mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-09-21 03:01:27 -03:00
feat(backend): CivitAI download support for other model types with subtype routing
This commit is contained in:
@@ -0,0 +1,565 @@
|
||||
"""DownloadManager support for the "other" model type (VAE, upscaler, ...).
|
||||
|
||||
Covers the Phase-2 scatter points from docs/plans/other-models-page.md §9.1:
|
||||
type map acceptance, existence gates consulting the other scanner (never
|
||||
falling through to the lora scanner), per-sub_type default roots, resume
|
||||
metadata and the archive extension set.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
from py.services import aria2_transfer_state
|
||||
from py.services import download_manager
|
||||
from py.services.download_manager import DownloadManager
|
||||
from py.services.service_registry import ServiceRegistry
|
||||
from py.services.settings_manager import SettingsManager, get_settings_manager
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_download_manager():
|
||||
"""Ensure each test operates on a fresh singleton."""
|
||||
DownloadManager._instance = None
|
||||
yield
|
||||
DownloadManager._instance = None
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def isolate_settings(monkeypatch, tmp_path):
|
||||
"""Point settings writes at a temporary directory to avoid touching real files."""
|
||||
manager = get_settings_manager()
|
||||
default_settings = manager._get_default_settings()
|
||||
default_settings.update(
|
||||
{
|
||||
"default_lora_root": str(tmp_path / "loras"),
|
||||
"default_checkpoint_root": str(tmp_path / "checkpoints"),
|
||||
"default_embedding_root": str(tmp_path / "embeddings"),
|
||||
"default_other_roots": {
|
||||
"vae": str(tmp_path / "vae"),
|
||||
"upscaler": str(tmp_path / "upscale_models"),
|
||||
"text_encoder": str(tmp_path / "text_encoders"),
|
||||
"clip_vision": str(tmp_path / "clip_vision"),
|
||||
},
|
||||
"download_path_templates": {
|
||||
"lora": "{base_model}/{first_tag}",
|
||||
"checkpoint": "{base_model}/{first_tag}",
|
||||
"embedding": "{base_model}/{first_tag}",
|
||||
"other": "",
|
||||
},
|
||||
"skip_previously_downloaded_model_versions": False,
|
||||
"download_skip_base_models": [],
|
||||
}
|
||||
)
|
||||
monkeypatch.setattr(manager, "settings", default_settings)
|
||||
monkeypatch.setattr(SettingsManager, "_save_settings", lambda self: None)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def isolate_aria2_state(monkeypatch, tmp_path):
|
||||
state_path = tmp_path / "cache" / "aria2" / "downloads.json"
|
||||
monkeypatch.setattr(
|
||||
aria2_transfer_state,
|
||||
"get_aria2_state_path",
|
||||
lambda: str(state_path),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def stub_metadata(monkeypatch):
|
||||
class _StubMetadata:
|
||||
def __init__(self, save_path: str):
|
||||
self.file_path = save_path
|
||||
self.sha256 = "sha256"
|
||||
self.file_name = Path(save_path).stem
|
||||
|
||||
def _make_class(name):
|
||||
@staticmethod
|
||||
def from_civitai_info(_version_info, _file_info, save_path):
|
||||
metadata = _StubMetadata(save_path)
|
||||
metadata.metadata_class = name
|
||||
return metadata
|
||||
|
||||
return type(name, (), {"from_civitai_info": from_civitai_info})
|
||||
|
||||
monkeypatch.setattr(download_manager, "LoraMetadata", _make_class("LoraMetadata"))
|
||||
monkeypatch.setattr(
|
||||
download_manager, "CheckpointMetadata", _make_class("CheckpointMetadata")
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
download_manager, "EmbeddingMetadata", _make_class("EmbeddingMetadata")
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
download_manager, "OtherModelMetadata", _make_class("OtherModelMetadata")
|
||||
)
|
||||
|
||||
|
||||
class DummyScanner:
|
||||
def __init__(self, exists: bool = False, raw_data=None):
|
||||
self.exists = exists
|
||||
self.calls = []
|
||||
self._cache = SimpleNamespace(raw_data=list(raw_data or []))
|
||||
|
||||
async def check_model_version_exists(self, version_id):
|
||||
self.calls.append(version_id)
|
||||
return self.exists
|
||||
|
||||
async def get_cached_data(self):
|
||||
return self._cache
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def scanners(monkeypatch):
|
||||
lora_scanner = DummyScanner()
|
||||
checkpoint_scanner = DummyScanner()
|
||||
embedding_scanner = DummyScanner()
|
||||
other_scanner = DummyScanner()
|
||||
|
||||
monkeypatch.setattr(
|
||||
ServiceRegistry, "get_lora_scanner", AsyncMock(return_value=lora_scanner)
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ServiceRegistry,
|
||||
"get_checkpoint_scanner",
|
||||
AsyncMock(return_value=checkpoint_scanner),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ServiceRegistry,
|
||||
"get_embedding_scanner",
|
||||
AsyncMock(return_value=embedding_scanner),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ServiceRegistry,
|
||||
"get_other_scanner",
|
||||
AsyncMock(return_value=other_scanner),
|
||||
)
|
||||
|
||||
return SimpleNamespace(
|
||||
lora=lora_scanner,
|
||||
checkpoint=checkpoint_scanner,
|
||||
embedding=embedding_scanner,
|
||||
other=other_scanner,
|
||||
)
|
||||
|
||||
|
||||
def _other_payload(civitai_type: str, *, files=None) -> dict:
|
||||
return {
|
||||
"id": 42,
|
||||
"model": {"type": civitai_type, "tags": ["utility"]},
|
||||
"baseModel": "SDXL 1.0",
|
||||
"creator": {"username": "Author"},
|
||||
"files": files
|
||||
or [
|
||||
{
|
||||
"type": "Model",
|
||||
"primary": True,
|
||||
"downloadUrl": "https://example.invalid/file.safetensors",
|
||||
"name": "file.safetensors",
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def metadata_provider(monkeypatch):
|
||||
class DummyProvider:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
self.payload = _other_payload("VAE")
|
||||
|
||||
async def get_model_version(self, model_id, model_version_id):
|
||||
self.calls.append((model_id, model_version_id))
|
||||
return self.payload
|
||||
|
||||
provider = DummyProvider()
|
||||
monkeypatch.setattr(
|
||||
download_manager,
|
||||
"get_default_metadata_provider",
|
||||
AsyncMock(return_value=provider),
|
||||
)
|
||||
return provider
|
||||
|
||||
|
||||
def _capture_execute(monkeypatch, captured):
|
||||
async def fake_execute_download(self, **kwargs):
|
||||
captured.update(kwargs)
|
||||
return {"success": True}
|
||||
|
||||
monkeypatch.setattr(
|
||||
DownloadManager, "_execute_download", fake_execute_download, raising=False
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"civitai_type",
|
||||
["VAE", "Upscaler", "TextEncoder", "CLIP", "CLIPVision", "Controlnet", "Other"],
|
||||
)
|
||||
async def test_download_accepts_other_model_types(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path, civitai_type
|
||||
):
|
||||
"""All VALID_OTHER_CIVITAI_TYPES route to model_type 'other'."""
|
||||
metadata_provider.payload = _other_payload(civitai_type)
|
||||
|
||||
captured = {}
|
||||
_capture_execute(monkeypatch, captured)
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=99, save_dir=str(tmp_path)
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert captured["model_type"] == "other"
|
||||
assert captured["metadata"].metadata_class == "OtherModelMetadata"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_rejects_unknown_model_type(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
metadata_provider.payload = _other_payload("Workflow")
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=99, save_dir=str(tmp_path)
|
||||
)
|
||||
|
||||
assert result["success"] is False
|
||||
assert result["error"].startswith("Model type")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_early_gate_checks_other_scanner(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
scanners.other.exists = True
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=101, save_dir=str(tmp_path)
|
||||
)
|
||||
|
||||
assert result["success"] is False
|
||||
assert result["error"] == "Model version already exists in other library"
|
||||
assert scanners.other.calls == [101]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scanner_dispatch_has_no_lora_fall_through(scanners):
|
||||
"""The Phase-2 trap: 'other' must reach the other scanner explicitly, and
|
||||
unknown types must raise instead of silently deduping against loras."""
|
||||
manager = DownloadManager()
|
||||
|
||||
scanner = await manager._get_scanner_for_model_type("other")
|
||||
assert scanner is scanners.other
|
||||
|
||||
scanner = await manager._get_scanner_for_model_type("lora")
|
||||
assert scanner is scanners.lora
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
await manager._get_scanner_for_model_type("bogus")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_explicit_file_gate_uses_other_scanner_not_lora(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
"""A matching local entry in the LORA library must not block an 'other'
|
||||
download; only the other scanner's library is consulted."""
|
||||
local_entry = {
|
||||
"file_name": "file",
|
||||
"sha256": "deadbeef",
|
||||
"civitai": {"id": 42},
|
||||
}
|
||||
scanners.lora._cache = SimpleNamespace(raw_data=[dict(local_entry)])
|
||||
scanners.other._cache = SimpleNamespace(raw_data=[])
|
||||
|
||||
metadata_provider.payload = _other_payload(
|
||||
"VAE",
|
||||
files=[
|
||||
{
|
||||
"id": 7,
|
||||
"type": "Model",
|
||||
"primary": True,
|
||||
"name": "file.safetensors",
|
||||
"hashes": {"SHA256": "deadbeef"},
|
||||
"downloadUrl": "https://example.invalid/file.safetensors",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
captured = {}
|
||||
_capture_execute(monkeypatch, captured)
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=42,
|
||||
save_dir=str(tmp_path),
|
||||
file_params={"id": 7, "type": "Model"},
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert captured["model_type"] == "other"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_explicit_file_gate_blocks_when_other_scanner_has_file(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
local_entry = {
|
||||
"file_name": "file",
|
||||
"sha256": "deadbeef",
|
||||
"civitai": {"id": 42},
|
||||
}
|
||||
scanners.other._cache = SimpleNamespace(raw_data=[dict(local_entry)])
|
||||
|
||||
metadata_provider.payload = _other_payload(
|
||||
"VAE",
|
||||
files=[
|
||||
{
|
||||
"id": 7,
|
||||
"type": "Model",
|
||||
"primary": True,
|
||||
"name": "file.safetensors",
|
||||
"hashes": {"SHA256": "deadbeef"},
|
||||
"downloadUrl": "https://example.invalid/file.safetensors",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
execute_mock = AsyncMock(return_value={"success": True})
|
||||
monkeypatch.setattr(DownloadManager, "_execute_download", execute_mock)
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=42,
|
||||
save_dir=str(tmp_path),
|
||||
file_params={"id": 7, "type": "Model"},
|
||||
)
|
||||
|
||||
assert result["success"] is False
|
||||
assert "already exists in other library" in result["error"]
|
||||
assert execute_mock.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_version_level_fallback_gate_checks_other_scanner(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
"""file_params that resolve to nothing fall back to the version-level
|
||||
gate, which must consult the other scanner."""
|
||||
scanners.other.exists = True
|
||||
|
||||
execute_mock = AsyncMock(return_value={"success": True})
|
||||
monkeypatch.setattr(DownloadManager, "_execute_download", execute_mock)
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=101,
|
||||
save_dir=str(tmp_path),
|
||||
file_params={"id": 999999, "type": "Model"},
|
||||
)
|
||||
|
||||
assert result["success"] is False
|
||||
assert result["error"] == "Model version already exists in other library"
|
||||
assert scanners.other.calls == [101]
|
||||
assert execute_mock.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_paths_use_per_sub_type_root(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
"""model.type VAE -> default_other_roots['vae']."""
|
||||
captured = {}
|
||||
_capture_execute(monkeypatch, captured)
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=99, use_default_paths=True
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert str(tmp_path / "vae") in str(captured["save_dir"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_paths_file_type_fallback_for_unmapped_model_type(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
"""model.type 'Other' maps to nothing; a 'Upscaler' file type decides."""
|
||||
metadata_provider.payload = _other_payload(
|
||||
"Other",
|
||||
files=[
|
||||
{
|
||||
"type": "Upscaler",
|
||||
"primary": True,
|
||||
"downloadUrl": "https://example.invalid/upscaler.safetensors",
|
||||
"name": "upscaler.safetensors",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
captured = {}
|
||||
_capture_execute(monkeypatch, captured)
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=99, use_default_paths=True
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert str(tmp_path / "upscale_models") in str(captured["save_dir"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_paths_explicit_file_pick_wins(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
"""An explicit pick of a bundled VAE component file routes to the vae
|
||||
root even though model.type maps to upscaler."""
|
||||
metadata_provider.payload = _other_payload(
|
||||
"Upscaler",
|
||||
files=[
|
||||
{
|
||||
"id": 1,
|
||||
"type": "Model",
|
||||
"primary": True,
|
||||
"downloadUrl": "https://example.invalid/model.safetensors",
|
||||
"name": "model.safetensors",
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"type": "VAE",
|
||||
"downloadUrl": "https://example.invalid/bundled-vae.safetensors",
|
||||
"name": "bundled-vae.safetensors",
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
captured = {}
|
||||
_capture_execute(monkeypatch, captured)
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=99,
|
||||
use_default_paths=True,
|
||||
file_params={"id": 2, "type": "VAE"},
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert str(tmp_path / "vae") in str(captured["save_dir"])
|
||||
assert captured["download_urls"] == [
|
||||
"https://example.invalid/bundled-vae.safetensors"
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_paths_errors_when_sub_type_root_unconfigured(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
"""controlnet has no configured default root in the fixture settings."""
|
||||
metadata_provider.payload = _other_payload("Controlnet")
|
||||
|
||||
execute_mock = AsyncMock(return_value={"success": True})
|
||||
monkeypatch.setattr(DownloadManager, "_execute_download", execute_mock)
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=99, use_default_paths=True
|
||||
)
|
||||
|
||||
assert result["success"] is False
|
||||
assert "controlnet" in result["error"]
|
||||
assert execute_mock.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_paths_errors_when_sub_type_undecidable(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
"""model.type 'Other' with only plain 'Model' files: never silently
|
||||
default to the vae folder — error and ask for an explicit folder."""
|
||||
metadata_provider.payload = _other_payload("Other")
|
||||
|
||||
execute_mock = AsyncMock(return_value={"success": True})
|
||||
monkeypatch.setattr(DownloadManager, "_execute_download", execute_mock)
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=99, use_default_paths=True
|
||||
)
|
||||
|
||||
assert result["success"] is False
|
||||
assert "sub-type" in result["error"]
|
||||
assert execute_mock.await_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_civarchive_source_same_payload_shape(
|
||||
monkeypatch, scanners, metadata_provider, tmp_path
|
||||
):
|
||||
"""CivArchive downloads walk the same path with the same payload shape."""
|
||||
metadata_provider.payload = _other_payload("TextEncoder")
|
||||
|
||||
captured = {}
|
||||
_capture_execute(monkeypatch, captured)
|
||||
|
||||
manager = DownloadManager()
|
||||
result = await manager.download_from_civitai(
|
||||
model_version_id=99, save_dir=str(tmp_path), source="civarchive"
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert captured["model_type"] == "other"
|
||||
|
||||
|
||||
def test_build_metadata_for_resume_uses_other_metadata():
|
||||
manager = DownloadManager()
|
||||
metadata = manager._build_metadata_for_resume(
|
||||
model_type="other",
|
||||
version_info={"model": {"type": "VAE"}},
|
||||
file_info={"name": "file.safetensors"},
|
||||
save_path="/tmp/file.safetensors",
|
||||
)
|
||||
assert metadata.metadata_class == "OtherModelMetadata"
|
||||
|
||||
|
||||
def test_other_extension_set_matches_checkpoint():
|
||||
manager = DownloadManager()
|
||||
extensions = manager._get_supported_extensions_for_type("other")
|
||||
assert extensions == manager._get_supported_extensions_for_type("checkpoint")
|
||||
assert ".gguf" in extensions
|
||||
assert ".safetensors" in extensions
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_downloaded_version_uses_other_scanner(monkeypatch, scanners):
|
||||
"""Update tracking for a downloaded other-model version consults the
|
||||
other scanner for local versions."""
|
||||
|
||||
class FakeUpdateService:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
async def update_in_library_versions(
|
||||
self, model_type, model_id, version_ids, version_info=None
|
||||
):
|
||||
self.calls.append((model_type, model_id, version_ids))
|
||||
|
||||
update_service = FakeUpdateService()
|
||||
monkeypatch.setattr(
|
||||
ServiceRegistry,
|
||||
"get_model_update_service",
|
||||
AsyncMock(return_value=update_service),
|
||||
)
|
||||
|
||||
manager = DownloadManager()
|
||||
await manager._sync_downloaded_version(
|
||||
"other", 7, {"id": 42, "model": {"id": 7}}
|
||||
)
|
||||
|
||||
assert update_service.calls == [("other", 7, [42])]
|
||||
@@ -38,3 +38,105 @@ def test_non_checkpoint_types_never_route_to_unet():
|
||||
def test_empty_inputs_stay_on_checkpoint_roots():
|
||||
assert not is_diffusion_model_download("checkpoint")
|
||||
assert not is_diffusion_model_download("checkpoint", file_types=[], base_model="")
|
||||
|
||||
|
||||
from py.services.download_routing import resolve_other_download_sub_type
|
||||
|
||||
|
||||
class TestResolveOtherDownloadSubType:
|
||||
"""Fixed priority: explicit file pick > model.type > file.type fallback."""
|
||||
|
||||
def test_explicit_file_pick_wins_over_model_type(self):
|
||||
"""User explicitly picked a VAE component file of a Checkpoint model —
|
||||
the picked file type wins."""
|
||||
assert (
|
||||
resolve_other_download_sub_type(
|
||||
"Checkpoint", file_types=["Model", "VAE"], selected_file_type="VAE"
|
||||
)
|
||||
== "vae"
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"selected,expected",
|
||||
[
|
||||
("VAE", "vae"),
|
||||
("Upscaler", "upscaler"),
|
||||
("Text Encoder", "text_encoder"),
|
||||
("Vision Encoder", "clip_vision"),
|
||||
("CLIPVision", "clip_vision"),
|
||||
("ControlNet", "controlnet"),
|
||||
],
|
||||
)
|
||||
def test_explicit_file_pick_maps_all_known_types(self, selected, expected):
|
||||
assert (
|
||||
resolve_other_download_sub_type("Other", selected_file_type=selected)
|
||||
== expected
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_type,expected",
|
||||
[
|
||||
("VAE", "vae"),
|
||||
("Upscaler", "upscaler"),
|
||||
("TextEncoder", "text_encoder"),
|
||||
("CLIP", "text_encoder"),
|
||||
("CLIPVision", "clip_vision"),
|
||||
("Controlnet", "controlnet"),
|
||||
],
|
||||
)
|
||||
def test_model_type_mapping(self, model_type, expected):
|
||||
assert resolve_other_download_sub_type(model_type) == expected
|
||||
|
||||
def test_model_type_beats_unmappable_file_pick(self):
|
||||
"""An explicit pick whose file type does not map (e.g. plain 'Model')
|
||||
falls through to model.type."""
|
||||
assert (
|
||||
resolve_other_download_sub_type(
|
||||
"TextEncoder", selected_file_type="Model"
|
||||
)
|
||||
== "text_encoder"
|
||||
)
|
||||
|
||||
def test_bundled_component_files_never_override_model_type(self):
|
||||
"""Anti-misrouting: a TextEncoder model bundling a VAE component file
|
||||
must stay text_encoder — file types are a fallback, not an override."""
|
||||
assert (
|
||||
resolve_other_download_sub_type(
|
||||
"TextEncoder", file_types=["Model", "VAE"]
|
||||
)
|
||||
== "text_encoder"
|
||||
)
|
||||
assert (
|
||||
resolve_other_download_sub_type(
|
||||
"Controlnet", file_types=["Model", "Text Encoder"]
|
||||
)
|
||||
== "controlnet"
|
||||
)
|
||||
|
||||
def test_file_type_fallback_when_model_type_unmapped(self):
|
||||
"""model.type 'Other' (or retired values) maps to nothing, so the
|
||||
first mappable file type decides."""
|
||||
assert (
|
||||
resolve_other_download_sub_type("Other", file_types=["Model", "Upscaler"])
|
||||
== "upscaler"
|
||||
)
|
||||
|
||||
def test_file_type_fallback_for_civarchive_payload(self):
|
||||
"""CivArchive-shaped payload: same fields, same decision path."""
|
||||
assert (
|
||||
resolve_other_download_sub_type(
|
||||
"Other",
|
||||
file_types=["Config", "Text Encoder"],
|
||||
)
|
||||
== "text_encoder"
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize("model_type", ["Other", "", "SomethingNew"])
|
||||
def test_undecidable_returns_none(self, model_type):
|
||||
assert (
|
||||
resolve_other_download_sub_type(model_type, file_types=["Model"]) is None
|
||||
)
|
||||
assert resolve_other_download_sub_type(model_type) is None
|
||||
|
||||
def test_model_type_matching_is_case_insensitive(self):
|
||||
assert resolve_other_download_sub_type("vAe") == "vae"
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
"""Example-images download dispatch accepts the "other" model type.
|
||||
|
||||
Covers the three scanner-dispatch sites from docs/plans/other-models-page.md
|
||||
§9.1: check_pending_models, _download_all_example_images and
|
||||
_download_specific_models_example_images_sync.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from py.services.settings_manager import get_settings_manager
|
||||
from py.utils import example_images_download_manager as download_module
|
||||
|
||||
|
||||
class StubScanner:
|
||||
"""Scanner double returning predetermined cache contents."""
|
||||
|
||||
def __init__(self, models: list[dict[str, Any]]) -> None:
|
||||
self._cache = SimpleNamespace(raw_data=models)
|
||||
|
||||
async def get_cached_data(self):
|
||||
return self._cache
|
||||
|
||||
|
||||
class RecordingWebSocketManager:
|
||||
def __init__(self) -> None:
|
||||
self.payloads: list[dict[str, Any]] = []
|
||||
|
||||
async def broadcast(self, payload: dict[str, Any]) -> None:
|
||||
self.payloads.append(payload)
|
||||
|
||||
|
||||
def _patch_all_scanners(monkeypatch: pytest.MonkeyPatch, **scanners) -> None:
|
||||
for name, getter in (
|
||||
("lora", "get_lora_scanner"),
|
||||
("checkpoint", "get_checkpoint_scanner"),
|
||||
("embedding", "get_embedding_scanner"),
|
||||
("other", "get_other_scanner"),
|
||||
):
|
||||
scanner = scanners.get(name) or StubScanner([])
|
||||
|
||||
async def _get_scanner(cls, _scanner=scanner):
|
||||
return _scanner
|
||||
|
||||
monkeypatch.setattr(
|
||||
download_module.ServiceRegistry,
|
||||
getter,
|
||||
classmethod(_get_scanner),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_pending_models_includes_other_scanner(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
tmp_path,
|
||||
settings_manager,
|
||||
):
|
||||
ws_manager = RecordingWebSocketManager()
|
||||
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||
|
||||
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||
|
||||
other_models = [{"sha256": "d" * 64, "model_name": "VAE Model"}]
|
||||
_patch_all_scanners(monkeypatch, other=StubScanner(other_models))
|
||||
|
||||
result = await manager.check_pending_models(["other"])
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["total_models"] == 1
|
||||
assert result["pending_count"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_all_example_images_processes_other_models(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
tmp_path,
|
||||
settings_manager,
|
||||
):
|
||||
ws_manager = RecordingWebSocketManager()
|
||||
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||
|
||||
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||
|
||||
other_models = [{"sha256": "e" * 64, "model_name": "Upscaler Model"}]
|
||||
_patch_all_scanners(monkeypatch, other=StubScanner(other_models))
|
||||
|
||||
async def fake_get_downloader():
|
||||
return object()
|
||||
|
||||
processed: list[tuple[str, dict[str, Any]]] = []
|
||||
|
||||
async def fake_process_model(self, scanner_type, model, scanner, *_args, **_kwargs):
|
||||
processed.append((scanner_type, model))
|
||||
return False
|
||||
|
||||
monkeypatch.setattr(download_module, "get_downloader", fake_get_downloader)
|
||||
monkeypatch.setattr(
|
||||
download_module.DownloadManager, "_process_model", fake_process_model
|
||||
)
|
||||
|
||||
# Simulate the running state that start_download establishes.
|
||||
manager._progress["status"] = "running"
|
||||
|
||||
await manager._download_all_example_images(
|
||||
str(tmp_path),
|
||||
optimize=False,
|
||||
model_types=["other"],
|
||||
delay=0,
|
||||
library_name="default",
|
||||
)
|
||||
|
||||
assert processed == [("other", other_models[0])]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_specific_models_example_images_processes_other_models(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
tmp_path,
|
||||
settings_manager,
|
||||
):
|
||||
ws_manager = RecordingWebSocketManager()
|
||||
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||
|
||||
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||
|
||||
model_hash = "f" * 64
|
||||
other_models = [{"sha256": model_hash, "model_name": "Text Encoder Model"}]
|
||||
_patch_all_scanners(monkeypatch, other=StubScanner(other_models))
|
||||
|
||||
async def fake_get_downloader():
|
||||
return object()
|
||||
|
||||
processed: list[tuple[str, dict[str, Any]]] = []
|
||||
|
||||
async def fake_process_specific_model(
|
||||
self, scanner_type, model, scanner, *_args, **_kwargs
|
||||
):
|
||||
processed.append((scanner_type, model))
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(download_module, "get_downloader", fake_get_downloader)
|
||||
monkeypatch.setattr(
|
||||
download_module.DownloadManager,
|
||||
"_process_specific_model",
|
||||
fake_process_specific_model,
|
||||
)
|
||||
|
||||
# Simulate the running state that start_force_download establishes.
|
||||
manager._progress["status"] = "running"
|
||||
|
||||
await manager._download_specific_models_example_images_sync(
|
||||
[model_hash],
|
||||
str(tmp_path),
|
||||
optimize=False,
|
||||
model_types=["other"],
|
||||
delay=0,
|
||||
library_name="default",
|
||||
)
|
||||
|
||||
assert processed == [("other", other_models[0])]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def settings_manager():
|
||||
return get_settings_manager()
|
||||
@@ -303,9 +303,13 @@ class TestOtherModelMetadataFromCivitai:
|
||||
def _build(self, civitai_type: str) -> OtherModelMetadata:
|
||||
return OtherModelMetadata.from_civitai_info(
|
||||
{
|
||||
"type": civitai_type,
|
||||
"baseModel": "SDXL",
|
||||
"model": {"name": "Model", "tags": ["tag"], "description": "desc"},
|
||||
"model": {
|
||||
"name": "Model",
|
||||
"tags": ["tag"],
|
||||
"description": "desc",
|
||||
"type": civitai_type,
|
||||
},
|
||||
},
|
||||
{"name": "model.safetensors", "sizeKB": 1, "hashes": {"SHA256": "AB"}},
|
||||
"/tmp/model.safetensors",
|
||||
@@ -329,6 +333,20 @@ class TestOtherModelMetadataFromCivitai:
|
||||
assert metadata.sha256 == "ab"
|
||||
assert metadata.tags == ["tag"]
|
||||
|
||||
def test_top_level_type_key_is_ignored(self):
|
||||
"""Regression: the CivitAI type lives at version["model"]["type"]; a
|
||||
top-level version["type"] key must not drive the mapping (#Phase-1 bug)."""
|
||||
metadata = OtherModelMetadata.from_civitai_info(
|
||||
{
|
||||
"type": "Upscaler",
|
||||
"baseModel": "SDXL",
|
||||
"model": {"name": "Model", "type": "VAE"},
|
||||
},
|
||||
{"name": "model.safetensors", "sizeKB": 1, "hashes": {"SHA256": "AB"}},
|
||||
"/tmp/model.safetensors",
|
||||
)
|
||||
assert metadata.sub_type == "vae"
|
||||
|
||||
|
||||
def test_page_type_maps_to_other():
|
||||
"""The WS progress page type for the other scanner is 'other'."""
|
||||
|
||||
@@ -1208,3 +1208,137 @@ def test_skip_previously_downloaded_model_versions_coerces_string_input(manager)
|
||||
|
||||
assert manager.get_skip_previously_downloaded_model_versions() is True
|
||||
assert manager.settings["skip_previously_downloaded_model_versions"] is True
|
||||
|
||||
|
||||
def test_default_other_roots_stay_empty_without_other_folders(manager):
|
||||
assert manager._get_default_settings()["default_other_roots"] == {}
|
||||
|
||||
manager.settings["default_other_roots"] = {}
|
||||
manager.settings["folder_paths"] = {}
|
||||
manager.settings["extra_folder_paths"] = {}
|
||||
|
||||
manager._auto_set_default_roots()
|
||||
|
||||
assert manager.get("default_other_roots") == {}
|
||||
|
||||
|
||||
def test_auto_set_default_other_roots(manager):
|
||||
manager.settings["default_other_roots"] = {}
|
||||
manager.settings["folder_paths"] = {
|
||||
"vae": ["/vae"],
|
||||
"upscale_models": ["/upscalers"],
|
||||
"clip_vision": ["/clip_vision"],
|
||||
}
|
||||
|
||||
manager._auto_set_default_roots()
|
||||
|
||||
roots = manager.get("default_other_roots")
|
||||
assert roots["vae"] == "/vae"
|
||||
assert roots["upscaler"] == "/upscalers"
|
||||
assert roots["clip_vision"] == "/clip_vision"
|
||||
# text_encoder has no configured folders -> no entry
|
||||
assert "text_encoder" not in roots
|
||||
assert "controlnet" not in roots
|
||||
|
||||
|
||||
def test_auto_set_default_other_roots_text_encoder_dual_key_union(manager):
|
||||
"""text_encoder candidates merge text_encoders and the legacy clip key."""
|
||||
manager.settings["default_other_roots"] = {}
|
||||
manager.settings["folder_paths"] = {
|
||||
"clip": ["/legacy-clip"],
|
||||
"text_encoders": ["/text-encoders"],
|
||||
}
|
||||
|
||||
manager._auto_set_default_roots()
|
||||
|
||||
roots = manager.get("default_other_roots")
|
||||
assert roots["text_encoder"] in {"/legacy-clip", "/text-encoders"}
|
||||
|
||||
# A value pointing at either key's root is considered valid
|
||||
manager.settings["default_other_roots"] = {"text_encoder": "/legacy-clip"}
|
||||
manager._auto_set_default_roots()
|
||||
assert manager.get("default_other_roots")["text_encoder"] == "/legacy-clip"
|
||||
|
||||
|
||||
def test_auto_set_default_other_roots_repairs_stale(manager):
|
||||
manager.settings["default_other_roots"] = {"vae": "/stale-vae"}
|
||||
manager.settings["folder_paths"] = {"vae": ["/vae"]}
|
||||
|
||||
manager._auto_set_default_roots()
|
||||
|
||||
assert manager.get("default_other_roots")["vae"] == "/vae"
|
||||
|
||||
|
||||
def test_auto_set_default_other_roots_uses_extra_folder_paths(manager):
|
||||
manager.settings["default_other_roots"] = {}
|
||||
manager.settings["folder_paths"] = {"vae": []}
|
||||
manager.settings["extra_folder_paths"] = {"vae": ["/extra-vae"]}
|
||||
|
||||
manager._auto_set_default_roots()
|
||||
|
||||
assert manager.get("default_other_roots")["vae"] == "/extra-vae"
|
||||
|
||||
|
||||
def test_set_default_other_roots_syncs_active_library(manager):
|
||||
manager.set("default_other_roots", {"vae": "/vae"})
|
||||
|
||||
libraries = manager.get_libraries()
|
||||
active = manager.get_active_library_name()
|
||||
assert libraries[active]["default_other_roots"] == {"vae": "/vae"}
|
||||
assert manager.get("default_other_roots") == {"vae": "/vae"}
|
||||
|
||||
|
||||
def test_set_default_other_roots_rejects_illegal_sub_type(manager):
|
||||
with pytest.raises(ValueError, match="Unknown other-model sub-type"):
|
||||
manager.set("default_other_roots", {"vae": "/vae", "lora": "/loras"})
|
||||
|
||||
|
||||
def test_set_default_other_roots_normalizes_values(manager):
|
||||
manager.set("default_other_roots", {"vae": " /vae ", "upscaler": ""})
|
||||
assert manager.get("default_other_roots") == {"vae": "/vae"}
|
||||
|
||||
|
||||
def test_upsert_library_passthrough_default_other_roots(manager, tmp_path):
|
||||
manager.upsert_library(
|
||||
"studio",
|
||||
folder_paths={"loras": ["/studio/loras"], "vae": ["/studio/vae"]},
|
||||
default_other_roots={"vae": "/studio/vae"},
|
||||
activate=True,
|
||||
)
|
||||
|
||||
libraries = manager.get_libraries()
|
||||
assert libraries["studio"]["default_other_roots"] == {"vae": "/studio/vae"}
|
||||
assert manager.get("default_other_roots") == {"vae": "/studio/vae"}
|
||||
|
||||
# Omitting the argument preserves the stored value
|
||||
manager.upsert_library("studio", folder_paths={"loras": ["/studio/loras"]})
|
||||
libraries = manager.get_libraries()
|
||||
assert libraries["studio"]["default_other_roots"] == {"vae": "/studio/vae"}
|
||||
|
||||
|
||||
def test_library_switch_restores_default_other_roots(manager):
|
||||
manager.set("default_other_roots", {"vae": "/default-vae"})
|
||||
manager.create_library(
|
||||
"studio",
|
||||
folder_paths={"loras": ["/studio/loras"]},
|
||||
default_other_roots={"vae": "/studio-vae"},
|
||||
)
|
||||
|
||||
manager.activate_library("studio")
|
||||
assert manager.get("default_other_roots") == {"vae": "/studio-vae"}
|
||||
|
||||
manager.activate_library("default")
|
||||
assert manager.get("default_other_roots") == {"vae": "/default-vae"}
|
||||
|
||||
|
||||
def test_migrate_sanitizes_legacy_libraries_includes_other_roots(tmp_path, monkeypatch):
|
||||
initial = {
|
||||
"libraries": {"legacy": "not-a-dict"},
|
||||
"active_library": "legacy",
|
||||
"folder_paths": {"loras": ["/old"]},
|
||||
}
|
||||
|
||||
manager = _create_manager_with_settings(tmp_path, monkeypatch, initial)
|
||||
|
||||
payload = manager.get_libraries()["legacy"]
|
||||
assert payload["default_other_roots"] == {}
|
||||
|
||||
Reference in New Issue
Block a user