mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-14 09:43:22 -03:00
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:
@@ -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(
|
||||
|
||||
@@ -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, '')]
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"}}
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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'
|
||||
|
||||
|
||||
@@ -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]
|
||||
Reference in New Issue
Block a user