mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-09-20 18:51:26 -03:00
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:
@@ -2113,6 +2113,8 @@ class ModelLibraryHandler:
|
|||||||
return "checkpoint"
|
return "checkpoint"
|
||||||
if normalized in {"embedding", "textualinversion"}:
|
if normalized in {"embedding", "textualinversion"}:
|
||||||
return "embedding"
|
return "embedding"
|
||||||
|
if normalized in VALID_OTHER_CIVITAI_TYPES:
|
||||||
|
return "other"
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def _get_scanner_for_type(self, model_type: str | 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()
|
return normalized_type, await self._service_registry.get_checkpoint_scanner()
|
||||||
if normalized_type == "embedding":
|
if normalized_type == "embedding":
|
||||||
return normalized_type, await self._service_registry.get_embedding_scanner()
|
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
|
return None, None
|
||||||
|
|
||||||
async def _get_download_history_service(self):
|
async def _get_download_history_service(self):
|
||||||
@@ -2214,6 +2223,11 @@ class ModelLibraryHandler:
|
|||||||
lora_scanner = await self._service_registry.get_lora_scanner()
|
lora_scanner = await self._service_registry.get_lora_scanner()
|
||||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||||
embedding_scanner = await self._service_registry.get_embedding_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:
|
if model_version_id_str:
|
||||||
try:
|
try:
|
||||||
@@ -2252,6 +2266,13 @@ class ModelLibraryHandler:
|
|||||||
exists = True
|
exists = True
|
||||||
model_type = "embedding"
|
model_type = "embedding"
|
||||||
matched_scanner = embedding_scanner
|
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:
|
if exists:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -2269,7 +2290,7 @@ class ModelLibraryHandler:
|
|||||||
history_service = await self._get_download_history_service()
|
history_service = await self._get_download_history_service()
|
||||||
has_been_downloaded = False
|
has_been_downloaded = False
|
||||||
history_type = None
|
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(
|
if await history_service.has_been_downloaded(
|
||||||
candidate_type,
|
candidate_type,
|
||||||
model_version_id,
|
model_version_id,
|
||||||
@@ -2291,6 +2312,7 @@ class ModelLibraryHandler:
|
|||||||
lora_versions = await lora_scanner.get_model_versions_by_id(model_id)
|
lora_versions = await lora_scanner.get_model_versions_by_id(model_id)
|
||||||
checkpoint_versions = []
|
checkpoint_versions = []
|
||||||
embedding_versions = []
|
embedding_versions = []
|
||||||
|
other_versions = []
|
||||||
if not lora_versions and checkpoint_scanner:
|
if not lora_versions and checkpoint_scanner:
|
||||||
checkpoint_versions = await checkpoint_scanner.get_model_versions_by_id(
|
checkpoint_versions = await checkpoint_scanner.get_model_versions_by_id(
|
||||||
model_id
|
model_id
|
||||||
@@ -2299,6 +2321,13 @@ class ModelLibraryHandler:
|
|||||||
embedding_versions = await embedding_scanner.get_model_versions_by_id(
|
embedding_versions = await embedding_scanner.get_model_versions_by_id(
|
||||||
model_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
|
model_type = None
|
||||||
versions = []
|
versions = []
|
||||||
@@ -2330,9 +2359,18 @@ class ModelLibraryHandler:
|
|||||||
"downloadedVersionIds": [],
|
"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()
|
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 = (
|
candidate_downloaded_version_ids = (
|
||||||
await history_service.get_downloaded_version_ids(
|
await history_service.get_downloaded_version_ids(
|
||||||
candidate_type,
|
candidate_type,
|
||||||
@@ -2387,6 +2425,11 @@ class ModelLibraryHandler:
|
|||||||
lora_scanner = await self._service_registry.get_lora_scanner()
|
lora_scanner = await self._service_registry.get_lora_scanner()
|
||||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||||
embedding_scanner = await self._service_registry.get_embedding_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]] = []
|
results: list[dict[str, Any]] = []
|
||||||
for model_id in model_ids:
|
for model_id in model_ids:
|
||||||
@@ -2422,6 +2465,17 @@ class ModelLibraryHandler:
|
|||||||
})
|
})
|
||||||
continue
|
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({
|
results.append({
|
||||||
"modelId": model_id,
|
"modelId": model_id,
|
||||||
"modelType": None,
|
"modelType": None,
|
||||||
|
|||||||
@@ -1534,6 +1534,9 @@ class DownloadManager:
|
|||||||
"Settings > Library before downloading VAE, upscaler, "
|
"Settings > Library before downloading VAE, upscaler, "
|
||||||
"text encoder or CLIP files."
|
"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"
|
model_type = "other"
|
||||||
else:
|
else:
|
||||||
@@ -1793,6 +1796,7 @@ class DownloadManager:
|
|||||||
f"disabled in settings. Please pick a destination "
|
f"disabled in settings. Please pick a destination "
|
||||||
f"folder explicitly instead of using default paths."
|
f"folder explicitly instead of using default paths."
|
||||||
),
|
),
|
||||||
|
"reason": "other_sub_type_disabled",
|
||||||
}
|
}
|
||||||
default_path = (
|
default_path = (
|
||||||
default_other_roots.get(other_sub_type)
|
default_other_roots.get(other_sub_type)
|
||||||
@@ -1805,17 +1809,20 @@ class DownloadManager:
|
|||||||
f"No default root configured for other-model "
|
f"No default root configured for other-model "
|
||||||
f"sub-type '{other_sub_type}'"
|
f"sub-type '{other_sub_type}'"
|
||||||
)
|
)
|
||||||
|
reason = "other_no_default_root"
|
||||||
else:
|
else:
|
||||||
detail = (
|
detail = (
|
||||||
"Could not determine the other-model sub-type "
|
"Could not determine the other-model sub-type "
|
||||||
"from the model metadata"
|
"from the model metadata"
|
||||||
)
|
)
|
||||||
|
reason = "other_sub_type_undecidable"
|
||||||
return {
|
return {
|
||||||
"success": False,
|
"success": False,
|
||||||
"error": (
|
"error": (
|
||||||
f"{detail}. Please pick a destination folder "
|
f"{detail}. Please pick a destination folder "
|
||||||
f"explicitly instead of using default paths."
|
f"explicitly instead of using default paths."
|
||||||
),
|
),
|
||||||
|
"reason": reason,
|
||||||
}
|
}
|
||||||
save_dir = default_path
|
save_dir = default_path
|
||||||
|
|
||||||
|
|||||||
@@ -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():
|
def test_create_handler_set_uses_provided_dependencies():
|
||||||
recorded_handlers: list[dict[str, Any]] = []
|
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 result["success"] is False
|
||||||
assert "disabled" in result["error"].lower()
|
assert "disabled" in result["error"].lower()
|
||||||
|
assert result["reason"] == "other_models_disabled"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -271,6 +272,7 @@ async def test_default_paths_reject_switched_off_sub_type(
|
|||||||
|
|
||||||
assert result["success"] is False
|
assert result["success"] is False
|
||||||
assert "disabled" in result["error"].lower()
|
assert "disabled" in result["error"].lower()
|
||||||
|
assert result["reason"] == "other_sub_type_disabled"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -516,6 +518,7 @@ async def test_default_paths_errors_when_sub_type_root_unconfigured(
|
|||||||
|
|
||||||
assert result["success"] is False
|
assert result["success"] is False
|
||||||
assert "controlnet" in result["error"]
|
assert "controlnet" in result["error"]
|
||||||
|
assert result["reason"] == "other_no_default_root"
|
||||||
assert execute_mock.await_count == 0
|
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 result["success"] is False
|
||||||
assert "sub-type" in result["error"]
|
assert "sub-type" in result["error"]
|
||||||
|
assert result["reason"] == "other_sub_type_undecidable"
|
||||||
assert execute_mock.await_count == 0
|
assert execute_mock.await_count == 0
|
||||||
|
|
||||||
|
|
||||||
@@ -559,6 +563,72 @@ async def test_civarchive_source_same_payload_shape(
|
|||||||
assert captured["model_type"] == "other"
|
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():
|
def test_build_metadata_for_resume_uses_other_metadata():
|
||||||
manager = DownloadManager()
|
manager = DownloadManager()
|
||||||
metadata = manager._build_metadata_for_resume(
|
metadata = manager._build_metadata_for_resume(
|
||||||
|
|||||||
Reference in New Issue
Block a user