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.
This commit is contained in:
Will Miao
2026-09-13 08:53:49 +08:00
parent 37f2cba72d
commit b5c1331911
4 changed files with 406 additions and 2 deletions
+273
View File
@@ -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]] = []
@@ -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(