fix(types): resolve pre-existing basedpyright errors in tests

Fix ~790 basedpyright errors across the test suite:
- Type stub subclasses of real production classes with super().__init__()
- Add missing generic type arguments and Dict[str, Any] annotations
- Add None guards before subscript/member access
- Adapt tests to production API changes (removed dead handlers,
  PersistentModelCache.get_default, _i18n_filter_added location)
This commit is contained in:
Will Miao
2026-08-08 20:12:59 +08:00
parent 8e724538bd
commit d2f955266d
95 changed files with 953 additions and 666 deletions
+7 -7
View File
@@ -6,7 +6,7 @@ from unittest.mock import AsyncMock
import pytest
from py.services.aria2_downloader import Aria2Downloader, Aria2Error
from py.services.aria2_downloader import Aria2Downloader, Aria2Error, Aria2Transfer
from py.services.aria2_transfer_state import Aria2TransferStateStore
from py.services import aria2_transfer_state
@@ -165,9 +165,9 @@ async def test_download_file_keeps_auth_headers_when_civitai_does_not_redirect(
@pytest.mark.asyncio
async def test_pause_resume_cancel_forward_to_rpc(monkeypatch):
downloader = Aria2Downloader()
downloader._transfers["download-1"] = type(
"Transfer", (), {"gid": "gid-1", "save_path": "/tmp/model.safetensors"}
)()
downloader._transfers["download-1"] = Aria2Transfer(
gid="gid-1", save_path="/tmp/model.safetensors"
)
calls = []
@@ -200,9 +200,9 @@ async def test_download_file_reuses_existing_transfer_without_add_uri(
downloader._rpc_secret = "secret"
save_path = tmp_path / "downloads" / "model.safetensors"
downloader._transfers["download-1"] = type(
"Transfer", (), {"gid": "gid-1", "save_path": str(save_path)}
)()
downloader._transfers["download-1"] = Aria2Transfer(
gid="gid-1", save_path=str(save_path)
)
rpc_calls = []
statuses = iter(
+17 -14
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
import json
from pathlib import Path
from typing import Any, Dict, List, Optional
from typing import Any, Dict, Iterator, List, Optional
import pytest
@@ -16,7 +16,7 @@ from py.services.persistent_model_cache import DEFAULT_LICENSE_FLAGS, Persistent
@pytest.fixture(autouse=True)
def reset_backfill_singleton() -> None:
def reset_backfill_singleton() -> Iterator[None]:
"""Reset the service singleton so every test starts from a fresh instance."""
Autov3BackfillService._instance = None
yield
@@ -66,7 +66,7 @@ class RecordingScanner:
self.model_type = model_type
self._persistent_cache = persistent_cache
self.entries: Dict[str, Dict[str, Any]] = {entry['file_path']: entry for entry in entries}
self.update_calls: List[tuple] = []
self.update_calls: List[tuple[str, str, str]] = []
async def update_autov3_for_model(self, model_type: str, file_path: str, autov3: str) -> bool:
self.update_calls.append((model_type, file_path, autov3))
@@ -114,7 +114,7 @@ async def test_backfill_updates_models_and_self_terminates(tmp_path: Path, monke
)
scanner = RecordingScanner('dummy', store, entries)
updated = await Autov3BackfillService.get_instance().backfill(scanner)
updated = await Autov3BackfillService.get_instance().backfill(scanner) # pyright: ignore[reportArgumentType]
# Non-safetensors files yield no embedded hash, so both are marked ''.
assert updated == 2
@@ -124,6 +124,7 @@ async def test_backfill_updates_models_and_self_terminates(tmp_path: Path, monke
assert store.get_models_missing_autov3('dummy') == []
persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data}
assert items[path_a]['autov3'] == ''
assert items[path_b]['autov3'] == ''
@@ -147,7 +148,7 @@ async def test_backfill_skips_missing_files_without_marking(tmp_path: Path, monk
)
scanner = RecordingScanner('dummy', store, entries)
updated = await Autov3BackfillService.get_instance().backfill(scanner)
updated = await Autov3BackfillService.get_instance().backfill(scanner) # pyright: ignore[reportArgumentType]
assert updated == 1
assert scanner.update_calls == [('dummy', existing, '')]
@@ -162,7 +163,7 @@ async def test_backfill_returns_zero_when_same_type_already_running(tmp_path: Pa
service = Autov3BackfillService.get_instance()
service._running_types = {'dummy'}
try:
assert await service.backfill(scanner) == 0
assert await service.backfill(scanner) == 0 # pyright: ignore[reportArgumentType]
finally:
service._running_types = set()
assert scanner.update_calls == []
@@ -195,7 +196,7 @@ async def test_backfill_runs_concurrently_for_different_model_types(tmp_path: Pa
try:
# The lora backfill must still run while checkpoint is in progress.
assert await service.backfill(lora_scanner) == 1
assert await service.backfill(lora_scanner) == 1 # pyright: ignore[reportArgumentType]
assert lora_scanner.update_calls == [('lora', lora_file, '')]
finally:
service._running_types = set()
@@ -213,7 +214,7 @@ async def test_backfill_never_raises_on_failure(tmp_path: Path, monkeypatch) ->
store.save_cache('dummy', entries, {'hash-boom': [existing]}, [])
scanner = RaisingScanner('dummy', store, entries)
updated = await Autov3BackfillService.get_instance().backfill(scanner)
updated = await Autov3BackfillService.get_instance().backfill(scanner) # pyright: ignore[reportArgumentType]
assert updated == 0
@@ -238,7 +239,7 @@ async def test_backfill_uses_default_cache_when_scanner_has_none(tmp_path: Path,
store.update_single_model(model_type, new_item, old_item)
return True
updated = await Autov3BackfillService.get_instance().backfill(BareScanner())
updated = await Autov3BackfillService.get_instance().backfill(BareScanner()) # pyright: ignore[reportArgumentType]
assert updated == 1
assert store.get_models_missing_autov3('dummy') == []
@@ -253,9 +254,9 @@ async def test_backfill_idempotent_second_run_is_noop(tmp_path: Path, monkeypatc
scanner = RecordingScanner('dummy', store, entries)
service = Autov3BackfillService.get_instance()
assert await service.backfill(scanner) == 1
assert await service.backfill(scanner) == 1 # pyright: ignore[reportArgumentType]
# A re-run has nothing left to do.
assert await service.backfill(scanner) == 0
assert await service.backfill(scanner) == 0 # pyright: ignore[reportArgumentType]
assert len(scanner.update_calls) == 1
@@ -275,7 +276,7 @@ async def test_backfill_end_to_end_through_scanner_lazy_import(tmp_path: Path, m
)
class RealScanner(ModelScanner):
def __init__(self) -> None:
def __init__(self) -> None: # pyright: ignore[reportMissingSuperCall]
self.model_type = 'dummy'
self._persistent_cache = store
self._cache = ModelCache(raw_data=[dict(e) for e in entries], folders=[])
@@ -285,6 +286,7 @@ async def test_backfill_end_to_end_through_scanner_lazy_import(tmp_path: Path, m
assert store.get_models_missing_autov3('dummy') == []
persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data}
assert items[path_a]['autov3'] == ''
assert items[path_b]['autov3'] == ''
@@ -314,12 +316,13 @@ async def test_backfill_prefers_civitai_autov3_from_sidecar(tmp_path: Path, monk
store.save_cache('dummy', [_entry(path, 'hash-ckpt')], {'hash-ckpt': [path]}, [])
scanner = RecordingScanner('dummy', store, [_entry(path, 'hash-ckpt')])
updated = await Autov3BackfillService.get_instance().backfill(scanner)
updated = await Autov3BackfillService.get_instance().backfill(scanner) # pyright: ignore[reportArgumentType]
assert updated == 1
assert scanner.update_calls == [('dummy', path, 'abcdef123456')]
persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data}
assert items[path]['autov3'] == 'abcdef123456'
# Self-terminating: the row is marked and the driving query empties.
@@ -344,7 +347,7 @@ async def test_backfill_falls_back_to_header_when_sidecar_has_no_match(tmp_path:
store.save_cache('dummy', [_entry(path, 'hash-plain')], {'hash-plain': [path]}, [])
scanner = RecordingScanner('dummy', store, [_entry(path, 'hash-plain')])
updated = await Autov3BackfillService.get_instance().backfill(scanner)
updated = await Autov3BackfillService.get_instance().backfill(scanner) # pyright: ignore[reportArgumentType]
assert updated == 1
assert scanner.update_calls == [('dummy', path, '')]
+2 -2
View File
@@ -209,8 +209,8 @@ async def test_model_update_service_migrates_legacy_snapshot_db(tmp_path, monkey
return str(legacy_db)
monkeypatch.setattr(
"py.services.persistent_model_cache.get_persistent_cache",
lambda *_args, **_kwargs: LegacyCache(),
"py.services.persistent_model_cache.PersistentModelCache.get_default",
lambda *args, **kwargs: LegacyCache(),
)
service = ModelUpdateService(settings_manager=DummySettingsManager())
+19 -14
View File
@@ -28,13 +28,14 @@ class DummyService(BaseModelService):
return model_data
class StubRepository:
class StubRepository(ModelCacheRepository):
def __init__(self, data):
super().__init__(scanner=object())
self._data = list(data)
self.parse_sort_calls = []
self.fetch_sorted_calls = []
def parse_sort(self, sort_by):
def parse_sort(self, sort_by): # pyright: ignore[reportIncompatibleMethodOverride]
params = ModelCacheRepository.parse_sort(sort_by)
self.parse_sort_calls.append(sort_by)
return params
@@ -44,8 +45,9 @@ class StubRepository:
return list(self._data)
class StubFilterSet:
class StubFilterSet(ModelFilterSet):
def __init__(self, result):
super().__init__(settings=StubSettings({}))
self.result = list(result)
self.calls = []
@@ -54,8 +56,9 @@ class StubFilterSet:
return list(self.result)
class StubSearchStrategy:
class StubSearchStrategy(SearchStrategy):
def __init__(self, search_result):
super().__init__()
self.search_result = list(search_result)
self.normalize_calls = []
self.apply_calls = []
@@ -67,7 +70,7 @@ class StubSearchStrategy:
normalized.update(options)
return normalized
def apply(self, data, search_term, options, fuzzy):
def apply(self, data, search_term, options, fuzzy=False):
self.apply_calls.append((list(data), search_term, options, fuzzy))
return list(self.search_result)
@@ -269,8 +272,9 @@ async def test_get_paginated_data_filters_and_searches_combination():
assert response["total_pages"] == 1
class PassThroughFilterSet:
class PassThroughFilterSet(ModelFilterSet):
def __init__(self):
super().__init__(settings=StubSettings({}))
self.calls = []
def apply(self, data, criteria):
@@ -278,8 +282,9 @@ class PassThroughFilterSet:
return list(data)
class NoSearchStrategy:
class NoSearchStrategy(SearchStrategy):
def __init__(self):
super().__init__()
self.normalize_calls = []
self.apply_called = False
@@ -355,7 +360,7 @@ async def test_get_paginated_data_filters_by_update_status():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -428,7 +433,7 @@ async def test_get_paginated_data_skips_items_when_update_check_fails():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -465,7 +470,7 @@ async def test_get_paginated_data_annotates_update_flags_with_bulk_dedup():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -561,7 +566,7 @@ async def test_version_grouping_same_base_prefers_matching_base():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -658,7 +663,7 @@ async def test_version_grouping_same_base_honors_latest_local_version():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -694,7 +699,7 @@ async def test_get_paginated_data_filters_update_available_only():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -1028,7 +1033,7 @@ def test_model_filter_set_supports_legacy_tag_arrays():
{"model_name": "AnimeOnly", "tags": ["anime"]},
]
criteria = FilterCriteria(tags=["style"])
criteria = FilterCriteria(tags=["style"]) # pyright: ignore[reportArgumentType]
result = filter_set.apply(data, criteria)
assert [item["model_name"] for item in result] == ["StyleOnly", "StyleAnime"]
+7 -7
View File
@@ -333,7 +333,7 @@ class TestBatchImportService:
)
service = BatchImportService(
analysis_service=analysis_service,
analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service,
ws_manager=ws_manager,
logger=logger,
@@ -445,8 +445,8 @@ class TestBatchImportServiceEdgeCases:
logger = logging.getLogger("test")
return BatchImportService(
analysis_service=analysis_service,
persistence_service=persistence_service,
analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service, # pyright: ignore[reportArgumentType]
ws_manager=ws_manager,
logger=logger,
)
@@ -506,8 +506,8 @@ class TestBatchImportServiceEdgeCases:
(tmp_path / "test.png").write_bytes(b"fake-image")
service = BatchImportService(
analysis_service=analysis_service,
persistence_service=persistence_service,
analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service, # pyright: ignore[reportArgumentType]
ws_manager=ws_manager,
logger=logger,
)
@@ -571,8 +571,8 @@ class TestInputValidation:
logger = logging.getLogger("test")
return BatchImportService(
analysis_service=analysis_service,
persistence_service=persistence_service,
analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service, # pyright: ignore[reportArgumentType]
ws_manager=ws_manager,
logger=logger,
)
+15 -6
View File
@@ -84,6 +84,7 @@ class TestCacheEntryValidator:
result = CacheEntryValidator.validate(entry, auto_repair=False)
assert result.is_valid is True
assert result.entry is not None
assert result.entry['sha256'] == ''
assert result.entry['hash_status'] == 'pending'
@@ -141,7 +142,7 @@ class TestCacheEntryValidator:
def test_validate_none_entry(self):
"""Test validation handles None entry"""
result = CacheEntryValidator.validate(None, auto_repair=False)
result = CacheEntryValidator.validate(None, auto_repair=False) # pyright: ignore[reportArgumentType]
assert result.is_valid is False
assert result.repaired is False
@@ -150,7 +151,7 @@ class TestCacheEntryValidator:
def test_validate_non_dict_entry(self):
"""Test validation handles non-dict entry"""
result = CacheEntryValidator.validate("not a dict", auto_repair=False)
result = CacheEntryValidator.validate("not a dict", auto_repair=False) # pyright: ignore[reportArgumentType]
assert result.is_valid is False
assert result.repaired is False
@@ -169,6 +170,7 @@ class TestCacheEntryValidator:
assert result.is_valid is True
assert result.repaired is True
assert result.entry is not None
assert result.entry['file_name'] == ''
assert result.entry['model_name'] == ''
assert result.entry['tags'] == []
@@ -186,6 +188,7 @@ class TestCacheEntryValidator:
assert result.is_valid is True
assert result.repaired is True
assert result.entry is not None
assert result.entry['size'] == 0 # Default value
assert result.entry['tags'] == [] # Default value
@@ -199,6 +202,7 @@ class TestCacheEntryValidator:
result = CacheEntryValidator.validate(entry, auto_repair=True)
assert result.is_valid is True
assert result.entry is not None
assert result.entry['sha256'] == 'abc123def456'
def test_validate_batch_all_valid(self):
@@ -262,8 +266,8 @@ class TestCacheEntryValidator:
def test_get_file_path_safe_not_dict(self):
"""Test safe file_path extraction from non-dict"""
assert CacheEntryValidator.get_file_path_safe(None) == ''
assert CacheEntryValidator.get_file_path_safe('string') == ''
assert CacheEntryValidator.get_file_path_safe(None) == '' # pyright: ignore[reportArgumentType]
assert CacheEntryValidator.get_file_path_safe('string') == '' # pyright: ignore[reportArgumentType]
def test_get_sha256_safe(self):
"""Test safe sha256 extraction"""
@@ -277,8 +281,8 @@ class TestCacheEntryValidator:
def test_get_sha256_safe_not_dict(self):
"""Test safe sha256 extraction from non-dict"""
assert CacheEntryValidator.get_sha256_safe(None) == ''
assert CacheEntryValidator.get_sha256_safe('string') == ''
assert CacheEntryValidator.get_sha256_safe(None) == '' # pyright: ignore[reportArgumentType]
assert CacheEntryValidator.get_sha256_safe('string') == '' # pyright: ignore[reportArgumentType]
def test_validate_with_all_optional_fields(self):
"""Test validation with all optional fields present"""
@@ -358,6 +362,7 @@ class TestAutov3Validation:
)
assert result.is_valid is True
assert result.entry is not None
assert result.entry['autov3'] == 'abcdef123456'
assert result.repaired is True
@@ -378,6 +383,7 @@ class TestAutov3Validation:
assert result.is_valid is True
assert result.repaired is False
assert result.entry is not None
assert result.entry['autov3'] is None
def test_validate_absent_autov3_is_valid_and_not_counted_as_repair(self):
@@ -386,6 +392,7 @@ class TestAutov3Validation:
assert result.is_valid is True
assert result.repaired is False
assert result.entry is not None
assert 'autov3' not in result.entry
def test_validate_short_autov3_still_valid_and_repaired_to_none(self):
@@ -396,6 +403,7 @@ class TestAutov3Validation:
)
assert result.is_valid is True
assert result.entry is not None
assert result.entry['autov3'] is None
assert result.repaired is True
@@ -407,5 +415,6 @@ class TestAutov3Validation:
)
assert result.is_valid is True
assert result.entry is not None
assert result.entry['autov3'] is None
assert result.repaired is True
+4 -3
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
import json
from types import SimpleNamespace
from typing import Any
import pytest
@@ -13,7 +14,7 @@ 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]) -> None:
def __init__(self, models: list[dict[str, Any]]) -> None:
self._cache = SimpleNamespace(raw_data=models)
async def get_cached_data(self):
@@ -58,9 +59,9 @@ class RecordingWebSocketManager:
"""Collects broadcast payloads for assertions."""
def __init__(self) -> None:
self.payloads: list[dict] = []
self.payloads: list[dict[str, Any]] = []
async def broadcast(self, payload: dict) -> None:
async def broadcast(self, payload: dict[str, Any]) -> None:
self.payloads.append(payload)
+3 -3
View File
@@ -4,7 +4,7 @@ import asyncio
import json
import os
from pathlib import Path
from typing import List
from typing import Any, List
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@@ -17,9 +17,9 @@ from py.utils.models import CheckpointMetadata
class RecordingWebSocketManager:
def __init__(self) -> None:
self.payloads: List[dict] = []
self.payloads: List[dict[str, Any]] = []
async def broadcast_init_progress(self, payload: dict) -> None:
async def broadcast_init_progress(self, payload: dict[str, Any]) -> None:
self.payloads.append(payload)
+3 -3
View File
@@ -1,6 +1,6 @@
import os
from pathlib import Path
from typing import List
from typing import Any, List
import pytest
@@ -12,9 +12,9 @@ from py.services.persistent_model_cache import PersistedCacheData
class RecordingWebSocketManager:
def __init__(self) -> None:
self.payloads: List[dict] = []
self.payloads: List[dict[str, Any]] = []
async def broadcast_init_progress(self, payload: dict) -> None:
async def broadcast_init_progress(self, payload: dict[str, Any]) -> None:
self.payloads.append(payload)
+6 -2
View File
@@ -1,4 +1,5 @@
import copy
from typing import Any, Dict
from unittest.mock import AsyncMock
import pytest
@@ -34,7 +35,7 @@ def downloader(monkeypatch):
return instance
def _base_civarchive_payload(version_id=1976567, *, trigger="mxpln", nsfw_level=31):
def _base_civarchive_payload(version_id=1976567, *, trigger="mxpln", nsfw_level=31) -> Dict[str, Any]:
version_name = "v2.0" if version_id != 1976567 else "v1.0"
file_sha = "e2b7a280d6539556f23f380b3f71e4e22bc4524445c4c96526e117c6005c6ad3"
return {
@@ -110,6 +111,7 @@ async def test_get_model_by_hash_transforms_payload(downloader):
result, error = await client.get_model_by_hash("abc")
assert error is None
assert result is not None
assert result["id"] == 1976567
assert result["nsfwLevel"] == 31
assert result["trainedWords"] == ["mxpln"]
@@ -131,7 +133,7 @@ async def test_get_model_versions_fetches_each_version(downloader):
base_payload = _base_civarchive_payload(version_id=2042594, trigger="mxpln-new", nsfw_level=5)
other_payload = _base_civarchive_payload()
responses = {
responses: Dict[Any, Dict[str, Any]] = {
(base_url, None): base_payload,
(base_url, (("modelVersionId", "2042594"),)): base_payload,
(base_url, (("modelVersionId", "1976567"),)): other_payload,
@@ -151,6 +153,7 @@ async def test_get_model_versions_fetches_each_version(downloader):
result = await client.get_model_versions("1746460")
assert result is not None
assert result["name"] == "Mixplin Style [Illustrious]"
assert result["type"] == "LORA"
versions = result["modelVersions"]
@@ -221,6 +224,7 @@ async def test_get_model_by_hash_uses_file_fallback(downloader, monkeypatch):
result, error = await client.get_model_by_hash("fallback")
assert error is None
assert result is not None
assert result["id"] == 1976567
assert result["model"]["name"] == "Mixplin Style [Illustrious]"
assert any("/models/1746460" in call["url"] for call in downloader.calls)
@@ -9,6 +9,8 @@ from py.services.civitai_base_model_service import CivitaiBaseModelService
class TestCivitaiBaseModelService:
"""Test suite for CivitaiBaseModelService."""
service = CivitaiBaseModelService()
@pytest.fixture(autouse=True)
def setup_service(self):
"""Create a fresh service instance for each test."""
@@ -46,7 +48,7 @@ class TestCivitaiBaseModelService:
def test_generate_abbreviation_edge_cases(self):
"""Test abbreviation generation edge cases."""
assert self.service.generate_abbreviation("") == "OTH"
assert self.service.generate_abbreviation(None) == "OTH"
assert self.service.generate_abbreviation(None) == "OTH" # pyright: ignore[reportArgumentType]
def test_cache_status_no_cache(self):
"""Test cache status when no cache exists."""
+7 -1
View File
@@ -97,6 +97,7 @@ async def test_get_model_by_hash_enriches_metadata(monkeypatch, downloader):
result, error = await client.get_model_by_hash("hash")
assert error is None
assert result is not None
assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"]
assert result["creator"] == {"username": "user"}
@@ -254,7 +255,7 @@ async def test_get_model_versions_bulk_success(monkeypatch, downloader):
client = await CivitaiClient.get_instance()
result = await client.get_model_versions_bulk([1, "2", 2])
result = await client.get_model_versions_bulk([1, "2", 2]) # pyright: ignore[reportArgumentType]
assert result == {
1: {
@@ -310,6 +311,7 @@ async def test_get_model_version_by_version_id(monkeypatch, downloader):
result = await client.get_model_version(version_id=7)
assert result is not None
assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"]
assert result["creator"] == {"username": "user"}
@@ -364,6 +366,7 @@ async def test_get_model_version_with_model_id_prefers_version_endpoint(monkeypa
result = await client.get_model_version(model_id=99, version_id=7)
assert result is not None
assert result["id"] == 7
assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"]
@@ -420,6 +423,7 @@ async def test_get_model_version_with_model_id_fallbacks_to_hash(monkeypatch, do
result = await client.get_model_version(model_id=99, version_id=7)
assert result is not None
assert result["id"] == 7
assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"]
@@ -461,6 +465,7 @@ async def test_get_model_version_with_model_id_builds_from_model_data(monkeypatc
result = await client.get_model_version(model_id=99, version_id=7)
assert result is not None
assert result["modelId"] == 99
assert result["model"]["name"] == "Model"
assert result["model"]["type"] == "LORA"
@@ -503,6 +508,7 @@ async def test_get_model_version_info_success(monkeypatch, downloader):
assert result == expected
assert error is None
assert result is not None
assert "comfy" not in result["images"][0]["meta"]
assert result["images"][0]["meta"]["other"] == "keep"
+2 -2
View File
@@ -91,7 +91,7 @@ async def test_parse_metadata_handles_nested_meta_and_lowercase_hashes(monkeypat
},
}
assert parser.is_metadata_matching(metadata)
assert parser.is_metadata_matching(metadata) # pyright: ignore[reportArgumentType]
result = await parser.parse_metadata(metadata)
@@ -272,7 +272,7 @@ async def test_parse_metadata_handles_modelVersionIds(monkeypatch):
"modelVersionIds": [2398829, 2398838],
}
assert parser.is_metadata_matching(metadata)
assert parser.is_metadata_matching(metadata) # pyright: ignore[reportArgumentType]
result = await parser.parse_metadata(metadata)
@@ -758,6 +758,7 @@ async def test_get_active_downloads_restores_orphaned_aria2_partial_as_paused(
downloads = await manager.get_active_downloads()
persisted = await manager._aria2_state_store.get("download-1")
assert persisted is not None
assert downloads["downloads"] == [
{
@@ -922,6 +923,7 @@ async def test_get_active_downloads_restores_persisted_aria2_without_initial_sav
downloads = await manager.get_active_downloads()
persisted = await manager._aria2_state_store.get("download-1")
assert persisted is not None
assert downloads["downloads"] == [
{
@@ -3,6 +3,7 @@
import os
from pathlib import Path
from types import SimpleNamespace
from typing import Optional
from unittest.mock import AsyncMock
import pytest
@@ -99,7 +100,7 @@ async def test_execute_download_uses_rewritten_civitai_preview(monkeypatch, tmp_
self.file_path = str(path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.preview_url: Optional[str] = None
self.autov3 = None
self.preview_nsfw_level = None
@@ -182,6 +183,7 @@ async def test_execute_download_uses_rewritten_civitai_preview(monkeypatch, tmp_
assert any("width=450,optimized=true" in url for url in preview_urls)
assert dummy_downloader.memory_calls == 0
assert optimize_called["value"] is False
assert metadata.preview_url is not None
assert metadata.preview_url.endswith(".jpeg")
assert metadata.preview_nsfw_level == 2
stored_preview = manager._active_downloads["dl"]["preview_path"]
@@ -204,7 +206,7 @@ async def test_execute_download_respects_blur_setting(monkeypatch, tmp_path):
self.file_path = str(path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.preview_url: Optional[str] = None
self.autov3 = None
self.preview_nsfw_level = None
@@ -326,7 +328,7 @@ async def test_execute_download_uses_auth_for_red_civitai_downloads(monkeypatch,
self.file_path = str(path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.preview_url: Optional[str] = None
self.autov3 = None
self.preview_nsfw_level = None
+14 -14
View File
@@ -75,7 +75,7 @@ async def test_execute_download_retries_urls(monkeypatch, tmp_path):
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -165,7 +165,7 @@ async def test_execute_download_uses_aria2_backend_for_model_files(monkeypatch,
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -272,7 +272,7 @@ async def test_execute_download_allows_anonymous_civitai_with_aria2(
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -350,7 +350,7 @@ async def test_execute_download_adjusts_checkpoint_sub_type(monkeypatch, tmp_pat
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
self.preview_nsfw_level = 0
self.sub_type = "checkpoint"
@@ -450,7 +450,7 @@ async def test_execute_download_extracts_zip_single_model(monkeypatch, tmp_path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -531,7 +531,7 @@ async def test_execute_download_extracts_zip_multiple_models(monkeypatch, tmp_pa
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -614,7 +614,7 @@ async def test_execute_download_extracts_zip_pt_embedding(monkeypatch, tmp_path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -1168,7 +1168,7 @@ async def test_execute_download_waits_for_paused_pre_transfer_gate(monkeypatch,
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -1278,7 +1278,7 @@ async def test_execute_download_reuses_existing_aria2_partial_path(monkeypatch,
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return "renamed.safetensors"
@@ -1356,7 +1356,7 @@ async def test_execute_download_rejects_conflicting_aria2_partial_path(tmp_path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
raise AssertionError("should not rename")
@@ -1408,7 +1408,7 @@ async def test_execute_download_reassigns_same_aria2_partial_to_new_download_id(
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
raise AssertionError("should not rename")
@@ -1460,9 +1460,9 @@ async def test_execute_download_reassigns_same_aria2_partial_to_new_download_id(
assert manager._active_downloads["new-download"]["file_path"] == str(target_path)
assert dummy_aria2.calls == [("reassign_transfer", "old-download", "new-download")]
assert await manager._aria2_state_store.get("old-download") is None
assert (await manager._aria2_state_store.get("new-download"))["save_path"] == str(
target_path
)
persisted = await manager._aria2_state_store.get("new-download")
assert persisted is not None
assert persisted["save_path"] == str(target_path)
def test_is_same_aria2_download_request_requires_version_id_match():
+17 -9
View File
@@ -9,7 +9,7 @@ from py.services.downloader import Downloader
class FakeStream:
def __init__(self, chunks: Sequence[Sequence] | Sequence[bytes]):
def __init__(self, chunks: Sequence[bytes | tuple[bytes, float]]):
self._chunks = list(chunks)
async def read(self, _chunk_size: int) -> bytes:
@@ -25,6 +25,7 @@ class FakeStream:
payload = item[0]
delay = item[1]
assert isinstance(payload, bytes)
await asyncio.sleep(delay)
return payload
@@ -84,11 +85,11 @@ def _build_downloader(responses, *, max_retries=0):
downloader.max_retries = max_retries
downloader.base_delay = 0
fake_session = FakeSession(responses)
downloader._session = fake_session
downloader._session = fake_session # pyright: ignore[reportAttributeAccessIssue]
downloader._session_created_at = datetime.now()
downloader._proxy_url = None
async def _noop_create_session():
downloader._session = fake_session
downloader._session = fake_session # pyright: ignore[reportAttributeAccessIssue]
downloader._session_created_at = datetime.now()
downloader._proxy_url = None
@@ -96,6 +97,13 @@ def _build_downloader(responses, *, max_retries=0):
return downloader
def _session(downloader: Downloader) -> FakeSession:
"""Return the injected fake session, asserting the runtime invariant."""
session = downloader._session
assert isinstance(session, FakeSession)
return session
@pytest.mark.asyncio
async def test_download_file_preserves_incomplete_part_when_size_mismatch(tmp_path):
target_path = tmp_path / "model" / "file.bin"
@@ -196,7 +204,7 @@ async def test_download_file_recovers_from_stall(tmp_path):
assert success is True
assert Path(result_path).read_bytes() == payload
assert downloader._session._get_calls == 2
assert _session(downloader)._get_calls == 2
assert not Path(str(target_path) + ".part").exists()
@@ -224,8 +232,8 @@ async def test_download_file_resumes_after_incomplete_integrity_check(tmp_path):
assert success is True
assert Path(result_path).read_bytes() == b"abcdef"
assert downloader._session._get_calls == 2
assert downloader._session.requests[1]["headers"]["Range"] == "bytes=3-"
assert _session(downloader)._get_calls == 2
assert _session(downloader).requests[1]["headers"]["Range"] == "bytes=3-"
assert not Path(str(target_path) + ".part").exists()
@@ -261,6 +269,6 @@ async def test_download_file_retries_redirected_url_when_range_not_honored(tmp_p
assert success is True
assert Path(result_path).read_bytes() == b"abcdef"
assert first_response.released is True
assert downloader._session.requests[0]["headers"]["Range"] == "bytes=3-"
assert downloader._session.requests[1]["url"] == redirected_url
assert downloader._session.requests[1]["headers"]["Range"] == "bytes=3-"
assert _session(downloader).requests[0]["headers"]["Range"] == "bytes=3-"
assert _session(downloader).requests[1]["url"] == redirected_url
assert _session(downloader).requests[1]["headers"]["Range"] == "bytes=3-"
@@ -54,7 +54,7 @@ async def test_cleanup_moves_empty_and_orphaned(tmp_path, monkeypatch):
result = await service.cleanup_example_image_folders()
deleted_bucket = Path(result['deleted_root'])
deleted_bucket = Path(str(result['deleted_root']))
assert result['success'] is True
assert result['moved_total'] == 2
assert not empty_folder.exists()
@@ -4,6 +4,7 @@ import asyncio
import json
from pathlib import Path
from types import SimpleNamespace
from typing import Any
import pytest
@@ -15,18 +16,18 @@ class RecordingWebSocketManager:
"""Collects broadcast payloads for assertions."""
def __init__(self) -> None:
self.payloads: list[dict] = []
self.payloads: list[dict[str, Any]] = []
async def broadcast(self, payload: dict) -> None:
async def broadcast(self, payload: dict[str, Any]) -> None:
self.payloads.append(payload)
class StubScanner:
"""Scanner double returning predetermined cache contents."""
def __init__(self, models: list[dict]) -> None:
def __init__(self, models: list[dict[str, Any]]) -> None:
self._cache = SimpleNamespace(raw_data=models)
self.sync_calls: list[tuple[str, dict]] = []
self.sync_calls: list[tuple[str, dict[str, Any]]] = []
async def get_cached_data(self):
return self._cache
@@ -39,7 +40,7 @@ class StubScanner:
break
return True
async def sync_cache_from_metadata(self, file_path: str, metadata: dict) -> bool:
async def sync_cache_from_metadata(self, file_path: str, metadata: dict[str, Any]) -> bool:
self.sync_calls.append((file_path, metadata))
for index, model in enumerate(self._cache.raw_data):
if model.get("file_path") == metadata.get("file_path"):
@@ -511,7 +512,7 @@ async def test_not_found_example_images_are_cleaned(
missing_url = "https://example.com/missing.png"
valid_url = "https://example.com/valid.png"
model_metadata = {
model_metadata: dict[str, Any] = {
"sha256": model_hash,
"model_name": "Missing Example",
"file_path": str(model_path),
+4 -3
View File
@@ -2,17 +2,18 @@ import asyncio
import json
import pytest
from pathlib import Path
from typing import Any
from py.services.settings_manager import get_settings_manager
from py.utils import example_images_download_manager as download_module
class RecordingWebSocketManager:
def __init__(self) -> None:
self.payloads: list[dict] = []
async def broadcast(self, payload: dict) -> None:
self.payloads: list[dict[str, Any]] = []
async def broadcast(self, payload: dict[str, Any]) -> None:
self.payloads.append(payload)
class StubScanner:
def __init__(self, models: list[dict]) -> None:
def __init__(self, models: list[dict[str, Any]]) -> None:
self.raw_data = models
async def get_cached_data(self):
class Cache:
+9 -1
View File
@@ -1,16 +1,24 @@
"""Tests for license-based filtering functionality."""
import pytest
from typing import Any
from unittest.mock import Mock, AsyncMock
from py.services.base_model_service import BaseModelService
from py.utils.civitai_utils import build_license_flags
from py.utils.models import BaseModelMetadata
class DummyModelService(BaseModelService):
"""Dummy implementation of BaseModelService for testing."""
def __init__(self):
super().__init__(
model_type="test",
scanner=Mock(),
metadata_class=BaseModelMetadata,
settings_provider=Mock(),
)
# Mock the required attributes
self.model_type = "test"
self.scanner = Mock()
@@ -28,7 +36,7 @@ class DummyModelService(BaseModelService):
self.scanner.get_cached_data = mock_get_cached_data
async def format_response(self, model_data: dict) -> dict:
async def format_response(self, model_data: dict[str, Any]) -> dict[str, Any]:
"""Required abstract method implementation."""
return model_data
@@ -1,10 +1,12 @@
"""Integration tests for license-based filtering in BaseModelService."""
import pytest
from typing import Any
from unittest.mock import Mock, AsyncMock
from py.services.base_model_service import BaseModelService
from py.utils.civitai_utils import build_license_flags
from py.utils.models import BaseModelMetadata
from py.services.model_query import ModelCacheRepository, ModelFilterSet, SearchStrategy, SortParams
@@ -12,6 +14,12 @@ class DummyModelService(BaseModelService):
"""Dummy implementation of BaseModelService for testing."""
def __init__(self):
super().__init__(
model_type="test",
scanner=Mock(),
metadata_class=BaseModelMetadata,
settings_provider=Mock(),
)
# Mock the required attributes
self.model_type = "test"
self.scanner = Mock()
@@ -33,7 +41,7 @@ class DummyModelService(BaseModelService):
self.scanner.get_cached_data = mock_get_cached_data
async def format_response(self, model_data: dict) -> dict:
async def format_response(self, model_data: dict[str, Any]) -> dict[str, Any]:
"""Required abstract method implementation."""
return model_data
+3
View File
@@ -57,6 +57,9 @@ class MockSession:
def __init__(self, response):
self._response = response
self.closed = False
self.last_url = None
self.last_json = None
self.last_headers = None
def post(self, url, json=None, headers=None):
self.last_url = url
+3 -2
View File
@@ -1,4 +1,5 @@
from types import SimpleNamespace
from typing import Optional
from unittest.mock import AsyncMock
import pytest
@@ -21,13 +22,13 @@ class DummyProvider(ModelMetadataProvider):
async def get_model_versions_bulk(self, model_ids):
return None
async def get_model_version(self, model_id: int = None, version_id: int = None):
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None):
return None
async def get_model_version_info(self, version_id: str):
return None, None
async def get_user_models(self, username: str):
async def get_user_models(self, username: str, cursor: Optional[str] = None):
return None
+7 -7
View File
@@ -10,7 +10,7 @@ from py.services.metadata_sync_service import MetadataSyncService
class DummySettings:
def __init__(self, values: dict | None = None) -> None:
def __init__(self, values: dict[str, Any] | None = None) -> None:
self._values = values or {}
def get(self, key: str, default=None):
@@ -19,7 +19,7 @@ class DummySettings:
def build_service(
*,
settings_values: dict | None = None,
settings_values: dict[str, Any] | None = None,
default_provider: SimpleNamespace | None = None,
provider_selector: AsyncMock | None = None,
):
@@ -43,7 +43,7 @@ def build_service(
service = MetadataSyncService(
metadata_manager=metadata_manager,
preview_service=preview_service,
settings=settings,
settings=settings, # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=default_provider_factory,
metadata_provider_selector=provider_selector,
)
@@ -194,7 +194,7 @@ async def test_fetch_and_update_model_success_updates_cache(tmp_path):
helpers.metadata_manager.hydrate_model_data.side_effect = hydrate
model_data = {
model_data: Dict[str, Any] = {
"model_name": "Local",
"folder": "root",
"file_path": str(model_path),
@@ -288,9 +288,9 @@ async def test_fetch_and_update_model_handles_missing_remote_metadata(tmp_path):
helpers.metadata_manager.hydrate_model_data.side_effect = hydrate
model_data = {
model_data: Dict[str, Any] = {
"model_name": "Local",
"folder": "sub",
"folder": "root",
"file_path": str(model_path),
}
@@ -659,7 +659,7 @@ async def test_fetch_and_update_model_does_not_overwrite_api_metadata_with_archi
helpers.default_provider.get_model_by_hash.return_value = (civarchive_payload, None)
model_path = tmp_path / "model.safetensors"
model_data = {
model_data: Dict[str, Any] = {
"model_name": "High Quality",
"metadata_source": "civitai_api",
"civitai": existing_civitai,
+40 -23
View File
@@ -1,6 +1,7 @@
import json
import os
from pathlib import Path
from typing import Any, Dict, cast
import pytest
@@ -9,6 +10,11 @@ from py.utils.metadata_manager import MetadataManager
from py.utils.models import LoraMetadata
async def _empty_metadata_loader(path: str) -> Dict[str, object]:
"""Default metadata loader for tests: return an empty payload."""
return {}
class ScannerWithRoots:
def __init__(self, roots):
self._roots = list(roots)
@@ -117,7 +123,7 @@ async def test_delete_model_rejects_path_outside_roots(tmp_path: Path):
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({"civitai": {"modelId": 1}}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
# Path within root should work (model file exists)
result = await service.delete_model(str(model))
@@ -133,7 +139,7 @@ async def test_delete_model_rejects_path_outside_roots(tmp_path: Path):
service2 = ModelLifecycleService(
scanner=scanner2,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="outside configured library"):
await service2.delete_model(str(outside))
@@ -148,7 +154,7 @@ async def test_rename_model_rejects_path_outside_roots(tmp_path: Path):
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
outside = tmp_path / "outside.safetensors"
outside.write_bytes(b"data")
@@ -170,7 +176,7 @@ async def test_bulk_delete_rejects_any_path_outside_roots(tmp_path: Path):
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="outside configured library"):
await service.bulk_delete_models([str(model_ok), str(outside)])
@@ -209,7 +215,7 @@ class VersionAwareScanner:
continue
candidate = civitai.get("modelId")
try:
normalized = int(candidate)
normalized = int(cast(Any, candidate))
except (TypeError, ValueError):
continue
if normalized != model_id:
@@ -297,6 +303,7 @@ async def test_rename_model_preserves_compound_extensions(tmp_path: Path):
assert expected_main.exists()
assert not model_path.exists()
assert isinstance(result["new_file_path"], str)
assert result["new_file_path"].endswith(f"{new_name}.safetensors")
assert expected_preview.exists()
assert not preview_path.exists()
@@ -343,7 +350,7 @@ async def test_delete_model_updates_update_service(tmp_path: Path):
scanner=scanner,
metadata_manager=metadata_manager,
metadata_loader=metadata_loader,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
result = await service.delete_model(model_path.as_posix())
@@ -396,6 +403,7 @@ async def test_rename_model_preserves_extension(tmp_path: Path):
assert expected_main.exists()
assert not model_path.exists()
assert isinstance(result["new_file_path"], str)
assert result["new_file_path"].endswith(f"{new_name}{old_extension}")
assert expected_preview.exists()
assert not preview_path.exists()
@@ -448,7 +456,12 @@ async def test_rename_model_with_dotted_basename(tmp_path: Path):
expected_main = tmp_path / f"{new_name}{old_extension}"
assert expected_main.exists()
assert result["new_file_path"] == expected_main.as_posix()
assert any(p.endswith(f"{new_name}{old_extension}") for p in result["renamed_files"])
renamed_files = result["renamed_files"]
assert isinstance(renamed_files, list)
assert any(
isinstance(p, str) and p.endswith(f"{new_name}{old_extension}")
for p in renamed_files
)
saved_metadata = json.loads((tmp_path / f"{new_name}.metadata.json").read_text())
assert saved_metadata["file_name"] == new_name
@@ -490,7 +503,9 @@ async def test_delete_model_removes_gguf_file(tmp_path: Path):
assert not model_path.exists()
assert not metadata_path.exists()
assert not preview_path.exists()
assert any(item.endswith("model.gguf") for item in result["deleted_files"])
deleted_files = result["deleted_files"]
assert isinstance(deleted_files, list)
assert any(isinstance(item, str) and item.endswith("model.gguf") for item in deleted_files)
# =============================================================================
@@ -531,10 +546,10 @@ async def test_exclude_model_marks_as_excluded(tmp_path: Path):
saved_metadata = []
class SavingMetadataManager:
async def save_metadata(self, path: str, metadata: dict):
async def save_metadata(self, path: str, metadata: Dict[str, Any]):
saved_metadata.append((path, metadata.copy()))
async def metadata_loader(path: str):
async def metadata_loader(path: str) -> Dict[str, Any]:
return metadata_payload.copy()
service = ModelLifecycleService(
@@ -546,6 +561,7 @@ async def test_exclude_model_marks_as_excluded(tmp_path: Path):
result = await service.exclude_model(str(model_path))
assert result["success"] is True
assert isinstance(result["message"], str)
assert "excluded" in result["message"].lower()
assert saved_metadata[0][1]["exclude"] is True
assert str(model_path) in scanner._excluded_models
@@ -581,7 +597,7 @@ async def test_exclude_model_updates_tag_counts(tmp_path: Path):
scanner = TagCountScanner(raw_data)
class DummyMetadataManagerLocal:
async def save_metadata(self, path: str, metadata: dict):
async def save_metadata(self, path: str, metadata: Dict[str, Any]):
pass
async def metadata_loader(path: str):
@@ -607,7 +623,7 @@ async def test_exclude_model_empty_path_raises_error():
service = ModelLifecycleService(
scanner=VersionAwareScanner([]),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="Model path is required"):
@@ -645,7 +661,7 @@ async def test_unexclude_model_restores_cache_entry(tmp_path: Path):
saved_metadata = []
class SavingMetadataManager:
async def save_metadata(self, path: str, metadata: dict):
async def save_metadata(self, path: str, metadata: Dict[str, Any]):
saved_metadata.append((path, metadata.copy()))
await MetadataManager.save_metadata(path, metadata)
@@ -663,6 +679,7 @@ async def test_unexclude_model_restores_cache_entry(tmp_path: Path):
result = await service.unexclude_model(str(model_path))
assert result["success"] is True
assert isinstance(result["message"], str)
assert "restored" in result["message"].lower()
assert scanner._excluded_models == []
assert saved_metadata[0][1]["exclude"] is False
@@ -700,7 +717,7 @@ async def test_bulk_delete_models_deletes_multiple_files(tmp_path: Path):
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
result = await service.bulk_delete_models(file_paths)
@@ -716,7 +733,7 @@ async def test_bulk_delete_models_empty_list_raises_error():
service = ModelLifecycleService(
scanner=VersionAwareScanner([]),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="No file paths provided"):
@@ -734,7 +751,7 @@ async def test_delete_model_empty_path_raises_error():
service = ModelLifecycleService(
scanner=VersionAwareScanner([]),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="Model path is required"):
@@ -747,7 +764,7 @@ async def test_rename_model_empty_path_raises_error():
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="required"):
@@ -763,7 +780,7 @@ async def test_rename_model_empty_name_raises_error(tmp_path: Path):
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="required"):
@@ -779,7 +796,7 @@ async def test_rename_model_invalid_characters_raises_error(tmp_path: Path):
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
invalid_names = [
@@ -817,7 +834,7 @@ async def test_rename_model_existing_file_raises_error(tmp_path: Path):
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="already exists"):
@@ -837,7 +854,7 @@ async def test_extract_model_id_from_civitai_payload():
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
# Test civitai.modelId
@@ -863,7 +880,7 @@ async def test_extract_model_id_returns_none_for_invalid_payload():
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
assert service._extract_model_id_from_payload({}) is None
@@ -879,7 +896,7 @@ async def test_extract_model_id_handles_string_values():
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
payload = {"civitai": {"modelId": "54321"}}
+40 -3
View File
@@ -6,11 +6,12 @@ from py.services import model_metadata_provider as provider_module
from py.services.errors import RateLimitError
from py.services.model_metadata_provider import (
FallbackMetadataProvider,
ModelMetadataProvider,
RateLimitRetryingProvider,
)
class RateLimitThenSuccessProvider:
class RateLimitThenSuccessProvider(ModelMetadataProvider):
def __init__(self) -> None:
self.calls = 0
@@ -20,8 +21,20 @@ class RateLimitThenSuccessProvider:
raise RateLimitError("limited", retry_after=1.0)
return {"id": "ok"}, None
async def get_model_versions(self, model_id: str):
return None
class AlwaysRateLimitedProvider:
async def get_model_version(self, model_id=None, version_id=None):
return None
async def get_model_version_info(self, version_id: str):
return None, None
async def get_user_models(self, username: str, cursor=None):
return None
class AlwaysRateLimitedProvider(ModelMetadataProvider):
def __init__(self) -> None:
self.calls = 0
@@ -29,8 +42,20 @@ class AlwaysRateLimitedProvider:
self.calls += 1
raise RateLimitError("limited")
async def get_model_versions(self, model_id: str):
return None
class TrackingProvider:
async def get_model_version(self, model_id=None, version_id=None):
return None
async def get_model_version_info(self, version_id: str):
return None, None
async def get_user_models(self, username: str, cursor=None):
return None
class TrackingProvider(ModelMetadataProvider):
def __init__(self) -> None:
self.calls = 0
@@ -38,6 +63,18 @@ class TrackingProvider:
self.calls += 1
return {"id": "secondary"}, None
async def get_model_versions(self, model_id: str):
return None
async def get_model_version(self, model_id=None, version_id=None):
return None
async def get_model_version_info(self, version_id: str):
return None, None
async def get_user_models(self, username: str, cursor=None):
return None
@pytest.mark.asyncio
async def test_fallback_retries_same_provider_on_rate_limit(monkeypatch):
+2 -2
View File
@@ -84,11 +84,11 @@ class TestResolveSubType:
def test_none_entry_returns_default(self):
"""None entry should return default."""
assert resolve_sub_type(None) == "LORA"
assert resolve_sub_type(None) == "LORA" # pyright: ignore[reportArgumentType]
def test_non_mapping_returns_default(self):
"""Non-mapping entry should return default."""
assert resolve_sub_type("invalid") == "LORA"
assert resolve_sub_type("invalid") == "LORA" # pyright: ignore[reportArgumentType]
class TestModelFilterSetWithSubType:
+8 -8
View File
@@ -2,7 +2,7 @@ import asyncio
import os
import sqlite3
from pathlib import Path
from typing import List
from typing import Any, Dict, List, Optional
from types import MethodType
import pytest
@@ -18,9 +18,9 @@ from py.utils.models import BaseModelMetadata
class RecordingWebSocketManager:
def __init__(self) -> None:
self.payloads: List[dict] = []
self.payloads: List[Dict[str, Any]] = []
async def broadcast_init_progress(self, payload: dict) -> None:
async def broadcast_init_progress(self, payload: Dict[str, Any]) -> None:
self.payloads.append(payload)
@@ -48,7 +48,7 @@ class DummyScanner(ModelScanner):
*,
hash_index: ModelHashIndex | None = None,
excluded_models: List[str] | None = None,
) -> dict:
) -> Optional[Dict[str, Any]]:
hash_index = hash_index or self._hash_index
excluded_models = excluded_models if excluded_models is not None else self._excluded_models
@@ -483,7 +483,7 @@ async def test_version_index_tracks_version_ids(tmp_path: Path):
assert cache.version_index[202]['file_path'] == second_path
assert await scanner.check_model_version_exists(101) is True
assert await scanner.check_model_version_exists('202') is True
assert await scanner.check_model_version_exists('202') is True # pyright: ignore[reportArgumentType]
assert await scanner.check_model_version_exists(999) is False
removed = await scanner._batch_update_cache_for_deleted_models([first_path])
@@ -530,7 +530,7 @@ async def test_reconcile_cache_applies_adjust_cached_entry(tmp_path: Path):
applied: List[str] = []
def _adjust(self, entry: dict) -> dict:
def _adjust(self, entry: Dict[str, Any]) -> Dict[str, Any]:
applied.append(entry["file_path"])
entry["custom_field"] = "adjusted"
return entry
@@ -607,7 +607,7 @@ async def test_reconcile_cache_removes_duplicate_alias_when_same_real_file_seen_
scanner = MultiRootDummyScanner([loras_root, extra_root])
await scanner._initialize_cache()
duplicate_entry = {
duplicate_entry: Dict[str, Any] = {
"file_path": _normalize_path(extra_root / "one.txt"),
"folder": "",
"sha256": "hash-one",
@@ -712,7 +712,7 @@ def test_cache_entries_differ_extra_key():
# ── sync_cache_from_metadata ─────────────────────────────────────────
def _make_cache_entry(**overrides) -> dict:
def _make_cache_entry(**overrides) -> Dict[str, Any]:
entry = {
"file_path": "/m/a.safetensors",
"model_name": "TestModel",
@@ -3,13 +3,19 @@ from types import SimpleNamespace
import pytest
from py.services.model_scanner import ModelScanner
from py.utils.models import BaseModelMetadata
class DummyScanner:
class DummyScanner(ModelScanner):
def __init__(self, raw_data):
super().__init__(
model_type="dummy",
model_class=BaseModelMetadata,
file_extensions={".safetensors"},
)
self._cache = SimpleNamespace(raw_data=raw_data)
async def get_cached_data(self):
async def get_cached_data(self, force_refresh: bool = False, rebuild_cache: bool = False):
return self._cache
@@ -391,6 +391,7 @@ async def test_update_in_library_versions_changes_update_state(tmp_path):
await service.update_in_library_versions("lora", 3, [31, 35])
record = await service.get_record("lora", 3)
assert record is not None
assert record.has_update() is False
+1 -1
View File
@@ -71,7 +71,7 @@ class StubLoraScanner:
def recipe_scanner(tmp_path, monkeypatch):
monkeypatch.setattr(config, "loras_roots", [str(tmp_path)])
stub = StubLoraScanner()
scanner = RecipeScanner(lora_scanner=stub)
scanner = RecipeScanner(lora_scanner=stub) # pyright: ignore[reportArgumentType]
return scanner
@pytest.mark.asyncio
@@ -1,4 +1,5 @@
from pathlib import Path
from typing import Any, Dict
import pytest
@@ -346,7 +347,7 @@ def test_update_single_model_update_hash(tmp_path: Path, monkeypatch):
# ── get_models_missing_autov3 ─────────────────────────────────────────
def _autov3_entry(file_path: str, sha256: str, autov3=None) -> dict:
def _autov3_entry(file_path: str, sha256: str, autov3=None) -> Dict[str, Any]:
"""Minimal model entry for the models table (autov3 tri-state preserved)."""
return {
'file_path': file_path,
+5 -5
View File
@@ -1,5 +1,5 @@
from pathlib import Path
from typing import Any
from typing import Any, Dict, List
import pytest
@@ -55,7 +55,7 @@ async def test_ensure_preview_prefers_rewritten_civitai_image(tmp_path):
exif_utils=exif_utils,
)
images = [
images: List[Dict[str, object]] = [
{
"url": "https://image.civitai.com/container/example/original=true/sample.jpeg",
"type": "image",
@@ -115,7 +115,7 @@ async def test_ensure_preview_falls_back_to_webp_when_rewrite_fails(tmp_path):
exif_utils=exif_utils,
)
images = [
images: List[Dict[str, object]] = [
{
"url": "https://image.civitai.com/container/example/original=true/sample.png",
"type": "image",
@@ -165,7 +165,7 @@ async def test_ensure_preview_rewrites_civitai_video(tmp_path):
exif_utils=RecordingExifUtils(),
)
images = [
images: List[Dict[str, object]] = [
{
"url": "https://image.civitai.com/container/example/original=true/sample.mp4",
"type": "video",
@@ -227,7 +227,7 @@ async def test_ensure_preview_respects_blur_setting(monkeypatch, tmp_path):
exif_utils=RecordingExifUtils(),
)
images = [
images: List[Dict[str, object]] = [
{
"url": "https://image.civitai.com/container/example/original=true/nsfw.jpeg",
"type": "image",
+3 -1
View File
@@ -1,4 +1,6 @@
import json
from typing import Any, Dict
import pytest
from py.recipes.parsers.recipe_format import RecipeFormatParser
@@ -83,7 +85,7 @@ async def test_recipe_format_parser_marks_lora_in_library_by_version(monkeypatch
fake_metadata_provider,
)
cached_entry = {
cached_entry: Dict[str, Any] = {
"file_path": "/loras/moriimee.safetensors",
"file_name": "MoriiMee Gothic Niji | LoRA Style",
"size": 4096,
+2 -1
View File
@@ -1,5 +1,6 @@
import pytest
import asyncio
from typing import Any, Dict
from unittest.mock import AsyncMock, MagicMock
from py.services.recipe_scanner import RecipeScanner
from types import SimpleNamespace
@@ -259,7 +260,7 @@ async def test_repair_all_recipes_strips_runtime_fields(setup_scanner):
recipe_scanner, mock_civitai_client, mock_metadata_provider = setup_scanner
# Recipe with runtime fields
recipe = {
recipe: Dict[str, Any] = {
"id": "r1",
"title": "Cleanup Test",
"checkpoint": {
+8 -7
View File
@@ -3,6 +3,7 @@ import json
import os
from pathlib import Path
from types import SimpleNamespace
from typing import Any, Dict
import pytest
@@ -24,7 +25,7 @@ class StubLoraScanner:
def __init__(self) -> None:
self._hash_index = StubHashIndex()
self._hash_meta: dict[str, dict[str, str]] = {}
self._models_by_name: dict[str, dict] = {}
self._models_by_name: dict[str, Dict[str, Any]] = {}
self._cache = SimpleNamespace(raw_data=[], version_index={})
async def get_cached_data(self):
@@ -44,7 +45,7 @@ class StubLoraScanner:
async def get_model_info_by_name(self, name: str):
return self._models_by_name.get(name)
def register_model(self, name: str, info: dict) -> None:
def register_model(self, name: str, info: Dict[str, Any]) -> None:
self._models_by_name[name] = info
hash_value = (info.get("sha256") or "").lower()
version_id = info.get("civitai", {}).get("id")
@@ -76,7 +77,7 @@ def recipe_scanner(tmp_path: Path, monkeypatch):
settings_manager_module.reset_settings_manager()
monkeypatch.setattr(config, "loras_roots", [str(tmp_path)])
stub = StubLoraScanner()
scanner = RecipeScanner(lora_scanner=stub)
scanner = RecipeScanner(lora_scanner=stub) # pyright: ignore[reportArgumentType]
async def _init():
await scanner.refresh_cache(force=True)
@@ -107,7 +108,7 @@ def test_recipes_dir_uses_custom_settings_path(tmp_path: Path, monkeypatch):
manager = settings_manager_module.get_settings_manager()
manager.set("recipes_path", str(custom_recipes))
scanner = RecipeScanner(lora_scanner=StubLoraScanner())
scanner = RecipeScanner(lora_scanner=StubLoraScanner()) # pyright: ignore[reportArgumentType]
resolved = scanner.recipes_dir
assert resolved == str((tmp_path / "custom_recipes").resolve())
@@ -123,7 +124,7 @@ def test_recipes_dir_falls_back_to_first_lora_root(tmp_path: Path, monkeypatch):
monkeypatch.setattr(config, "loras_roots", [str(tmp_path / "alpha")])
scanner = RecipeScanner(lora_scanner=StubLoraScanner())
scanner = RecipeScanner(lora_scanner=StubLoraScanner()) # pyright: ignore[reportArgumentType]
resolved = scanner.recipes_dir
assert resolved == str(tmp_path / "alpha" / "recipes")
@@ -719,7 +720,7 @@ async def test_initialize_waits_for_lora_scanner(monkeypatch):
ready_flag.set()
lora_scanner = StubLoraScanner()
scanner = RecipeScanner(lora_scanner=lora_scanner)
scanner = RecipeScanner(lora_scanner=lora_scanner) # pyright: ignore[reportArgumentType]
await scanner.initialize_in_background()
@@ -736,7 +737,7 @@ async def test_invalid_model_version_marked_deleted_and_not_retried(
recipes_dir = Path(config.loras_roots[0]) / "recipes"
recipes_dir.mkdir(parents=True, exist_ok=True)
recipe = {
recipe: Dict[str, Any] = {
"id": "invalid-version",
"file_path": str(recipes_dir / "invalid-version.webp"),
"title": "Invalid",
+10 -4
View File
@@ -4,8 +4,9 @@ import os
from io import BytesIO
from pathlib import Path
from types import SimpleNamespace
from typing import Any, Dict
import piexif
import piexif # pyright: ignore[reportMissingTypeStubs]
import pytest
from PIL import Image, PngImagePlugin
@@ -463,12 +464,17 @@ async def test_save_recipe_preserves_workflow_when_png_is_converted_to_webp(tmp_
image_path = Path(result.payload["image_path"])
exif_dict = piexif.load(str(image_path))
assert exif_dict is not None
exif_0th = exif_dict["0th"]
assert exif_0th is not None
assert (
exif_dict["0th"][piexif.ImageIFD.ImageDescription].decode("utf-8")
exif_0th[piexif.ImageIFD.ImageDescription].decode("utf-8")
== 'Workflow:{"nodes":[{"id":1}]}'
)
user_comment = exif_dict["Exif"][piexif.ExifIFD.UserComment]
exif_section = exif_dict["Exif"]
assert exif_section is not None
user_comment = exif_section[piexif.ExifIFD.UserComment]
decoded_comment = user_comment[8:].decode("utf-16be")
assert "prompt text" in decoded_comment
assert "Recipe metadata:" in decoded_comment
@@ -705,7 +711,7 @@ async def test_move_recipe_updates_paths(tmp_path):
matches = list(Path(self.recipes_dir).rglob(f"{target_id}.recipe.json"))
return str(matches[0]) if matches else None
async def update_recipe_metadata(self, target_id: str, metadata: dict):
async def update_recipe_metadata(self, target_id: str, metadata: Dict[str, Any]):
if target_id != recipe_id:
return False
self.recipe.update(metadata)
+14 -9
View File
@@ -2,6 +2,7 @@ import pytest
from py.services.model_query import ModelFilterSet, FilterCriteria
from py.services.recipe_scanner import RecipeScanner
from types import SimpleNamespace
from typing import Any, cast
# Mock settings
@@ -193,9 +194,9 @@ async def test_recipe_scanner_root_recursive_true():
async def get_cached_data(self):
return SimpleNamespace(raw_data=[])
scanner = RecipeScanner(lora_scanner=StubLoraScanner())
scanner = RecipeScanner(lora_scanner=StubLoraScanner()) # pyright: ignore[reportArgumentType]
# Manually populate cache for testing get_paginated_data logic
scanner._cache = SimpleNamespace(
scanner._cache = cast(Any, SimpleNamespace(
raw_data=[
{
"id": "r1",
@@ -234,13 +235,15 @@ async def test_recipe_scanner_root_recursive_true():
],
sorted_by_name=[],
version_index={},
)
))
result = await scanner.get_paginated_data(
page=1, page_size=10, folder="", recursive=True
)
assert len(result["items"]) == 2
items = result["items"]
assert isinstance(items, list)
assert len(items) == 2
@pytest.mark.asyncio
@@ -250,8 +253,8 @@ async def test_recipe_scanner_root_recursive_false():
async def get_cached_data(self):
return SimpleNamespace(raw_data=[])
scanner = RecipeScanner(lora_scanner=StubLoraScanner())
scanner._cache = SimpleNamespace(
scanner = RecipeScanner(lora_scanner=StubLoraScanner()) # pyright: ignore[reportArgumentType]
scanner._cache = cast(Any, SimpleNamespace(
raw_data=[
{
"id": "r1",
@@ -290,11 +293,13 @@ async def test_recipe_scanner_root_recursive_false():
],
sorted_by_name=[],
version_index={},
)
))
result = await scanner.get_paginated_data(
page=1, page_size=10, folder="", recursive=False
)
assert len(result["items"]) == 1
assert result["items"][0]["id"] == "r1"
items = result["items"]
assert isinstance(items, list)
assert len(items) == 1
assert items[0]["id"] == "r1"
@@ -51,10 +51,10 @@ class DummyProvider:
def __init__(self, payload: Dict[str, Any]) -> None:
self.payload = payload
async def get_model_by_hash(self, sha256: str):
async def get_model_by_hash(self, model_hash: str):
return self.payload, None
async def get_model_version(self, model_id: int, model_version_id: int | None):
async def get_model_version(self, model_id: Any = None, version_id: Any = None):
return self.payload
@@ -77,7 +77,7 @@ def test_metadata_sync_merges_remote_fields(tmp_path: Path) -> None:
service = MetadataSyncService(
metadata_manager=manager,
preview_service=preview,
settings=DummySettings(),
settings=DummySettings(), # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=lambda: asyncio.sleep(0, result=provider),
metadata_provider_selector=lambda _name=None: asyncio.sleep(0, result=provider),
)
@@ -112,7 +112,7 @@ def test_metadata_sync_fetch_and_update_updates_cache(tmp_path: Path) -> None:
service = MetadataSyncService(
metadata_manager=manager,
preview_service=preview,
settings=DummySettings(),
settings=DummySettings(), # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=lambda: asyncio.sleep(0, result=provider),
metadata_provider_selector=lambda _name=None: asyncio.sleep(0, result=provider),
)
+1 -1
View File
@@ -52,7 +52,7 @@ async def test_lazy_loaded_scanners(monkeypatch, method_name, module_path, class
async def test_lazy_loaded_websocket_manager(monkeypatch):
fake_manager = object()
module = types.ModuleType("py.services.websocket_manager")
module.ws_manager = fake_manager
setattr(module, "ws_manager", fake_manager)
monkeypatch.setitem(sys.modules, "py.services.websocket_manager", module)
first = await ServiceRegistry.get_websocket_manager()
+2 -2
View File
@@ -482,7 +482,7 @@ def test_model_name_display_setting_notifies_scanners(tmp_path, monkeypatch):
manager = _create_manager_with_settings(tmp_path, monkeypatch, initial)
loop = asyncio.new_event_loop()
loop._thread_id = 1
setattr(loop, "_thread_id", 1)
class DummyScanner:
def __init__(self):
@@ -530,7 +530,7 @@ def test_model_name_display_setting_notifies_scanners(tmp_path, monkeypatch):
assert dummy_scanner.calls == ["file_name"]
assert dispatched_loops == [dummy_scanner.loop]
finally:
loop._thread_id = None
setattr(loop, "_thread_id", None)
loop.close()
@@ -8,6 +8,8 @@ from py.recipes.parsers import SuiImageParamsParser
class TestSuiImageParamsParser:
"""Test cases for SuiImageParamsParser."""
parser: SuiImageParamsParser = SuiImageParamsParser()
def setup_method(self):
"""Set up test fixtures."""
self.parser = SuiImageParamsParser()
@@ -116,6 +118,7 @@ class TestSuiImageParamsParser:
result = await self.parser.parse_metadata(metadata_str)
loras = result.get('loras')
assert isinstance(loras, list)
assert len(loras) == 1
assert loras[0]['type'] == 'lora'
assert loras[0]['name'] == 'test_lora'
@@ -142,6 +145,7 @@ class TestSuiImageParamsParser:
result = await self.parser.parse_metadata(metadata_str)
loras = result.get('loras')
assert isinstance(loras, list)
assert len(loras) == 1
assert loras[0]['type'] == 'lora'
+25 -10
View File
@@ -1,11 +1,14 @@
import asyncio
import logging
from dataclasses import dataclass
from types import SimpleNamespace
from typing import Any, Dict, List, Optional
import pytest
from py.services.model_file_service import AutoOrganizeResult
from py.services.download_coordinator import DownloadCoordinator
from py.services.metadata_sync_service import MetadataSyncService
from py.services.model_file_service import AutoOrganizeResult, ModelFileService
from py.services.use_cases import (
AutoOrganizeInProgressError,
AutoOrganizeUseCase,
@@ -26,6 +29,7 @@ from py.utils.example_images_download_manager import (
)
from py.utils.example_images_processor import (
ExampleImagesImportError,
ExampleImagesProcessor,
ExampleImagesValidationError,
)
from py.utils.metadata_manager import MetadataManager
@@ -44,13 +48,13 @@ class StubLockProvider:
return self._lock
class StubFileService:
class StubFileService(ModelFileService):
def __init__(self) -> None:
super().__init__(scanner=None, model_type="lora")
self.calls: List[Dict[str, Any]] = []
async def auto_organize_models(
self,
*,
file_paths: Optional[List[str]] = None,
progress_callback=None,
exclusion_patterns=None,
@@ -65,8 +69,15 @@ class StubFileService:
return result
class StubMetadataSync:
class StubMetadataSync(MetadataSyncService):
def __init__(self) -> None:
super().__init__(
metadata_manager=object(),
preview_service=object(),
settings=StubSettings(), # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=lambda: asyncio.sleep(0, result=None), # pyright: ignore[reportArgumentType]
metadata_provider_selector=lambda _name=None: asyncio.sleep(0, result=None), # pyright: ignore[reportArgumentType]
)
self.calls: List[Dict[str, Any]] = []
async def fetch_and_update_model(self, **kwargs: Any):
@@ -94,8 +105,12 @@ class ProgressCollector:
self.events.append(payload)
class StubDownloadCoordinator:
class StubDownloadCoordinator(DownloadCoordinator):
def __init__(self, *, error: Optional[str] = None) -> None:
super().__init__(
ws_manager=SimpleNamespace(generate_download_id=lambda: "abc123"),
download_manager_factory=lambda: asyncio.sleep(0, result=None),
)
self.error = error
self.payloads: List[Dict[str, Any]] = []
@@ -125,13 +140,13 @@ class StubExampleImagesDownloadManager:
return {"success": True, "message": "ok"}
class StubExampleImagesProcessor:
class StubExampleImagesProcessor(ExampleImagesProcessor):
def __init__(self) -> None:
self.calls: List[Dict[str, Any]] = []
self.error: Optional[str] = None
self.response: Dict[str, Any] = {"success": True}
async def import_images(self, model_hash: str, files: List[str]) -> Dict[str, Any]:
async def import_images(self, model_hash: str, files: List[str]) -> Dict[str, Any]: # pyright: ignore[reportIncompatibleMethodOverride]
self.calls.append({"model_hash": model_hash, "files": files})
if self.error == "validation":
raise ExampleImagesValidationError("missing")
@@ -464,7 +479,7 @@ async def test_import_example_images_use_case_delegates() -> None:
use_case = ImportExampleImagesUseCase(processor=processor)
request = DummyJsonRequest({"model_hash": "abc", "file_paths": ["/tmp/file"]})
result = await use_case.execute(request)
result = await use_case.execute(request) # pyright: ignore[reportArgumentType]
assert processor.calls == [{"model_hash": "abc", "files": ["/tmp/file"]}]
assert result == {"success": True}
@@ -477,7 +492,7 @@ async def test_import_example_images_use_case_maps_validation_error() -> None:
request = DummyJsonRequest({"model_hash": None, "file_paths": []})
with pytest.raises(ImportExampleImagesValidationError):
await use_case.execute(request)
await use_case.execute(request) # pyright: ignore[reportArgumentType]
async def test_import_example_images_use_case_propagates_generic_error() -> None:
@@ -487,4 +502,4 @@ async def test_import_example_images_use_case_propagates_generic_error() -> None
request = DummyJsonRequest({"model_hash": "abc", "file_paths": ["/tmp/file"]})
with pytest.raises(ExampleImagesImportError):
await use_case.execute(request)
await use_case.execute(request) # pyright: ignore[reportArgumentType]