From b5c1331911bdcdd04954ea363682a6698a39511f Mon Sep 17 00:00:00 2001 From: Will Miao Date: Sun, 13 Sep 2026 08:53:49 +0800 Subject: [PATCH] feat(backend): make model existence checks other-aware Implements the backend slice (B1-B7) of lm-civitai-extension/docs/other-models-support.md, which lets the companion browser extension detect, badge and download the opt-in Other Models types (VAE / upscaler / text encoder / CLIP vision / ControlNet). ModelLibraryHandler: - _normalize_model_type() learns the CivitAI other aliases (vae, upscaler, textencoder, clip, clipvision, controlnet, other) and maps them to "other". - _get_scanner_for_type() resolves "other" through the other scanner, but only while enable_other_models is on, so model-versions-status and model-version-download-status keep their legacy 400 when the feature is off. - check_model_exists() / check_models_exist() consult the other scanner last (lora -> checkpoint -> embedding -> other) and report modelType "other". With the feature disabled both endpoints stay byte-identical to before and the other scanner is never touched. DownloadManager: - The four other-type default-path failures now carry a machine-readable "reason" (contract C4): other_models_disabled, other_sub_type_disabled, other_no_default_root, other_sub_type_undecidable. The user-facing "error" strings are unchanged; the key is additive and reaches the client because both download endpoints pass the result dict through verbatim. Tests cover the opt-in on/off branches for both existence endpoints, mixed lora + other ids in the batch endpoint, the CivitAI alias acceptance and the 400 regression for unknown types, and the exact reason/error pairs for all four download failure modes. --- py/routes/handlers/misc_handlers.py | 58 +++- py/services/download_manager.py | 7 + tests/routes/test_misc_routes.py | 273 ++++++++++++++++++ tests/services/test_download_manager_other.py | 70 +++++ 4 files changed, 406 insertions(+), 2 deletions(-) diff --git a/py/routes/handlers/misc_handlers.py b/py/routes/handlers/misc_handlers.py index 711e4eb9..e88c198f 100644 --- a/py/routes/handlers/misc_handlers.py +++ b/py/routes/handlers/misc_handlers.py @@ -2113,6 +2113,8 @@ class ModelLibraryHandler: return "checkpoint" if normalized in {"embedding", "textualinversion"}: return "embedding" + if normalized in VALID_OTHER_CIVITAI_TYPES: + return "other" return None async def _get_scanner_for_type(self, model_type: str | None): @@ -2123,6 +2125,13 @@ class ModelLibraryHandler: return normalized_type, await self._service_registry.get_checkpoint_scanner() if normalized_type == "embedding": return normalized_type, await self._service_registry.get_embedding_scanner() + if normalized_type == "other": + # Opt-in feature: the other scanner only resolves while the master + # switch is on, so callers keep returning the legacy "required" + # error (400) when it is off. + if not get_settings_manager().is_other_models_enabled(): + return None, None + return normalized_type, await self._service_registry.get_other_scanner() return None, None async def _get_download_history_service(self): @@ -2214,6 +2223,11 @@ class ModelLibraryHandler: lora_scanner = await self._service_registry.get_lora_scanner() checkpoint_scanner = await self._service_registry.get_checkpoint_scanner() embedding_scanner = await self._service_registry.get_embedding_scanner() + # Opt-in: probe the other scanner only while Other Models is enabled, + # so the disabled behaviour stays byte-identical to the legacy one. + other_scanner = None + if get_settings_manager().is_other_models_enabled(): + other_scanner = await self._service_registry.get_other_scanner() if model_version_id_str: try: @@ -2252,6 +2266,13 @@ class ModelLibraryHandler: exists = True model_type = "embedding" matched_scanner = embedding_scanner + elif ( + other_scanner + and await other_scanner.check_model_version_exists(model_version_id) + ): + exists = True + model_type = "other" + matched_scanner = other_scanner if exists: return web.json_response( @@ -2269,7 +2290,7 @@ class ModelLibraryHandler: history_service = await self._get_download_history_service() has_been_downloaded = False history_type = None - for candidate_type in ("lora", "checkpoint", "embedding"): + for candidate_type in ("lora", "checkpoint", "embedding", "other"): if await history_service.has_been_downloaded( candidate_type, model_version_id, @@ -2291,6 +2312,7 @@ class ModelLibraryHandler: lora_versions = await lora_scanner.get_model_versions_by_id(model_id) checkpoint_versions = [] embedding_versions = [] + other_versions = [] if not lora_versions and checkpoint_scanner: checkpoint_versions = await checkpoint_scanner.get_model_versions_by_id( model_id @@ -2299,6 +2321,13 @@ class ModelLibraryHandler: embedding_versions = await embedding_scanner.get_model_versions_by_id( model_id ) + if ( + not lora_versions + and not checkpoint_versions + and not embedding_versions + and other_scanner + ): + other_versions = await other_scanner.get_model_versions_by_id(model_id) model_type = None versions = [] @@ -2330,9 +2359,18 @@ class ModelLibraryHandler: "downloadedVersionIds": [], } ) + if other_versions: + return web.json_response( + { + "success": True, + "modelType": "other", + "versions": self._with_downloaded_flag(other_versions), + "downloadedVersionIds": [], + } + ) history_service = await self._get_download_history_service() - for candidate_type in ("lora", "checkpoint", "embedding"): + for candidate_type in ("lora", "checkpoint", "embedding", "other"): candidate_downloaded_version_ids = ( await history_service.get_downloaded_version_ids( candidate_type, @@ -2387,6 +2425,11 @@ class ModelLibraryHandler: lora_scanner = await self._service_registry.get_lora_scanner() checkpoint_scanner = await self._service_registry.get_checkpoint_scanner() embedding_scanner = await self._service_registry.get_embedding_scanner() + # Opt-in: keep the other probe last so model cards for lora / + # checkpoint / embedding ids are unaffected by the extra scanner. + other_scanner = None + if get_settings_manager().is_other_models_enabled(): + other_scanner = await self._service_registry.get_other_scanner() results: list[dict[str, Any]] = [] for model_id in model_ids: @@ -2422,6 +2465,17 @@ class ModelLibraryHandler: }) continue + if other_scanner: + other_versions = await other_scanner.get_model_versions_by_id(model_id) + if other_versions: + results.append({ + "modelId": model_id, + "modelType": "other", + "versions": self._with_downloaded_flag(other_versions), + "downloadedVersionIds": [], + }) + continue + results.append({ "modelId": model_id, "modelType": None, diff --git a/py/services/download_manager.py b/py/services/download_manager.py index 0bae9e75..7b95fe69 100644 --- a/py/services/download_manager.py +++ b/py/services/download_manager.py @@ -1534,6 +1534,9 @@ class DownloadManager: "Settings > Library before downloading VAE, upscaler, " "text encoder or CLIP files." ), + # Machine-readable failure code consumed by the companion + # browser extension (docs/other-models-support.md C4). + "reason": "other_models_disabled", } model_type = "other" else: @@ -1793,6 +1796,7 @@ class DownloadManager: f"disabled in settings. Please pick a destination " f"folder explicitly instead of using default paths." ), + "reason": "other_sub_type_disabled", } default_path = ( default_other_roots.get(other_sub_type) @@ -1805,17 +1809,20 @@ class DownloadManager: f"No default root configured for other-model " f"sub-type '{other_sub_type}'" ) + reason = "other_no_default_root" else: detail = ( "Could not determine the other-model sub-type " "from the model metadata" ) + reason = "other_sub_type_undecidable" return { "success": False, "error": ( f"{detail}. Please pick a destination folder " f"explicitly instead of using default paths." ), + "reason": reason, } save_dir = default_path diff --git a/tests/routes/test_misc_routes.py b/tests/routes/test_misc_routes.py index 62908587..2c2b49f5 100644 --- a/tests/routes/test_misc_routes.py +++ b/tests/routes/test_misc_routes.py @@ -1775,6 +1775,279 @@ async def test_model_version_download_status_endpoints(): } +class OtherRecordingScanner: + """Other-scanner stub recording both probe kinds.""" + + def __init__(self, versions_by_model_id=None, version_ids=()): + self.versions_by_model_id = versions_by_model_id or {} + self.version_ids = set(version_ids) + self.version_calls: list[int] = [] + + async def get_model_versions_by_id(self, model_id): + self.version_calls.append(model_id) + return list(self.versions_by_model_id.get(model_id, [])) + + async def check_model_version_exists(self, version_id): + return version_id in self.version_ids + + +def _set_other_models_enabled(enabled: bool) -> None: + from py.services.settings_manager import get_settings_manager + + get_settings_manager().set("enable_other_models", enabled) + + +@pytest.mark.asyncio +async def test_check_model_exists_with_other_models_enabled(): + """An other-type version resolves through the other scanner when opted in.""" + _set_other_models_enabled(True) + other_scanner = OtherRecordingScanner(version_ids={400}) + + async def other_factory(): + return other_scanner + + handler = ModelLibraryHandler( + ServiceRegistryAdapter( + get_lora_scanner=fake_scanner_factory, + get_checkpoint_scanner=fake_scanner_factory, + get_embedding_scanner=fake_scanner_factory, + get_other_scanner=other_factory, + get_downloaded_version_history_service=fake_download_history_service_factory, + ), + metadata_provider_factory=fake_metadata_provider_factory, + ) + + response = await handler.check_model_exists( + FakeRequest(query={"modelId": "5", "modelVersionId": "400"}) # pyright: ignore[reportArgumentType] + ) + payload = _json_payload(response) + + assert payload == { + "success": True, + "exists": True, + "modelType": "other", + "hasBeenDownloaded": False, + "downloadedFiles": [], + } + + +@pytest.mark.asyncio +async def test_check_model_exists_skips_other_scanner_when_disabled(): + """Opt-out stays byte-identical: no other probe, modelType stays null.""" + _set_other_models_enabled(False) + other_scanner = OtherRecordingScanner(versions_by_model_id={5: [{"versionId": 400}]}) + + async def other_factory(): + return other_scanner + + handler = ModelLibraryHandler( + ServiceRegistryAdapter( + get_lora_scanner=fake_scanner_factory, + get_checkpoint_scanner=fake_scanner_factory, + get_embedding_scanner=fake_scanner_factory, + get_other_scanner=other_factory, + get_downloaded_version_history_service=fake_download_history_service_factory, + ), + metadata_provider_factory=fake_metadata_provider_factory, + ) + + response = await handler.check_model_exists( + FakeRequest(query={"modelId": "5"}) # pyright: ignore[reportArgumentType] + ) + payload = _json_payload(response) + + assert payload == { + "success": True, + "modelType": None, + "versions": [], + "downloadedVersionIds": [], + } + assert other_scanner.version_calls == [] + + +@pytest.mark.asyncio +async def test_check_models_exist_resolves_other_ids(): + """Mixed lora + vae ids resolve independently in the batch endpoint.""" + _set_other_models_enabled(True) + lora_scanner = OtherRecordingScanner( + versions_by_model_id={5: [{"versionId": 11, "name": "v1"}]} + ) + other_scanner = OtherRecordingScanner( + versions_by_model_id={6: [{"versionId": 400, "name": "vae-v1"}]} + ) + + async def lora_factory(): + return lora_scanner + + async def other_factory(): + return other_scanner + + handler = ModelLibraryHandler( + ServiceRegistryAdapter( + get_lora_scanner=lora_factory, + get_checkpoint_scanner=fake_scanner_factory, + get_embedding_scanner=fake_scanner_factory, + get_other_scanner=other_factory, + get_downloaded_version_history_service=fake_download_history_service_factory, + ), + metadata_provider_factory=fake_metadata_provider_factory, + ) + + response = await handler.check_models_exist( + FakeRequest(query={"modelIds": "5,6"}) # pyright: ignore[reportArgumentType] + ) + payload = _json_payload(response) + + assert payload["success"] is True + results = {item["modelId"]: item for item in payload["results"]} + assert results[5]["modelType"] == "lora" + assert results[5]["versions"] == [ + {"versionId": 11, "name": "v1", "hasBeenDownloaded": True} + ] + assert results[6]["modelType"] == "other" + assert results[6]["versions"] == [ + {"versionId": 400, "name": "vae-v1", "hasBeenDownloaded": True} + ] + + +@pytest.mark.asyncio +async def test_check_models_exist_ignores_other_scanner_when_disabled(): + _set_other_models_enabled(False) + other_scanner = OtherRecordingScanner(versions_by_model_id={5: [{"versionId": 400}]}) + + async def other_factory(): + return other_scanner + + handler = ModelLibraryHandler( + ServiceRegistryAdapter( + get_lora_scanner=fake_scanner_factory, + get_checkpoint_scanner=fake_scanner_factory, + get_embedding_scanner=fake_scanner_factory, + get_other_scanner=other_factory, + get_downloaded_version_history_service=fake_download_history_service_factory, + ), + metadata_provider_factory=fake_metadata_provider_factory, + ) + + response = await handler.check_models_exist( + FakeRequest(query={"modelIds": "6"}) # pyright: ignore[reportArgumentType] + ) + payload = _json_payload(response) + + assert payload["results"] == [ + { + "modelId": 6, + "modelType": None, + "versions": [], + "downloadedVersionIds": [], + } + ] + assert other_scanner.version_calls == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model_type", ["vae", "textencoder", "clip", "other"]) +async def test_model_version_download_status_accepts_other_types_when_enabled( + model_type, +): + _set_other_models_enabled(True) + history_service = FakeDownloadHistoryService() + + async def history_factory(): + return history_service + + handler = ModelLibraryHandler( + ServiceRegistryAdapter( + get_lora_scanner=fake_scanner_factory, + get_checkpoint_scanner=fake_scanner_factory, + get_embedding_scanner=fake_scanner_factory, + get_other_scanner=fake_scanner_factory, + get_downloaded_version_history_service=history_factory, + ), + metadata_provider_factory=fake_metadata_provider_factory, + ) + + response = await handler.get_model_version_download_status( + FakeRequest( # pyright: ignore[reportArgumentType] + query={"modelType": model_type, "modelVersionId": "400"} + ) + ) + payload = _json_payload(response) + + assert response.status == 200 + assert payload == { + "success": True, + "modelType": "other", + "modelVersionId": 400, + "hasBeenDownloaded": False, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model_type", ["vae", "textencoder", "clip", "other"]) +async def test_model_version_download_status_rejects_other_types_when_disabled( + model_type, +): + _set_other_models_enabled(False) + + async def history_factory(): + return FakeDownloadHistoryService() + + handler = ModelLibraryHandler( + ServiceRegistryAdapter( + get_lora_scanner=fake_scanner_factory, + get_checkpoint_scanner=fake_scanner_factory, + get_embedding_scanner=fake_scanner_factory, + get_other_scanner=fake_scanner_factory, + get_downloaded_version_history_service=history_factory, + ), + metadata_provider_factory=fake_metadata_provider_factory, + ) + + response = await handler.get_model_version_download_status( + FakeRequest( # pyright: ignore[reportArgumentType] + query={"modelType": model_type, "modelVersionId": "400"} + ) + ) + payload = _json_payload(response) + + assert response.status == 400 + assert payload == { + "success": False, + "error": "Parameter modelType is required", + } + + +@pytest.mark.asyncio +async def test_model_version_download_status_rejects_unknown_type(): + """Regression: garbage modelType keeps the legacy 400 error.""" + _set_other_models_enabled(True) + + handler = ModelLibraryHandler( + ServiceRegistryAdapter( + get_lora_scanner=fake_scanner_factory, + get_checkpoint_scanner=fake_scanner_factory, + get_embedding_scanner=fake_scanner_factory, + get_other_scanner=fake_scanner_factory, + get_downloaded_version_history_service=fake_download_history_service_factory, + ), + metadata_provider_factory=fake_metadata_provider_factory, + ) + + response = await handler.get_model_version_download_status( + FakeRequest( # pyright: ignore[reportArgumentType] + query={"modelType": "garbage", "modelVersionId": "400"} + ) + ) + payload = _json_payload(response) + + assert response.status == 400 + assert payload == { + "success": False, + "error": "Parameter modelType is required", + } + + def test_create_handler_set_uses_provided_dependencies(): recorded_handlers: list[dict[str, Any]] = [] diff --git a/tests/services/test_download_manager_other.py b/tests/services/test_download_manager_other.py index 12367523..6ddc3b5e 100644 --- a/tests/services/test_download_manager_other.py +++ b/tests/services/test_download_manager_other.py @@ -254,6 +254,7 @@ async def test_download_rejects_other_when_feature_disabled( assert result["success"] is False assert "disabled" in result["error"].lower() + assert result["reason"] == "other_models_disabled" @pytest.mark.asyncio @@ -271,6 +272,7 @@ async def test_default_paths_reject_switched_off_sub_type( assert result["success"] is False assert "disabled" in result["error"].lower() + assert result["reason"] == "other_sub_type_disabled" @pytest.mark.asyncio @@ -516,6 +518,7 @@ async def test_default_paths_errors_when_sub_type_root_unconfigured( assert result["success"] is False assert "controlnet" in result["error"] + assert result["reason"] == "other_no_default_root" assert execute_mock.await_count == 0 @@ -537,6 +540,7 @@ async def test_default_paths_errors_when_sub_type_undecidable( assert result["success"] is False assert "sub-type" in result["error"] + assert result["reason"] == "other_sub_type_undecidable" assert execute_mock.await_count == 0 @@ -559,6 +563,72 @@ async def test_civarchive_source_same_payload_shape( assert captured["model_type"] == "other" +@pytest.mark.asyncio +async def test_other_failure_reasons_are_machine_readable( + monkeypatch, scanners, metadata_provider, tmp_path +): + """Contract C4: every other-type default-path failure carries a ``reason``. + + The companion browser extension binds to ``reason`` and only falls back to + substring matching for backends that predate the field, so the exact values + below must not drift. + """ + expected_reasons = { + "disabled": "other_models_disabled", + "sub_type_disabled": "other_sub_type_disabled", + "no_default_root": "other_no_default_root", + "undecidable": "other_sub_type_undecidable", + } + reasons: dict[str, str] = {} + + manager = DownloadManager() + + # 1. Master switch off. + metadata_provider.payload = _other_payload("VAE") + get_settings_manager().settings["enable_other_models"] = False + disabled = await manager.download_from_civitai( + model_version_id=99, save_dir=str(tmp_path) + ) + reasons["disabled"] = disabled["reason"] + assert disabled["error"].strip() + + get_settings_manager().settings["enable_other_models"] = True + + # 2. Resolved sub_type not enabled. + get_settings_manager().settings["enabled_other_sub_types"] = ["upscaler"] + sub_type_disabled = await manager.download_from_civitai( + model_version_id=99, use_default_paths=True + ) + reasons["sub_type_disabled"] = sub_type_disabled["reason"] + assert sub_type_disabled["error"].strip() + + get_settings_manager().settings["enabled_other_sub_types"] = [ + "vae", + "upscaler", + "text_encoder", + "clip_vision", + "controlnet", + ] + + # 3. Sub_type resolved but no default root configured. + metadata_provider.payload = _other_payload("Controlnet") + no_default_root = await manager.download_from_civitai( + model_version_id=99, use_default_paths=True + ) + reasons["no_default_root"] = no_default_root["reason"] + assert no_default_root["error"].strip() + + # 4. Neither model.type nor file types map to a sub_type. + metadata_provider.payload = _other_payload("Other") + undecidable = await manager.download_from_civitai( + model_version_id=99, use_default_paths=True + ) + reasons["undecidable"] = undecidable["reason"] + assert undecidable["error"].strip() + + assert reasons == expected_reasons + + def test_build_metadata_for_resume_uses_other_metadata(): manager = DownloadManager() metadata = manager._build_metadata_for_resume(