From cf0da2e6a33754e9f1f933ea5d15aa9d5d01d3bc Mon Sep 17 00:00:00 2001 From: SEPURI-SAI-KRISHNA Date: Tue, 4 Aug 2026 12:56:41 +0530 Subject: [PATCH] fix(models): honor model_name on /models/load and /models/unload (fixes #977) --- backend/models.py | 13 ++ backend/routes/models.py | 75 ++++++++-- backend/tests/test_models_load_by_name.py | 134 ++++++++++++++++++ .../docs/developer/model-management.mdx | 6 +- 4 files changed, 218 insertions(+), 10 deletions(-) create mode 100644 backend/tests/test_models_load_by_name.py diff --git a/backend/models.py b/backend/models.py index a42f3b7d..0f6d2dac 100644 --- a/backend/models.py +++ b/backend/models.py @@ -489,6 +489,19 @@ class ModelDownloadRequest(BaseModel): model_name: str +class ModelLoadRequest(BaseModel): + """Request model for loading or unloading a model by name. + + ``model_name`` is one of the ids returned by ``GET /models/status`` + (e.g. ``"kokoro"``, ``"qwen-tts-0.6B"``, ``"whisper-turbo"``). + ``model_size`` is the legacy Qwen-only selector, kept so existing + ``POST /models/load?model_size=0.6B`` callers keep working. + """ + + model_name: Optional[str] = None + model_size: Optional[str] = None + + class ModelMigrateRequest(BaseModel): """Request model for migrating models to a new directory.""" diff --git a/backend/routes/models.py b/backend/routes/models.py index fcad0d35..600e0221 100644 --- a/backend/routes/models.py +++ b/backend/routes/models.py @@ -47,29 +47,88 @@ def _copy_with_progress(src: Path, dst: Path, progress_manager, copied_so_far: i return copied_so_far +DEFAULT_QWEN_MODEL_SIZE = "1.7B" + + +def _resolve_model_config(model_name: str): + """Look up a model config by name or raise a 400 listing the valid ids.""" + from ..backends import get_all_model_configs, get_model_config + + config = get_model_config(model_name) + if config is None: + known = sorted(cfg.model_name for cfg in get_all_model_configs()) + raise HTTPException( + status_code=400, + detail=f"Unknown model: {model_name}. Available: {known}", + ) + return config + + @router.post("/models/load") -async def load_model(model_size: str = "1.7B"): - """Manually load TTS model.""" +async def load_model( + request: models.ModelLoadRequest | None = None, + model_size: str | None = None, +): + """Load a model into memory. + + Pass ``{"model_name": "kokoro"}`` to target any registered engine. Without + a name this falls back to the default Qwen TTS backend, selected by the + ``model_size`` query parameter, which is what callers did before + ``model_name`` existed. + """ + from ..backends import get_model_load_func from ..services import tts + requested_name = request.model_name if request else None + requested_size = (request.model_size if request else None) or model_size + + if requested_name: + config = _resolve_model_config(requested_name) + try: + result = get_model_load_func(config)() + if asyncio.iscoroutine(result): + await result + except Exception as e: + raise HTTPException(status_code=500, detail=str(e)) from e + return {"message": f"Model {config.model_name} loaded successfully"} + + size = requested_size or DEFAULT_QWEN_MODEL_SIZE try: tts_model = tts.get_tts_model() - await tts_model.load_model_async(model_size) - return {"message": f"Model {model_size} loaded successfully"} + await tts_model.load_model_async(size) + return {"message": f"Model {size} loaded successfully"} except Exception as e: - raise HTTPException(status_code=500, detail=str(e)) + raise HTTPException(status_code=500, detail=str(e)) from e @router.post("/models/unload") -async def unload_model(): - """Unload the default Qwen TTS model to free memory.""" +async def unload_model(request: models.ModelLoadRequest | None = None): + """Unload a model from memory. + + Pass ``{"model_name": "chatterbox-tts"}`` to target a specific engine — + same behaviour as ``POST /models/{model_name}/unload``. Without a name + this unloads the default Qwen TTS model, as it always has. + """ + from ..backends import unload_model_by_config from ..services import tts + requested_name = request.model_name if request else None + + if requested_name: + config = _resolve_model_config(requested_name) + try: + was_loaded = unload_model_by_config(config) + except Exception as e: + raise HTTPException(status_code=500, detail=str(e)) from e + if not was_loaded: + return {"message": f"Model {config.model_name} is not loaded"} + return {"message": f"Model {config.model_name} unloaded successfully"} + try: tts.unload_tts_model() return {"message": "Model unloaded successfully"} except Exception as e: - raise HTTPException(status_code=500, detail=str(e)) + raise HTTPException(status_code=500, detail=str(e)) from e @router.post("/models/{model_name}/unload") diff --git a/backend/tests/test_models_load_by_name.py b/backend/tests/test_models_load_by_name.py new file mode 100644 index 00000000..4773c76a --- /dev/null +++ b/backend/tests/test_models_load_by_name.py @@ -0,0 +1,134 @@ +""" +Regression tests for POST /models/load and POST /models/unload targeting a +model by name (issue #977). + +``/models/load`` only accepted a ``model_size`` query parameter and always +dispatched to the default Qwen TTS backend. A caller following the documented +form — ``POST /models/load {"model_name": "kokoro"}`` — got a 200 back while +the server silently started a 3.6 GB Qwen 1.7B download instead. ``/models/unload`` +had the same shape: the documented ``{"model_name": ...}`` body was ignored and +the default Qwen model was unloaded regardless. + +Usage: + python -m pytest backend/tests/test_models_load_by_name.py -v +""" + +import sys +from pathlib import Path + +import pytest +from fastapi import FastAPI +from starlette.testclient import TestClient + +# Repo root on sys.path so ``backend`` imports as a package (the routes use +# package-relative imports). +sys.path.insert(0, str(Path(__file__).parent.parent.parent)) + +from backend import backends +from backend.routes.models import router as models_router +from backend.services import tts + + +@pytest.fixture +def loaded(monkeypatch): + """Record what the route asked the registry to load/unload.""" + calls: dict[str, list] = {"by_config": [], "qwen_size": [], "qwen_unload": 0} + + def fake_load_func(config): + async def _load(): + calls["by_config"].append(config.model_name) + + return _load + + def fake_unload_by_config(config): + calls["by_config"].append(f"unload:{config.model_name}") + return True + + class FakeQwenBackend: + async def load_model_async(self, model_size): + calls["qwen_size"].append(model_size) + + def fake_unload_tts_model(): + calls["qwen_unload"] += 1 + + monkeypatch.setattr(backends, "get_model_load_func", fake_load_func) + monkeypatch.setattr(backends, "unload_model_by_config", fake_unload_by_config) + monkeypatch.setattr(tts, "get_tts_model", lambda: FakeQwenBackend()) + monkeypatch.setattr(tts, "unload_tts_model", fake_unload_tts_model) + return calls + + +@pytest.fixture +def client(): + app = FastAPI() + app.include_router(models_router) + with TestClient(app) as test_client: + yield test_client + + +def test_load_by_model_name_targets_that_model(client, loaded): + """The whole point of #977 — asking for kokoro must not load Qwen.""" + response = client.post("/models/load", json={"model_name": "kokoro"}) + + assert response.status_code == 200 + assert loaded["by_config"] == ["kokoro"] + assert loaded["qwen_size"] == [], "must not touch the default Qwen backend" + assert "kokoro" in response.json()["message"] + + +@pytest.mark.parametrize( + "model_name", + ["qwen-tts-0.6B", "chatterbox-tts", "whisper-turbo", "qwen3-0.6b"], +) +def test_load_dispatches_across_engine_families(client, loaded, model_name): + """TTS, STT, and LLM entries all resolve through the same registry.""" + assert client.post("/models/load", json={"model_name": model_name}).status_code == 200 + assert loaded["by_config"] == [model_name] + + +def test_load_unknown_model_is_rejected(client, loaded): + """An unknown name is a 400 that names the valid ids, not a silent Qwen load.""" + response = client.post("/models/load", json={"model_name": "not-a-model"}) + + assert response.status_code == 400 + assert "kokoro" in response.json()["detail"] + assert loaded["by_config"] == [] + assert loaded["qwen_size"] == [] + + +def test_load_without_body_keeps_the_qwen_default(client, loaded): + """Pre-existing callers post no body at all.""" + assert client.post("/models/load").status_code == 200 + assert loaded["qwen_size"] == ["1.7B"] + + +def test_load_still_honors_the_model_size_query_param(client, loaded): + """The legacy ?model_size= form must keep selecting the Qwen variant.""" + assert client.post("/models/load?model_size=0.6B").status_code == 200 + assert loaded["qwen_size"] == ["0.6B"] + + +def test_load_model_size_in_body_also_works(client, loaded): + assert client.post("/models/load", json={"model_size": "0.6B"}).status_code == 200 + assert loaded["qwen_size"] == ["0.6B"] + + +def test_unload_by_model_name_targets_that_model(client, loaded): + response = client.post("/models/unload", json={"model_name": "chatterbox-tts"}) + + assert response.status_code == 200 + assert loaded["by_config"] == ["unload:chatterbox-tts"] + assert loaded["qwen_unload"] == 0, "must not unload the default Qwen model instead" + + +def test_unload_without_body_keeps_the_qwen_default(client, loaded): + assert client.post("/models/unload").status_code == 200 + assert loaded["qwen_unload"] == 1 + assert loaded["by_config"] == [] + + +def test_unload_unknown_model_is_rejected(client, loaded): + response = client.post("/models/unload", json={"model_name": "not-a-model"}) + + assert response.status_code == 400 + assert loaded["qwen_unload"] == 0 diff --git a/docs/content/docs/developer/model-management.mdx b/docs/content/docs/developer/model-management.mdx index 93c87e57..757e24e8 100644 --- a/docs/content/docs/developer/model-management.mdx +++ b/docs/content/docs/developer/model-management.mdx @@ -117,7 +117,9 @@ POST /models/load } ``` -The route looks up the config, dispatches to `get_model_load_func(config)`, and returns once the model is ready. +The route looks up the config, dispatches to `get_model_load_func(config)`, and returns once the model is ready. `model_name` is any id from `GET /models/status` — TTS, Whisper, and LLM entries all resolve through the same registry. + +Omitting `model_name` falls back to the default Qwen TTS backend, selected by an optional `model_size` (`POST /models/load?model_size=0.6B`). This is the pre-`model_name` form and still works. ### Unload @@ -128,7 +130,7 @@ POST /models/unload } ``` -Calls `unload_model_by_config(config)`, which routes to the right backend's `unload_model()` and frees GPU memory. +Calls `unload_model_by_config(config)`, which routes to the right backend's `unload_model()` and frees GPU memory. `POST /models/{name}/unload` is equivalent. Omitting `model_name` unloads the default Qwen TTS model. ### Download