feat(backend): CivitAI download support for other model types with subtype routing

This commit is contained in:
Will Miao
2026-09-12 15:56:47 +08:00
parent 57729375b6
commit f2a7297cb9
18 changed files with 1492 additions and 16 deletions
@@ -91,3 +91,61 @@ async def test_invalid_json_rejected():
FakeRequest(json.JSONDecodeError("bad", "", 0))
)
assert response.status == 400
@pytest.mark.asyncio
async def test_other_model_type_returns_sub_type():
handler = DownloadRoutingHandler()
response = await handler.get_download_routing(
FakeRequest({"model_type": "TextEncoder", "file_types": ["Model"]})
)
payload = json.loads(response.text)
assert response.status == 200
assert payload == {"success": True, "root_kind": "other", "sub_type": "text_encoder"}
@pytest.mark.asyncio
async def test_other_explicit_file_pick_wins():
handler = DownloadRoutingHandler()
response = await handler.get_download_routing(
FakeRequest(
{
"model_type": "Other",
"file_types": ["Model"],
"selected_file_type": "VAE",
}
)
)
payload = json.loads(response.text)
assert payload["root_kind"] == "other"
assert payload["sub_type"] == "vae"
@pytest.mark.asyncio
async def test_other_file_type_fallback_when_model_type_unmapped():
handler = DownloadRoutingHandler()
response = await handler.get_download_routing(
FakeRequest({"model_type": "Other", "file_types": ["Model", "Upscaler"]})
)
payload = json.loads(response.text)
assert payload["sub_type"] == "upscaler"
@pytest.mark.asyncio
async def test_other_undecidable_sub_type_is_none():
handler = DownloadRoutingHandler()
response = await handler.get_download_routing(
FakeRequest({"model_type": "Other", "file_types": ["Model"]})
)
payload = json.loads(response.text)
assert response.status == 200
assert payload == {"success": True, "root_kind": "other", "sub_type": None}
@pytest.mark.asyncio
async def test_other_invalid_selected_file_type_rejected():
handler = DownloadRoutingHandler()
response = await handler.get_download_routing(
FakeRequest({"model_type": "VAE", "selected_file_type": 123})
)
assert response.status == 400
+31 -3
View File
@@ -1133,8 +1133,8 @@ async def test_get_civitai_user_models_marks_library_versions():
},
{
"id": 4,
"name": "Unsupported",
"type": "Other",
"name": "VAE Model",
"type": "VAE",
"modelVersions": [
{
"id": 400,
@@ -1142,6 +1142,17 @@ async def test_get_civitai_user_models_marks_library_versions():
}
],
},
{
"id": 5,
"name": "Unsupported",
"type": "Wildcard",
"modelVersions": [
{
"id": 500,
"name": "v1",
}
],
},
]
provider = FakeUserModelsProvider(models)
@@ -1152,6 +1163,7 @@ async def test_get_civitai_user_models_marks_library_versions():
lora_scanner = FakeExistenceScanner({101})
checkpoint_scanner = FakeExistenceScanner()
embedding_scanner = FakeExistenceScanner({202})
other_scanner = FakeExistenceScanner({400})
async def lora_factory():
return lora_scanner
@@ -1162,11 +1174,15 @@ async def test_get_civitai_user_models_marks_library_versions():
async def embedding_factory():
return embedding_scanner
async def other_factory():
return other_scanner
handler = ModelLibraryHandler(
ServiceRegistryAdapter(
get_lora_scanner=lora_factory,
get_checkpoint_scanner=checkpoint_factory,
get_embedding_scanner=embedding_factory,
get_other_scanner=other_factory,
get_downloaded_version_history_service=lambda: fake_download_history_service_factory(),
),
metadata_provider_factory=provider_factory,
@@ -1240,6 +1256,18 @@ async def test_get_civitai_user_models_marks_library_versions():
"inLibrary": False,
"hasBeenDownloaded": False,
},
{
"modelId": 4,
"versionId": 400,
"modelName": "VAE Model",
"versionName": "v1",
"type": "VAE",
"tags": [],
"baseModel": None,
"thumbnailUrl": None,
"inLibrary": True,
"hasBeenDownloaded": False,
},
]
assert provider.received_usernames == ["pixel"]
@@ -1351,7 +1379,7 @@ async def test_get_civitai_user_models_returns_pagination_fields():
{
"id": 2,
"name": "Unsupported",
"type": "Other",
"type": "Wildcard",
"modelVersions": [{"id": 200, "name": "v1"}],
},
]
+47
View File
@@ -116,3 +116,50 @@ async def test_initialize_services_builds_other_model_service(monkeypatch):
assert isinstance(handler.service, OtherModelService)
assert handler.service.model_type == "other"
assert handler.service.scanner is sentinel_scanner
def test_roots_by_subtype_route_registered():
app = web.Application()
OtherRoutes().setup_routes(app)
registered = {(route.method, route.resource.canonical) for route in app.router.routes()}
assert ("GET", "/api/lm/other/roots_by_subtype") in registered
async def test_get_roots_by_subtype_aggregates_folder_keys(monkeypatch):
"""text_encoders and the legacy clip key both land under text_encoder."""
from py.config import config
monkeypatch.setattr(
config,
"other_folder_roots",
{
"vae": ["/models/vae", "/models/vae2"],
"text_encoders": ["/models/text_encoders"],
"clip": ["/models/clip_legacy"],
"upscale_models": ["/models/upscale"],
"unknown_key": ["/models/ignored"],
},
)
response = await OtherRoutes().get_roots_by_subtype(DummyRequest())
payload = json.loads(response.text)
assert payload["success"] is True
assert payload["roots_by_subtype"] == {
"vae": ["/models/vae", "/models/vae2"],
"text_encoder": ["/models/text_encoders", "/models/clip_legacy"],
"upscaler": ["/models/upscale"],
}
async def test_get_roots_by_subtype_empty_config(monkeypatch):
from py.config import config
monkeypatch.setattr(config, "other_folder_roots", {})
response = await OtherRoutes().get_roots_by_subtype(DummyRequest())
payload = json.loads(response.text)
assert payload == {"success": True, "roots_by_subtype": {}}