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
+16 -4
View File
@@ -1,5 +1,5 @@
import logging import logging
from typing import Dict, Iterable, List from typing import Any, Dict, Iterable, List
import pytest import pytest
@@ -146,6 +146,8 @@ def test_save_paths_repairs_empty_default_roots(monkeypatch: pytest.MonkeyPatch,
class FakeSettingsService: class FakeSettingsService:
active_library = "comfyui" active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self): def get_libraries(self):
return { return {
@@ -183,6 +185,8 @@ def test_save_paths_repairs_stale_default_roots(monkeypatch: pytest.MonkeyPatch,
class FakeSettingsService: class FakeSettingsService:
active_library = "comfyui" active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self): def get_libraries(self):
return { return {
@@ -220,6 +224,8 @@ def test_save_paths_keeps_valid_default_roots(monkeypatch: pytest.MonkeyPatch, t
class FakeSettingsService: class FakeSettingsService:
active_library = "comfyui" active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self): def get_libraries(self):
return { return {
@@ -357,6 +363,8 @@ def test_save_paths_keeps_default_roots_in_extra_paths(monkeypatch: pytest.Monke
class FakeSettingsService: class FakeSettingsService:
active_library = "comfyui" active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self): def get_libraries(self):
return { return {
@@ -409,6 +417,8 @@ def test_save_paths_keeps_default_roots_in_extra_paths_with_windows_slash_mismat
class FakeSettingsService: class FakeSettingsService:
active_library = "comfyui" active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self): def get_libraries(self):
return { return {
@@ -460,6 +470,8 @@ def test_save_paths_repairs_empty_default_roots_to_extra_paths_when_primary_miss
class FakeSettingsService: class FakeSettingsService:
active_library = "comfyui" active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self): def get_libraries(self):
return { return {
@@ -577,7 +589,7 @@ def test_apply_library_settings_merges_extra_paths(monkeypatch, tmp_path):
assert str(extra_loras_dir) in config_instance.extra_loras_roots assert str(extra_loras_dir) in config_instance.extra_loras_roots
assert str(checkpoints_dir) in config_instance.base_models_roots assert str(checkpoints_dir) in config_instance.base_models_roots
assert str(extra_checkpoints_dir) in config_instance.extra_checkpoints_roots assert str(extra_checkpoints_dir) in config_instance.extra_checkpoints_roots
assert str(embeddings_dir) in config_instance.embeddings_roots assert str(embeddings_dir) in (config_instance.embeddings_roots or [])
assert str(extra_embeddings_dir) in config_instance.extra_embeddings_roots assert str(extra_embeddings_dir) in config_instance.extra_embeddings_roots
@@ -609,7 +621,7 @@ def test_apply_library_settings_without_extra_paths(monkeypatch, tmp_path):
assert config_instance.extra_loras_roots == [] assert config_instance.extra_loras_roots == []
assert str(checkpoints_dir) in config_instance.base_models_roots assert str(checkpoints_dir) in config_instance.base_models_roots
assert config_instance.extra_checkpoints_roots == [] assert config_instance.extra_checkpoints_roots == []
assert str(embeddings_dir) in config_instance.embeddings_roots assert str(embeddings_dir) in (config_instance.embeddings_roots or [])
assert config_instance.extra_embeddings_roots == [] assert config_instance.extra_embeddings_roots == []
@@ -858,7 +870,7 @@ def test_save_paths_removes_stale_empty_default_when_comfyui_exists(
# dict order, returning "default". # dict order, returning "default".
self.active_library = "default" self.active_library = "default"
self.delete_calls: list[str] = [] self.delete_calls: list[str] = []
self.upsert_calls: list[tuple[str, dict]] = [] self.upsert_calls: list[tuple[str, dict[str, Any]]] = []
def get_libraries(self): def get_libraries(self):
return dict(self.libraries) return dict(self.libraries)
+17 -17
View File
@@ -63,23 +63,23 @@ sys.modules.setdefault("py_local", _repo_package)
# Mock ComfyUI modules before any imports from the main project # Mock ComfyUI modules before any imports from the main project
server_mock = MockModule("server") server_mock = MockModule("server")
server_mock.PromptServer = mock.MagicMock() setattr(server_mock, "PromptServer", mock.MagicMock())
sys.modules['server'] = server_mock sys.modules['server'] = server_mock
folder_paths_mock = MockModule("folder_paths") folder_paths_mock = MockModule("folder_paths")
folder_paths_mock.get_folder_paths = mock.MagicMock(return_value=[]) setattr(folder_paths_mock, "get_folder_paths", mock.MagicMock(return_value=[]))
folder_paths_mock.folder_names_and_paths = {} setattr(folder_paths_mock, "folder_names_and_paths", {})
sys.modules['folder_paths'] = folder_paths_mock sys.modules['folder_paths'] = folder_paths_mock
# Mock other ComfyUI modules that might be imported # Mock other ComfyUI modules that might be imported
comfy_mock = MockModule("comfy") comfy_mock = MockModule("comfy")
comfy_mock.utils = MockModule("comfy.utils") setattr(comfy_mock, "utils", MockModule("comfy.utils"))
comfy_mock.utils.load_torch_file = mock.MagicMock(return_value={}) setattr(comfy_mock.utils, "load_torch_file", mock.MagicMock(return_value={}))
comfy_mock.sd = MockModule("comfy.sd") setattr(comfy_mock, "sd", MockModule("comfy.sd"))
comfy_mock.sd.load_lora_for_models = mock.MagicMock(return_value=(None, None)) setattr(comfy_mock.sd, "load_lora_for_models", mock.MagicMock(return_value=(None, None)))
comfy_mock.model_management = MockModule("comfy.model_management") setattr(comfy_mock, "model_management", MockModule("comfy.model_management"))
comfy_mock.comfy_types = MockModule("comfy.comfy_types") setattr(comfy_mock, "comfy_types", MockModule("comfy.comfy_types"))
comfy_mock.comfy_types.IO = mock.MagicMock() setattr(comfy_mock.comfy_types, "IO", mock.MagicMock())
sys.modules['comfy'] = comfy_mock sys.modules['comfy'] = comfy_mock
sys.modules['comfy.utils'] = comfy_mock.utils sys.modules['comfy.utils'] = comfy_mock.utils
sys.modules['comfy.sd'] = comfy_mock.sd sys.modules['comfy.sd'] = comfy_mock.sd
@@ -88,14 +88,14 @@ sys.modules['comfy.comfy_types'] = comfy_mock.comfy_types
sys.modules['comfy.hooks'] = MockModule("comfy.hooks") sys.modules['comfy.hooks'] = MockModule("comfy.hooks")
execution_mock = MockModule("execution") execution_mock = MockModule("execution")
execution_mock.PromptExecutor = mock.MagicMock() setattr(execution_mock, "PromptExecutor", mock.MagicMock())
sys.modules['execution'] = execution_mock sys.modules['execution'] = execution_mock
# Mock ComfyUI nodes module # Mock ComfyUI nodes module
nodes_mock = MockModule("nodes") nodes_mock = MockModule("nodes")
nodes_mock.LoraLoader = mock.MagicMock() setattr(nodes_mock, "LoraLoader", mock.MagicMock())
nodes_mock.SaveImage = mock.MagicMock() setattr(nodes_mock, "SaveImage", mock.MagicMock())
nodes_mock.NODE_CLASS_MAPPINGS = {} setattr(nodes_mock, "NODE_CLASS_MAPPINGS", {})
sys.modules['nodes'] = nodes_mock sys.modules['nodes'] = nodes_mock
@@ -347,7 +347,7 @@ def reset_singletons():
# Reset ServiceRegistry # Reset ServiceRegistry
ServiceRegistry._services = {} ServiceRegistry._services = {}
ServiceRegistry._initialized = False ServiceRegistry._initialized = False # pyright: ignore[reportAttributeAccessIssue]
# Reset ModelScanner instances # Reset ModelScanner instances
if hasattr(ModelScanner, '_instances'): if hasattr(ModelScanner, '_instances'):
@@ -356,14 +356,14 @@ def reset_singletons():
# Reset SettingsManager # Reset SettingsManager
settings_manager = get_settings_manager() settings_manager = get_settings_manager()
if hasattr(settings_manager, '_reset'): if hasattr(settings_manager, '_reset'):
settings_manager._reset() settings_manager._reset() # pyright: ignore[reportAttributeAccessIssue]
yield yield
# Cleanup after test # Cleanup after test
DownloadManager._instance = None DownloadManager._instance = None
ServiceRegistry._services = {} ServiceRegistry._services = {}
ServiceRegistry._initialized = False ServiceRegistry._initialized = False # pyright: ignore[reportAttributeAccessIssue]
if hasattr(ModelScanner, '_instances'): if hasattr(ModelScanner, '_instances'):
ModelScanner._instances.clear() ModelScanner._instances.clear()
+7 -7
View File
@@ -12,7 +12,7 @@ from __future__ import annotations
import json import json
import re import re
from pathlib import Path from pathlib import Path
from typing import Dict, Iterable, Set from typing import Any, Dict, Iterable, Set
import pytest import pytest
@@ -105,9 +105,9 @@ HTML_TRANSLATION_PATTERN = (
@pytest.fixture(scope="module") @pytest.fixture(scope="module")
def loaded_locales() -> Dict[str, dict]: def loaded_locales() -> Dict[str, Any]:
"""Load locale JSON once per test module.""" """Load locale JSON once per test module."""
locales: Dict[str, dict] = {} locales: Dict[str, Any] = {}
for locale in EXPECTED_LOCALES: for locale in EXPECTED_LOCALES:
path = LOCALES_DIR / f"{locale}.json" path = LOCALES_DIR / f"{locale}.json"
@@ -131,7 +131,7 @@ def loaded_locales() -> Dict[str, dict]:
@pytest.fixture(scope="module") @pytest.fixture(scope="module")
def english_translation_keys(loaded_locales: Dict[str, dict]) -> Set[str]: def english_translation_keys(loaded_locales: Dict[str, Any]) -> Set[str]:
return collect_translation_keys(loaded_locales["en"]) return collect_translation_keys(loaded_locales["en"])
@@ -140,7 +140,7 @@ def static_code_translation_keys() -> Set[str]:
return gather_static_translation_keys() return gather_static_translation_keys()
def collect_translation_keys(data: dict, prefix: str = "") -> Set[str]: def collect_translation_keys(data: Dict[str, Any], prefix: str = "") -> Set[str]:
"""Recursively collect translation keys from a locale dictionary.""" """Recursively collect translation keys from a locale dictionary."""
keys: Set[str] = set() keys: Set[str] = set()
@@ -212,7 +212,7 @@ def extract_i18n_keys_from_html(file_path: Path) -> Set[str]:
@pytest.mark.parametrize("locale", EXPECTED_LOCALES) @pytest.mark.parametrize("locale", EXPECTED_LOCALES)
def test_locale_files_have_expected_structure(locale: str, loaded_locales: Dict[str, dict]) -> None: def test_locale_files_have_expected_structure(locale: str, loaded_locales: Dict[str, Any]) -> None:
"""Every locale must contain the required sections.""" """Every locale must contain the required sections."""
data = loaded_locales[locale] data = loaded_locales[locale]
missing_sections = sorted(REQUIRED_SECTIONS - data.keys()) missing_sections = sorted(REQUIRED_SECTIONS - data.keys())
@@ -221,7 +221,7 @@ def test_locale_files_have_expected_structure(locale: str, loaded_locales: Dict[
@pytest.mark.parametrize("locale", EXPECTED_LOCALES[1:]) @pytest.mark.parametrize("locale", EXPECTED_LOCALES[1:])
def test_locale_keys_match_english( def test_locale_keys_match_english(
locale: str, loaded_locales: Dict[str, dict], english_translation_keys: Set[str] locale: str, loaded_locales: Dict[str, Any], english_translation_keys: Set[str]
) -> None: ) -> None:
"""Locales must expose the same translation keys as English.""" """Locales must expose the same translation keys as English."""
locale_keys = collect_translation_keys(loaded_locales[locale]) locale_keys = collect_translation_keys(loaded_locales[locale])
+2 -2
View File
@@ -124,7 +124,7 @@ def mock_metadata_manager():
"""Provide a mock metadata manager.""" """Provide a mock metadata manager."""
class MockMetadataManager: class MockMetadataManager:
def __init__(self): def __init__(self):
self.saved_metadata: List[tuple] = [] self.saved_metadata: List[tuple[str, Any]] = []
self.loaded_payloads: Dict[str, Dict[str, Any]] = {} self.loaded_payloads: Dict[str, Dict[str, Any]] = {}
async def save_metadata(self, file_path: str, metadata: Dict[str, Any]) -> None: async def save_metadata(self, file_path: str, metadata: Dict[str, Any]) -> None:
@@ -194,7 +194,7 @@ async def test_http_server(
site = web.TCPSite(runner, "127.0.0.1", 0) site = web.TCPSite(runner, "127.0.0.1", 0)
await site.start() await site.start()
port = site._server.sockets[0].getsockname()[1] port = site._server.sockets[0].getsockname()[1] # pyright: ignore[reportAttributeAccessIssue, reportOptionalMemberAccess]
base_url = f"http://127.0.0.1:{port}" base_url = f"http://127.0.0.1:{port}"
yield base_url, port yield base_url, port
+1 -1
View File
@@ -193,7 +193,7 @@ class TestDownloadRouteIntegration:
assert response.status == 400 assert response.status == 400
# Response might be JSON or text, check both # Response might be JSON or text, check both
if hasattr(response, 'text'): if hasattr(response, 'text') and response.text is not None:
error_text = response.text.lower() error_text = response.text.lower()
else: else:
body = response.body body = response.body
+2 -1
View File
@@ -75,6 +75,7 @@ class TestRecipeFlowIntegration:
# Verify update # Verify update
loaded = cache.load_cache() loaded = cache.load_cache()
assert loaded is not None
loaded_recipe = loaded.raw_data[0] loaded_recipe = loaded.raw_data[0]
assert loaded_recipe["title"] == "Updated Recipe Title" assert loaded_recipe["title"] == "Updated Recipe Title"
@@ -243,7 +244,7 @@ steps: 20
cfg: 7.0""" cfg: 7.0"""
# Basic parsing logic for testing # Basic parsing logic for testing
def parse_simple_metadata(text: str) -> dict: def parse_simple_metadata(text: str) -> Dict[str, str]:
result = {} result = {}
for line in text.strip().split('\n'): for line in text.strip().split('\n'):
if ':' in line: if ':' in line:
+7 -7
View File
@@ -44,13 +44,13 @@ def populated_registry(metadata_registry):
# Direct assignment to avoid scanner.py false positive # Direct assignment to avoid scanner.py false positive
# (scanner.py matches _CLASS_MAPPINGS.update({...}) pattern) # (scanner.py matches _CLASS_MAPPINGS.update({...}) pattern)
nodes.NODE_CLASS_MAPPINGS["TSC_EfficientLoader"] = TSC_EfficientLoader nodes.NODE_CLASS_MAPPINGS["TSC_EfficientLoader"] = TSC_EfficientLoader # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["SamplerCustomAdvanced"] = SamplerCustomAdvanced nodes.NODE_CLASS_MAPPINGS["SamplerCustomAdvanced"] = SamplerCustomAdvanced # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["BasicScheduler"] = BasicScheduler nodes.NODE_CLASS_MAPPINGS["BasicScheduler"] = BasicScheduler # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["KSamplerSelect"] = KSamplerSelect nodes.NODE_CLASS_MAPPINGS["KSamplerSelect"] = KSamplerSelect # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["CFGGuider"] = CFGGuider nodes.NODE_CLASS_MAPPINGS["CFGGuider"] = CFGGuider # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["CLIPTextEncode"] = CLIPTextEncode nodes.NODE_CLASS_MAPPINGS["CLIPTextEncode"] = CLIPTextEncode # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["VAEDecode"] = VAEDecode nodes.NODE_CLASS_MAPPINGS["VAEDecode"] = VAEDecode # pyright: ignore[reportAttributeAccessIssue]
prompt_graph = { prompt_graph = {
"loader": {"class_type": "TSC_EfficientLoader", "inputs": {}}, "loader": {"class_type": "TSC_EfficientLoader", "inputs": {}},
@@ -1,6 +1,7 @@
import sys import sys
import types import types
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any, Dict
from py.metadata_collector import metadata_processor from py.metadata_collector import metadata_processor
from py.metadata_collector.metadata_hook import MetadataHook from py.metadata_collector.metadata_hook import MetadataHook
@@ -44,6 +45,7 @@ def test_metadata_hook_installs_and_traces_execution(monkeypatch, metadata_regis
class FakeNode: class FakeNode:
FUNCTION = "run" FUNCTION = "run"
unique_id: str = ""
node = FakeNode() node = FakeNode()
node.unique_id = "node-1" node.unique_id = "node-1"
@@ -702,7 +704,7 @@ def test_lora_manager_cache_updates_when_loras_removed(metadata_registry):
class LoraLoaderLM: # type: ignore[too-many-ancestors] class LoraLoaderLM: # type: ignore[too-many-ancestors]
__name__ = "LoraLoaderLM" __name__ = "LoraLoaderLM"
nodes.NODE_CLASS_MAPPINGS["LoraLoaderLM"] = LoraLoaderLM nodes.NODE_CLASS_MAPPINGS["LoraLoaderLM"] = LoraLoaderLM # pyright: ignore[reportAttributeAccessIssue]
prompt_graph = { prompt_graph = {
"lora_node": {"class_type": "LoraLoaderLM", "inputs": {}}, "lora_node": {"class_type": "LoraLoaderLM", "inputs": {}},
@@ -883,7 +885,7 @@ def test_metadata_overwrite_extractor_empty_inputs(metadata_registry):
from py.metadata_collector.constants import CLIP_SKIP_SENTINEL from py.metadata_collector.constants import CLIP_SKIP_SENTINEL
inputs = {key: "" for key in METADATA_OVERWRITE_FIELDS} inputs: Dict[str, Any] = {key: "" for key in METADATA_OVERWRITE_FIELDS}
inputs.update({"seed": 0, "steps": 0, "cfg_scale": 0.0, "clip_skip": CLIP_SKIP_SENTINEL}) inputs.update({"seed": 0, "steps": 0, "cfg_scale": 0.0, "clip_skip": CLIP_SKIP_SENTINEL})
MetadataOverwriteExtractor.extract("ow-2", inputs, None, metadata) MetadataOverwriteExtractor.extract("ow-2", inputs, None, metadata)
+6 -5
View File
@@ -9,6 +9,7 @@ Mock targets must match where imports are resolved inside each function
from __future__ import annotations from __future__ import annotations
from typing import Any
from unittest import mock from unittest import mock
import pytest import pytest
@@ -28,14 +29,14 @@ from py.metadata_ops import (
class MockCache: class MockCache:
def __init__(self, raw_data: list[dict] | None = None): def __init__(self, raw_data: list[dict[str, Any]] | None = None):
self.raw_data = raw_data or [] self.raw_data = raw_data or []
class MockScanner: class MockScanner:
"""Simulates a ModelScanner for testing.""" """Simulates a ModelScanner for testing."""
def __init__(self, raw_data: list[dict] | None = None): def __init__(self, raw_data: list[dict[str, Any]] | None = None):
self._raw_data = raw_data or [] self._raw_data = raw_data or []
self.update_single_model_cache = mock.AsyncMock(return_value=True) self.update_single_model_cache = mock.AsyncMock(return_value=True)
@@ -599,7 +600,7 @@ class TestExtractGalleryTableImages:
""" """
@staticmethod @staticmethod
def _extract(md: str, repo: str = _REPO, existing: set | None = None): def _extract(md: str, repo: str = _REPO, existing: set[str] | None = None):
from py.services.agent.skills.enrich_hf_metadata.readme_processor import \ from py.services.agent.skills.enrich_hf_metadata.readme_processor import \
extract_gallery_table_images extract_gallery_table_images
return extract_gallery_table_images(md, repo, existing_urls=existing) return extract_gallery_table_images(md, repo, existing_urls=existing)
@@ -643,7 +644,7 @@ class TestExtractGalleryTableImages:
class TestCleanReadmeForLlm: class TestCleanReadmeForLlm:
@staticmethod @staticmethod
def _clean(md: str, max_length: int = 6000) -> str: def _clean(md: str | None, max_length: int = 6000) -> str:
from py.services.agent.skills.enrich_hf_metadata.readme_processor import \ from py.services.agent.skills.enrich_hf_metadata.readme_processor import \
clean_readme_for_llm clean_readme_for_llm
return clean_readme_for_llm(md, max_length=max_length) return clean_readme_for_llm(md, max_length=max_length)
@@ -651,7 +652,7 @@ class TestCleanReadmeForLlm:
# -- basic guards -------------------------------------------------------- # -- basic guards --------------------------------------------------------
def test_none_returns_empty(self): def test_none_returns_empty(self):
assert self._clean(None) == "" # type: ignore[arg-type] assert self._clean(None) == ""
def test_empty_returns_empty(self): def test_empty_returns_empty(self):
assert self._clean("") == "" assert self._clean("") == ""
@@ -19,6 +19,8 @@ _MODULE_PATH = Path(__file__).parents[2] / "py" / "services" / "agent" / "skills
def R(): def R():
"""Load the ``readme_processor`` module once per session.""" """Load the ``readme_processor`` module once per session."""
spec = importlib.util.spec_from_file_location("readme_processor", str(_MODULE_PATH)) spec = importlib.util.spec_from_file_location("readme_processor", str(_MODULE_PATH))
assert spec is not None
assert spec.loader is not None
mod = importlib.util.module_from_spec(spec) mod = importlib.util.module_from_spec(spec)
spec.loader.exec_module(mod) spec.loader.exec_module(mod)
return mod return mod
+1 -1
View File
@@ -32,7 +32,7 @@ def _parse_directives(header: str) -> dict[str, list[str]]:
async def _invoke_middleware( async def _invoke_middleware(
path: str, response: web.Response, csp_header: str | None = DEFAULT_CSP path: str, response: web.Response, csp_header: str | None = DEFAULT_CSP
) -> web.Response: ) -> web.StreamResponse:
async def handler(_request: web.Request) -> web.Response: async def handler(_request: web.Request) -> web.Response:
if csp_header is not None: if csp_header is not None:
response.headers["Content-Security-Policy"] = csp_header response.headers["Content-Security-Policy"] = csp_header
+2 -2
View File
@@ -31,7 +31,7 @@ class _DummyModel:
return self return self
def test_nunchaku_load_lora_legacy_fallback(monkeypatch, caplog): def test_nunchaku_load_lora_legacy_fallback(monkeypatch, caplog):
import folder_paths import folder_paths # pyright: ignore[reportMissingImports]
import copy import copy
dummy_model = _DummyModel() dummy_model = _DummyModel()
@@ -59,7 +59,7 @@ def test_nunchaku_load_lora_legacy_fallback(monkeypatch, caplog):
assert result_model.model.diffusion_model.loras[0][1] == 0.8 assert result_model.model.diffusion_model.loras[0][1] == 0.8
def test_nunchaku_load_lora_new_logic(monkeypatch): def test_nunchaku_load_lora_new_logic(monkeypatch):
import folder_paths import folder_paths # pyright: ignore[reportMissingImports]
import os import os
dummy_model = _DummyModel() dummy_model = _DummyModel()
+4 -2
View File
@@ -1,5 +1,7 @@
from __future__ import annotations from __future__ import annotations
from typing import Any, cast
from py.nodes.prompt import PromptLM from py.nodes.prompt import PromptLM
from py.nodes.text import TextLM from py.nodes.text import TextLM
@@ -49,7 +51,7 @@ def test_prompt_lm_input_types_expose_input_only_seed():
assert seed_type == "INT" assert seed_type == "INT"
assert seed_options["forceInput"] is True assert seed_options["forceInput"] is True
assert "wildcard generation" in seed_options["tooltip"] assert "wildcard generation" in cast(Any, seed_options)["tooltip"]
def test_text_lm_input_types_expose_input_only_seed(): def test_text_lm_input_types_expose_input_only_seed():
@@ -58,7 +60,7 @@ def test_text_lm_input_types_expose_input_only_seed():
assert seed_type == "INT" assert seed_type == "INT"
assert seed_options["forceInput"] is True assert seed_options["forceInput"] is True
assert "wildcard generation" in seed_options["tooltip"] assert "wildcard generation" in cast(Any, seed_options)["tooltip"]
def test_text_lm_is_changed_forces_rerun_without_seed_when_text_is_dynamic(): def test_text_lm_is_changed_forces_rerun_without_seed_when_text_is_dynamic():
+24 -12
View File
@@ -1,8 +1,9 @@
import json import json
import os import os
from typing import Any, cast
import numpy as np import numpy as np
import piexif import piexif # pyright: ignore[reportMissingTypeStubs]
from PIL import Image from PIL import Image
from py.services.service_registry import ServiceRegistry from py.services.service_registry import ServiceRegistry
@@ -136,7 +137,8 @@ def test_save_image_skips_jpeg_metadata_when_disabled(monkeypatch, tmp_path):
image_path = tmp_path / "sample_00001_.jpg" image_path = tmp_path / "sample_00001_.jpg"
exif_dict = piexif.load(str(image_path)) exif_dict = piexif.load(str(image_path))
assert piexif.ExifIFD.UserComment not in exif_dict.get("Exif", {}) exif_ifd = exif_dict.get("Exif", {}) or {}
assert piexif.ExifIFD.UserComment not in exif_ifd
def test_save_image_skips_webp_metadata_when_disabled(monkeypatch, tmp_path): def test_save_image_skips_webp_metadata_when_disabled(monkeypatch, tmp_path):
@@ -154,7 +156,8 @@ def test_save_image_skips_webp_metadata_when_disabled(monkeypatch, tmp_path):
image_path = tmp_path / "sample_00001_.webp" image_path = tmp_path / "sample_00001_.webp"
exif_dict = piexif.load(str(image_path)) exif_dict = piexif.load(str(image_path))
assert piexif.ExifIFD.UserComment not in exif_dict.get("Exif", {}) exif_ifd = exif_dict.get("Exif", {}) or {}
assert piexif.ExifIFD.UserComment not in exif_ifd
def test_process_image_returns_passthrough_result_and_ui_images(monkeypatch, tmp_path): def test_process_image_returns_passthrough_result_and_ui_images(monkeypatch, tmp_path):
@@ -474,25 +477,34 @@ class TestParameterDefaultConsistency:
input_types = SaveImageLM.INPUT_TYPES() input_types = SaveImageLM.INPUT_TYPES()
optional = input_types["optional"] optional = input_types["optional"]
assert optional["webp_method"][1]["default"] == 6 widget_spec = cast(Any, optional["webp_method"])
assert SaveImageLM.save_images.__defaults__[4] == 6 # positional: webp_method=6 is at index 4 assert widget_spec[1]["default"] == 6
assert SaveImageLM.process_image.__defaults__[6] == 6 save_defaults = cast(tuple[Any, ...], SaveImageLM.save_images.__defaults__ or ())
process_defaults = cast(tuple[Any, ...], SaveImageLM.process_image.__defaults__ or ())
assert save_defaults[4] == 6 # positional: webp_method=6 is at index 4
assert process_defaults[6] == 6
def test_jpeg_subsampling_defaults_are_consistent(self): def test_jpeg_subsampling_defaults_are_consistent(self):
input_types = SaveImageLM.INPUT_TYPES() input_types = SaveImageLM.INPUT_TYPES()
optional = input_types["optional"] optional = input_types["optional"]
assert optional["jpeg_subsampling"][1]["default"] == 0 widget_spec = cast(Any, optional["jpeg_subsampling"])
assert SaveImageLM.save_images.__defaults__[5] == 0 assert widget_spec[1]["default"] == 0
assert SaveImageLM.process_image.__defaults__[7] == 0 save_defaults = cast(tuple[Any, ...], SaveImageLM.save_images.__defaults__ or ())
process_defaults = cast(tuple[Any, ...], SaveImageLM.process_image.__defaults__ or ())
assert save_defaults[5] == 0
assert process_defaults[7] == 0
def test_add_loras_to_prompt_defaults_are_consistent(self): def test_add_loras_to_prompt_defaults_are_consistent(self):
input_types = SaveImageLM.INPUT_TYPES() input_types = SaveImageLM.INPUT_TYPES()
optional = input_types["optional"] optional = input_types["optional"]
assert optional["add_loras_to_prompt"][1]["default"] is False widget_spec = cast(Any, optional["add_loras_to_prompt"])
assert SaveImageLM.save_images.__defaults__[-1] is False assert widget_spec[1]["default"] is False
assert SaveImageLM.process_image.__defaults__[-1] is False save_defaults = cast(tuple[Any, ...], SaveImageLM.save_images.__defaults__ or ())
process_defaults = cast(tuple[Any, ...], SaveImageLM.process_image.__defaults__ or ())
assert save_defaults[-1] is False
assert process_defaults[-1] is False
def test_png_does_not_pass_webp_method_or_jpeg_subsampling(monkeypatch, tmp_path): def test_png_does_not_pass_webp_method_or_jpeg_subsampling(monkeypatch, tmp_path):
+1 -1
View File
@@ -34,7 +34,7 @@ class _DummyModel:
def test_nunchaku_load_lora_skips_missing_lora(monkeypatch, caplog): def test_nunchaku_load_lora_skips_missing_lora(monkeypatch, caplog):
import folder_paths import folder_paths # pyright: ignore[reportMissingImports]
dummy_model = _DummyModel() dummy_model = _DummyModel()
+20 -10
View File
@@ -8,6 +8,8 @@ from __future__ import annotations
import random import random
import string import string
from typing import Any, Dict, cast
import pytest import pytest
from py.services.model_hash_index import ModelHashIndex from py.services.model_hash_index import ModelHashIndex
@@ -22,9 +24,11 @@ class TestHashIndexPerformance:
def test_hash_index_lookup_small(self, benchmark): def test_hash_index_lookup_small(self, benchmark):
"""Benchmark hash index lookup with 100 models.""" """Benchmark hash index lookup with 100 models."""
index, target_hash = self._create_hash_index_with_n_models( index, target_hash = cast(
100, return_target=True tuple[ModelHashIndex, str | None],
self._create_hash_index_with_n_models(100, return_target=True),
) )
assert target_hash is not None
def lookup(): def lookup():
return index.get_path(target_hash) return index.get_path(target_hash)
@@ -34,9 +38,11 @@ class TestHashIndexPerformance:
def test_hash_index_lookup_medium(self, benchmark): def test_hash_index_lookup_medium(self, benchmark):
"""Benchmark hash index lookup with 1,000 models.""" """Benchmark hash index lookup with 1,000 models."""
index, target_hash = self._create_hash_index_with_n_models( index, target_hash = cast(
1000, return_target=True tuple[ModelHashIndex, str | None],
self._create_hash_index_with_n_models(1000, return_target=True),
) )
assert target_hash is not None
def lookup(): def lookup():
return index.get_path(target_hash) return index.get_path(target_hash)
@@ -46,9 +52,11 @@ class TestHashIndexPerformance:
def test_hash_index_lookup_large(self, benchmark): def test_hash_index_lookup_large(self, benchmark):
"""Benchmark hash index lookup with 10,000 models.""" """Benchmark hash index lookup with 10,000 models."""
index, target_hash = self._create_hash_index_with_n_models( index, target_hash = cast(
10000, return_target=True tuple[ModelHashIndex, str | None],
self._create_hash_index_with_n_models(10000, return_target=True),
) )
assert target_hash is not None
def lookup(): def lookup():
return index.get_path(target_hash) return index.get_path(target_hash)
@@ -58,7 +66,7 @@ class TestHashIndexPerformance:
def test_hash_index_add_entry_small(self, benchmark): def test_hash_index_add_entry_small(self, benchmark):
"""Benchmark adding entries to hash index with 100 existing models.""" """Benchmark adding entries to hash index with 100 existing models."""
index = self._create_hash_index_with_n_models(100) index = cast(ModelHashIndex, self._create_hash_index_with_n_models(100))
new_hash = f"new_hash_{self._random_string(16)}" new_hash = f"new_hash_{self._random_string(16)}"
new_path = "/path/to/new_model.safetensors" new_path = "/path/to/new_model.safetensors"
@@ -69,7 +77,7 @@ class TestHashIndexPerformance:
def test_hash_index_add_entry_large(self, benchmark): def test_hash_index_add_entry_large(self, benchmark):
"""Benchmark adding entries to hash index with 10,000 existing models.""" """Benchmark adding entries to hash index with 10,000 existing models."""
index = self._create_hash_index_with_n_models(10000) index = cast(ModelHashIndex, self._create_hash_index_with_n_models(10000))
new_hash = f"new_hash_{self._random_string(16)}" new_hash = f"new_hash_{self._random_string(16)}"
new_path = "/path/to/new_model.safetensors" new_path = "/path/to/new_model.safetensors"
@@ -78,7 +86,9 @@ class TestHashIndexPerformance:
benchmark(add_entry) benchmark(add_entry)
def _create_hash_index_with_n_models(self, n: int, return_target: bool = False): def _create_hash_index_with_n_models(
self, n: int, return_target: bool = False
) -> ModelHashIndex | tuple[ModelHashIndex, str | None]:
"""Create a hash index with n mock models. """Create a hash index with n mock models.
Args: Args:
@@ -170,7 +180,7 @@ class TestRecipeFingerprintPerformance:
benchmark(calculate) benchmark(calculate)
def _create_loras(self, n: int) -> list: def _create_loras(self, n: int) -> list[Dict[str, Any]]:
"""Create a list of n mock LoRA dictionaries.""" """Create a list of n mock LoRA dictionaries."""
loras = [] loras = []
for i in range(n): for i in range(n):
+43 -17
View File
@@ -8,8 +8,10 @@ response schemas.
from __future__ import annotations from __future__ import annotations
import json import json
import pytest
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any
import pytest
from syrupy import SnapshotAssertion from syrupy import SnapshotAssertion
from py.routes.handlers.misc_handlers import ( from py.routes.handlers.misc_handlers import (
@@ -54,13 +56,35 @@ async def noop_async(*_args, **_kwargs):
return None return None
class FakeDownloader:
"""Minimal downloader stub satisfying DownloaderProtocol."""
async def refresh_session(self) -> None:
return None
async def fake_downloader_factory() -> FakeDownloader:
return FakeDownloader()
async def fake_metadata_provider_factory():
return None
def json_payload(response) -> Any:
"""Decode the JSON body of a web.Response, asserting it is not null."""
text = response.text
assert text is not None
return json.loads(text)
class FakePromptServer: class FakePromptServer:
"""Fake prompt server for testing.""" """Fake prompt server for testing."""
sent = [] sent = []
class Instance: class Instance:
sockets: dict = {} sockets: dict[str, Any] = {}
def send_sync(self, event, payload, sid=None): def send_sync(self, event, payload, sid=None):
FakePromptServer.sent.append((event, payload)) FakePromptServer.sent.append((event, payload))
@@ -103,11 +127,11 @@ class TestSettingsHandlerSnapshots:
handler = SettingsHandler( handler = SettingsHandler(
settings_service=settings_service, settings_service=settings_service,
metadata_provider_updater=noop_async, metadata_provider_updater=noop_async,
downloader_factory=lambda: None, downloader_factory=fake_downloader_factory,
) )
response = await handler.get_settings(FakeRequest()) response = await handler.get_settings(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = json_payload(response)
assert payload == snapshot assert payload == snapshot
@@ -118,12 +142,12 @@ class TestSettingsHandlerSnapshots:
handler = SettingsHandler( handler = SettingsHandler(
settings_service=settings_service, settings_service=settings_service,
metadata_provider_updater=noop_async, metadata_provider_updater=noop_async,
downloader_factory=lambda: None, downloader_factory=fake_downloader_factory,
) )
request = FakeRequest(json_data={"language": "zh"}) request = FakeRequest(json_data={"language": "zh"})
response = await handler.update_settings(request) response = await handler.update_settings(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = json_payload(response)
assert payload == snapshot assert payload == snapshot
@@ -137,7 +161,7 @@ class TestNodeRegistryHandlerSnapshots:
node_registry = NodeRegistry() node_registry = NodeRegistry()
handler = NodeRegistryHandler( handler = NodeRegistryHandler(
node_registry=node_registry, node_registry=node_registry,
prompt_server=FakePromptServer, prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False, standalone_mode=False,
) )
@@ -155,8 +179,8 @@ class TestNodeRegistryHandlerSnapshots:
} }
) )
response = await handler.register_nodes(request) response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = json_payload(response)
assert payload == snapshot assert payload == snapshot
@@ -166,13 +190,13 @@ class TestNodeRegistryHandlerSnapshots:
node_registry = NodeRegistry() node_registry = NodeRegistry()
handler = NodeRegistryHandler( handler = NodeRegistryHandler(
node_registry=node_registry, node_registry=node_registry,
prompt_server=FakePromptServer, prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False, standalone_mode=False,
) )
request = FakeRequest(json_data={"nodes": [], "client_id": "test-client-1"}) request = FakeRequest(json_data={"nodes": [], "client_id": "test-client-1"})
response = await handler.register_nodes(request) response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = json_payload(response)
assert payload == snapshot assert payload == snapshot
@@ -249,10 +273,12 @@ class TestModelLibraryHandlerSnapshots:
get_embedding_scanner=scanner_factory, get_embedding_scanner=scanner_factory,
get_downloaded_version_history_service=fake_download_history_service_factory, get_downloaded_version_history_service=fake_download_history_service_factory,
), ),
metadata_provider_factory=lambda: None, metadata_provider_factory=fake_metadata_provider_factory,
) )
response = await handler.check_model_exists(FakeRequest(query={"modelId": "1"})) response = await handler.check_model_exists(
payload = json.loads(response.text) FakeRequest(query={"modelId": "1"}) # pyright: ignore[reportArgumentType]
)
payload = json_payload(response)
assert payload == snapshot assert payload == snapshot
+12 -10
View File
@@ -6,10 +6,10 @@ from pathlib import Path
import types import types
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Optional from typing import Any, Optional
folder_paths_stub = types.SimpleNamespace(get_folder_paths=lambda *_: []) folder_paths_stub = types.SimpleNamespace(get_folder_paths=lambda *_: [])
sys.modules.setdefault("folder_paths", folder_paths_stub) sys.modules.setdefault("folder_paths", folder_paths_stub) # pyright: ignore[reportArgumentType]
import pytest import pytest
from aiohttp import FormData, web from aiohttp import FormData, web
@@ -38,7 +38,9 @@ class DummyRoutes(BaseModelRoutes):
def __init__(self, service=None): def __init__(self, service=None):
super().__init__(service) super().__init__(service)
self.set_model_update_service(NullModelUpdateService()) self.set_model_update_service(
NullModelUpdateService() # pyright: ignore[reportArgumentType]
)
@dataclass @dataclass
@@ -110,7 +112,7 @@ class NullModelUpdateService:
return None return None
async def create_test_client(service) -> TestClient: async def create_test_client(service) -> TestClient[Any, Any]:
routes = DummyRoutes(service) routes = DummyRoutes(service)
app = web.Application() app = web.Application()
routes.setup_routes(app, "test-models") routes.setup_routes(app, "test-models")
@@ -457,18 +459,18 @@ def test_fetch_civitai_hydrates_metadata_before_sync(
mock_scanner._cache.raw_data = [minimal_cache_entry] mock_scanner._cache.raw_data = [minimal_cache_entry]
class FakeMetadata: class FakeMetadata:
def __init__(self, payload: dict) -> None: def __init__(self, payload: dict[str, Any]) -> None:
self._payload = payload self._payload = payload
self._unknown_fields = {"legacy_field": "legacy"} self._unknown_fields = {"legacy_field": "legacy"}
def to_dict(self) -> dict: def to_dict(self) -> dict[str, Any]:
return self._payload.copy() return self._payload.copy()
async def fake_load_metadata(path: str, *_args, **_kwargs): async def fake_load_metadata(path: str, *_args, **_kwargs):
assert path == str(model_path) assert path == str(model_path)
return FakeMetadata(existing_metadata), False return FakeMetadata(existing_metadata), False
async def fake_save_metadata(path: str, metadata: dict) -> bool: async def fake_save_metadata(path: str, metadata: dict[str, Any]) -> bool:
save_calls.append((path, json.loads(json.dumps(metadata)))) save_calls.append((path, json.loads(json.dumps(metadata))))
return True return True
@@ -477,7 +479,7 @@ def test_fetch_civitai_hydrates_metadata_before_sync(
*, *,
sha256: str, sha256: str,
file_path: str, file_path: str,
model_data: dict, model_data: dict[str, Any],
update_cache_func, update_cache_func,
): ):
captured["model_data"] = json.loads(json.dumps(model_data)) captured["model_data"] = json.loads(json.dumps(model_data))
@@ -490,8 +492,8 @@ def test_fetch_civitai_hydrates_metadata_before_sync(
await update_cache_func(file_path, file_path, model_data) await update_cache_func(file_path, file_path, model_data)
return True, None return True, None
save_calls: list[tuple[str, dict]] = [] save_calls: list[tuple[str, dict[str, Any]]] = []
captured: dict[str, dict] = {} captured: dict[str, dict[str, Any]] = {}
monkeypatch.setattr( monkeypatch.setattr(
MetadataManager, "load_metadata", staticmethod(fake_load_metadata) MetadataManager, "load_metadata", staticmethod(fake_load_metadata)
+1 -1
View File
@@ -24,7 +24,7 @@ class StubEmbeddingService:
@pytest.fixture @pytest.fixture
def routes(): def routes():
handler = EmbeddingRoutes() handler = EmbeddingRoutes()
handler.service = StubEmbeddingService() handler.service = StubEmbeddingService() # pyright: ignore[reportAttributeAccessIssue]
return handler return handler
@@ -3,9 +3,9 @@ from __future__ import annotations
import json import json
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Dict from typing import Any, AsyncGenerator, Dict
from aiohttp import web from aiohttp import ClientResponse, web
from aiohttp.test_utils import TestClient, TestServer from aiohttp.test_utils import TestClient, TestServer
from py.routes.example_images_route_registrar import ExampleImagesRouteRegistrar from py.routes.example_images_route_registrar import ExampleImagesRouteRegistrar
@@ -140,14 +140,14 @@ class StubFileManager:
@dataclass @dataclass
class RegistrarHarness: class RegistrarHarness:
client: TestClient client: TestClient[Any, Any]
download_use_case: StubDownloadUseCase download_use_case: StubDownloadUseCase
download_manager: StubDownloadManager download_manager: StubDownloadManager
import_use_case: StubImportUseCase import_use_case: StubImportUseCase
@asynccontextmanager @asynccontextmanager
async def registrar_app() -> RegistrarHarness: async def registrar_app() -> AsyncGenerator[RegistrarHarness, None]:
app = web.Application() app = web.Application()
download_use_case = StubDownloadUseCase() download_use_case = StubDownloadUseCase()
@@ -158,8 +158,15 @@ async def registrar_app() -> RegistrarHarness:
file_manager = StubFileManager() file_manager = StubFileManager()
handler_set = ExampleImagesHandlerSet( handler_set = ExampleImagesHandlerSet(
download=ExampleImagesDownloadHandler(download_use_case, download_manager), download=ExampleImagesDownloadHandler(
management=ExampleImagesManagementHandler(import_use_case, processor, cleanup_service), download_use_case, # pyright: ignore[reportArgumentType]
download_manager,
),
management=ExampleImagesManagementHandler(
import_use_case, # pyright: ignore[reportArgumentType]
processor,
cleanup_service,
),
files=ExampleImagesFileHandler(file_manager), files=ExampleImagesFileHandler(file_manager),
) )
@@ -181,7 +188,7 @@ async def registrar_app() -> RegistrarHarness:
await client.close() await client.close()
async def _json(response: web.StreamResponse) -> Dict[str, Any]: async def _json(response: ClientResponse) -> Dict[str, Any]:
text = await response.text() text = await response.text()
return json.loads(text) if text else {} return json.loads(text) if text else {}
@@ -358,7 +365,7 @@ async def test_check_example_images_needed_returns_error_on_exception():
# Actually, we need to make the method raise an exception # Actually, we need to make the method raise an exception
original_method = harness.download_manager.check_pending_models original_method = harness.download_manager.check_pending_models
async def failing_check(_model_types): async def failing_check(model_types):
raise RuntimeError("Database connection failed") raise RuntimeError("Database connection failed")
harness.download_manager.check_pending_models = failing_check harness.download_manager.check_pending_models = failing_check
+81 -46
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
import json import json
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Dict, List, Tuple from typing import Any, AsyncGenerator, Dict, List, Tuple
from aiohttp import web from aiohttp import web
from aiohttp.test_utils import TestClient, TestServer from aiohttp.test_utils import TestClient, TestServer
@@ -19,11 +19,16 @@ from py.routes.handlers.example_images_handlers import (
) )
def _json_response(response) -> Any:
"""Decode the JSON body of a handler response."""
return json.loads(response.text)
@dataclass @dataclass
class ExampleImagesHarness: class ExampleImagesHarness:
"""Container exposing the aiohttp client and stubbed collaborators.""" """Container exposing the aiohttp client and stubbed collaborators."""
client: TestClient client: TestClient[Any, Any]
download_manager: "StubDownloadManager" download_manager: "StubDownloadManager"
processor: "StubExampleImagesProcessor" processor: "StubExampleImagesProcessor"
file_manager: "StubExampleImagesFileManager" file_manager: "StubExampleImagesFileManager"
@@ -35,27 +40,27 @@ class StubDownloadManager:
def __init__(self) -> None: def __init__(self) -> None:
self.calls: List[Tuple[str, Any]] = [] self.calls: List[Tuple[str, Any]] = []
async def start_download(self, payload: Any) -> dict: async def start_download(self, payload: Any) -> Dict[str, Any]:
self.calls.append(("start_download", payload)) self.calls.append(("start_download", payload))
return {"operation": "start_download", "payload": payload} return {"operation": "start_download", "payload": payload}
async def get_status(self, request: web.Request) -> dict: async def get_status(self, request: web.Request) -> Dict[str, Any]:
self.calls.append(("get_status", dict(request.query))) self.calls.append(("get_status", dict(request.query)))
return {"operation": "get_status"} return {"operation": "get_status"}
async def pause_download(self, request: web.Request) -> dict: async def pause_download(self, request: web.Request) -> Dict[str, Any]:
self.calls.append(("pause_download", None)) self.calls.append(("pause_download", None))
return {"operation": "pause_download"} return {"operation": "pause_download"}
async def resume_download(self, request: web.Request) -> dict: async def resume_download(self, request: web.Request) -> Dict[str, Any]:
self.calls.append(("resume_download", None)) self.calls.append(("resume_download", None))
return {"operation": "resume_download"} return {"operation": "resume_download"}
async def stop_download(self, request: web.Request) -> dict: async def stop_download(self, request: web.Request) -> Dict[str, Any]:
self.calls.append(("stop_download", None)) self.calls.append(("stop_download", None))
return {"operation": "stop_download"} return {"operation": "stop_download"}
async def start_force_download(self, payload: Any) -> dict: async def start_force_download(self, payload: Any) -> Dict[str, Any]:
self.calls.append(("start_force_download", payload)) self.calls.append(("start_force_download", payload))
return {"operation": "start_force_download", "payload": payload} return {"operation": "start_force_download", "payload": payload}
@@ -64,7 +69,7 @@ class StubExampleImagesProcessor:
def __init__(self) -> None: def __init__(self) -> None:
self.calls: List[Tuple[str, Any]] = [] self.calls: List[Tuple[str, Any]] = []
async def import_images(self, model_hash: str, files: List[str]) -> dict: async def import_images(self, model_hash: str, files: List[str]) -> Dict[str, Any]:
payload = {"model_hash": model_hash, "file_paths": files} payload = {"model_hash": model_hash, "file_paths": files}
self.calls.append(("import_images", payload)) self.calls.append(("import_images", payload))
return {"operation": "import_images", "payload": payload} return {"operation": "import_images", "payload": payload}
@@ -122,7 +127,7 @@ class StubWebSocketManager:
@asynccontextmanager @asynccontextmanager
async def example_images_app() -> ExampleImagesHarness: async def example_images_app() -> AsyncGenerator[ExampleImagesHarness, None]:
"""Yield an ExampleImagesRoutes app wired with stubbed collaborators.""" """Yield an ExampleImagesRoutes app wired with stubbed collaborators."""
download_manager = StubDownloadManager() download_manager = StubDownloadManager()
@@ -133,10 +138,10 @@ async def example_images_app() -> ExampleImagesHarness:
controller = ExampleImagesRoutes( controller = ExampleImagesRoutes(
ws_manager=ws_manager, ws_manager=ws_manager,
download_manager=download_manager, download_manager=download_manager, # pyright: ignore[reportArgumentType]
processor=processor, processor=processor,
file_manager=file_manager, file_manager=file_manager, # pyright: ignore[reportArgumentType]
cleanup_service=cleanup_service, cleanup_service=cleanup_service, # pyright: ignore[reportArgumentType]
) )
app = web.Application() app = web.Application()
@@ -323,23 +328,23 @@ async def test_download_handler_methods_delegate() -> None:
def __init__(self) -> None: def __init__(self) -> None:
self.calls: List[Tuple[str, Any]] = [] self.calls: List[Tuple[str, Any]] = []
async def get_status(self, request) -> dict: async def get_status(self, request) -> Dict[str, Any]:
self.calls.append(("get_status", request)) self.calls.append(("get_status", request))
return {"status": "ok"} return {"status": "ok"}
async def pause_download(self, request) -> dict: async def pause_download(self, request) -> Dict[str, Any]:
self.calls.append(("pause_download", request)) self.calls.append(("pause_download", request))
return {"status": "paused"} return {"status": "paused"}
async def resume_download(self, request) -> dict: async def resume_download(self, request) -> Dict[str, Any]:
self.calls.append(("resume_download", request)) self.calls.append(("resume_download", request))
return {"status": "running"} return {"status": "running"}
async def stop_download(self, request) -> dict: async def stop_download(self, request) -> Dict[str, Any]:
self.calls.append(("stop_download", request)) self.calls.append(("stop_download", request))
return {"status": "stopping"} return {"status": "stopping"}
async def start_force_download(self, payload) -> dict: async def start_force_download(self, payload) -> Dict[str, Any]:
self.calls.append(("start_force_download", payload)) self.calls.append(("start_force_download", payload))
return {"status": "force", "payload": payload} return {"status": "force", "payload": payload}
@@ -347,35 +352,50 @@ async def test_download_handler_methods_delegate() -> None:
def __init__(self) -> None: def __init__(self) -> None:
self.payloads: List[Any] = [] self.payloads: List[Any] = []
async def execute(self, payload: dict) -> dict: async def execute(self, payload: Dict[str, Any]) -> Dict[str, Any]:
self.payloads.append(payload) self.payloads.append(payload)
return {"status": "started", "payload": payload} return {"status": "started", "payload": payload}
class DummyRequest: class DummyRequest:
def __init__(self, payload: dict) -> None: def __init__(self, payload: Dict[str, Any]) -> None:
self._payload = payload self._payload = payload
self.query = {} self.query = {}
async def json(self) -> dict: async def json(self) -> Dict[str, Any]:
return self._payload return self._payload
recorder = Recorder() recorder = Recorder()
use_case = StubDownloadUseCase() use_case = StubDownloadUseCase()
handler = ExampleImagesDownloadHandler(use_case, recorder) handler = ExampleImagesDownloadHandler(
use_case, # pyright: ignore[reportArgumentType]
recorder,
)
request = DummyRequest({"foo": "bar"}) request = DummyRequest({"foo": "bar"})
download_response = await handler.download_example_images(request) download_response = await handler.download_example_images(
assert json.loads(download_response.text) == {"status": "started", "payload": {"foo": "bar"}} request # pyright: ignore[reportArgumentType]
status_response = await handler.get_example_images_status(request) )
assert json.loads(status_response.text) == {"status": "ok"} assert _json_response(download_response) == {"status": "started", "payload": {"foo": "bar"}}
pause_response = await handler.pause_example_images(request) status_response = await handler.get_example_images_status(
assert json.loads(pause_response.text) == {"status": "paused"} request # pyright: ignore[reportArgumentType]
resume_response = await handler.resume_example_images(request) )
assert json.loads(resume_response.text) == {"status": "running"} assert _json_response(status_response) == {"status": "ok"}
stop_response = await handler.stop_example_images(request) pause_response = await handler.pause_example_images(
assert json.loads(stop_response.text) == {"status": "stopping"} request # pyright: ignore[reportArgumentType]
force_response = await handler.force_download_example_images(request) )
assert json.loads(force_response.text) == {"status": "force", "payload": {"foo": "bar"}} assert _json_response(pause_response) == {"status": "paused"}
resume_response = await handler.resume_example_images(
request # pyright: ignore[reportArgumentType]
)
assert _json_response(resume_response) == {"status": "running"}
stop_response = await handler.stop_example_images(
request # pyright: ignore[reportArgumentType]
)
assert _json_response(stop_response) == {"status": "stopping"}
force_response = await handler.force_download_example_images(
request # pyright: ignore[reportArgumentType]
)
assert _json_response(force_response) == {"status": "force", "payload": {"foo": "bar"}}
assert use_case.payloads == [{"foo": "bar"}] assert use_case.payloads == [{"foo": "bar"}]
assert recorder.calls == [ assert recorder.calls == [
@@ -393,7 +413,7 @@ async def test_management_handler_methods_delegate() -> None:
def __init__(self) -> None: def __init__(self) -> None:
self.requests: List[Any] = [] self.requests: List[Any] = []
async def execute(self, request: Any) -> dict: async def execute(self, request: Any) -> Dict[str, Any]:
self.requests.append(request) self.requests.append(request)
return {"status": "imported"} return {"status": "imported"}
@@ -412,15 +432,23 @@ async def test_management_handler_methods_delegate() -> None:
recorder = Recorder() recorder = Recorder()
cleanup_service = StubExampleImagesCleanupService() cleanup_service = StubExampleImagesCleanupService()
use_case = StubImportUseCase() use_case = StubImportUseCase()
handler = ExampleImagesManagementHandler(use_case, recorder, cleanup_service) handler = ExampleImagesManagementHandler(
use_case, # pyright: ignore[reportArgumentType]
recorder,
cleanup_service,
)
request = object() request = object()
import_response = await handler.import_example_images(request) import_response = await handler.import_example_images(
assert json.loads(import_response.text) == {"status": "imported"} request # pyright: ignore[reportArgumentType]
assert await handler.delete_example_image(request) == "delete" )
assert _json_response(import_response) == {"status": "imported"}
assert await handler.delete_example_image(request) == "delete" # pyright: ignore[reportArgumentType]
cleanup_service.result = {"success": True} cleanup_service.result = {"success": True}
cleanup_response = await handler.cleanup_example_image_folders(request) cleanup_response = await handler.cleanup_example_image_folders(
assert json.loads(cleanup_response.text) == {"success": True} request # pyright: ignore[reportArgumentType]
)
assert _json_response(cleanup_response) == {"success": True}
assert use_case.requests == [request] assert use_case.requests == [request]
assert recorder.calls == [("delete_custom_image", request)] assert recorder.calls == [("delete_custom_image", request)]
assert len(cleanup_service.calls) == 1 assert len(cleanup_service.calls) == 1
@@ -448,9 +476,9 @@ async def test_file_handler_methods_delegate() -> None:
handler = ExampleImagesFileHandler(recorder) handler = ExampleImagesFileHandler(recorder)
request = object() request = object()
assert await handler.open_example_images_folder(request) == "open" assert await handler.open_example_images_folder(request) == "open" # pyright: ignore[reportArgumentType]
assert await handler.get_example_image_files(request) == "files" assert await handler.get_example_image_files(request) == "files" # pyright: ignore[reportArgumentType]
assert await handler.has_example_images(request) == "has" assert await handler.has_example_images(request) == "has" # pyright: ignore[reportArgumentType]
assert recorder.calls == [ assert recorder.calls == [
("open_folder", request), ("open_folder", request),
("get_files", request), ("get_files", request),
@@ -483,9 +511,16 @@ def test_handler_set_route_mapping_includes_all_handlers() -> None:
async def set_example_image_nsfw_level(self, request): async def set_example_image_nsfw_level(self, request):
return {} return {}
download = ExampleImagesDownloadHandler(DummyUseCase(), DummyManager()) download = ExampleImagesDownloadHandler(
DummyUseCase(), # pyright: ignore[reportArgumentType]
DummyManager(),
)
cleanup_service = StubExampleImagesCleanupService() cleanup_service = StubExampleImagesCleanupService()
management = ExampleImagesManagementHandler(DummyUseCase(), DummyProcessor(), cleanup_service) management = ExampleImagesManagementHandler(
DummyUseCase(), # pyright: ignore[reportArgumentType]
DummyProcessor(),
cleanup_service,
)
files = ExampleImagesFileHandler(object()) files = ExampleImagesFileHandler(object())
handler_set = ExampleImagesHandlerSet( handler_set = ExampleImagesHandlerSet(
download=download, download=download,
+2 -1
View File
@@ -4,6 +4,7 @@ import asyncio
import logging import logging
from pathlib import Path from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any
import pytest import pytest
from aiohttp import web from aiohttp import web
@@ -160,7 +161,7 @@ async def test_lora_manager_lifecycle(monkeypatch: pytest.MonkeyPatch, tmp_path:
) )
original_create_task = asyncio.create_task original_create_task = asyncio.create_task
scheduled_tasks: list[asyncio.Task] = [] scheduled_tasks: list[asyncio.Task[Any]] = []
def track_create_task(coro, *, name=None): def track_create_task(coro, *, name=None):
task = original_create_task(coro, name=name) task = original_create_task(coro, name=name)
+2 -100
View File
@@ -5,7 +5,7 @@ from unittest.mock import MagicMock
import pytest import pytest
from py.routes.lora_routes import LoraRoutes from py.routes.lora_routes import LoraRoutes
from server import PromptServer from server import PromptServer # pyright: ignore[reportMissingImports]
class DummyRequest: class DummyRequest:
@@ -20,14 +20,8 @@ class DummyRequest:
class StubLoraService: class StubLoraService:
def __init__(self): def __init__(self):
self.notes = {}
self.trigger_words = {} self.trigger_words = {}
self.usage_tips = {} self.usage_tips = {}
self.previews = {}
self.civitai = {}
async def get_lora_notes(self, name):
return self.notes.get(name)
async def get_lora_trigger_words(self, name): async def get_lora_trigger_words(self, name):
return self.trigger_words.get(name, []) return self.trigger_words.get(name, [])
@@ -35,57 +29,14 @@ class StubLoraService:
async def get_lora_usage_tips_by_relative_path(self, path): async def get_lora_usage_tips_by_relative_path(self, path):
return self.usage_tips.get(path) return self.usage_tips.get(path)
async def get_lora_preview_url(self, name):
return self.previews.get(name)
async def get_lora_civitai_url(self, name):
return self.civitai.get(name, {"civitai_url": ""})
@pytest.fixture @pytest.fixture
def routes(): def routes():
handler = LoraRoutes() handler = LoraRoutes()
handler.service = StubLoraService() handler.service = StubLoraService() # pyright: ignore[reportAttributeAccessIssue]
return handler return handler
async def test_get_lora_notes_success(routes):
routes.service.notes["demo"] = "Great notes"
request = DummyRequest(query={"name": "demo"})
response = await routes.get_lora_notes(request)
payload = json.loads(response.text)
assert payload == {"success": True, "notes": "Great notes"}
async def test_get_lora_notes_missing_name(routes):
response = await routes.get_lora_notes(DummyRequest())
assert response.status == 400
assert response.text == "Lora file name is required"
async def test_get_lora_notes_not_found(routes):
response = await routes.get_lora_notes(DummyRequest(query={"name": "missing"}))
payload = json.loads(response.text)
assert response.status == 404
assert payload == {"success": False, "error": "LoRA not found in cache"}
async def test_get_lora_notes_error(routes, monkeypatch):
async def failing(*_args, **_kwargs):
raise RuntimeError("boom")
routes.service.get_lora_notes = failing
response = await routes.get_lora_notes(DummyRequest(query={"name": "demo"}))
payload = json.loads(response.text)
assert response.status == 500
assert payload["success"] is False
assert payload["error"] == "boom"
async def test_get_lora_trigger_words_success(routes): async def test_get_lora_trigger_words_success(routes):
routes.service.trigger_words["demo"] = ["trigger"] routes.service.trigger_words["demo"] = ["trigger"]
response = await routes.get_lora_trigger_words(DummyRequest(query={"name": "demo"})) response = await routes.get_lora_trigger_words(DummyRequest(query={"name": "demo"}))
@@ -133,55 +84,6 @@ async def test_get_usage_tips_error(routes):
assert payload["success"] is False assert payload["success"] is False
async def test_get_preview_url_success(routes):
routes.service.previews["demo"] = "http://preview"
response = await routes.get_lora_preview_url(DummyRequest(query={"name": "demo"}))
payload = json.loads(response.text)
assert payload == {"success": True, "preview_url": "http://preview"}
async def test_get_preview_url_missing(routes):
response = await routes.get_lora_preview_url(DummyRequest())
assert response.status == 400
async def test_get_preview_url_not_found(routes):
response = await routes.get_lora_preview_url(DummyRequest(query={"name": "missing"}))
payload = json.loads(response.text)
assert response.status == 404
assert payload["success"] is False
async def test_get_civitai_url_success(routes):
routes.service.civitai["demo"] = {"civitai_url": "https://civitai.com"}
response = await routes.get_lora_civitai_url(DummyRequest(query={"name": "demo"}))
payload = json.loads(response.text)
assert payload == {"success": True, "civitai_url": "https://civitai.com"}
async def test_get_civitai_url_missing(routes):
response = await routes.get_lora_civitai_url(DummyRequest())
assert response.status == 400
async def test_get_civitai_url_not_found(routes):
response = await routes.get_lora_civitai_url(DummyRequest(query={"name": "missing"}))
payload = json.loads(response.text)
assert response.status == 404
assert payload["success"] is False
async def test_get_civitai_url_error(routes):
async def failing(*_args, **_kwargs):
raise RuntimeError("oops")
routes.service.get_lora_civitai_url = failing
response = await routes.get_lora_civitai_url(DummyRequest(query={"name": "demo"}))
payload = json.loads(response.text)
assert response.status == 500
assert payload["success"] is False
async def test_get_trigger_words_broadcasts(monkeypatch, routes): async def test_get_trigger_words_broadcasts(monkeypatch, routes):
send_mock = MagicMock() send_mock = MagicMock()
PromptServer.instance = SimpleNamespace(send_sync=send_mock) PromptServer.instance = SimpleNamespace(send_sync=send_mock)
+114 -102
View File
@@ -5,6 +5,7 @@ import os
import subprocess import subprocess
import zipfile import zipfile
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any
from unittest.mock import patch, MagicMock from unittest.mock import patch, MagicMock
import pytest import pytest
@@ -34,6 +35,13 @@ from py.routes.misc_route_registrar import MISC_ROUTE_DEFINITIONS, MiscRouteRegi
from py.routes.misc_routes import MiscRoutes from py.routes.misc_routes import MiscRoutes
def _json_payload(response) -> dict[str, Any]:
"""Decode the JSON body of a web.Response, asserting it is not null."""
text = response.text
assert text is not None
return json.loads(text)
class FakeRequest: class FakeRequest:
def __init__(self, *, json_data=None, query=None, method="POST"): def __init__(self, *, json_data=None, query=None, method="POST"):
self._json_data = json_data or {} self._json_data = json_data or {}
@@ -129,8 +137,8 @@ async def test_get_settings_excludes_no_sync_keys():
downloader_factory=dummy_downloader_factory, downloader_factory=dummy_downloader_factory,
) )
response = await handler.get_settings(FakeRequest()) response = await handler.get_settings(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
# Regular settings should be synced # Regular settings should be synced
@@ -155,8 +163,8 @@ async def test_update_settings_rejects_missing_example_path(tmp_path):
missing_path = tmp_path / "does-not-exist" missing_path = tmp_path / "does-not-exist"
request = FakeRequest(json_data={"example_images_path": str(missing_path)}) request = FakeRequest(json_data={"example_images_path": str(missing_path)})
response = await handler.update_settings(request) response = await handler.update_settings(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is False assert payload["success"] is False
assert "Path does not exist" in payload["error"] assert "Path does not exist" in payload["error"]
@@ -181,9 +189,9 @@ async def test_doctor_handler_reports_key_cache_and_ui_issues():
) )
response = await handler.get_doctor_diagnostics( response = await handler.get_doctor_diagnostics(
FakeRequest(query={"clientVersion": "1.2.2-client"}, method="GET") FakeRequest(query={"clientVersion": "1.2.2-client"}, method="GET") # pyright: ignore[reportArgumentType]
) )
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
assert payload["summary"]["status"] == "error" assert payload["summary"]["status"] == "error"
@@ -209,8 +217,8 @@ async def test_doctor_handler_can_repair_cache():
scanner_factories=(("lora", "LoRAs", scanner_factory),), scanner_factories=(("lora", "LoRAs", scanner_factory),),
) )
response = await handler.repair_doctor_cache(FakeRequest()) response = await handler.repair_doctor_cache(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert response.status == 200 assert response.status == 200
assert payload["success"] is True assert payload["success"] is True
@@ -230,7 +238,7 @@ async def test_doctor_handler_exports_support_bundle():
) )
response = await handler.export_doctor_bundle( response = await handler.export_doctor_bundle(
FakeRequest( FakeRequest( # pyright: ignore[reportArgumentType]
json_data={ json_data={
"summary": {"status": "warning"}, "summary": {"status": "warning"},
"diagnostics": [{"id": "cache_health", "status": "warning"}], "diagnostics": [{"id": "cache_health", "status": "warning"}],
@@ -241,6 +249,7 @@ async def test_doctor_handler_exports_support_bundle():
) )
assert response.status == 200 assert response.status == 200
assert isinstance(response.body, bytes)
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive: with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
names = set(archive.namelist()) names = set(archive.namelist())
assert "doctor-report.json" in names assert "doctor-report.json" in names
@@ -263,7 +272,7 @@ async def test_doctor_handler_redacts_string_secrets_in_bundle():
) )
response = await handler.export_doctor_bundle( response = await handler.export_doctor_bundle(
FakeRequest( FakeRequest( # pyright: ignore[reportArgumentType]
json_data={ json_data={
"frontend_logs": [ "frontend_logs": [
{ {
@@ -276,6 +285,7 @@ async def test_doctor_handler_redacts_string_secrets_in_bundle():
) )
assert response.status == 200 assert response.status == 200
assert isinstance(response.body, bytes)
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive: with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
frontend_logs = archive.read("frontend-console.json").decode("utf-8") frontend_logs = archive.read("frontend-console.json").decode("utf-8")
assert "abcdef123456" not in frontend_logs assert "abcdef123456" not in frontend_logs
@@ -308,7 +318,7 @@ async def test_doctor_handler_redacts_json_shaped_string_secrets_in_bundle():
} }
response = await handler.export_doctor_bundle( response = await handler.export_doctor_bundle(
FakeRequest( FakeRequest( # pyright: ignore[reportArgumentType]
json_data={ json_data={
"frontend_logs": [ "frontend_logs": [
{ {
@@ -321,6 +331,7 @@ async def test_doctor_handler_redacts_json_shaped_string_secrets_in_bundle():
) )
assert response.status == 200 assert response.status == 200
assert isinstance(response.body, bytes)
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive: with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
frontend_logs = archive.read("frontend-console.json").decode("utf-8") frontend_logs = archive.read("frontend-console.json").decode("utf-8")
backend_logs = archive.read("backend-logs.txt").decode("utf-8") backend_logs = archive.read("backend-logs.txt").decode("utf-8")
@@ -359,9 +370,10 @@ async def test_doctor_handler_exports_backend_session_logs_from_helper():
"notes": [], "notes": [],
} }
response = await handler.export_doctor_bundle(FakeRequest(json_data={})) response = await handler.export_doctor_bundle(FakeRequest(json_data={})) # pyright: ignore[reportArgumentType]
assert response.status == 200 assert response.status == 200
assert isinstance(response.body, bytes)
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive: with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
backend_logs = archive.read("backend-logs.txt").decode("utf-8") backend_logs = archive.read("backend-logs.txt").decode("utf-8")
backend_source = json.loads( backend_source = json.loads(
@@ -449,14 +461,14 @@ async def test_backup_handler_returns_status_and_exports(monkeypatch):
handler = BackupHandler(backup_service_factory=factory) handler = BackupHandler(backup_service_factory=factory)
status_response = await handler.get_backup_status(FakeRequest()) status_response = await handler.get_backup_status(FakeRequest()) # pyright: ignore[reportArgumentType]
status_payload = json.loads(status_response.text) status_payload = _json_payload(status_response)
assert status_payload["success"] is True assert status_payload["success"] is True
assert status_payload["status"]["backupDir"] == "/tmp/backups" assert status_payload["status"]["backupDir"] == "/tmp/backups"
assert status_payload["status"]["enabled"] is True assert status_payload["status"]["enabled"] is True
assert status_payload["snapshots"][0]["name"] == "backup.zip" assert status_payload["snapshots"][0]["name"] == "backup.zip"
export_response = await handler.export_backup(FakeRequest()) export_response = await handler.export_backup(FakeRequest()) # pyright: ignore[reportArgumentType]
assert export_response.status == 200 assert export_response.status == 200
assert export_response.body == b"zip-bytes" assert export_response.body == b"zip-bytes"
@@ -476,8 +488,8 @@ async def test_backup_handler_rejects_missing_import_archive():
async def read(self): async def read(self):
return b"" return b""
response = await handler.import_backup(EmptyRequest()) response = await handler.import_backup(EmptyRequest()) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert response.status == 400 assert response.status == 400
assert payload["success"] is False assert payload["success"] is False
@@ -504,8 +516,8 @@ async def test_open_backup_location_uses_settings_directory(tmp_path, monkeypatc
monkeypatch.setattr("py.routes.handlers.misc_handlers._is_docker", lambda: False) monkeypatch.setattr("py.routes.handlers.misc_handlers._is_docker", lambda: False)
monkeypatch.setattr("py.routes.handlers.misc_handlers._is_wsl", lambda: False) monkeypatch.setattr("py.routes.handlers.misc_handlers._is_wsl", lambda: False)
response = await handler.open_backup_location(FakeRequest()) response = await handler.open_backup_location(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert response.status == 200 assert response.status == 200
assert payload["success"] is True assert payload["success"] is True
@@ -535,8 +547,8 @@ async def test_open_wildcards_location_creates_and_opens_directory(tmp_path, mon
else str(wildcards_dir), else str(wildcards_dir),
) )
response = await handler.open_wildcards_location(FakeRequest()) response = await handler.open_wildcards_location(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert response.status == 200 assert response.status == 200
assert payload["success"] is True assert payload["success"] is True
@@ -564,7 +576,7 @@ class RecordingRouter:
def test_misc_route_registrar_registers_all_routes(): def test_misc_route_registrar_registers_all_routes():
app = SimpleNamespace(router=RecordingRouter()) app = SimpleNamespace(router=RecordingRouter())
registrar = MiscRouteRegistrar(app) # type: ignore[arg-type] registrar = MiscRouteRegistrar(app) # pyright: ignore[reportArgumentType]
async def dummy_handler(_request): async def dummy_handler(_request):
return web.Response() return web.Response()
@@ -586,7 +598,7 @@ class FakePromptServer:
sent = [] sent = []
class Instance: class Instance:
sockets: dict = {} sockets: dict[str, Any] = {}
def send_sync(self, event, payload, sid=None): def send_sync(self, event, payload, sid=None):
FakePromptServer.sent.append((event, payload)) FakePromptServer.sent.append((event, payload))
@@ -599,7 +611,7 @@ async def test_register_nodes_requires_graph_id():
node_registry = NodeRegistry() node_registry = NodeRegistry()
handler = NodeRegistryHandler( handler = NodeRegistryHandler(
node_registry=node_registry, node_registry=node_registry,
prompt_server=FakePromptServer, prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False, standalone_mode=False,
) )
@@ -609,8 +621,8 @@ async def test_register_nodes_requires_graph_id():
"client_id": "test-client-1", "client_id": "test-client-1",
} }
) )
response = await handler.register_nodes(request) response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert response.status == 400 assert response.status == 400
assert payload["success"] is False assert payload["success"] is False
@@ -622,7 +634,7 @@ async def test_register_nodes_stores_graph_identifier():
node_registry = NodeRegistry() node_registry = NodeRegistry()
handler = NodeRegistryHandler( handler = NodeRegistryHandler(
node_registry=node_registry, node_registry=node_registry,
prompt_server=FakePromptServer, prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False, standalone_mode=False,
) )
@@ -641,8 +653,8 @@ async def test_register_nodes_stores_graph_identifier():
} }
) )
response = await handler.register_nodes(request) response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
@@ -659,7 +671,7 @@ async def test_register_nodes_defaults_graph_name_to_none():
node_registry = NodeRegistry() node_registry = NodeRegistry()
handler = NodeRegistryHandler( handler = NodeRegistryHandler(
node_registry=node_registry, node_registry=node_registry,
prompt_server=FakePromptServer, prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False, standalone_mode=False,
) )
@@ -677,8 +689,8 @@ async def test_register_nodes_defaults_graph_name_to_none():
} }
) )
response = await handler.register_nodes(request) response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
@@ -692,7 +704,7 @@ async def test_register_nodes_includes_capabilities():
node_registry = NodeRegistry() node_registry = NodeRegistry()
handler = NodeRegistryHandler( handler = NodeRegistryHandler(
node_registry=node_registry, node_registry=node_registry,
prompt_server=FakePromptServer, prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False, standalone_mode=False,
) )
@@ -714,8 +726,8 @@ async def test_register_nodes_includes_capabilities():
} }
) )
response = await handler.register_nodes(request) response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
@@ -734,7 +746,7 @@ async def test_register_nodes_accepts_compound_node_ids():
node_registry = NodeRegistry() node_registry = NodeRegistry()
handler = NodeRegistryHandler( handler = NodeRegistryHandler(
node_registry=node_registry, node_registry=node_registry,
prompt_server=FakePromptServer, prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False, standalone_mode=False,
) )
@@ -758,8 +770,8 @@ async def test_register_nodes_accepts_compound_node_ids():
} }
) )
response = await handler.register_nodes(request) response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert response.status == 200 assert response.status == 200
assert payload["success"] is True assert payload["success"] is True
@@ -778,11 +790,11 @@ async def test_register_nodes_accepts_compound_node_ids():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_update_node_widget_sends_payload(): async def test_update_node_widget_sends_payload():
send_calls: list[tuple[str, dict]] = [] send_calls: list[tuple[str, dict[str, Any]]] = []
class RecordingPromptServer: class RecordingPromptServer:
class Instance: class Instance:
sockets: dict = {} sockets: dict[str, Any] = {}
def send_sync(self, event, payload, sid=None): def send_sync(self, event, payload, sid=None):
send_calls.append((event, payload)) send_calls.append((event, payload))
@@ -791,7 +803,7 @@ async def test_update_node_widget_sends_payload():
handler = NodeRegistryHandler( handler = NodeRegistryHandler(
node_registry=NodeRegistry(), node_registry=NodeRegistry(),
prompt_server=RecordingPromptServer, prompt_server=RecordingPromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False, standalone_mode=False,
) )
@@ -803,8 +815,8 @@ async def test_update_node_widget_sends_payload():
} }
) )
response = await handler.update_node_widget(request) response = await handler.update_node_widget(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert response.status == 200 assert response.status == 200
assert payload["success"] is True assert payload["success"] is True
@@ -824,18 +836,18 @@ async def test_update_node_widget_sends_payload():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_update_lora_code_includes_graph_identifier(): async def test_update_lora_code_includes_graph_identifier():
send_calls: list[tuple[str, dict]] = [] send_calls: list[tuple[str, dict[str, Any]]] = []
class RecordingPromptServer: class RecordingPromptServer:
class Instance: class Instance:
sockets: dict = {} sockets: dict[str, Any] = {}
def send_sync(self, event, payload, sid=None): def send_sync(self, event, payload, sid=None):
send_calls.append((event, payload)) send_calls.append((event, payload))
instance = Instance() instance = Instance()
handler = LoraCodeHandler(RecordingPromptServer) handler = LoraCodeHandler(RecordingPromptServer) # pyright: ignore[reportArgumentType]
request = FakeRequest( request = FakeRequest(
json_data={ json_data={
@@ -845,8 +857,8 @@ async def test_update_lora_code_includes_graph_identifier():
} }
) )
response = await handler.update_lora_code(request) response = await handler.update_lora_code(request) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
assert payload["results"] == [ assert payload["results"] == [
@@ -913,7 +925,7 @@ class FakeUserModelsProvider(FakeMetadataProvider):
self.next_cursor = next_cursor self.next_cursor = next_cursor
self.estimated_total = estimated_total self.estimated_total = estimated_total
self.received_usernames: list[str] = [] self.received_usernames: list[str] = []
self.received_cursors: list = [] self.received_cursors: list[Any] = []
async def get_user_models(self, username, cursor=None): async def get_user_models(self, username, cursor=None):
self.received_usernames.append(username) self.received_usernames.append(username)
@@ -924,7 +936,7 @@ class FakeUserModelsProvider(FakeMetadataProvider):
return self.estimated_total return self.estimated_total
async def fake_metadata_provider_factory(): async def fake_metadata_provider_factory() -> Any:
return FakeMetadataProvider() return FakeMetadataProvider()
@@ -949,8 +961,8 @@ async def fake_metadata_archive_manager_factory():
class FakeDownloadHistoryService: class FakeDownloadHistoryService:
def __init__(self, downloaded_by_type=None): def __init__(self, downloaded_by_type=None):
self.downloaded_by_type = downloaded_by_type or {} self.downloaded_by_type = downloaded_by_type or {}
self.marked_downloaded: list[tuple] = [] self.marked_downloaded: list[tuple[Any, ...]] = []
self.marked_not_downloaded: list[tuple] = [] self.marked_not_downloaded: list[tuple[Any, ...]] = []
async def has_been_downloaded(self, model_type, version_id): async def has_been_downloaded(self, model_type, version_id):
return version_id in self.downloaded_by_type.get(model_type, set()) return version_id in self.downloaded_by_type.get(model_type, set())
@@ -1012,20 +1024,20 @@ async def test_misc_routes_bind_produces_expected_handlers():
controller = MiscRoutes( controller = MiscRoutes(
settings_service=DummySettings(), settings_service=DummySettings(),
usage_stats_factory=lambda: SimpleNamespace( usage_stats_factory=lambda: SimpleNamespace( # pyright: ignore[reportArgumentType]
process_execution=noop_async, get_stats=noop_async process_execution=noop_async, get_stats=noop_async
), ),
prompt_server=FakePromptServer, prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
service_registry_adapter=service_registry_adapter, service_registry_adapter=service_registry_adapter,
metadata_provider_factory=fake_metadata_provider_factory, metadata_provider_factory=fake_metadata_provider_factory,
metadata_archive_manager_factory=fake_metadata_archive_manager_factory, metadata_archive_manager_factory=fake_metadata_archive_manager_factory,
metadata_provider_updater=noop_async, metadata_provider_updater=noop_async,
downloader_factory=dummy_downloader_factory, downloader_factory=dummy_downloader_factory,
registrar_factory=registrar_factory, registrar_factory=registrar_factory, # pyright: ignore[reportArgumentType]
) )
app = SimpleNamespace(router=RecordingRouter()) app = SimpleNamespace(router=RecordingRouter())
controller.bind(app) # type: ignore[arg-type] controller.bind(app) # pyright: ignore[reportArgumentType]
assert recorded_registrars, "Expected registrar to be created" assert recorded_registrars, "Expected registrar to be created"
mapping = recorded_registrars[0].registered_mapping mapping = recorded_registrars[0].registered_mapping
@@ -1106,7 +1118,7 @@ async def test_get_civitai_user_models_marks_library_versions():
provider = FakeUserModelsProvider(models) provider = FakeUserModelsProvider(models)
async def provider_factory(): async def provider_factory() -> Any:
return provider return provider
lora_scanner = FakeExistenceScanner({101}) lora_scanner = FakeExistenceScanner({101})
@@ -1133,9 +1145,9 @@ async def test_get_civitai_user_models_marks_library_versions():
) )
response = await handler.get_civitai_user_models( response = await handler.get_civitai_user_models(
FakeRequest(query={"username": "pixel"}) FakeRequest(query={"username": "pixel"}) # pyright: ignore[reportArgumentType]
) )
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
assert payload["username"] == "pixel" assert payload["username"] == "pixel"
@@ -1239,7 +1251,7 @@ async def test_get_civitai_user_models_rewrites_civitai_previews():
provider = FakeUserModelsProvider(models) provider = FakeUserModelsProvider(models)
async def provider_factory(): async def provider_factory() -> Any:
return provider return provider
handler = ModelLibraryHandler( handler = ModelLibraryHandler(
@@ -1253,9 +1265,9 @@ async def test_get_civitai_user_models_rewrites_civitai_previews():
) )
response = await handler.get_civitai_user_models( response = await handler.get_civitai_user_models(
FakeRequest(query={"username": "pixel"}) FakeRequest(query={"username": "pixel"}) # pyright: ignore[reportArgumentType]
) )
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
previews_by_version = { previews_by_version = {
@@ -1275,7 +1287,7 @@ async def test_get_civitai_user_models_rewrites_civitai_previews():
async def test_get_civitai_user_models_requires_username(): async def test_get_civitai_user_models_requires_username():
provider = FakeUserModelsProvider([]) provider = FakeUserModelsProvider([])
async def provider_factory(): async def provider_factory() -> Any:
return provider return provider
handler = ModelLibraryHandler( handler = ModelLibraryHandler(
@@ -1288,8 +1300,8 @@ async def test_get_civitai_user_models_requires_username():
metadata_provider_factory=provider_factory, metadata_provider_factory=provider_factory,
) )
response = await handler.get_civitai_user_models(FakeRequest()) response = await handler.get_civitai_user_models(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert response.status == 400 assert response.status == 400
assert payload["success"] is False assert payload["success"] is False
@@ -1318,7 +1330,7 @@ async def test_get_civitai_user_models_returns_pagination_fields():
provider = FakeUserModelsProvider(models, next_cursor="cursor-token", estimated_total=2140) provider = FakeUserModelsProvider(models, next_cursor="cursor-token", estimated_total=2140)
async def provider_factory(): async def provider_factory() -> Any:
return provider return provider
handler = ModelLibraryHandler( handler = ModelLibraryHandler(
@@ -1332,9 +1344,9 @@ async def test_get_civitai_user_models_returns_pagination_fields():
) )
response = await handler.get_civitai_user_models( response = await handler.get_civitai_user_models(
FakeRequest(query={"username": "pixel"}) FakeRequest(query={"username": "pixel"}) # pyright: ignore[reportArgumentType]
) )
payload = json.loads(response.text) payload = _json_payload(response)
assert response.status == 200 assert response.status == 200
assert payload["success"] is True assert payload["success"] is True
@@ -1351,7 +1363,7 @@ async def test_get_civitai_user_models_returns_pagination_fields():
async def test_get_civitai_user_models_passes_cursor_and_omits_estimate(): async def test_get_civitai_user_models_passes_cursor_and_omits_estimate():
provider = FakeUserModelsProvider([], next_cursor=None, estimated_total=999) provider = FakeUserModelsProvider([], next_cursor=None, estimated_total=999)
async def provider_factory(): async def provider_factory() -> Any:
return provider return provider
handler = ModelLibraryHandler( handler = ModelLibraryHandler(
@@ -1365,9 +1377,9 @@ async def test_get_civitai_user_models_passes_cursor_and_omits_estimate():
) )
response = await handler.get_civitai_user_models( response = await handler.get_civitai_user_models(
FakeRequest(query={"username": "pixel", "cursor": "opaque-token"}) FakeRequest(query={"username": "pixel", "cursor": "opaque-token"}) # pyright: ignore[reportArgumentType]
) )
payload = json.loads(response.text) payload = _json_payload(response)
assert response.status == 200 assert response.status == 200
assert payload["success"] is True assert payload["success"] is True
@@ -1391,10 +1403,10 @@ def test_ensure_handler_mapping_caches_result():
controller = MiscRoutes( controller = MiscRoutes(
settings_service=DummySettings(), settings_service=DummySettings(),
usage_stats_factory=lambda: SimpleNamespace( usage_stats_factory=lambda: SimpleNamespace( # pyright: ignore[reportArgumentType]
process_execution=noop_async, get_stats=noop_async process_execution=noop_async, get_stats=noop_async
), ),
prompt_server=FakePromptServer, prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
service_registry_adapter=ServiceRegistryAdapter( service_registry_adapter=ServiceRegistryAdapter(
get_lora_scanner=fake_scanner_factory, get_lora_scanner=fake_scanner_factory,
get_checkpoint_scanner=fake_scanner_factory, get_checkpoint_scanner=fake_scanner_factory,
@@ -1405,7 +1417,7 @@ def test_ensure_handler_mapping_caches_result():
metadata_archive_manager_factory=fake_metadata_archive_manager_factory, metadata_archive_manager_factory=fake_metadata_archive_manager_factory,
metadata_provider_updater=noop_async, metadata_provider_updater=noop_async,
downloader_factory=dummy_downloader_factory, downloader_factory=dummy_downloader_factory,
handler_set_factory=RecordingHandlerSet, handler_set_factory=RecordingHandlerSet, # pyright: ignore[reportArgumentType]
) )
first_mapping = controller._ensure_handler_mapping() first_mapping = controller._ensure_handler_mapping()
@@ -1447,8 +1459,8 @@ async def test_check_model_exists_returns_local_versions():
metadata_provider_factory=fake_metadata_provider_factory, metadata_provider_factory=fake_metadata_provider_factory,
) )
response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"})) response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"})) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
assert payload["modelType"] == "lora" assert payload["modelType"] == "lora"
@@ -1461,7 +1473,7 @@ async def test_check_model_exists_returns_local_versions():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_check_model_exists_model_id_only_does_not_call_metadata_provider(): async def test_check_model_exists_model_id_only_does_not_call_metadata_provider():
async def metadata_provider_factory(): async def metadata_provider_factory() -> Any:
raise AssertionError("metadata provider should not be called for modelId-only checks") raise AssertionError("metadata provider should not be called for modelId-only checks")
handler = ModelLibraryHandler( handler = ModelLibraryHandler(
@@ -1474,8 +1486,8 @@ async def test_check_model_exists_model_id_only_does_not_call_metadata_provider(
metadata_provider_factory=metadata_provider_factory, metadata_provider_factory=metadata_provider_factory,
) )
response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"})) response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"})) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert payload == { assert payload == {
"success": True, "success": True,
@@ -1503,9 +1515,9 @@ async def test_check_model_exists_returns_download_history_when_file_missing():
) )
response = await handler.check_model_exists( response = await handler.check_model_exists(
FakeRequest(query={"modelId": "5", "modelVersionId": "999"}) FakeRequest(query={"modelId": "5", "modelVersionId": "999"}) # pyright: ignore[reportArgumentType]
) )
payload = json.loads(response.text) payload = _json_payload(response)
assert payload == { assert payload == {
"success": True, "success": True,
@@ -1533,9 +1545,9 @@ async def test_model_version_download_status_endpoints():
) )
get_response = await handler.get_model_version_download_status( get_response = await handler.get_model_version_download_status(
FakeRequest(query={"modelType": "lora", "modelVersionId": "123"}) FakeRequest(query={"modelType": "lora", "modelVersionId": "123"}) # pyright: ignore[reportArgumentType]
) )
get_payload = json.loads(get_response.text) get_payload = _json_payload(get_response)
assert get_payload == { assert get_payload == {
"success": True, "success": True,
"modelType": "lora", "modelType": "lora",
@@ -1544,7 +1556,7 @@ async def test_model_version_download_status_endpoints():
} }
set_response = await handler.set_model_version_download_status( set_response = await handler.set_model_version_download_status(
FakeRequest( FakeRequest( # pyright: ignore[reportArgumentType]
json_data={ json_data={
"modelType": "checkpoint", "modelType": "checkpoint",
"modelVersionId": 456, "modelVersionId": 456,
@@ -1554,7 +1566,7 @@ async def test_model_version_download_status_endpoints():
} }
) )
) )
set_payload = json.loads(set_response.text) set_payload = _json_payload(set_response)
assert set_payload == { assert set_payload == {
"success": True, "success": True,
"modelType": "checkpoint", "modelType": "checkpoint",
@@ -1566,7 +1578,7 @@ async def test_model_version_download_status_endpoints():
] ]
set_get_response = await handler.set_model_version_download_status( set_get_response = await handler.set_model_version_download_status(
FakeRequest( FakeRequest( # pyright: ignore[reportArgumentType]
method="GET", method="GET",
query={ query={
"modelType": "embedding", "modelType": "embedding",
@@ -1576,7 +1588,7 @@ async def test_model_version_download_status_endpoints():
}, },
) )
) )
set_get_payload = json.loads(set_get_response.text) set_get_payload = _json_payload(set_get_response)
assert set_get_payload == { assert set_get_payload == {
"success": True, "success": True,
"modelType": "embedding", "modelType": "embedding",
@@ -1586,7 +1598,7 @@ async def test_model_version_download_status_endpoints():
def test_create_handler_set_uses_provided_dependencies(): def test_create_handler_set_uses_provided_dependencies():
recorded_handlers: list[dict] = [] recorded_handlers: list[dict[str, Any]] = []
class RecordingHandlerSet: class RecordingHandlerSet:
def __init__(self, **handlers): def __init__(self, **handlers):
@@ -1609,8 +1621,8 @@ def test_create_handler_set_uses_provided_dependencies():
controller = MiscRoutes( controller = MiscRoutes(
settings_service=DummySettings(), settings_service=DummySettings(),
usage_stats_factory=lambda: FakeUsageStats(), usage_stats_factory=lambda: FakeUsageStats(), # pyright: ignore[reportArgumentType]
prompt_server=CustomPromptServer, prompt_server=CustomPromptServer, # pyright: ignore[reportArgumentType]
service_registry_adapter=ServiceRegistryAdapter( service_registry_adapter=ServiceRegistryAdapter(
get_lora_scanner=fake_scanner_factory, get_lora_scanner=fake_scanner_factory,
get_checkpoint_scanner=fake_scanner_factory, get_checkpoint_scanner=fake_scanner_factory,
@@ -1621,8 +1633,8 @@ def test_create_handler_set_uses_provided_dependencies():
metadata_archive_manager_factory=fake_metadata_archive_manager_factory, metadata_archive_manager_factory=fake_metadata_archive_manager_factory,
metadata_provider_updater=noop_async, metadata_provider_updater=noop_async,
downloader_factory=dummy_downloader_factory, downloader_factory=dummy_downloader_factory,
handler_set_factory=RecordingHandlerSet, handler_set_factory=RecordingHandlerSet, # pyright: ignore[reportArgumentType]
node_registry=fake_node_registry, node_registry=fake_node_registry, # pyright: ignore[reportArgumentType]
standalone_mode_flag=True, standalone_mode_flag=True,
) )
@@ -1665,7 +1677,7 @@ def test_is_wsl_returns_false_on_read_error():
assert result is False assert result is False
def test_is_wsl_returns_false_on_read_error(): def test_is_wsl_returns_false_on_read_error_builtins():
with patch("builtins.open", side_effect=OSError()): with patch("builtins.open", side_effect=OSError()):
result = _is_wsl() result = _is_wsl()
assert result is False assert result is False
@@ -1688,7 +1700,7 @@ def test_wsl_to_windows_path_returns_none_on_error():
assert result is None assert result is None
def test_wsl_to_windows_path_returns_none_on_subprocess_error(): def test_wsl_to_windows_path_returns_none_on_subprocess_error_plain():
with patch( with patch(
"subprocess.run", side_effect=subprocess.CalledProcessError(1, "wslpath") "subprocess.run", side_effect=subprocess.CalledProcessError(1, "wslpath")
): ):
@@ -1769,8 +1781,8 @@ async def test_check_filename_conflicts_returns_ok_when_no_duplicates():
scanner_factories=(("lora", "LoRAs", scanner_factory),), scanner_factories=(("lora", "LoRAs", scanner_factory),),
) )
response = await handler.get_doctor_diagnostics(FakeRequest(method="GET")) response = await handler.get_doctor_diagnostics(FakeRequest(method="GET")) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
diagnostic_map = {item["id"]: item for item in payload["diagnostics"]} diagnostic_map = {item["id"]: item for item in payload["diagnostics"]}
assert diagnostic_map["filename_conflicts"]["status"] == "ok" assert diagnostic_map["filename_conflicts"]["status"] == "ok"
@@ -1798,8 +1810,8 @@ async def test_check_filename_conflicts_detects_duplicates():
scanner_factories=(("lora", "LoRAs", scanner_factory),), scanner_factories=(("lora", "LoRAs", scanner_factory),),
) )
response = await handler.get_doctor_diagnostics(FakeRequest(method="GET")) response = await handler.get_doctor_diagnostics(FakeRequest(method="GET")) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
diagnostic_map = {item["id"]: item for item in payload["diagnostics"]} diagnostic_map = {item["id"]: item for item in payload["diagnostics"]}
conflict_diag = diagnostic_map["filename_conflicts"] conflict_diag = diagnostic_map["filename_conflicts"]
@@ -1826,8 +1838,8 @@ async def test_resolve_filename_conflicts_returns_renamed_list():
scanner_factories=(("lora", "LoRAs", scanner_factory),), scanner_factories=(("lora", "LoRAs", scanner_factory),),
) )
response = await handler.resolve_filename_conflicts(FakeRequest(method="POST")) response = await handler.resolve_filename_conflicts(FakeRequest(method="POST")) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
# Files don't exist on disk, so nothing gets renamed # Files don't exist on disk, so nothing gets renamed
@@ -1850,8 +1862,8 @@ async def test_resolve_filename_conflicts_handles_scanner_error_gracefully():
scanner_factories=(("lora", "LoRAs", scanner_factory),), scanner_factories=(("lora", "LoRAs", scanner_factory),),
) )
response = await handler.resolve_filename_conflicts(FakeRequest(method="POST")) response = await handler.resolve_filename_conflicts(FakeRequest(method="POST")) # pyright: ignore[reportArgumentType]
payload = json.loads(response.text) payload = _json_payload(response)
assert payload["success"] is True assert payload["success"] is True
assert payload["count"] == 0 assert payload["count"] == 0
+5 -7
View File
@@ -1,5 +1,6 @@
from __future__ import annotations from __future__ import annotations
import logging
from types import SimpleNamespace from types import SimpleNamespace
import jinja2 import jinja2
@@ -48,19 +49,16 @@ async def test_model_page_view_reads_version_per_request():
template_env=template_env, template_env=template_env,
template_name="dummy.html", template_name="dummy.html",
service=DummyService(), service=DummyService(),
settings_service=DummySettings(), settings_service=DummySettings(), # pyright: ignore[reportArgumentType]
server_i18n=DummyI18n(), server_i18n=DummyI18n(),
logger=SimpleNamespace( logger=logging.getLogger("test_model_page_view"),
debug=lambda *_args, **_kwargs: None,
error=lambda *_args, **_kwargs: None,
),
) )
view._get_app_version = lambda: "1.0.2-old" view._get_app_version = lambda: "1.0.2-old"
first = await view.handle(SimpleNamespace()) first = await view.handle(SimpleNamespace()) # pyright: ignore[reportArgumentType]
view._get_app_version = lambda: "1.0.2-new" view._get_app_version = lambda: "1.0.2-new"
second = await view.handle(SimpleNamespace()) second = await view.handle(SimpleNamespace()) # pyright: ignore[reportArgumentType]
assert first.text == "1.0.2-old" assert first.text == "1.0.2-old"
assert second.text == "1.0.2-new" assert second.text == "1.0.2-new"
+18 -7
View File
@@ -21,8 +21,12 @@ async def test_model_query_handler_accepts_limit_zero_for_base_models():
service = DummyService() service = DummyService()
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__)) handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
response = await handler.get_base_models(SimpleNamespace(query={"limit": "0"})) response = await handler.get_base_models(
payload = json.loads(response.text) SimpleNamespace(query={"limit": "0"}) # pyright: ignore[reportArgumentType]
)
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True assert payload["success"] is True
assert service.received_limit == 0 assert service.received_limit == 0
@@ -33,7 +37,9 @@ async def test_model_query_handler_rejects_negative_limit_for_base_models():
service = DummyService() service = DummyService()
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__)) handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
await handler.get_base_models(SimpleNamespace(query={"limit": "-1"})) await handler.get_base_models(
SimpleNamespace(query={"limit": "-1"}) # pyright: ignore[reportArgumentType]
)
assert service.received_limit == 20 assert service.received_limit == 20
@@ -58,9 +64,11 @@ async def test_model_query_handler_search_tags_passes_query_and_limit():
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__)) handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
response = await handler.search_tags( response = await handler.search_tags(
SimpleNamespace(query={"q": "ani", "limit": "50"}) SimpleNamespace(query={"q": "ani", "limit": "50"}) # pyright: ignore[reportArgumentType]
) )
payload = json.loads(response.text) text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True assert payload["success"] is True
assert payload["tags"] == [{"tag": "anime", "count": 3}] assert payload["tags"] == [{"tag": "anime", "count": 3}]
@@ -73,7 +81,8 @@ async def test_model_query_handler_search_tags_defaults_limit_to_20():
service = DummySearchTagsService() service = DummySearchTagsService()
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__)) handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
await handler.search_tags(SimpleNamespace(query={})) await handler.search_tags(SimpleNamespace(query={}) # pyright: ignore[reportArgumentType]
)
assert service.received_limit == 20 assert service.received_limit == 20
@@ -83,6 +92,8 @@ async def test_model_query_handler_search_tags_clamps_negative_limit():
service = DummySearchTagsService() service = DummySearchTagsService()
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__)) handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
await handler.search_tags(SimpleNamespace(query={"limit": "-5"})) await handler.search_tags(
SimpleNamespace(query={"limit": "-5"}) # pyright: ignore[reportArgumentType]
)
assert service.received_limit == 20 assert service.received_limit == 20
+38 -19
View File
@@ -2,6 +2,7 @@ import copy
import json import json
import logging import logging
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any
import pytest import pytest
@@ -198,22 +199,24 @@ async def test_get_civitai_versions_degrades_when_download_history_unavailable(m
handler = ModelCivitaiHandler( handler = ModelCivitaiHandler(
service=service, service=service,
settings_service=SimpleNamespace(get=lambda *_: False), settings_service=SimpleNamespace(get=lambda *_: False), # pyright: ignore[reportArgumentType]
ws_manager=SimpleNamespace(), ws_manager=SimpleNamespace(), # pyright: ignore[reportArgumentType]
logger=logging.getLogger(__name__), logger=logging.getLogger(__name__),
metadata_provider_factory=metadata_provider_factory, metadata_provider_factory=metadata_provider_factory,
validate_model_type=lambda *_: True, validate_model_type=lambda *_: True,
expected_model_types=lambda: "LoRA", expected_model_types=lambda: "LoRA",
find_model_file=lambda *_: None, find_model_file=lambda *_: None,
metadata_sync=SimpleNamespace(), metadata_sync=SimpleNamespace(), # pyright: ignore[reportArgumentType]
metadata_refresh_use_case=SimpleNamespace(), metadata_refresh_use_case=SimpleNamespace(), # pyright: ignore[reportArgumentType]
metadata_progress_callback=lambda *_args, **_kwargs: None, metadata_progress_callback=lambda *_args, **_kwargs: None, # pyright: ignore[reportArgumentType]
) )
response = await handler.get_civitai_versions( response = await handler.get_civitai_versions(
SimpleNamespace(match_info={"model_id": "42"}) SimpleNamespace(match_info={"model_id": "42"}) # pyright: ignore[reportArgumentType]
) )
payload = json.loads(response.text) text = response.text
assert text is not None
payload = json.loads(text)
assert response.status == 200 assert response.status == 200
assert payload[0]["id"] == 7 assert payload[0]["id"] == 7
@@ -284,10 +287,14 @@ async def test_refresh_model_updates_filters_records_without_updates():
async def json(self): async def json(self):
return {} return {}
response = await handler.refresh_model_updates(DummyRequest()) response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200 assert response.status == 200
payload = json.loads(response.text) text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True assert payload["success"] is True
assert len(payload["records"]) == 1 assert len(payload["records"]) == 1
assert payload["records"][0]["modelId"] == 1 assert payload["records"][0]["modelId"] == 1
@@ -347,7 +354,9 @@ async def test_refresh_model_updates_with_target_ids():
async def json(self): async def json(self):
return {"modelIds": [1, "2", None]} return {"modelIds": [1, "2", None]}
response = await handler.refresh_model_updates(DummyRequest()) response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200 assert response.status == 200
call = update_service.calls[0] call = update_service.calls[0]
@@ -399,7 +408,9 @@ async def test_refresh_model_updates_accepts_snake_case_ids():
async def json(self): async def json(self):
return {"model_ids": [3, "4", "abc", None]} return {"model_ids": [3, "4", "abc", None]}
response = await handler.refresh_model_updates(DummyRequest()) response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200 assert response.status == 200
call = update_service.calls[0] call = update_service.calls[0]
@@ -429,9 +440,9 @@ async def test_fetch_missing_license_data_updates_metadata(monkeypatch):
return None, False return None, False
return SimpleNamespace(to_dict=lambda: copy.deepcopy(data)), False return SimpleNamespace(to_dict=lambda: copy.deepcopy(data)), False
saved: list[tuple[str, dict]] = [] saved: list[tuple[str, dict[str, Any]]] = []
async def fake_save(path: str, metadata: dict): async def fake_save(path: str, metadata: dict[str, Any]):
saved.append((path, copy.deepcopy(metadata))) saved.append((path, copy.deepcopy(metadata)))
return True return True
@@ -479,10 +490,14 @@ async def test_fetch_missing_license_data_updates_metadata(monkeypatch):
async def json(self): async def json(self):
return {} return {}
response = await handler.fetch_missing_civitai_license_data(DummyRequest()) response = await handler.fetch_missing_civitai_license_data(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200 assert response.status == 200
payload = json.loads(response.text) text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True assert payload["success"] is True
assert len(payload["updated"]) == 3 assert len(payload["updated"]) == 3
assert provider_calls == [[10, 20]] assert provider_calls == [[10, 20]]
@@ -516,9 +531,9 @@ async def test_fetch_missing_license_data_filters_model_ids(monkeypatch):
return None, False return None, False
return SimpleNamespace(to_dict=lambda: copy.deepcopy(data)), False return SimpleNamespace(to_dict=lambda: copy.deepcopy(data)), False
saved: list[tuple[str, dict]] = [] saved: list[tuple[str, dict[str, Any]]] = []
async def fake_save(path: str, metadata: dict): async def fake_save(path: str, metadata: dict[str, Any]):
saved.append((path, copy.deepcopy(metadata))) saved.append((path, copy.deepcopy(metadata)))
return True return True
@@ -566,10 +581,14 @@ async def test_fetch_missing_license_data_filters_model_ids(monkeypatch):
async def json(self): async def json(self):
return {"modelIds": [20]} return {"modelIds": [20]}
response = await handler.fetch_missing_civitai_license_data(DummyRequest()) response = await handler.fetch_missing_civitai_license_data(
DummyRequest() # pyright: ignore[reportArgumentType]
)
assert response.status == 200 assert response.status == 200
payload = json.loads(response.text) text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True assert payload["success"] is True
assert len(payload["updated"]) == 1 assert len(payload["updated"]) == 1
assert provider_calls == [[20]] assert provider_calls == [[20]]
+1 -1
View File
@@ -31,7 +31,7 @@ class StubLoraService:
@pytest.fixture @pytest.fixture
def routes(): def routes():
handler = LoraRoutes() handler = LoraRoutes()
handler.service = StubLoraService() handler.service = StubLoraService() # pyright: ignore[reportAttributeAccessIssue]
return handler return handler
+6 -2
View File
@@ -34,8 +34,12 @@ async def test_recipe_query_handler_base_models_limit_zero_returns_all():
logger=logging.getLogger(__name__), logger=logging.getLogger(__name__),
) )
response = await handler.get_base_models(SimpleNamespace(query={"limit": "0"})) response = await handler.get_base_models(
payload = json.loads(response.text) SimpleNamespace(query={"limit": "0"}) # pyright: ignore[reportArgumentType]
)
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True assert payload["success"] is True
assert payload["base_models"] == [ assert payload["base_models"] == [
@@ -124,7 +124,7 @@ def test_to_route_mapping_uses_handler_set():
super().__init__() super().__init__()
self.created = 0 self.created = 0
def _create_handler_set(self): # noqa: D401 - simple override for test def _create_handler_set(self): # noqa: D401 - simple override for test # pyright: ignore[reportIncompatibleMethodOverride]
self.created += 1 self.created += 1
return DummyHandlerSet() return DummyHandlerSet()
@@ -162,10 +162,12 @@ def test_recipe_route_registrar_binds_every_route():
self.router = FakeRouter() self.router = FakeRouter()
app = FakeApp() app = FakeApp()
registrar = recipe_route_registrar.RecipeRouteRegistrar(app) registrar = recipe_route_registrar.RecipeRouteRegistrar(
app # pyright: ignore[reportArgumentType]
)
handler_mapping = { handler_mapping = {
definition.handler_name: object() definition.handler_name: lambda _request: None
for definition in recipe_route_registrar.ROUTE_DEFINITIONS for definition in recipe_route_registrar.ROUTE_DEFINITIONS
} }
+10 -4
View File
@@ -26,7 +26,7 @@ from py.services.service_registry import ServiceRegistry
class RecipeRouteHarness: class RecipeRouteHarness:
"""Container exposing the aiohttp client and stubbed collaborators.""" """Container exposing the aiohttp client and stubbed collaborators."""
client: TestClient client: TestClient[Any, Any]
scanner: "StubRecipeScanner" scanner: "StubRecipeScanner"
analysis: "StubAnalysisService" analysis: "StubAnalysisService"
persistence: "StubPersistenceService" persistence: "StubPersistenceService"
@@ -92,6 +92,9 @@ class StubRecipeScanner:
candidate = Path(self.recipes_dir) / f"{recipe_id}.recipe.json" candidate = Path(self.recipes_dir) / f"{recipe_id}.recipe.json"
return str(candidate) if candidate.exists() else None return str(candidate) if candidate.exists() else None
async def get_recipe_syntax_tokens(self, recipe_id: str) -> List[str]:
return self.recipes.get(recipe_id, {}).get("syntax", []) # pragma: no cover - overridden per test
async def remove_recipe(self, recipe_id: str) -> None: async def remove_recipe(self, recipe_id: str) -> None:
self.removed.append(recipe_id) self.removed.append(recipe_id)
self.recipes.pop(recipe_id, None) self.recipes.pop(recipe_id, None)
@@ -110,7 +113,7 @@ class StubAnalysisService:
self.remote_calls: List[Optional[str]] = [] self.remote_calls: List[Optional[str]] = []
self.local_calls: List[Optional[str]] = [] self.local_calls: List[Optional[str]] = []
self.result = SimpleNamespace(payload={"loras": []}, status=200) self.result = SimpleNamespace(payload={"loras": []}, status=200)
self._recipe_parser_factory = None self._recipe_parser_factory: Any = None
StubAnalysisService.instances.append(self) StubAnalysisService.instances.append(self)
async def analyze_uploaded_image( async def analyze_uploaded_image(
@@ -456,7 +459,7 @@ async def test_list_recipes_offloads_dimensions_to_thread(
harness.scanner.cached_raw = list(harness.scanner.listing_items) harness.scanner.cached_raw = list(harness.scanner.listing_items)
real_to_thread = asyncio.to_thread real_to_thread = asyncio.to_thread
to_thread_calls: list[tuple] = [] to_thread_calls: list[tuple[Any, ...]] = []
async def counting_to_thread(fn, *args, **kwargs): async def counting_to_thread(fn, *args, **kwargs):
to_thread_calls.append((fn, args, kwargs)) to_thread_calls.append((fn, args, kwargs))
@@ -536,6 +539,7 @@ async def test_list_recipes_passes_checkpoint_hash_filter(
assert response.status == 200 assert response.status == 200
assert payload["items"] == [] assert payload["items"] == []
assert harness.scanner.last_paginated_params is not None
assert harness.scanner.last_paginated_params["checkpoint_hash"] == "ckpt123" assert harness.scanner.last_paginated_params["checkpoint_hash"] == "ckpt123"
@@ -1010,7 +1014,9 @@ async def test_get_recipe_syntax(monkeypatch, tmp_path: Path) -> None:
return ["<lora:lora1:0.5>"] return ["<lora:lora1:0.5>"]
raise RecipeNotFoundError(f"Recipe {rid} not found") raise RecipeNotFoundError(f"Recipe {rid} not found")
harness.scanner.get_recipe_syntax_tokens = fake_get_recipe_syntax_tokens harness.scanner.get_recipe_syntax_tokens = ( # pyright: ignore[reportAttributeAccessIssue]
fake_get_recipe_syntax_tokens
)
response = await harness.client.get(f"/api/lm/recipe/{recipe_id}/syntax") response = await harness.client.get(f"/api/lm/recipe/{recipe_id}/syntax")
payload = await response.json() payload = await response.json()
+3 -3
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import asyncio import asyncio
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from types import SimpleNamespace from types import SimpleNamespace
from typing import AsyncIterator, Dict, Iterable, List, Sequence from typing import Any, AsyncIterator, Dict, Iterable, List, Sequence
from aiohttp import web from aiohttp import web
from aiohttp.test_utils import TestClient, TestServer from aiohttp.test_utils import TestClient, TestServer
@@ -34,7 +34,7 @@ class IntegrationCache:
class IntegrationScanner: class IntegrationScanner:
"""Scanner double that registers with ServiceRegistry expectations.""" """Scanner double that registers with ServiceRegistry expectations."""
def __init__(self, items: Iterable[Dict[str, object]]) -> None: def __init__(self, items: Iterable[Dict[str, Any]]) -> None:
self.model_type = "lora" self.model_type = "lora"
self._cache = IntegrationCache(list(items)) self._cache = IntegrationCache(list(items))
self._hash_index = SimpleNamespace( self._hash_index = SimpleNamespace(
@@ -68,7 +68,7 @@ class IntegrationScanner:
@asynccontextmanager @asynccontextmanager
async def aiohttp_client(app: web.Application) -> AsyncIterator[TestClient]: async def aiohttp_client(app: web.Application) -> AsyncIterator[TestClient[Any, Any]]:
"""Spin up a TestClient with lifecycle management.""" """Spin up a TestClient with lifecycle management."""
server = TestServer(app) server = TestServer(app)
+24 -10
View File
@@ -1,4 +1,5 @@
import json import json
from typing import Any, Optional
import pytest import pytest
@@ -17,7 +18,7 @@ class FakeRequest:
class DummySettings: class DummySettings:
def __init__(self): def __init__(self):
self.activated = None self.activated = None
self.should_raise = None self.should_raise: Optional[Exception] = None
def activate_library(self, name): def activate_library(self, name):
if self.should_raise: if self.should_raise:
@@ -25,6 +26,13 @@ class DummySettings:
self.activated = name self.activated = name
def json_payload(response) -> Any:
"""Decode the JSON body of a web.Response, asserting it is not null."""
text = response.text
assert text is not None
return json.loads(text)
class DummyDownloader: class DummyDownloader:
async def refresh_session(self): # pragma: no cover - helper async def refresh_session(self): # pragma: no cover - helper
return None return None
@@ -53,7 +61,7 @@ async def test_get_libraries_returns_registry(monkeypatch, handler):
monkeypatch.setattr(config, "get_library_registry_snapshot", lambda: registry) monkeypatch.setattr(config, "get_library_registry_snapshot", lambda: registry)
response = await handler.get_libraries(FakeRequest()) response = await handler.get_libraries(FakeRequest())
payload = json.loads(response.text) payload = json_payload(response)
assert response.status == 200 assert response.status == 200
assert payload == { assert payload == {
@@ -71,7 +79,7 @@ async def test_get_libraries_handles_errors(monkeypatch, handler):
monkeypatch.setattr(config, "get_library_registry_snapshot", boom) monkeypatch.setattr(config, "get_library_registry_snapshot", boom)
response = await handler.get_libraries(FakeRequest()) response = await handler.get_libraries(FakeRequest())
payload = json.loads(response.text) payload = json_payload(response)
assert response.status == 500 assert response.status == 500
assert payload["success"] is False assert payload["success"] is False
@@ -90,8 +98,10 @@ async def test_activate_library_success(monkeypatch):
registry = {"libraries": {"alpha": {"name": "Alpha"}}, "active_library": "alpha"} registry = {"libraries": {"alpha": {"name": "Alpha"}}, "active_library": "alpha"}
monkeypatch.setattr(config, "get_library_registry_snapshot", lambda: registry) monkeypatch.setattr(config, "get_library_registry_snapshot", lambda: registry)
response = await handler.activate_library(FakeRequest(json_data={"library": "alpha"})) response = await handler.activate_library(
payload = json.loads(response.text) FakeRequest(json_data={"library": "alpha"}) # pyright: ignore[reportArgumentType]
)
payload = json_payload(response)
assert response.status == 200 assert response.status == 200
assert payload == { assert payload == {
@@ -105,7 +115,7 @@ async def test_activate_library_success(monkeypatch):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_activate_library_requires_name(handler): async def test_activate_library_requires_name(handler):
response = await handler.activate_library(FakeRequest(json_data={})) response = await handler.activate_library(FakeRequest(json_data={}))
payload = json.loads(response.text) payload = json_payload(response)
assert response.status == 400 assert response.status == 400
assert payload["success"] is False assert payload["success"] is False
@@ -122,8 +132,10 @@ async def test_activate_library_unknown_returns_404(monkeypatch):
downloader_factory=dummy_downloader_factory, downloader_factory=dummy_downloader_factory,
) )
response = await handler.activate_library(FakeRequest(json_data={"library": "ghost"})) response = await handler.activate_library(
payload = json.loads(response.text) FakeRequest(json_data={"library": "ghost"}) # pyright: ignore[reportArgumentType]
)
payload = json_payload(response)
assert response.status == 404 assert response.status == 404
assert payload["success"] is False assert payload["success"] is False
@@ -140,8 +152,10 @@ async def test_activate_library_unexpected_error_returns_500(monkeypatch):
downloader_factory=dummy_downloader_factory, downloader_factory=dummy_downloader_factory,
) )
response = await handler.activate_library(FakeRequest(json_data={"library": "broken"})) response = await handler.activate_library(
payload = json.loads(response.text) FakeRequest(json_data={"library": "broken"}) # pyright: ignore[reportArgumentType]
)
payload = json_payload(response)
assert response.status == 500 assert response.status == 500
assert payload["success"] is False assert payload["success"] is False
+1 -1
View File
@@ -341,7 +341,7 @@ async def test_handle_stats_page_renders_template(stats_routes):
assert response.status == 200 assert response.status == 200
assert response.text == "rendered" assert response.text == "rendered"
assert stats_routes.server_i18n.locale_calls[-1] == "ja" assert stats_routes.server_i18n.locale_calls[-1] == "ja"
assert stats_routes.routes.template_env._i18n_filter_added is True assert stats_routes.routes._i18n_filter_added is True
assert "t" in stats_routes.routes.template_env.filters assert "t" in stats_routes.routes.template_env.filters
assert stats_routes.routes.template_env.filters["t"]("greeting") == "translated:greeting" assert stats_routes.routes.template_env.filters["t"]("greeting") == "translated:greeting"
assert template_context["is_initializing"] is False assert template_context["is_initializing"] is False
+3 -1
View File
@@ -7,9 +7,10 @@ from aiohttp.test_utils import TestClient, TestServer
import sys import sys
import types import types
from typing import Any
folder_paths_stub = types.SimpleNamespace(get_folder_paths=lambda *_: []) folder_paths_stub = types.SimpleNamespace(get_folder_paths=lambda *_: [])
sys.modules.setdefault("folder_paths", folder_paths_stub) sys.modules.setdefault("folder_paths", folder_paths_stub) # pyright: ignore[reportArgumentType]
from py.routes.handlers.model_handlers import ModelListingHandler from py.routes.handlers.model_handlers import ModelListingHandler
@@ -19,6 +20,7 @@ class MockService:
def __init__(self): def __init__(self):
self.model_type = "test-model" self.model_type = "test-model"
self.last_call_kwargs: dict[str, Any] = {}
async def get_paginated_data(self, **kwargs): async def get_paginated_data(self, **kwargs):
# Store the kwargs for verification # Store the kwargs for verification
+2 -2
View File
@@ -24,7 +24,7 @@ def _fake_request(body=None, query_params=None):
async def _json(): async def _json():
return body or {} return body or {}
req.json = _json req.json = _json # pyright: ignore[reportAttributeAccessIssue]
return req return req
@@ -131,7 +131,7 @@ async def test_perform_git_update_preserves_user_dirs(monkeypatch, tmp_path):
class FakeHeads: class FakeHeads:
def __getitem__(self, name): def __getitem__(self, name):
class Head: class Head:
def checkout(self_inner): def checkout(self):
calls.append(("head-checkout", (name,))) calls.append(("head-checkout", (name,)))
return Head() return Head()
+10 -4
View File
@@ -32,9 +32,11 @@ async def test_search_wildcards_returns_results():
handler = WildcardsHandler(service=StubService()) handler = WildcardsHandler(service=StubService())
response = await handler.search_wildcards( response = await handler.search_wildcards(
FakeRequest(query={"search": "cat", "limit": "25", "offset": "2"}) FakeRequest(query={"search": "cat", "limit": "25", "offset": "2"}) # pyright: ignore[reportArgumentType]
) )
payload = json.loads(response.text) text = response.text
assert text is not None
payload = json.loads(text)
assert response.status == 200 assert response.status == 200
assert payload == { assert payload == {
@@ -62,8 +64,12 @@ async def test_search_wildcards_handles_errors():
raise RuntimeError("boom") raise RuntimeError("boom")
handler = WildcardsHandler(service=StubService()) handler = WildcardsHandler(service=StubService())
response = await handler.search_wildcards(FakeRequest(query={"search": "cat"})) response = await handler.search_wildcards(
payload = json.loads(response.text) FakeRequest(query={"search": "cat"}) # pyright: ignore[reportArgumentType]
)
text = response.text
assert text is not None
payload = json.loads(text)
assert response.status == 500 assert response.status == 500
assert payload["error"] == "boom" assert payload["error"] == "boom"
+7 -7
View File
@@ -6,7 +6,7 @@ from unittest.mock import AsyncMock
import pytest 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.aria2_transfer_state import Aria2TransferStateStore
from py.services import aria2_transfer_state 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 @pytest.mark.asyncio
async def test_pause_resume_cancel_forward_to_rpc(monkeypatch): async def test_pause_resume_cancel_forward_to_rpc(monkeypatch):
downloader = Aria2Downloader() downloader = Aria2Downloader()
downloader._transfers["download-1"] = type( downloader._transfers["download-1"] = Aria2Transfer(
"Transfer", (), {"gid": "gid-1", "save_path": "/tmp/model.safetensors"} gid="gid-1", save_path="/tmp/model.safetensors"
)() )
calls = [] calls = []
@@ -200,9 +200,9 @@ async def test_download_file_reuses_existing_transfer_without_add_uri(
downloader._rpc_secret = "secret" downloader._rpc_secret = "secret"
save_path = tmp_path / "downloads" / "model.safetensors" save_path = tmp_path / "downloads" / "model.safetensors"
downloader._transfers["download-1"] = type( downloader._transfers["download-1"] = Aria2Transfer(
"Transfer", (), {"gid": "gid-1", "save_path": str(save_path)} gid="gid-1", save_path=str(save_path)
)() )
rpc_calls = [] rpc_calls = []
statuses = iter( statuses = iter(
+17 -14
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
import json import json
from pathlib import Path from pathlib import Path
from typing import Any, Dict, List, Optional from typing import Any, Dict, Iterator, List, Optional
import pytest import pytest
@@ -16,7 +16,7 @@ from py.services.persistent_model_cache import DEFAULT_LICENSE_FLAGS, Persistent
@pytest.fixture(autouse=True) @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.""" """Reset the service singleton so every test starts from a fresh instance."""
Autov3BackfillService._instance = None Autov3BackfillService._instance = None
yield yield
@@ -66,7 +66,7 @@ class RecordingScanner:
self.model_type = model_type self.model_type = model_type
self._persistent_cache = persistent_cache self._persistent_cache = persistent_cache
self.entries: Dict[str, Dict[str, Any]] = {entry['file_path']: entry for entry in entries} 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: 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)) 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) 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 ''. # Non-safetensors files yield no embedded hash, so both are marked ''.
assert updated == 2 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') == [] assert store.get_models_missing_autov3('dummy') == []
persisted = store.load_cache('dummy') persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data} items = {item['file_path']: item for item in persisted.raw_data}
assert items[path_a]['autov3'] == '' assert items[path_a]['autov3'] == ''
assert items[path_b]['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) 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 updated == 1
assert scanner.update_calls == [('dummy', existing, '')] 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 = Autov3BackfillService.get_instance()
service._running_types = {'dummy'} service._running_types = {'dummy'}
try: try:
assert await service.backfill(scanner) == 0 assert await service.backfill(scanner) == 0 # pyright: ignore[reportArgumentType]
finally: finally:
service._running_types = set() service._running_types = set()
assert scanner.update_calls == [] assert scanner.update_calls == []
@@ -195,7 +196,7 @@ async def test_backfill_runs_concurrently_for_different_model_types(tmp_path: Pa
try: try:
# The lora backfill must still run while checkpoint is in progress. # 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, '')] assert lora_scanner.update_calls == [('lora', lora_file, '')]
finally: finally:
service._running_types = set() 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]}, []) store.save_cache('dummy', entries, {'hash-boom': [existing]}, [])
scanner = RaisingScanner('dummy', store, entries) 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 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) store.update_single_model(model_type, new_item, old_item)
return True return True
updated = await Autov3BackfillService.get_instance().backfill(BareScanner()) updated = await Autov3BackfillService.get_instance().backfill(BareScanner()) # pyright: ignore[reportArgumentType]
assert updated == 1 assert updated == 1
assert store.get_models_missing_autov3('dummy') == [] 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) scanner = RecordingScanner('dummy', store, entries)
service = Autov3BackfillService.get_instance() 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. # 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 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): class RealScanner(ModelScanner):
def __init__(self) -> None: def __init__(self) -> None: # pyright: ignore[reportMissingSuperCall]
self.model_type = 'dummy' self.model_type = 'dummy'
self._persistent_cache = store self._persistent_cache = store
self._cache = ModelCache(raw_data=[dict(e) for e in entries], folders=[]) 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') == [] assert store.get_models_missing_autov3('dummy') == []
persisted = store.load_cache('dummy') persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data} items = {item['file_path']: item for item in persisted.raw_data}
assert items[path_a]['autov3'] == '' assert items[path_a]['autov3'] == ''
assert items[path_b]['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]}, []) store.save_cache('dummy', [_entry(path, 'hash-ckpt')], {'hash-ckpt': [path]}, [])
scanner = RecordingScanner('dummy', store, [_entry(path, 'hash-ckpt')]) 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 updated == 1
assert scanner.update_calls == [('dummy', path, 'abcdef123456')] assert scanner.update_calls == [('dummy', path, 'abcdef123456')]
persisted = store.load_cache('dummy') persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data} items = {item['file_path']: item for item in persisted.raw_data}
assert items[path]['autov3'] == 'abcdef123456' assert items[path]['autov3'] == 'abcdef123456'
# Self-terminating: the row is marked and the driving query empties. # 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]}, []) store.save_cache('dummy', [_entry(path, 'hash-plain')], {'hash-plain': [path]}, [])
scanner = RecordingScanner('dummy', store, [_entry(path, 'hash-plain')]) 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 updated == 1
assert scanner.update_calls == [('dummy', path, '')] 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) return str(legacy_db)
monkeypatch.setattr( monkeypatch.setattr(
"py.services.persistent_model_cache.get_persistent_cache", "py.services.persistent_model_cache.PersistentModelCache.get_default",
lambda *_args, **_kwargs: LegacyCache(), lambda *args, **kwargs: LegacyCache(),
) )
service = ModelUpdateService(settings_manager=DummySettingsManager()) service = ModelUpdateService(settings_manager=DummySettingsManager())
+19 -14
View File
@@ -28,13 +28,14 @@ class DummyService(BaseModelService):
return model_data return model_data
class StubRepository: class StubRepository(ModelCacheRepository):
def __init__(self, data): def __init__(self, data):
super().__init__(scanner=object())
self._data = list(data) self._data = list(data)
self.parse_sort_calls = [] self.parse_sort_calls = []
self.fetch_sorted_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) params = ModelCacheRepository.parse_sort(sort_by)
self.parse_sort_calls.append(sort_by) self.parse_sort_calls.append(sort_by)
return params return params
@@ -44,8 +45,9 @@ class StubRepository:
return list(self._data) return list(self._data)
class StubFilterSet: class StubFilterSet(ModelFilterSet):
def __init__(self, result): def __init__(self, result):
super().__init__(settings=StubSettings({}))
self.result = list(result) self.result = list(result)
self.calls = [] self.calls = []
@@ -54,8 +56,9 @@ class StubFilterSet:
return list(self.result) return list(self.result)
class StubSearchStrategy: class StubSearchStrategy(SearchStrategy):
def __init__(self, search_result): def __init__(self, search_result):
super().__init__()
self.search_result = list(search_result) self.search_result = list(search_result)
self.normalize_calls = [] self.normalize_calls = []
self.apply_calls = [] self.apply_calls = []
@@ -67,7 +70,7 @@ class StubSearchStrategy:
normalized.update(options) normalized.update(options)
return normalized 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)) self.apply_calls.append((list(data), search_term, options, fuzzy))
return list(self.search_result) return list(self.search_result)
@@ -269,8 +272,9 @@ async def test_get_paginated_data_filters_and_searches_combination():
assert response["total_pages"] == 1 assert response["total_pages"] == 1
class PassThroughFilterSet: class PassThroughFilterSet(ModelFilterSet):
def __init__(self): def __init__(self):
super().__init__(settings=StubSettings({}))
self.calls = [] self.calls = []
def apply(self, data, criteria): def apply(self, data, criteria):
@@ -278,8 +282,9 @@ class PassThroughFilterSet:
return list(data) return list(data)
class NoSearchStrategy: class NoSearchStrategy(SearchStrategy):
def __init__(self): def __init__(self):
super().__init__()
self.normalize_calls = [] self.normalize_calls = []
self.apply_called = False self.apply_called = False
@@ -355,7 +360,7 @@ async def test_get_paginated_data_filters_by_update_status():
filter_set=filter_set, filter_set=filter_set,
search_strategy=search_strategy, search_strategy=search_strategy,
settings_provider=settings, settings_provider=settings,
update_service=update_service, update_service=update_service, # pyright: ignore[reportArgumentType]
) )
response = await service.get_paginated_data( 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, filter_set=filter_set,
search_strategy=search_strategy, search_strategy=search_strategy,
settings_provider=settings, settings_provider=settings,
update_service=update_service, update_service=update_service, # pyright: ignore[reportArgumentType]
) )
response = await service.get_paginated_data( 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, filter_set=filter_set,
search_strategy=search_strategy, search_strategy=search_strategy,
settings_provider=settings, settings_provider=settings,
update_service=update_service, update_service=update_service, # pyright: ignore[reportArgumentType]
) )
response = await service.get_paginated_data( response = await service.get_paginated_data(
@@ -561,7 +566,7 @@ async def test_version_grouping_same_base_prefers_matching_base():
filter_set=filter_set, filter_set=filter_set,
search_strategy=search_strategy, search_strategy=search_strategy,
settings_provider=settings, settings_provider=settings,
update_service=update_service, update_service=update_service, # pyright: ignore[reportArgumentType]
) )
response = await service.get_paginated_data( 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, filter_set=filter_set,
search_strategy=search_strategy, search_strategy=search_strategy,
settings_provider=settings, settings_provider=settings,
update_service=update_service, update_service=update_service, # pyright: ignore[reportArgumentType]
) )
response = await service.get_paginated_data( response = await service.get_paginated_data(
@@ -694,7 +699,7 @@ async def test_get_paginated_data_filters_update_available_only():
filter_set=filter_set, filter_set=filter_set,
search_strategy=search_strategy, search_strategy=search_strategy,
settings_provider=settings, settings_provider=settings,
update_service=update_service, update_service=update_service, # pyright: ignore[reportArgumentType]
) )
response = await service.get_paginated_data( response = await service.get_paginated_data(
@@ -1028,7 +1033,7 @@ def test_model_filter_set_supports_legacy_tag_arrays():
{"model_name": "AnimeOnly", "tags": ["anime"]}, {"model_name": "AnimeOnly", "tags": ["anime"]},
] ]
criteria = FilterCriteria(tags=["style"]) criteria = FilterCriteria(tags=["style"]) # pyright: ignore[reportArgumentType]
result = filter_set.apply(data, criteria) result = filter_set.apply(data, criteria)
assert [item["model_name"] for item in result] == ["StyleOnly", "StyleAnime"] assert [item["model_name"] for item in result] == ["StyleOnly", "StyleAnime"]
+7 -7
View File
@@ -333,7 +333,7 @@ class TestBatchImportService:
) )
service = BatchImportService( service = BatchImportService(
analysis_service=analysis_service, analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service, persistence_service=persistence_service,
ws_manager=ws_manager, ws_manager=ws_manager,
logger=logger, logger=logger,
@@ -445,8 +445,8 @@ class TestBatchImportServiceEdgeCases:
logger = logging.getLogger("test") logger = logging.getLogger("test")
return BatchImportService( return BatchImportService(
analysis_service=analysis_service, analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service, persistence_service=persistence_service, # pyright: ignore[reportArgumentType]
ws_manager=ws_manager, ws_manager=ws_manager,
logger=logger, logger=logger,
) )
@@ -506,8 +506,8 @@ class TestBatchImportServiceEdgeCases:
(tmp_path / "test.png").write_bytes(b"fake-image") (tmp_path / "test.png").write_bytes(b"fake-image")
service = BatchImportService( service = BatchImportService(
analysis_service=analysis_service, analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service, persistence_service=persistence_service, # pyright: ignore[reportArgumentType]
ws_manager=ws_manager, ws_manager=ws_manager,
logger=logger, logger=logger,
) )
@@ -571,8 +571,8 @@ class TestInputValidation:
logger = logging.getLogger("test") logger = logging.getLogger("test")
return BatchImportService( return BatchImportService(
analysis_service=analysis_service, analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service, persistence_service=persistence_service, # pyright: ignore[reportArgumentType]
ws_manager=ws_manager, ws_manager=ws_manager,
logger=logger, logger=logger,
) )
+15 -6
View File
@@ -84,6 +84,7 @@ class TestCacheEntryValidator:
result = CacheEntryValidator.validate(entry, auto_repair=False) result = CacheEntryValidator.validate(entry, auto_repair=False)
assert result.is_valid is True assert result.is_valid is True
assert result.entry is not None
assert result.entry['sha256'] == '' assert result.entry['sha256'] == ''
assert result.entry['hash_status'] == 'pending' assert result.entry['hash_status'] == 'pending'
@@ -141,7 +142,7 @@ class TestCacheEntryValidator:
def test_validate_none_entry(self): def test_validate_none_entry(self):
"""Test validation handles None entry""" """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.is_valid is False
assert result.repaired is False assert result.repaired is False
@@ -150,7 +151,7 @@ class TestCacheEntryValidator:
def test_validate_non_dict_entry(self): def test_validate_non_dict_entry(self):
"""Test validation handles non-dict entry""" """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.is_valid is False
assert result.repaired is False assert result.repaired is False
@@ -169,6 +170,7 @@ class TestCacheEntryValidator:
assert result.is_valid is True assert result.is_valid is True
assert result.repaired is True assert result.repaired is True
assert result.entry is not None
assert result.entry['file_name'] == '' assert result.entry['file_name'] == ''
assert result.entry['model_name'] == '' assert result.entry['model_name'] == ''
assert result.entry['tags'] == [] assert result.entry['tags'] == []
@@ -186,6 +188,7 @@ class TestCacheEntryValidator:
assert result.is_valid is True assert result.is_valid is True
assert result.repaired is True assert result.repaired is True
assert result.entry is not None
assert result.entry['size'] == 0 # Default value assert result.entry['size'] == 0 # Default value
assert result.entry['tags'] == [] # Default value assert result.entry['tags'] == [] # Default value
@@ -199,6 +202,7 @@ class TestCacheEntryValidator:
result = CacheEntryValidator.validate(entry, auto_repair=True) result = CacheEntryValidator.validate(entry, auto_repair=True)
assert result.is_valid is True assert result.is_valid is True
assert result.entry is not None
assert result.entry['sha256'] == 'abc123def456' assert result.entry['sha256'] == 'abc123def456'
def test_validate_batch_all_valid(self): def test_validate_batch_all_valid(self):
@@ -262,8 +266,8 @@ class TestCacheEntryValidator:
def test_get_file_path_safe_not_dict(self): def test_get_file_path_safe_not_dict(self):
"""Test safe file_path extraction from non-dict""" """Test safe file_path extraction from non-dict"""
assert CacheEntryValidator.get_file_path_safe(None) == '' assert CacheEntryValidator.get_file_path_safe(None) == '' # pyright: ignore[reportArgumentType]
assert CacheEntryValidator.get_file_path_safe('string') == '' assert CacheEntryValidator.get_file_path_safe('string') == '' # pyright: ignore[reportArgumentType]
def test_get_sha256_safe(self): def test_get_sha256_safe(self):
"""Test safe sha256 extraction""" """Test safe sha256 extraction"""
@@ -277,8 +281,8 @@ class TestCacheEntryValidator:
def test_get_sha256_safe_not_dict(self): def test_get_sha256_safe_not_dict(self):
"""Test safe sha256 extraction from non-dict""" """Test safe sha256 extraction from non-dict"""
assert CacheEntryValidator.get_sha256_safe(None) == '' assert CacheEntryValidator.get_sha256_safe(None) == '' # pyright: ignore[reportArgumentType]
assert CacheEntryValidator.get_sha256_safe('string') == '' assert CacheEntryValidator.get_sha256_safe('string') == '' # pyright: ignore[reportArgumentType]
def test_validate_with_all_optional_fields(self): def test_validate_with_all_optional_fields(self):
"""Test validation with all optional fields present""" """Test validation with all optional fields present"""
@@ -358,6 +362,7 @@ class TestAutov3Validation:
) )
assert result.is_valid is True assert result.is_valid is True
assert result.entry is not None
assert result.entry['autov3'] == 'abcdef123456' assert result.entry['autov3'] == 'abcdef123456'
assert result.repaired is True assert result.repaired is True
@@ -378,6 +383,7 @@ class TestAutov3Validation:
assert result.is_valid is True assert result.is_valid is True
assert result.repaired is False assert result.repaired is False
assert result.entry is not None
assert result.entry['autov3'] is None assert result.entry['autov3'] is None
def test_validate_absent_autov3_is_valid_and_not_counted_as_repair(self): 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.is_valid is True
assert result.repaired is False assert result.repaired is False
assert result.entry is not None
assert 'autov3' not in result.entry assert 'autov3' not in result.entry
def test_validate_short_autov3_still_valid_and_repaired_to_none(self): 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.is_valid is True
assert result.entry is not None
assert result.entry['autov3'] is None assert result.entry['autov3'] is None
assert result.repaired is True assert result.repaired is True
@@ -407,5 +415,6 @@ class TestAutov3Validation:
) )
assert result.is_valid is True assert result.is_valid is True
assert result.entry is not None
assert result.entry['autov3'] is None assert result.entry['autov3'] is None
assert result.repaired is True assert result.repaired is True
+4 -3
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
import json import json
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any
import pytest import pytest
@@ -13,7 +14,7 @@ from py.utils import example_images_download_manager as download_module
class StubScanner: class StubScanner:
"""Scanner double returning predetermined cache contents.""" """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._cache = SimpleNamespace(raw_data=models)
async def get_cached_data(self): async def get_cached_data(self):
@@ -58,9 +59,9 @@ class RecordingWebSocketManager:
"""Collects broadcast payloads for assertions.""" """Collects broadcast payloads for assertions."""
def __init__(self) -> None: 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) self.payloads.append(payload)
+3 -3
View File
@@ -4,7 +4,7 @@ import asyncio
import json import json
import os import os
from pathlib import Path from pathlib import Path
from typing import List from typing import Any, List
from unittest.mock import AsyncMock, MagicMock, patch from unittest.mock import AsyncMock, MagicMock, patch
import pytest import pytest
@@ -17,9 +17,9 @@ from py.utils.models import CheckpointMetadata
class RecordingWebSocketManager: class RecordingWebSocketManager:
def __init__(self) -> None: 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) self.payloads.append(payload)
+3 -3
View File
@@ -1,6 +1,6 @@
import os import os
from pathlib import Path from pathlib import Path
from typing import List from typing import Any, List
import pytest import pytest
@@ -12,9 +12,9 @@ from py.services.persistent_model_cache import PersistedCacheData
class RecordingWebSocketManager: class RecordingWebSocketManager:
def __init__(self) -> None: 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) self.payloads.append(payload)
+6 -2
View File
@@ -1,4 +1,5 @@
import copy import copy
from typing import Any, Dict
from unittest.mock import AsyncMock from unittest.mock import AsyncMock
import pytest import pytest
@@ -34,7 +35,7 @@ def downloader(monkeypatch):
return instance 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" version_name = "v2.0" if version_id != 1976567 else "v1.0"
file_sha = "e2b7a280d6539556f23f380b3f71e4e22bc4524445c4c96526e117c6005c6ad3" file_sha = "e2b7a280d6539556f23f380b3f71e4e22bc4524445c4c96526e117c6005c6ad3"
return { return {
@@ -110,6 +111,7 @@ async def test_get_model_by_hash_transforms_payload(downloader):
result, error = await client.get_model_by_hash("abc") result, error = await client.get_model_by_hash("abc")
assert error is None assert error is None
assert result is not None
assert result["id"] == 1976567 assert result["id"] == 1976567
assert result["nsfwLevel"] == 31 assert result["nsfwLevel"] == 31
assert result["trainedWords"] == ["mxpln"] 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) base_payload = _base_civarchive_payload(version_id=2042594, trigger="mxpln-new", nsfw_level=5)
other_payload = _base_civarchive_payload() other_payload = _base_civarchive_payload()
responses = { responses: Dict[Any, Dict[str, Any]] = {
(base_url, None): base_payload, (base_url, None): base_payload,
(base_url, (("modelVersionId", "2042594"),)): base_payload, (base_url, (("modelVersionId", "2042594"),)): base_payload,
(base_url, (("modelVersionId", "1976567"),)): other_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") result = await client.get_model_versions("1746460")
assert result is not None
assert result["name"] == "Mixplin Style [Illustrious]" assert result["name"] == "Mixplin Style [Illustrious]"
assert result["type"] == "LORA" assert result["type"] == "LORA"
versions = result["modelVersions"] 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") result, error = await client.get_model_by_hash("fallback")
assert error is None assert error is None
assert result is not None
assert result["id"] == 1976567 assert result["id"] == 1976567
assert result["model"]["name"] == "Mixplin Style [Illustrious]" assert result["model"]["name"] == "Mixplin Style [Illustrious]"
assert any("/models/1746460" in call["url"] for call in downloader.calls) 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: class TestCivitaiBaseModelService:
"""Test suite for CivitaiBaseModelService.""" """Test suite for CivitaiBaseModelService."""
service = CivitaiBaseModelService()
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
def setup_service(self): def setup_service(self):
"""Create a fresh service instance for each test.""" """Create a fresh service instance for each test."""
@@ -46,7 +48,7 @@ class TestCivitaiBaseModelService:
def test_generate_abbreviation_edge_cases(self): def test_generate_abbreviation_edge_cases(self):
"""Test abbreviation generation edge cases.""" """Test abbreviation generation edge cases."""
assert self.service.generate_abbreviation("") == "OTH" 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): def test_cache_status_no_cache(self):
"""Test cache status when no cache exists.""" """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") result, error = await client.get_model_by_hash("hash")
assert error is None assert error is None
assert result is not None
assert result["model"]["description"] == "desc" assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"] assert result["model"]["tags"] == ["tag"]
assert result["creator"] == {"username": "user"} assert result["creator"] == {"username": "user"}
@@ -254,7 +255,7 @@ async def test_get_model_versions_bulk_success(monkeypatch, downloader):
client = await CivitaiClient.get_instance() 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 == { assert result == {
1: { 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) result = await client.get_model_version(version_id=7)
assert result is not None
assert result["model"]["description"] == "desc" assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"] assert result["model"]["tags"] == ["tag"]
assert result["creator"] == {"username": "user"} 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) result = await client.get_model_version(model_id=99, version_id=7)
assert result is not None
assert result["id"] == 7 assert result["id"] == 7
assert result["model"]["description"] == "desc" assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"] 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) result = await client.get_model_version(model_id=99, version_id=7)
assert result is not None
assert result["id"] == 7 assert result["id"] == 7
assert result["model"]["description"] == "desc" assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"] 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) result = await client.get_model_version(model_id=99, version_id=7)
assert result is not None
assert result["modelId"] == 99 assert result["modelId"] == 99
assert result["model"]["name"] == "Model" assert result["model"]["name"] == "Model"
assert result["model"]["type"] == "LORA" assert result["model"]["type"] == "LORA"
@@ -503,6 +508,7 @@ async def test_get_model_version_info_success(monkeypatch, downloader):
assert result == expected assert result == expected
assert error is None assert error is None
assert result is not None
assert "comfy" not in result["images"][0]["meta"] assert "comfy" not in result["images"][0]["meta"]
assert result["images"][0]["meta"]["other"] == "keep" 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) result = await parser.parse_metadata(metadata)
@@ -272,7 +272,7 @@ async def test_parse_metadata_handles_modelVersionIds(monkeypatch):
"modelVersionIds": [2398829, 2398838], "modelVersionIds": [2398829, 2398838],
} }
assert parser.is_metadata_matching(metadata) assert parser.is_metadata_matching(metadata) # pyright: ignore[reportArgumentType]
result = await parser.parse_metadata(metadata) 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() downloads = await manager.get_active_downloads()
persisted = await manager._aria2_state_store.get("download-1") persisted = await manager._aria2_state_store.get("download-1")
assert persisted is not None
assert downloads["downloads"] == [ 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() downloads = await manager.get_active_downloads()
persisted = await manager._aria2_state_store.get("download-1") persisted = await manager._aria2_state_store.get("download-1")
assert persisted is not None
assert downloads["downloads"] == [ assert downloads["downloads"] == [
{ {
@@ -3,6 +3,7 @@
import os import os
from pathlib import Path from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from typing import Optional
from unittest.mock import AsyncMock from unittest.mock import AsyncMock
import pytest import pytest
@@ -99,7 +100,7 @@ async def test_execute_download_uses_rewritten_civitai_preview(monkeypatch, tmp_
self.file_path = str(path) self.file_path = str(path)
self.sha256 = "sha256" self.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url: Optional[str] = None
self.autov3 = None self.autov3 = None
self.preview_nsfw_level = 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 any("width=450,optimized=true" in url for url in preview_urls)
assert dummy_downloader.memory_calls == 0 assert dummy_downloader.memory_calls == 0
assert optimize_called["value"] is False assert optimize_called["value"] is False
assert metadata.preview_url is not None
assert metadata.preview_url.endswith(".jpeg") assert metadata.preview_url.endswith(".jpeg")
assert metadata.preview_nsfw_level == 2 assert metadata.preview_nsfw_level == 2
stored_preview = manager._active_downloads["dl"]["preview_path"] 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.file_path = str(path)
self.sha256 = "sha256" self.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url: Optional[str] = None
self.autov3 = None self.autov3 = None
self.preview_nsfw_level = 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.file_path = str(path)
self.sha256 = "sha256" self.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url: Optional[str] = None
self.autov3 = None self.autov3 = None
self.preview_nsfw_level = 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.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs): def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path) 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.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs): def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path) 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.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs): def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path) 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.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
self.preview_nsfw_level = 0 self.preview_nsfw_level = 0
self.sub_type = "checkpoint" self.sub_type = "checkpoint"
@@ -450,7 +450,7 @@ async def test_execute_download_extracts_zip_single_model(monkeypatch, tmp_path)
self.sha256 = "sha256" self.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs): def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path) 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.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs): def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path) 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.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs): def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path) 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.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs): def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path) 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.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs): def generate_unique_filename(self, *_args, **_kwargs):
return "renamed.safetensors" return "renamed.safetensors"
@@ -1356,7 +1356,7 @@ async def test_execute_download_rejects_conflicting_aria2_partial_path(tmp_path)
self.sha256 = "sha256" self.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs): def generate_unique_filename(self, *_args, **_kwargs):
raise AssertionError("should not rename") 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.sha256 = "sha256"
self.file_name = path.stem self.file_name = path.stem
self.preview_url = None self.preview_url = None
self.autov3 = None self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs): def generate_unique_filename(self, *_args, **_kwargs):
raise AssertionError("should not rename") 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 manager._active_downloads["new-download"]["file_path"] == str(target_path)
assert dummy_aria2.calls == [("reassign_transfer", "old-download", "new-download")] 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("old-download") is None
assert (await manager._aria2_state_store.get("new-download"))["save_path"] == str( persisted = await manager._aria2_state_store.get("new-download")
target_path assert persisted is not None
) assert persisted["save_path"] == str(target_path)
def test_is_same_aria2_download_request_requires_version_id_match(): 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: class FakeStream:
def __init__(self, chunks: Sequence[Sequence] | Sequence[bytes]): def __init__(self, chunks: Sequence[bytes | tuple[bytes, float]]):
self._chunks = list(chunks) self._chunks = list(chunks)
async def read(self, _chunk_size: int) -> bytes: async def read(self, _chunk_size: int) -> bytes:
@@ -25,6 +25,7 @@ class FakeStream:
payload = item[0] payload = item[0]
delay = item[1] delay = item[1]
assert isinstance(payload, bytes)
await asyncio.sleep(delay) await asyncio.sleep(delay)
return payload return payload
@@ -84,11 +85,11 @@ def _build_downloader(responses, *, max_retries=0):
downloader.max_retries = max_retries downloader.max_retries = max_retries
downloader.base_delay = 0 downloader.base_delay = 0
fake_session = FakeSession(responses) fake_session = FakeSession(responses)
downloader._session = fake_session downloader._session = fake_session # pyright: ignore[reportAttributeAccessIssue]
downloader._session_created_at = datetime.now() downloader._session_created_at = datetime.now()
downloader._proxy_url = None downloader._proxy_url = None
async def _noop_create_session(): async def _noop_create_session():
downloader._session = fake_session downloader._session = fake_session # pyright: ignore[reportAttributeAccessIssue]
downloader._session_created_at = datetime.now() downloader._session_created_at = datetime.now()
downloader._proxy_url = None downloader._proxy_url = None
@@ -96,6 +97,13 @@ def _build_downloader(responses, *, max_retries=0):
return downloader 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 @pytest.mark.asyncio
async def test_download_file_preserves_incomplete_part_when_size_mismatch(tmp_path): async def test_download_file_preserves_incomplete_part_when_size_mismatch(tmp_path):
target_path = tmp_path / "model" / "file.bin" 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 success is True
assert Path(result_path).read_bytes() == payload 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() 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 success is True
assert Path(result_path).read_bytes() == b"abcdef" assert Path(result_path).read_bytes() == b"abcdef"
assert downloader._session._get_calls == 2 assert _session(downloader)._get_calls == 2
assert downloader._session.requests[1]["headers"]["Range"] == "bytes=3-" assert _session(downloader).requests[1]["headers"]["Range"] == "bytes=3-"
assert not Path(str(target_path) + ".part").exists() 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 success is True
assert Path(result_path).read_bytes() == b"abcdef" assert Path(result_path).read_bytes() == b"abcdef"
assert first_response.released is True assert first_response.released is True
assert downloader._session.requests[0]["headers"]["Range"] == "bytes=3-" assert _session(downloader).requests[0]["headers"]["Range"] == "bytes=3-"
assert downloader._session.requests[1]["url"] == redirected_url assert _session(downloader).requests[1]["url"] == redirected_url
assert downloader._session.requests[1]["headers"]["Range"] == "bytes=3-" 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() 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['success'] is True
assert result['moved_total'] == 2 assert result['moved_total'] == 2
assert not empty_folder.exists() assert not empty_folder.exists()
@@ -4,6 +4,7 @@ import asyncio
import json import json
from pathlib import Path from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any
import pytest import pytest
@@ -15,18 +16,18 @@ class RecordingWebSocketManager:
"""Collects broadcast payloads for assertions.""" """Collects broadcast payloads for assertions."""
def __init__(self) -> None: 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) self.payloads.append(payload)
class StubScanner: class StubScanner:
"""Scanner double returning predetermined cache contents.""" """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._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): async def get_cached_data(self):
return self._cache return self._cache
@@ -39,7 +40,7 @@ class StubScanner:
break break
return True 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)) self.sync_calls.append((file_path, metadata))
for index, model in enumerate(self._cache.raw_data): for index, model in enumerate(self._cache.raw_data):
if model.get("file_path") == metadata.get("file_path"): 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" missing_url = "https://example.com/missing.png"
valid_url = "https://example.com/valid.png" valid_url = "https://example.com/valid.png"
model_metadata = { model_metadata: dict[str, Any] = {
"sha256": model_hash, "sha256": model_hash,
"model_name": "Missing Example", "model_name": "Missing Example",
"file_path": str(model_path), "file_path": str(model_path),
+4 -3
View File
@@ -2,17 +2,18 @@ import asyncio
import json import json
import pytest import pytest
from pathlib import Path from pathlib import Path
from typing import Any
from py.services.settings_manager import get_settings_manager from py.services.settings_manager import get_settings_manager
from py.utils import example_images_download_manager as download_module from py.utils import example_images_download_manager as download_module
class RecordingWebSocketManager: class RecordingWebSocketManager:
def __init__(self) -> None: 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) self.payloads.append(payload)
class StubScanner: class StubScanner:
def __init__(self, models: list[dict]) -> None: def __init__(self, models: list[dict[str, Any]]) -> None:
self.raw_data = models self.raw_data = models
async def get_cached_data(self): async def get_cached_data(self):
class Cache: class Cache:
+9 -1
View File
@@ -1,16 +1,24 @@
"""Tests for license-based filtering functionality.""" """Tests for license-based filtering functionality."""
import pytest import pytest
from typing import Any
from unittest.mock import Mock, AsyncMock from unittest.mock import Mock, AsyncMock
from py.services.base_model_service import BaseModelService from py.services.base_model_service import BaseModelService
from py.utils.civitai_utils import build_license_flags from py.utils.civitai_utils import build_license_flags
from py.utils.models import BaseModelMetadata
class DummyModelService(BaseModelService): class DummyModelService(BaseModelService):
"""Dummy implementation of BaseModelService for testing.""" """Dummy implementation of BaseModelService for testing."""
def __init__(self): def __init__(self):
super().__init__(
model_type="test",
scanner=Mock(),
metadata_class=BaseModelMetadata,
settings_provider=Mock(),
)
# Mock the required attributes # Mock the required attributes
self.model_type = "test" self.model_type = "test"
self.scanner = Mock() self.scanner = Mock()
@@ -28,7 +36,7 @@ class DummyModelService(BaseModelService):
self.scanner.get_cached_data = mock_get_cached_data 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.""" """Required abstract method implementation."""
return model_data return model_data
@@ -1,10 +1,12 @@
"""Integration tests for license-based filtering in BaseModelService.""" """Integration tests for license-based filtering in BaseModelService."""
import pytest import pytest
from typing import Any
from unittest.mock import Mock, AsyncMock from unittest.mock import Mock, AsyncMock
from py.services.base_model_service import BaseModelService from py.services.base_model_service import BaseModelService
from py.utils.civitai_utils import build_license_flags 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 from py.services.model_query import ModelCacheRepository, ModelFilterSet, SearchStrategy, SortParams
@@ -12,6 +14,12 @@ class DummyModelService(BaseModelService):
"""Dummy implementation of BaseModelService for testing.""" """Dummy implementation of BaseModelService for testing."""
def __init__(self): def __init__(self):
super().__init__(
model_type="test",
scanner=Mock(),
metadata_class=BaseModelMetadata,
settings_provider=Mock(),
)
# Mock the required attributes # Mock the required attributes
self.model_type = "test" self.model_type = "test"
self.scanner = Mock() self.scanner = Mock()
@@ -33,7 +41,7 @@ class DummyModelService(BaseModelService):
self.scanner.get_cached_data = mock_get_cached_data 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.""" """Required abstract method implementation."""
return model_data return model_data
+3
View File
@@ -57,6 +57,9 @@ class MockSession:
def __init__(self, response): def __init__(self, response):
self._response = response self._response = response
self.closed = False self.closed = False
self.last_url = None
self.last_json = None
self.last_headers = None
def post(self, url, json=None, headers=None): def post(self, url, json=None, headers=None):
self.last_url = url self.last_url = url
+3 -2
View File
@@ -1,4 +1,5 @@
from types import SimpleNamespace from types import SimpleNamespace
from typing import Optional
from unittest.mock import AsyncMock from unittest.mock import AsyncMock
import pytest import pytest
@@ -21,13 +22,13 @@ class DummyProvider(ModelMetadataProvider):
async def get_model_versions_bulk(self, model_ids): async def get_model_versions_bulk(self, model_ids):
return None 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 return None
async def get_model_version_info(self, version_id: str): async def get_model_version_info(self, version_id: str):
return None, None 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 return None
+7 -7
View File
@@ -10,7 +10,7 @@ from py.services.metadata_sync_service import MetadataSyncService
class DummySettings: class DummySettings:
def __init__(self, values: dict | None = None) -> None: def __init__(self, values: dict[str, Any] | None = None) -> None:
self._values = values or {} self._values = values or {}
def get(self, key: str, default=None): def get(self, key: str, default=None):
@@ -19,7 +19,7 @@ class DummySettings:
def build_service( def build_service(
*, *,
settings_values: dict | None = None, settings_values: dict[str, Any] | None = None,
default_provider: SimpleNamespace | None = None, default_provider: SimpleNamespace | None = None,
provider_selector: AsyncMock | None = None, provider_selector: AsyncMock | None = None,
): ):
@@ -43,7 +43,7 @@ def build_service(
service = MetadataSyncService( service = MetadataSyncService(
metadata_manager=metadata_manager, metadata_manager=metadata_manager,
preview_service=preview_service, preview_service=preview_service,
settings=settings, settings=settings, # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=default_provider_factory, default_metadata_provider_factory=default_provider_factory,
metadata_provider_selector=provider_selector, 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 helpers.metadata_manager.hydrate_model_data.side_effect = hydrate
model_data = { model_data: Dict[str, Any] = {
"model_name": "Local", "model_name": "Local",
"folder": "root", "folder": "root",
"file_path": str(model_path), "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 helpers.metadata_manager.hydrate_model_data.side_effect = hydrate
model_data = { model_data: Dict[str, Any] = {
"model_name": "Local", "model_name": "Local",
"folder": "sub", "folder": "root",
"file_path": str(model_path), "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) helpers.default_provider.get_model_by_hash.return_value = (civarchive_payload, None)
model_path = tmp_path / "model.safetensors" model_path = tmp_path / "model.safetensors"
model_data = { model_data: Dict[str, Any] = {
"model_name": "High Quality", "model_name": "High Quality",
"metadata_source": "civitai_api", "metadata_source": "civitai_api",
"civitai": existing_civitai, "civitai": existing_civitai,
+40 -23
View File
@@ -1,6 +1,7 @@
import json import json
import os import os
from pathlib import Path from pathlib import Path
from typing import Any, Dict, cast
import pytest import pytest
@@ -9,6 +10,11 @@ from py.utils.metadata_manager import MetadataManager
from py.utils.models import LoraMetadata 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: class ScannerWithRoots:
def __init__(self, roots): def __init__(self, roots):
self._roots = list(roots) self._roots = list(roots)
@@ -117,7 +123,7 @@ async def test_delete_model_rejects_path_outside_roots(tmp_path: Path):
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=scanner, scanner=scanner,
metadata_manager=DummyMetadataManager({"civitai": {"modelId": 1}}), metadata_manager=DummyMetadataManager({"civitai": {"modelId": 1}}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
# Path within root should work (model file exists) # Path within root should work (model file exists)
result = await service.delete_model(str(model)) 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( service2 = ModelLifecycleService(
scanner=scanner2, scanner=scanner2,
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
with pytest.raises(ValueError, match="outside configured library"): with pytest.raises(ValueError, match="outside configured library"):
await service2.delete_model(str(outside)) await service2.delete_model(str(outside))
@@ -148,7 +154,7 @@ async def test_rename_model_rejects_path_outside_roots(tmp_path: Path):
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=scanner, scanner=scanner,
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
outside = tmp_path / "outside.safetensors" outside = tmp_path / "outside.safetensors"
outside.write_bytes(b"data") outside.write_bytes(b"data")
@@ -170,7 +176,7 @@ async def test_bulk_delete_rejects_any_path_outside_roots(tmp_path: Path):
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=scanner, scanner=scanner,
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
with pytest.raises(ValueError, match="outside configured library"): with pytest.raises(ValueError, match="outside configured library"):
await service.bulk_delete_models([str(model_ok), str(outside)]) await service.bulk_delete_models([str(model_ok), str(outside)])
@@ -209,7 +215,7 @@ class VersionAwareScanner:
continue continue
candidate = civitai.get("modelId") candidate = civitai.get("modelId")
try: try:
normalized = int(candidate) normalized = int(cast(Any, candidate))
except (TypeError, ValueError): except (TypeError, ValueError):
continue continue
if normalized != model_id: if normalized != model_id:
@@ -297,6 +303,7 @@ async def test_rename_model_preserves_compound_extensions(tmp_path: Path):
assert expected_main.exists() assert expected_main.exists()
assert not model_path.exists() assert not model_path.exists()
assert isinstance(result["new_file_path"], str)
assert result["new_file_path"].endswith(f"{new_name}.safetensors") assert result["new_file_path"].endswith(f"{new_name}.safetensors")
assert expected_preview.exists() assert expected_preview.exists()
assert not preview_path.exists() assert not preview_path.exists()
@@ -343,7 +350,7 @@ async def test_delete_model_updates_update_service(tmp_path: Path):
scanner=scanner, scanner=scanner,
metadata_manager=metadata_manager, metadata_manager=metadata_manager,
metadata_loader=metadata_loader, metadata_loader=metadata_loader,
update_service=update_service, update_service=update_service, # pyright: ignore[reportArgumentType]
) )
result = await service.delete_model(model_path.as_posix()) 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 expected_main.exists()
assert not model_path.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 result["new_file_path"].endswith(f"{new_name}{old_extension}")
assert expected_preview.exists() assert expected_preview.exists()
assert not preview_path.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}" expected_main = tmp_path / f"{new_name}{old_extension}"
assert expected_main.exists() assert expected_main.exists()
assert result["new_file_path"] == expected_main.as_posix() 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()) saved_metadata = json.loads((tmp_path / f"{new_name}.metadata.json").read_text())
assert saved_metadata["file_name"] == new_name 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 model_path.exists()
assert not metadata_path.exists() assert not metadata_path.exists()
assert not preview_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 = [] saved_metadata = []
class SavingMetadataManager: 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())) 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() return metadata_payload.copy()
service = ModelLifecycleService( 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)) result = await service.exclude_model(str(model_path))
assert result["success"] is True assert result["success"] is True
assert isinstance(result["message"], str)
assert "excluded" in result["message"].lower() assert "excluded" in result["message"].lower()
assert saved_metadata[0][1]["exclude"] is True assert saved_metadata[0][1]["exclude"] is True
assert str(model_path) in scanner._excluded_models 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) scanner = TagCountScanner(raw_data)
class DummyMetadataManagerLocal: class DummyMetadataManagerLocal:
async def save_metadata(self, path: str, metadata: dict): async def save_metadata(self, path: str, metadata: Dict[str, Any]):
pass pass
async def metadata_loader(path: str): async def metadata_loader(path: str):
@@ -607,7 +623,7 @@ async def test_exclude_model_empty_path_raises_error():
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=VersionAwareScanner([]), scanner=VersionAwareScanner([]),
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
with pytest.raises(ValueError, match="Model path is required"): 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 = [] saved_metadata = []
class SavingMetadataManager: 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())) saved_metadata.append((path, metadata.copy()))
await MetadataManager.save_metadata(path, metadata) 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)) result = await service.unexclude_model(str(model_path))
assert result["success"] is True assert result["success"] is True
assert isinstance(result["message"], str)
assert "restored" in result["message"].lower() assert "restored" in result["message"].lower()
assert scanner._excluded_models == [] assert scanner._excluded_models == []
assert saved_metadata[0][1]["exclude"] is False 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( service = ModelLifecycleService(
scanner=scanner, scanner=scanner,
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
result = await service.bulk_delete_models(file_paths) result = await service.bulk_delete_models(file_paths)
@@ -716,7 +733,7 @@ async def test_bulk_delete_models_empty_list_raises_error():
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=VersionAwareScanner([]), scanner=VersionAwareScanner([]),
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
with pytest.raises(ValueError, match="No file paths provided"): with pytest.raises(ValueError, match="No file paths provided"):
@@ -734,7 +751,7 @@ async def test_delete_model_empty_path_raises_error():
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=VersionAwareScanner([]), scanner=VersionAwareScanner([]),
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
with pytest.raises(ValueError, match="Model path is required"): with pytest.raises(ValueError, match="Model path is required"):
@@ -747,7 +764,7 @@ async def test_rename_model_empty_path_raises_error():
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=DummyScanner(), scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
with pytest.raises(ValueError, match="required"): with pytest.raises(ValueError, match="required"):
@@ -763,7 +780,7 @@ async def test_rename_model_empty_name_raises_error(tmp_path: Path):
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=DummyScanner(), scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
with pytest.raises(ValueError, match="required"): with pytest.raises(ValueError, match="required"):
@@ -779,7 +796,7 @@ async def test_rename_model_invalid_characters_raises_error(tmp_path: Path):
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=DummyScanner(), scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
invalid_names = [ invalid_names = [
@@ -817,7 +834,7 @@ async def test_rename_model_existing_file_raises_error(tmp_path: Path):
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=DummyScanner(), scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
with pytest.raises(ValueError, match="already exists"): with pytest.raises(ValueError, match="already exists"):
@@ -837,7 +854,7 @@ async def test_extract_model_id_from_civitai_payload():
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=DummyScanner(), scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
# Test civitai.modelId # Test civitai.modelId
@@ -863,7 +880,7 @@ async def test_extract_model_id_returns_none_for_invalid_payload():
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=DummyScanner(), scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
assert service._extract_model_id_from_payload({}) is None assert service._extract_model_id_from_payload({}) is None
@@ -879,7 +896,7 @@ async def test_extract_model_id_handles_string_values():
service = ModelLifecycleService( service = ModelLifecycleService(
scanner=DummyScanner(), scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}), metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {}, metadata_loader=_empty_metadata_loader,
) )
payload = {"civitai": {"modelId": "54321"}} 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.errors import RateLimitError
from py.services.model_metadata_provider import ( from py.services.model_metadata_provider import (
FallbackMetadataProvider, FallbackMetadataProvider,
ModelMetadataProvider,
RateLimitRetryingProvider, RateLimitRetryingProvider,
) )
class RateLimitThenSuccessProvider: class RateLimitThenSuccessProvider(ModelMetadataProvider):
def __init__(self) -> None: def __init__(self) -> None:
self.calls = 0 self.calls = 0
@@ -20,8 +21,20 @@ class RateLimitThenSuccessProvider:
raise RateLimitError("limited", retry_after=1.0) raise RateLimitError("limited", retry_after=1.0)
return {"id": "ok"}, None 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: def __init__(self) -> None:
self.calls = 0 self.calls = 0
@@ -29,8 +42,20 @@ class AlwaysRateLimitedProvider:
self.calls += 1 self.calls += 1
raise RateLimitError("limited") 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: def __init__(self) -> None:
self.calls = 0 self.calls = 0
@@ -38,6 +63,18 @@ class TrackingProvider:
self.calls += 1 self.calls += 1
return {"id": "secondary"}, None 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 @pytest.mark.asyncio
async def test_fallback_retries_same_provider_on_rate_limit(monkeypatch): 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): def test_none_entry_returns_default(self):
"""None entry should return default.""" """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): def test_non_mapping_returns_default(self):
"""Non-mapping entry should return default.""" """Non-mapping entry should return default."""
assert resolve_sub_type("invalid") == "LORA" assert resolve_sub_type("invalid") == "LORA" # pyright: ignore[reportArgumentType]
class TestModelFilterSetWithSubType: class TestModelFilterSetWithSubType:
+8 -8
View File
@@ -2,7 +2,7 @@ import asyncio
import os import os
import sqlite3 import sqlite3
from pathlib import Path from pathlib import Path
from typing import List from typing import Any, Dict, List, Optional
from types import MethodType from types import MethodType
import pytest import pytest
@@ -18,9 +18,9 @@ from py.utils.models import BaseModelMetadata
class RecordingWebSocketManager: class RecordingWebSocketManager:
def __init__(self) -> None: 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) self.payloads.append(payload)
@@ -48,7 +48,7 @@ class DummyScanner(ModelScanner):
*, *,
hash_index: ModelHashIndex | None = None, hash_index: ModelHashIndex | None = None,
excluded_models: List[str] | None = None, excluded_models: List[str] | None = None,
) -> dict: ) -> Optional[Dict[str, Any]]:
hash_index = hash_index or self._hash_index hash_index = hash_index or self._hash_index
excluded_models = excluded_models if excluded_models is not None else self._excluded_models 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 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(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 assert await scanner.check_model_version_exists(999) is False
removed = await scanner._batch_update_cache_for_deleted_models([first_path]) 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] = [] applied: List[str] = []
def _adjust(self, entry: dict) -> dict: def _adjust(self, entry: Dict[str, Any]) -> Dict[str, Any]:
applied.append(entry["file_path"]) applied.append(entry["file_path"])
entry["custom_field"] = "adjusted" entry["custom_field"] = "adjusted"
return entry 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]) scanner = MultiRootDummyScanner([loras_root, extra_root])
await scanner._initialize_cache() await scanner._initialize_cache()
duplicate_entry = { duplicate_entry: Dict[str, Any] = {
"file_path": _normalize_path(extra_root / "one.txt"), "file_path": _normalize_path(extra_root / "one.txt"),
"folder": "", "folder": "",
"sha256": "hash-one", "sha256": "hash-one",
@@ -712,7 +712,7 @@ def test_cache_entries_differ_extra_key():
# ── sync_cache_from_metadata ───────────────────────────────────────── # ── sync_cache_from_metadata ─────────────────────────────────────────
def _make_cache_entry(**overrides) -> dict: def _make_cache_entry(**overrides) -> Dict[str, Any]:
entry = { entry = {
"file_path": "/m/a.safetensors", "file_path": "/m/a.safetensors",
"model_name": "TestModel", "model_name": "TestModel",
@@ -3,13 +3,19 @@ from types import SimpleNamespace
import pytest import pytest
from py.services.model_scanner import ModelScanner from py.services.model_scanner import ModelScanner
from py.utils.models import BaseModelMetadata
class DummyScanner: class DummyScanner(ModelScanner):
def __init__(self, raw_data): def __init__(self, raw_data):
super().__init__(
model_type="dummy",
model_class=BaseModelMetadata,
file_extensions={".safetensors"},
)
self._cache = SimpleNamespace(raw_data=raw_data) 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 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]) await service.update_in_library_versions("lora", 3, [31, 35])
record = await service.get_record("lora", 3) record = await service.get_record("lora", 3)
assert record is not None
assert record.has_update() is False assert record.has_update() is False
+1 -1
View File
@@ -71,7 +71,7 @@ class StubLoraScanner:
def recipe_scanner(tmp_path, monkeypatch): def recipe_scanner(tmp_path, monkeypatch):
monkeypatch.setattr(config, "loras_roots", [str(tmp_path)]) monkeypatch.setattr(config, "loras_roots", [str(tmp_path)])
stub = StubLoraScanner() stub = StubLoraScanner()
scanner = RecipeScanner(lora_scanner=stub) scanner = RecipeScanner(lora_scanner=stub) # pyright: ignore[reportArgumentType]
return scanner return scanner
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -1,4 +1,5 @@
from pathlib import Path from pathlib import Path
from typing import Any, Dict
import pytest import pytest
@@ -346,7 +347,7 @@ def test_update_single_model_update_hash(tmp_path: Path, monkeypatch):
# ── get_models_missing_autov3 ───────────────────────────────────────── # ── 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).""" """Minimal model entry for the models table (autov3 tri-state preserved)."""
return { return {
'file_path': file_path, 'file_path': file_path,
+5 -5
View File
@@ -1,5 +1,5 @@
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any, Dict, List
import pytest import pytest
@@ -55,7 +55,7 @@ async def test_ensure_preview_prefers_rewritten_civitai_image(tmp_path):
exif_utils=exif_utils, exif_utils=exif_utils,
) )
images = [ images: List[Dict[str, object]] = [
{ {
"url": "https://image.civitai.com/container/example/original=true/sample.jpeg", "url": "https://image.civitai.com/container/example/original=true/sample.jpeg",
"type": "image", "type": "image",
@@ -115,7 +115,7 @@ async def test_ensure_preview_falls_back_to_webp_when_rewrite_fails(tmp_path):
exif_utils=exif_utils, exif_utils=exif_utils,
) )
images = [ images: List[Dict[str, object]] = [
{ {
"url": "https://image.civitai.com/container/example/original=true/sample.png", "url": "https://image.civitai.com/container/example/original=true/sample.png",
"type": "image", "type": "image",
@@ -165,7 +165,7 @@ async def test_ensure_preview_rewrites_civitai_video(tmp_path):
exif_utils=RecordingExifUtils(), exif_utils=RecordingExifUtils(),
) )
images = [ images: List[Dict[str, object]] = [
{ {
"url": "https://image.civitai.com/container/example/original=true/sample.mp4", "url": "https://image.civitai.com/container/example/original=true/sample.mp4",
"type": "video", "type": "video",
@@ -227,7 +227,7 @@ async def test_ensure_preview_respects_blur_setting(monkeypatch, tmp_path):
exif_utils=RecordingExifUtils(), exif_utils=RecordingExifUtils(),
) )
images = [ images: List[Dict[str, object]] = [
{ {
"url": "https://image.civitai.com/container/example/original=true/nsfw.jpeg", "url": "https://image.civitai.com/container/example/original=true/nsfw.jpeg",
"type": "image", "type": "image",
+3 -1
View File
@@ -1,4 +1,6 @@
import json import json
from typing import Any, Dict
import pytest import pytest
from py.recipes.parsers.recipe_format import RecipeFormatParser 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, fake_metadata_provider,
) )
cached_entry = { cached_entry: Dict[str, Any] = {
"file_path": "/loras/moriimee.safetensors", "file_path": "/loras/moriimee.safetensors",
"file_name": "MoriiMee Gothic Niji | LoRA Style", "file_name": "MoriiMee Gothic Niji | LoRA Style",
"size": 4096, "size": 4096,
+2 -1
View File
@@ -1,5 +1,6 @@
import pytest import pytest
import asyncio import asyncio
from typing import Any, Dict
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
from py.services.recipe_scanner import RecipeScanner from py.services.recipe_scanner import RecipeScanner
from types import SimpleNamespace 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_scanner, mock_civitai_client, mock_metadata_provider = setup_scanner
# Recipe with runtime fields # Recipe with runtime fields
recipe = { recipe: Dict[str, Any] = {
"id": "r1", "id": "r1",
"title": "Cleanup Test", "title": "Cleanup Test",
"checkpoint": { "checkpoint": {
+8 -7
View File
@@ -3,6 +3,7 @@ import json
import os import os
from pathlib import Path from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any, Dict
import pytest import pytest
@@ -24,7 +25,7 @@ class StubLoraScanner:
def __init__(self) -> None: def __init__(self) -> None:
self._hash_index = StubHashIndex() self._hash_index = StubHashIndex()
self._hash_meta: dict[str, dict[str, str]] = {} 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={}) self._cache = SimpleNamespace(raw_data=[], version_index={})
async def get_cached_data(self): async def get_cached_data(self):
@@ -44,7 +45,7 @@ class StubLoraScanner:
async def get_model_info_by_name(self, name: str): async def get_model_info_by_name(self, name: str):
return self._models_by_name.get(name) 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 self._models_by_name[name] = info
hash_value = (info.get("sha256") or "").lower() hash_value = (info.get("sha256") or "").lower()
version_id = info.get("civitai", {}).get("id") version_id = info.get("civitai", {}).get("id")
@@ -76,7 +77,7 @@ def recipe_scanner(tmp_path: Path, monkeypatch):
settings_manager_module.reset_settings_manager() settings_manager_module.reset_settings_manager()
monkeypatch.setattr(config, "loras_roots", [str(tmp_path)]) monkeypatch.setattr(config, "loras_roots", [str(tmp_path)])
stub = StubLoraScanner() stub = StubLoraScanner()
scanner = RecipeScanner(lora_scanner=stub) scanner = RecipeScanner(lora_scanner=stub) # pyright: ignore[reportArgumentType]
async def _init(): async def _init():
await scanner.refresh_cache(force=True) 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 = settings_manager_module.get_settings_manager()
manager.set("recipes_path", str(custom_recipes)) manager.set("recipes_path", str(custom_recipes))
scanner = RecipeScanner(lora_scanner=StubLoraScanner()) scanner = RecipeScanner(lora_scanner=StubLoraScanner()) # pyright: ignore[reportArgumentType]
resolved = scanner.recipes_dir resolved = scanner.recipes_dir
assert resolved == str((tmp_path / "custom_recipes").resolve()) 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")]) 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 resolved = scanner.recipes_dir
assert resolved == str(tmp_path / "alpha" / "recipes") assert resolved == str(tmp_path / "alpha" / "recipes")
@@ -719,7 +720,7 @@ async def test_initialize_waits_for_lora_scanner(monkeypatch):
ready_flag.set() ready_flag.set()
lora_scanner = StubLoraScanner() lora_scanner = StubLoraScanner()
scanner = RecipeScanner(lora_scanner=lora_scanner) scanner = RecipeScanner(lora_scanner=lora_scanner) # pyright: ignore[reportArgumentType]
await scanner.initialize_in_background() 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 = Path(config.loras_roots[0]) / "recipes"
recipes_dir.mkdir(parents=True, exist_ok=True) recipes_dir.mkdir(parents=True, exist_ok=True)
recipe = { recipe: Dict[str, Any] = {
"id": "invalid-version", "id": "invalid-version",
"file_path": str(recipes_dir / "invalid-version.webp"), "file_path": str(recipes_dir / "invalid-version.webp"),
"title": "Invalid", "title": "Invalid",
+10 -4
View File
@@ -4,8 +4,9 @@ import os
from io import BytesIO from io import BytesIO
from pathlib import Path from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any, Dict
import piexif import piexif # pyright: ignore[reportMissingTypeStubs]
import pytest import pytest
from PIL import Image, PngImagePlugin 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"]) image_path = Path(result.payload["image_path"])
exif_dict = piexif.load(str(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 ( assert (
exif_dict["0th"][piexif.ImageIFD.ImageDescription].decode("utf-8") exif_0th[piexif.ImageIFD.ImageDescription].decode("utf-8")
== 'Workflow:{"nodes":[{"id":1}]}' == '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") decoded_comment = user_comment[8:].decode("utf-16be")
assert "prompt text" in decoded_comment assert "prompt text" in decoded_comment
assert "Recipe metadata:" 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")) matches = list(Path(self.recipes_dir).rglob(f"{target_id}.recipe.json"))
return str(matches[0]) if matches else None 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: if target_id != recipe_id:
return False return False
self.recipe.update(metadata) 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.model_query import ModelFilterSet, FilterCriteria
from py.services.recipe_scanner import RecipeScanner from py.services.recipe_scanner import RecipeScanner
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any, cast
# Mock settings # Mock settings
@@ -193,9 +194,9 @@ async def test_recipe_scanner_root_recursive_true():
async def get_cached_data(self): async def get_cached_data(self):
return SimpleNamespace(raw_data=[]) 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 # Manually populate cache for testing get_paginated_data logic
scanner._cache = SimpleNamespace( scanner._cache = cast(Any, SimpleNamespace(
raw_data=[ raw_data=[
{ {
"id": "r1", "id": "r1",
@@ -234,13 +235,15 @@ async def test_recipe_scanner_root_recursive_true():
], ],
sorted_by_name=[], sorted_by_name=[],
version_index={}, version_index={},
) ))
result = await scanner.get_paginated_data( result = await scanner.get_paginated_data(
page=1, page_size=10, folder="", recursive=True 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 @pytest.mark.asyncio
@@ -250,8 +253,8 @@ async def test_recipe_scanner_root_recursive_false():
async def get_cached_data(self): async def get_cached_data(self):
return SimpleNamespace(raw_data=[]) return SimpleNamespace(raw_data=[])
scanner = RecipeScanner(lora_scanner=StubLoraScanner()) scanner = RecipeScanner(lora_scanner=StubLoraScanner()) # pyright: ignore[reportArgumentType]
scanner._cache = SimpleNamespace( scanner._cache = cast(Any, SimpleNamespace(
raw_data=[ raw_data=[
{ {
"id": "r1", "id": "r1",
@@ -290,11 +293,13 @@ async def test_recipe_scanner_root_recursive_false():
], ],
sorted_by_name=[], sorted_by_name=[],
version_index={}, version_index={},
) ))
result = await scanner.get_paginated_data( result = await scanner.get_paginated_data(
page=1, page_size=10, folder="", recursive=False page=1, page_size=10, folder="", recursive=False
) )
assert len(result["items"]) == 1 items = result["items"]
assert result["items"][0]["id"] == "r1" 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: def __init__(self, payload: Dict[str, Any]) -> None:
self.payload = payload 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 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 return self.payload
@@ -77,7 +77,7 @@ def test_metadata_sync_merges_remote_fields(tmp_path: Path) -> None:
service = MetadataSyncService( service = MetadataSyncService(
metadata_manager=manager, metadata_manager=manager,
preview_service=preview, preview_service=preview,
settings=DummySettings(), settings=DummySettings(), # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=lambda: asyncio.sleep(0, result=provider), default_metadata_provider_factory=lambda: asyncio.sleep(0, result=provider),
metadata_provider_selector=lambda _name=None: 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( service = MetadataSyncService(
metadata_manager=manager, metadata_manager=manager,
preview_service=preview, preview_service=preview,
settings=DummySettings(), settings=DummySettings(), # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=lambda: asyncio.sleep(0, result=provider), default_metadata_provider_factory=lambda: asyncio.sleep(0, result=provider),
metadata_provider_selector=lambda _name=None: 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): async def test_lazy_loaded_websocket_manager(monkeypatch):
fake_manager = object() fake_manager = object()
module = types.ModuleType("py.services.websocket_manager") 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) monkeypatch.setitem(sys.modules, "py.services.websocket_manager", module)
first = await ServiceRegistry.get_websocket_manager() 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) manager = _create_manager_with_settings(tmp_path, monkeypatch, initial)
loop = asyncio.new_event_loop() loop = asyncio.new_event_loop()
loop._thread_id = 1 setattr(loop, "_thread_id", 1)
class DummyScanner: class DummyScanner:
def __init__(self): 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 dummy_scanner.calls == ["file_name"]
assert dispatched_loops == [dummy_scanner.loop] assert dispatched_loops == [dummy_scanner.loop]
finally: finally:
loop._thread_id = None setattr(loop, "_thread_id", None)
loop.close() loop.close()
@@ -8,6 +8,8 @@ from py.recipes.parsers import SuiImageParamsParser
class TestSuiImageParamsParser: class TestSuiImageParamsParser:
"""Test cases for SuiImageParamsParser.""" """Test cases for SuiImageParamsParser."""
parser: SuiImageParamsParser = SuiImageParamsParser()
def setup_method(self): def setup_method(self):
"""Set up test fixtures.""" """Set up test fixtures."""
self.parser = SuiImageParamsParser() self.parser = SuiImageParamsParser()
@@ -116,6 +118,7 @@ class TestSuiImageParamsParser:
result = await self.parser.parse_metadata(metadata_str) result = await self.parser.parse_metadata(metadata_str)
loras = result.get('loras') loras = result.get('loras')
assert isinstance(loras, list)
assert len(loras) == 1 assert len(loras) == 1
assert loras[0]['type'] == 'lora' assert loras[0]['type'] == 'lora'
assert loras[0]['name'] == 'test_lora' assert loras[0]['name'] == 'test_lora'
@@ -142,6 +145,7 @@ class TestSuiImageParamsParser:
result = await self.parser.parse_metadata(metadata_str) result = await self.parser.parse_metadata(metadata_str)
loras = result.get('loras') loras = result.get('loras')
assert isinstance(loras, list)
assert len(loras) == 1 assert len(loras) == 1
assert loras[0]['type'] == 'lora' assert loras[0]['type'] == 'lora'
+25 -10
View File
@@ -1,11 +1,14 @@
import asyncio import asyncio
import logging import logging
from dataclasses import dataclass from dataclasses import dataclass
from types import SimpleNamespace
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
import pytest 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 ( from py.services.use_cases import (
AutoOrganizeInProgressError, AutoOrganizeInProgressError,
AutoOrganizeUseCase, AutoOrganizeUseCase,
@@ -26,6 +29,7 @@ from py.utils.example_images_download_manager import (
) )
from py.utils.example_images_processor import ( from py.utils.example_images_processor import (
ExampleImagesImportError, ExampleImagesImportError,
ExampleImagesProcessor,
ExampleImagesValidationError, ExampleImagesValidationError,
) )
from py.utils.metadata_manager import MetadataManager from py.utils.metadata_manager import MetadataManager
@@ -44,13 +48,13 @@ class StubLockProvider:
return self._lock return self._lock
class StubFileService: class StubFileService(ModelFileService):
def __init__(self) -> None: def __init__(self) -> None:
super().__init__(scanner=None, model_type="lora")
self.calls: List[Dict[str, Any]] = [] self.calls: List[Dict[str, Any]] = []
async def auto_organize_models( async def auto_organize_models(
self, self,
*,
file_paths: Optional[List[str]] = None, file_paths: Optional[List[str]] = None,
progress_callback=None, progress_callback=None,
exclusion_patterns=None, exclusion_patterns=None,
@@ -65,8 +69,15 @@ class StubFileService:
return result return result
class StubMetadataSync: class StubMetadataSync(MetadataSyncService):
def __init__(self) -> None: 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]] = [] self.calls: List[Dict[str, Any]] = []
async def fetch_and_update_model(self, **kwargs: Any): async def fetch_and_update_model(self, **kwargs: Any):
@@ -94,8 +105,12 @@ class ProgressCollector:
self.events.append(payload) self.events.append(payload)
class StubDownloadCoordinator: class StubDownloadCoordinator(DownloadCoordinator):
def __init__(self, *, error: Optional[str] = None) -> None: 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.error = error
self.payloads: List[Dict[str, Any]] = [] self.payloads: List[Dict[str, Any]] = []
@@ -125,13 +140,13 @@ class StubExampleImagesDownloadManager:
return {"success": True, "message": "ok"} return {"success": True, "message": "ok"}
class StubExampleImagesProcessor: class StubExampleImagesProcessor(ExampleImagesProcessor):
def __init__(self) -> None: def __init__(self) -> None:
self.calls: List[Dict[str, Any]] = [] self.calls: List[Dict[str, Any]] = []
self.error: Optional[str] = None self.error: Optional[str] = None
self.response: Dict[str, Any] = {"success": True} 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}) self.calls.append({"model_hash": model_hash, "files": files})
if self.error == "validation": if self.error == "validation":
raise ExampleImagesValidationError("missing") raise ExampleImagesValidationError("missing")
@@ -464,7 +479,7 @@ async def test_import_example_images_use_case_delegates() -> None:
use_case = ImportExampleImagesUseCase(processor=processor) use_case = ImportExampleImagesUseCase(processor=processor)
request = DummyJsonRequest({"model_hash": "abc", "file_paths": ["/tmp/file"]}) 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 processor.calls == [{"model_hash": "abc", "files": ["/tmp/file"]}]
assert result == {"success": True} 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": []}) request = DummyJsonRequest({"model_hash": None, "file_paths": []})
with pytest.raises(ImportExampleImagesValidationError): 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: 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"]}) request = DummyJsonRequest({"model_hash": "abc", "file_paths": ["/tmp/file"]})
with pytest.raises(ExampleImagesImportError): with pytest.raises(ExampleImagesImportError):
await use_case.execute(request) await use_case.execute(request) # pyright: ignore[reportArgumentType]
+7 -7
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import json import json
from pathlib import Path from pathlib import Path
from types import ModuleType, SimpleNamespace from types import ModuleType, SimpleNamespace
from typing import List, Tuple from typing import Any, Generator, List, Tuple
import pytest import pytest
from aiohttp import web from aiohttp import web
@@ -13,11 +13,11 @@ from aiohttp import web
from py.utils.settings_paths import ensure_settings_file from py.utils.settings_paths import ensure_settings_file
ROUTE_CALLS_KEY: web.AppKey[List[Tuple[str, dict]]] = web.AppKey("route_calls") ROUTE_CALLS_KEY: web.AppKey[List[Tuple[str, dict[str, Any]]]] = web.AppKey("route_calls")
@pytest.fixture @pytest.fixture
def standalone_module(monkeypatch) -> ModuleType: def standalone_module(monkeypatch) -> Generator[ModuleType, None, None]:
"""Load the ``standalone`` module with a lightweight ``LoraManager`` stub.""" """Load the ``standalone`` module with a lightweight ``LoraManager`` stub."""
import importlib import importlib
@@ -41,7 +41,7 @@ def standalone_module(monkeypatch) -> ModuleType:
async def _cleanup(cls, app): # pragma: no cover - compatibility shim async def _cleanup(cls, app): # pragma: no cover - compatibility shim
return None return None
stub_module.LoraManager = _StubLoraManager stub_module.LoraManager = _StubLoraManager # pyright: ignore[reportAttributeAccessIssue]
sys.modules["py.lora_manager"] = stub_module sys.modules["py.lora_manager"] = stub_module
module = importlib.import_module("standalone") module = importlib.import_module("standalone")
@@ -54,7 +54,7 @@ def standalone_module(monkeypatch) -> ModuleType:
sys.modules.pop("py.lora_manager", None) sys.modules.pop("py.lora_manager", None)
def _write_settings(contents: dict) -> Path: def _write_settings(contents: dict[str, Any]) -> Path:
"""Persist *contents* into the isolated settings.json.""" """Persist *contents* into the isolated settings.json."""
settings_path = Path(ensure_settings_file()) settings_path = Path(ensure_settings_file())
@@ -118,7 +118,7 @@ def test_standalone_lora_manager_registers_routes(monkeypatch, tmp_path, standal
"""``StandaloneLoraManager.add_routes`` registers static and websocket routes.""" """``StandaloneLoraManager.add_routes`` registers static and websocket routes."""
app = web.Application() app = web.Application()
route_calls: List[Tuple[str, dict]] = [] route_calls: List[Tuple[str, dict[str, Any]]] = []
app[ROUTE_CALLS_KEY] = route_calls app[ROUTE_CALLS_KEY] = route_calls
locales_dir = tmp_path / "locales" locales_dir = tmp_path / "locales"
@@ -213,7 +213,7 @@ def test_standalone_lora_manager_registers_routes(monkeypatch, tmp_path, standal
assert "/locales" in canonical_routes assert "/locales" in canonical_routes
assert "/loras_static" in canonical_routes assert "/loras_static" in canonical_routes
websocket_paths = {route.resource.canonical for route in app.router.routes() if "ws" in route.resource.canonical} websocket_paths = {route.resource.canonical for route in app.router.routes() if "ws" in route.resource.canonical} # pyright: ignore[reportOptionalMemberAccess]
assert { assert {
"/ws/fetch-progress", "/ws/fetch-progress",
"/ws/download-progress", "/ws/download-progress",
+4 -4
View File
@@ -4,7 +4,7 @@ import os
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "py")) sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "py"))
from services.auto_tag_service import extract_auto_tags, AUTO_TAG_CATEGORIES from services.auto_tag_service import extract_auto_tags, AUTO_TAG_CATEGORIES # pyright: ignore[reportMissingImports]
class TestExtractAutoTags: class TestExtractAutoTags:
@@ -208,18 +208,18 @@ class TestAutoTagCategories:
re.compile(pattern, re.IGNORECASE) re.compile(pattern, re.IGNORECASE)
def test_mode_group_tags(self): def test_mode_group_tags(self):
from services.auto_tag_service import MODE_TAGS from services.auto_tag_service import MODE_TAGS # pyright: ignore[reportMissingImports]
assert "HIGH" in MODE_TAGS assert "HIGH" in MODE_TAGS
assert "LOW" in MODE_TAGS assert "LOW" in MODE_TAGS
def test_video_group_tags(self): def test_video_group_tags(self):
from services.auto_tag_service import VIDEO_MODE_TAGS from services.auto_tag_service import VIDEO_MODE_TAGS # pyright: ignore[reportMissingImports]
assert "I2V" in VIDEO_MODE_TAGS assert "I2V" in VIDEO_MODE_TAGS
assert "T2V" in VIDEO_MODE_TAGS assert "T2V" in VIDEO_MODE_TAGS
assert "TI2V" in VIDEO_MODE_TAGS assert "TI2V" in VIDEO_MODE_TAGS
def test_default_enabled_groups(self): def test_default_enabled_groups(self):
from services.auto_tag_service import DEFAULT_ENABLED_GROUPS from services.auto_tag_service import DEFAULT_ENABLED_GROUPS # pyright: ignore[reportMissingImports]
assert "mode" in DEFAULT_ENABLED_GROUPS assert "mode" in DEFAULT_ENABLED_GROUPS
assert "video" in DEFAULT_ENABLED_GROUPS assert "video" in DEFAULT_ENABLED_GROUPS
assert "speed" not in DEFAULT_ENABLED_GROUPS assert "speed" not in DEFAULT_ENABLED_GROUPS
+10 -2
View File
@@ -3,7 +3,7 @@
import json import json
import os import os
import tempfile import tempfile
from typing import Dict, List from typing import Any, Dict, List
import pytest import pytest
@@ -27,7 +27,7 @@ def temp_db_path():
@pytest.fixture @pytest.fixture
def sample_recipes() -> List[Dict]: def sample_recipes() -> List[Dict[str, Any]]:
"""Create sample recipe data.""" """Create sample recipe data."""
return [ return [
{ {
@@ -133,6 +133,7 @@ class TestPersistentRecipeCache:
# Load and verify # Load and verify
loaded = cache.load_cache() loaded = cache.load_cache()
assert loaded is not None
r1 = next(r for r in loaded.raw_data if r["id"] == "recipe-001") r1 = next(r for r in loaded.raw_data if r["id"] == "recipe-001")
assert r1["title"] == "Updated Title" assert r1["title"] == "Updated Title"
assert r1["favorite"] is False assert r1["favorite"] is False
@@ -147,6 +148,7 @@ class TestPersistentRecipeCache:
# Load and verify # Load and verify
loaded = cache.load_cache() loaded = cache.load_cache()
assert loaded is not None
assert len(loaded.raw_data) == 1 assert len(loaded.raw_data) == 1
assert loaded.raw_data[0]["id"] == "recipe-002" assert loaded.raw_data[0]["id"] == "recipe-002"
@@ -203,6 +205,7 @@ class TestPersistentRecipeCache:
cache.save_cache(recipes) cache.save_cache(recipes)
loaded = cache.load_cache() loaded = cache.load_cache()
assert loaded is not None
assert len(loaded.raw_data) == 1 assert len(loaded.raw_data) == 1
assert loaded.raw_data[0]["id"] == "valid-001" assert loaded.raw_data[0]["id"] == "valid-001"
@@ -251,6 +254,7 @@ class TestPersistentRecipeCache:
cache.save_cache(recipes) cache.save_cache(recipes)
loaded = cache.load_cache() loaded = cache.load_cache()
assert loaded is not None
loras = loaded.raw_data[0]["loras"] loras = loaded.raw_data[0]["loras"]
assert len(loras) == 2 assert len(loras) == 2
assert loras[0]["modelVersionId"] == 12345 assert loras[0]["modelVersionId"] == 12345
@@ -509,6 +513,7 @@ class TestPersistentRecipeCache:
cache.save_cache(sample_recipes) cache.save_cache(sample_recipes)
loaded = cache.load_cache() loaded = cache.load_cache()
assert loaded is not None
assert loaded.image_id_map == {} assert loaded.image_id_map == {}
def test_image_id_map_survives_recipe_update(self, temp_db_path, sample_recipes): def test_image_id_map_survives_recipe_update(self, temp_db_path, sample_recipes):
@@ -522,6 +527,7 @@ class TestPersistentRecipeCache:
cache.update_recipe(updated) cache.update_recipe(updated)
loaded = cache.load_cache() loaded = cache.load_cache()
assert loaded is not None
assert loaded.image_id_map == {"123": "recipe-alpha"} assert loaded.image_id_map == {"123": "recipe-alpha"}
def test_save_image_id_map_persists_without_full_save(self, temp_db_path, sample_recipes): def test_save_image_id_map_persists_without_full_save(self, temp_db_path, sample_recipes):
@@ -532,6 +538,7 @@ class TestPersistentRecipeCache:
cache.save_image_id_map({"555": "new-recipe", "666": "another-recipe"}) cache.save_image_id_map({"555": "new-recipe", "666": "another-recipe"})
loaded = cache.load_cache() loaded = cache.load_cache()
assert loaded is not None
assert loaded.image_id_map == {"555": "new-recipe", "666": "another-recipe"} assert loaded.image_id_map == {"555": "new-recipe", "666": "another-recipe"}
def test_save_image_id_map_overwrites_previous(self, temp_db_path, sample_recipes): def test_save_image_id_map_overwrites_previous(self, temp_db_path, sample_recipes):
@@ -542,4 +549,5 @@ class TestPersistentRecipeCache:
cache.save_image_id_map({"222": "new-only"}) cache.save_image_id_map({"222": "new-only"})
loaded = cache.load_cache() loaded = cache.load_cache()
assert loaded is not None
assert loaded.image_id_map == {"222": "new-only"} assert loaded.image_id_map == {"222": "new-only"}
+2 -2
View File
@@ -2,7 +2,7 @@
import os import os
import tempfile import tempfile
from typing import Dict, List from typing import Any, Dict, List
import pytest import pytest
@@ -25,7 +25,7 @@ def temp_db_path():
@pytest.fixture @pytest.fixture
def sample_recipes() -> List[Dict]: def sample_recipes() -> List[Dict[str, Any]]:
"""Create sample recipe data for FTS indexing.""" """Create sample recipe data for FTS indexing."""
return [ return [
{ {
+3 -2
View File
@@ -1,6 +1,7 @@
import importlib import importlib
import json import json
from pathlib import Path from pathlib import Path
from typing import Any
import pytest import pytest
@@ -21,12 +22,12 @@ def reset_settings(tmp_path, monkeypatch):
reset_settings_manager() reset_settings_manager()
def read_settings_file(path: Path) -> dict: def read_settings_file(path: Path) -> dict[str, Any]:
with path.open('r', encoding='utf-8') as handle: with path.open('r', encoding='utf-8') as handle:
return json.load(handle) return json.load(handle)
def read_example_settings() -> dict: def read_example_settings() -> dict[str, Any]:
example_path = Path(__file__).resolve().parents[1] / "settings.json.example" example_path = Path(__file__).resolve().parents[1] / "settings.json.example"
with example_path.open('r', encoding='utf-8') as handle: with example_path.open('r', encoding='utf-8') as handle:
return json.load(handle) return json.load(handle)
@@ -76,6 +76,7 @@ class TestRewritePreviewUrl:
for url in test_cases: for url in test_cases:
result, was_rewritten = rewrite_preview_url(url, "image") result, was_rewritten = rewrite_preview_url(url, "image")
assert was_rewritten is True assert was_rewritten is True
assert result is not None
assert "width=450,optimized=true" in result assert "width=450,optimized=true" in result
def test_handles_urls_with_explicit_port(self): def test_handles_urls_with_explicit_port(self):
@@ -83,6 +84,7 @@ class TestRewritePreviewUrl:
url = "https://image.civitai.com:443/checkpoints/original=true" url = "https://image.civitai.com:443/checkpoints/original=true"
result, was_rewritten = rewrite_preview_url(url, "image") result, was_rewritten = rewrite_preview_url(url, "image")
assert was_rewritten is True assert was_rewritten is True
assert result is not None
assert "width=450,optimized=true" in result assert "width=450,optimized=true" in result
# Port is preserved in the URL (this is acceptable behavior) # Port is preserved in the URL (this is acceptable behavior)
assert ":443" in result assert ":443" in result
@@ -104,6 +106,8 @@ class TestRewritePreviewUrl:
result2, was2 = rewrite_preview_url(url, "Video") result2, was2 = rewrite_preview_url(url, "Video")
assert was1 is True assert was1 is True
assert was2 is True assert was2 is True
assert result1 is not None
assert result2 is not None
assert "transcode=true" in result1 assert "transcode=true" in result1
assert "transcode=true" in result2 assert "transcode=true" in result2
@@ -119,6 +123,7 @@ class TestRewritePreviewUrl:
url = "https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/abc123/original=true/12345.png" url = "https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/abc123/original=true/12345.png"
result, was_rewritten = rewrite_preview_url(url, "image") result, was_rewritten = rewrite_preview_url(url, "image")
assert was_rewritten is True assert was_rewritten is True
assert result is not None
assert result.startswith( assert result.startswith(
"https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/abc123/" "https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/abc123/"
) )
@@ -129,6 +134,7 @@ class TestRewritePreviewUrl:
url = "https://image.civitai.com/original=true/test.png" url = "https://image.civitai.com/original=true/test.png"
result, was_rewritten = rewrite_preview_url(url, None) result, was_rewritten = rewrite_preview_url(url, None)
assert was_rewritten is True assert was_rewritten is True
assert result is not None
assert "transcode=true" not in result assert "transcode=true" not in result
assert "width=450,optimized=true" in result assert "width=450,optimized=true" in result
@@ -2,7 +2,7 @@ from __future__ import annotations
import asyncio import asyncio
import time import time
from typing import Any, Dict from typing import Any, Dict, Generator
import pytest import pytest
@@ -19,7 +19,7 @@ class RecordingWebSocketManager:
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
def restore_settings() -> None: def restore_settings() -> Generator[None, None, None]:
manager = get_settings_manager() manager = get_settings_manager()
original = manager.settings.copy() original = manager.settings.copy()
try: try:
@@ -45,7 +45,9 @@ async def test_start_download_requires_configured_path(
result = await manager.start_download({"auto_mode": True}) result = await manager.start_download({"auto_mode": True})
assert result["success"] is True assert result["success"] is True
assert "skipping auto download" in result["message"] message = result["message"]
assert isinstance(message, str)
assert "skipping auto download" in message
async def test_start_download_bootstraps_progress_and_task( async def test_start_download_bootstraps_progress_and_task(
+16 -11
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
import json import json
import os import os
import subprocess import subprocess
from typing import Any, Dict from typing import Any, Dict, Generator
import pytest import pytest
@@ -20,8 +20,13 @@ class JsonRequest:
return self._payload return self._payload
def _parse_json(response) -> Dict[str, Any]:
assert response.text is not None
return json.loads(response.text)
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
def restore_settings() -> None: def restore_settings() -> Generator[None, None, None]:
manager = get_settings_manager() manager = get_settings_manager()
original = manager.settings.copy() original = manager.settings.copy()
try: try:
@@ -54,7 +59,7 @@ async def test_open_folder_requires_existing_model_directory(monkeypatch: pytest
request = JsonRequest({"model_hash": model_hash}) request = JsonRequest({"model_hash": model_hash})
response = await ExampleImagesFileManager.open_folder(request) response = await ExampleImagesFileManager.open_folder(request)
body = json.loads(response.text) body = _parse_json(response)
assert body["success"] is True assert body["success"] is True
# On Windows, os.startfile is used; on other platforms, subprocess.Popen # On Windows, os.startfile is used; on other platforms, subprocess.Popen
@@ -89,7 +94,7 @@ async def test_open_folder_returns_clipboard_mode_with_mapped_local_path(
request = JsonRequest({"model_hash": model_hash}) request = JsonRequest({"model_hash": model_hash})
response = await ExampleImagesFileManager.open_folder(request) response = await ExampleImagesFileManager.open_folder(request)
body = json.loads(response.text) body = _parse_json(response)
assert response.status == 200 assert response.status == 200
assert body == { assert body == {
@@ -126,7 +131,7 @@ async def test_open_folder_returns_uri_mode_with_rendered_template(
request = JsonRequest({"model_hash": model_hash}) request = JsonRequest({"model_hash": model_hash})
response = await ExampleImagesFileManager.open_folder(request) response = await ExampleImagesFileManager.open_folder(request)
body = json.loads(response.text) body = _parse_json(response)
assert response.status == 200 assert response.status == 200
assert body["success"] is True assert body["success"] is True
@@ -150,7 +155,7 @@ async def test_open_folder_rejects_missing_uri_template(monkeypatch: pytest.Monk
(model_folder / "image.png").write_text("data", encoding="utf-8") (model_folder / "image.png").write_text("data", encoding="utf-8")
response = await ExampleImagesFileManager.open_folder(JsonRequest({"model_hash": model_hash})) response = await ExampleImagesFileManager.open_folder(JsonRequest({"model_hash": model_hash}))
body = json.loads(response.text) body = _parse_json(response)
assert response.status == 400 assert response.status == 400
assert body["success"] is False assert body["success"] is False
@@ -168,7 +173,7 @@ async def test_open_folder_rejects_invalid_paths(monkeypatch: pytest.MonkeyPatch
request = JsonRequest({"model_hash": "a" * 64}) request = JsonRequest({"model_hash": "a" * 64})
response = await ExampleImagesFileManager.open_folder(request) response = await ExampleImagesFileManager.open_folder(request)
body = json.loads(response.text) body = _parse_json(response)
assert response.status == 400 assert response.status == 400
assert body["success"] is False assert body["success"] is False
@@ -186,7 +191,7 @@ async def test_get_files_lists_supported_media(tmp_path) -> None:
request = JsonRequest({}, {"model_hash": model_hash}) request = JsonRequest({}, {"model_hash": model_hash})
response = await ExampleImagesFileManager.get_files(request) response = await ExampleImagesFileManager.get_files(request)
body = json.loads(response.text) body = _parse_json(response)
assert response.status == 200 assert response.status == 200
names = {entry["name"] for entry in body["files"]} names = {entry["name"] for entry in body["files"]}
@@ -203,18 +208,18 @@ async def test_has_images_reports_presence(tmp_path) -> None:
request = JsonRequest({}, {"model_hash": model_hash}) request = JsonRequest({}, {"model_hash": model_hash})
response = await ExampleImagesFileManager.has_images(request) response = await ExampleImagesFileManager.has_images(request)
body = json.loads(response.text) body = _parse_json(response)
assert body["has_images"] is True assert body["has_images"] is True
empty_request = JsonRequest({}, {"model_hash": "missing"}) empty_request = JsonRequest({}, {"model_hash": "missing"})
empty_response = await ExampleImagesFileManager.has_images(empty_request) empty_response = await ExampleImagesFileManager.has_images(empty_request)
empty_body = json.loads(empty_response.text) empty_body = _parse_json(empty_response)
assert empty_body["has_images"] is False assert empty_body["has_images"] is False
async def test_has_images_requires_model_hash() -> None: async def test_has_images_requires_model_hash() -> None:
response = await ExampleImagesFileManager.has_images(JsonRequest({}, {})) response = await ExampleImagesFileManager.has_images(JsonRequest({}, {}))
body = json.loads(response.text) body = _parse_json(response)
assert response.status == 400 assert response.status == 400
assert body["success"] is False assert body["success"] is False
+2 -2
View File
@@ -102,7 +102,7 @@ async def test_update_metadata_after_import_preserves_existing_metadata(
model_file.write_text("content", encoding="utf-8") model_file.write_text("content", encoding="utf-8")
metadata_path = tmp_path / "preserve.metadata.json" metadata_path = tmp_path / "preserve.metadata.json"
existing_payload = { existing_payload: Dict[str, Any] = {
"model_name": "Example", "model_name": "Example",
"file_path": str(model_file), "file_path": str(model_file),
"civitai": { "civitai": {
@@ -200,7 +200,7 @@ async def test_update_metadata_from_local_examples_generates_entries(monkeypatch
model_dir = tmp_path / model_hash model_dir = tmp_path / model_hash
model_dir.mkdir() model_dir.mkdir()
(model_dir / "image.png").write_text("data", encoding="utf-8") (model_dir / "image.png").write_text("data", encoding="utf-8")
model_data = {"model_name": "Local", "civitai": {}, "file_path": str(tmp_path / "model.safetensors")} model_data: Dict[str, Any] = {"model_name": "Local", "civitai": {}, "file_path": str(tmp_path / "model.safetensors")}
async def fake_save(path, metadata): async def fake_save(path, metadata):
return True return True
@@ -4,7 +4,7 @@ import json
import os import os
from pathlib import Path from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any, Dict, Tuple from typing import Any, Dict, Generator, Tuple
import pytest import pytest
@@ -15,7 +15,7 @@ from py.utils.example_images_paths import get_model_folder
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
def restore_settings() -> None: def restore_settings() -> Generator[None, None, None]:
manager = get_settings_manager() manager = get_settings_manager()
original = manager.settings.copy() original = manager.settings.copy()
try: try:
@@ -206,7 +206,7 @@ async def test_import_images_creates_hash_directory(monkeypatch: pytest.MonkeyPa
monkeypatch.setattr(processor_module.MetadataUpdater, "update_metadata_after_import", staticmethod(fake_update_metadata)) monkeypatch.setattr(processor_module.MetadataUpdater, "update_metadata_after_import", staticmethod(fake_update_metadata))
result = await processor_module.ExampleImagesProcessor.import_images("a" * 64, [str(source_file)]) result: Dict[str, Any] = await processor_module.ExampleImagesProcessor.import_images("a" * 64, [str(source_file)])
assert result["success"] is True assert result["success"] is True
assert result["files"][0]["name"].startswith("custom_short") assert result["files"][0]["name"].startswith("custom_short")
@@ -255,7 +255,7 @@ async def test_delete_custom_image_preserves_existing_metadata(monkeypatch: pyte
model_file.write_text("content", encoding="utf-8") model_file.write_text("content", encoding="utf-8")
metadata_path = tmp_path / "keep.metadata.json" metadata_path = tmp_path / "keep.metadata.json"
existing_metadata = { existing_metadata: Dict[str, Any] = {
"model_name": "Keep", "model_name": "Keep",
"file_path": str(model_file), "file_path": str(model_file),
"civitai": { "civitai": {
@@ -313,6 +313,7 @@ async def test_delete_custom_image_preserves_existing_metadata(monkeypatch: pyte
) )
assert response.status == 200 assert response.status == 200
assert response.text is not None
body = json.loads(response.text) body = json.loads(response.text)
assert body["success"] is True assert body["success"] is True
assert body["custom_images"] == [] assert body["custom_images"] == []
+11 -4
View File
@@ -1,6 +1,7 @@
import json import json
from typing import Any, Dict
import piexif import piexif # pyright: ignore[reportMissingTypeStubs]
from PIL import Image, PngImagePlugin from PIL import Image, PngImagePlugin
from py.utils.exif_utils import ExifUtils from py.utils.exif_utils import ExifUtils
@@ -84,10 +85,12 @@ def test_optimize_image_preserves_workflow_when_converting_png_to_webp(tmp_path)
optimized_path.write_bytes(optimized_data) optimized_path.write_bytes(optimized_data)
exif_dict = piexif.load(str(optimized_path)) exif_dict = piexif.load(str(optimized_path))
assert exif_dict["0th"] is not None
assert ( assert (
exif_dict["0th"][piexif.ImageIFD.ImageDescription].decode("utf-8") exif_dict["0th"][piexif.ImageIFD.ImageDescription].decode("utf-8")
== 'Workflow:{"nodes": [{"id": 1}]}' == 'Workflow:{"nodes": [{"id": 1}]}'
) )
assert exif_dict["Exif"] is not None
user_comment = exif_dict["Exif"][piexif.ExifIFD.UserComment] user_comment = exif_dict["Exif"][piexif.ExifIFD.UserComment]
assert user_comment.startswith(b"UNICODE\0") assert user_comment.startswith(b"UNICODE\0")
assert user_comment[8:].decode("utf-16be") == "prompt text\nSteps: 20" assert user_comment[8:].decode("utf-16be") == "prompt text\nSteps: 20"
@@ -113,10 +116,12 @@ def test_update_image_metadata_preserves_webp_workflow(tmp_path):
) )
updated_exif = piexif.load(str(image_path)) updated_exif = piexif.load(str(image_path))
assert updated_exif["0th"] is not None
assert ( assert (
updated_exif["0th"][piexif.ImageIFD.ImageDescription].decode("utf-8") updated_exif["0th"][piexif.ImageIFD.ImageDescription].decode("utf-8")
== 'Workflow:{"nodes":[{"id":1}]}' == 'Workflow:{"nodes":[{"id":1}]}'
) )
assert updated_exif["Exif"] is not None
updated_comment = updated_exif["Exif"][piexif.ExifIFD.UserComment] updated_comment = updated_exif["Exif"][piexif.ExifIFD.UserComment]
assert ( assert (
updated_comment[8:].decode("utf-16be") updated_comment[8:].decode("utf-16be")
@@ -147,10 +152,10 @@ def test_update_image_metadata_preserves_png_workflow(tmp_path):
import struct import struct
import brotli import brotli # pyright: ignore[reportMissingTypeStubs]
def _build_jxl_with_brob(payload_json: dict) -> bytes: def _build_jxl_with_brob(payload_json: Dict[str, Any]) -> bytes:
"""Build a minimal JXL container with a brob box containing brotli-compressed JSON.""" """Build a minimal JXL container with a brob box containing brotli-compressed JSON."""
# ISOBMFF box 1: JXL signature box (size=12, type='JXL ', signature) # ISOBMFF box 1: JXL signature box (size=12, type='JXL ', signature)
box1 = struct.pack(">I", 12) + b"JXL " + bytes([0x0d, 0x0a, 0x87, 0x0a]) box1 = struct.pack(">I", 12) + b"JXL " + bytes([0x0d, 0x0a, 0x87, 0x0a])
@@ -163,7 +168,7 @@ def _build_jxl_with_brob(payload_json: dict) -> bytes:
return box1 + box2 + box3 return box1 + box2 + box3
def _build_avif_with_brob(payload_json: dict) -> bytes: def _build_avif_with_brob(payload_json: Dict[str, Any]) -> bytes:
"""Build a minimal AVIF container with a brob box containing brotli-compressed JSON.""" """Build a minimal AVIF container with a brob box containing brotli-compressed JSON."""
compressed = brotli.compress(json.dumps(payload_json).encode("utf-8")) compressed = brotli.compress(json.dumps(payload_json).encode("utf-8"))
brob_payload = b"comf" + compressed brob_payload = b"comf" + compressed
@@ -262,6 +267,7 @@ class TestIsobmffBrotliExtraction:
path.write_bytes(data) path.write_bytes(data)
result = ExifUtils._load_structured_metadata(str(path)) result = ExifUtils._load_structured_metadata(str(path))
assert result["prompt"] is not None
assert json.loads(result["prompt"]) == {"text": "hello", "negative": "bad"} assert json.loads(result["prompt"]) == {"text": "hello", "negative": "bad"}
def test_extract_workflow_as_list(self, tmp_path): def test_extract_workflow_as_list(self, tmp_path):
@@ -272,6 +278,7 @@ class TestIsobmffBrotliExtraction:
path.write_bytes(data) path.write_bytes(data)
result = ExifUtils._load_structured_metadata(str(path)) result = ExifUtils._load_structured_metadata(str(path))
assert result["workflow"] is not None
assert json.loads(result["workflow"]) == [{"id": 1}, {"id": 2}] assert json.loads(result["workflow"]) == [{"id": 1}, {"id": 2}]
def test_over_decompressed_size_limit(self, tmp_path, monkeypatch): def test_over_decompressed_size_limit(self, tmp_path, monkeypatch):
+5 -3
View File
@@ -1,5 +1,7 @@
"""Tests for model sub_type field refactoring.""" """Tests for model sub_type field refactoring."""
from typing import Any, Dict
import pytest import pytest
from py.utils.models import ( from py.utils.models import (
BaseModelMetadata, BaseModelMetadata,
@@ -44,7 +46,7 @@ class TestCheckpointMetadataSubType:
def test_checkpoint_from_civitai_info_uses_sub_type(self): def test_checkpoint_from_civitai_info_uses_sub_type(self):
"""from_civitai_info should use sub_type from version_info.""" """from_civitai_info should use sub_type from version_info."""
version_info = { version_info: Dict[str, Any] = {
"baseModel": "SDXL", "baseModel": "SDXL",
"model": {"name": "Test", "description": "", "tags": []}, "model": {"name": "Test", "description": "", "tags": []},
"files": [{"name": "model.safetensors", "sizeKB": 1000, "hashes": {"SHA256": "abc123"}, "primary": True}], "files": [{"name": "model.safetensors", "sizeKB": 1000, "hashes": {"SHA256": "abc123"}, "primary": True}],
@@ -79,7 +81,7 @@ class TestEmbeddingMetadataSubType:
def test_embedding_from_civitai_info_uses_sub_type(self): def test_embedding_from_civitai_info_uses_sub_type(self):
"""from_civitai_info should use sub_type from version_info.""" """from_civitai_info should use sub_type from version_info."""
version_info = { version_info: Dict[str, Any] = {
"baseModel": "SD1.5", "baseModel": "SD1.5",
"model": {"name": "Test", "description": "", "tags": []}, "model": {"name": "Test", "description": "", "tags": []},
"files": [{"name": "model.pt", "sizeKB": 1000, "hashes": {"SHA256": "abc123"}, "primary": True}], "files": [{"name": "model.pt", "sizeKB": 1000, "hashes": {"SHA256": "abc123"}, "primary": True}],
@@ -113,7 +115,7 @@ class TestLoraMetadataConsistency:
def test_lora_from_civitai_info_extracts_type(self): def test_lora_from_civitai_info_extracts_type(self):
"""from_civitai_info should extract type from civitai data.""" """from_civitai_info should extract type from civitai data."""
version_info = { version_info: Dict[str, Any] = {
"baseModel": "SDXL", "baseModel": "SDXL",
"model": {"name": "Test", "description": "", "tags": [], "type": "Lora"}, "model": {"name": "Test", "description": "", "tags": [], "type": "Lora"},
"files": [{"name": "model.safetensors", "sizeKB": 1000, "hashes": {"SHA256": "abc123"}, "primary": True}], "files": [{"name": "model.safetensors", "sizeKB": 1000, "hashes": {"SHA256": "abc123"}, "primary": True}],
+2
View File
@@ -12,6 +12,7 @@ def test_select_preview_returns_first_when_blur_disabled():
selected, level = select_preview_media(images, blur_mature_content=False) selected, level = select_preview_media(images, blur_mature_content=False)
assert selected is not None
assert selected["url"] == "nsfw" assert selected["url"] == "nsfw"
assert level == 32 assert level == 32
@@ -40,6 +41,7 @@ def test_select_preview_respects_configurable_threshold(threshold_name, expected
mature_threshold=NSFW_LEVELS[threshold_name], mature_threshold=NSFW_LEVELS[threshold_name],
) )
assert selected is not None
assert selected["url"] == expected_url assert selected["url"] == expected_url
assert level == next(item["nsfwLevel"] for item in images if item["url"] == expected_url) assert level == next(item["nsfwLevel"] for item in images if item["url"] == expected_url)
+7 -3
View File
@@ -6,6 +6,8 @@ property-based testing to catch edge cases and ensure correctness.
from __future__ import annotations from __future__ import annotations
from typing import Any, Dict
import pytest import pytest
from hypothesis import given, settings, strategies as st from hypothesis import given, settings, strategies as st
@@ -80,6 +82,8 @@ class TestNormalizePath:
@given(st.text(alphabet='abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789._-/\\') | st.none()) @given(st.text(alphabet='abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789._-/\\') | st.none())
def test_normalize_path_is_idempotent_for_ascii(self, path: str | None): def test_normalize_path_is_idempotent_for_ascii(self, path: str | None):
"""Normalizing an already normalized ASCII path should not change it.""" """Normalizing an already normalized ASCII path should not change it."""
if path is None:
return
normalized = normalize_path(path) normalized = normalize_path(path)
renormalized = normalize_path(normalized) renormalized = normalize_path(normalized)
assert normalized == renormalized assert normalized == renormalized
@@ -160,14 +164,14 @@ class TestCalculateRecipeFingerprint:
"""Property-based tests for calculate_recipe_fingerprint function.""" """Property-based tests for calculate_recipe_fingerprint function."""
@given(st.lists(st.dictionaries(st.text(), st.text() | st.integers() | st.floats(), min_size=1), min_size=0, max_size=50)) @given(st.lists(st.dictionaries(st.text(), st.text() | st.integers() | st.floats(), min_size=1), min_size=0, max_size=50))
def test_fingerprint_is_deterministic(self, loras: list): def test_fingerprint_is_deterministic(self, loras: list[Dict[str, Any]]):
"""Same input should always produce same fingerprint.""" """Same input should always produce same fingerprint."""
fp1 = calculate_recipe_fingerprint(loras) fp1 = calculate_recipe_fingerprint(loras)
fp2 = calculate_recipe_fingerprint(loras) fp2 = calculate_recipe_fingerprint(loras)
assert fp1 == fp2 assert fp1 == fp2
@given(st.lists(st.dictionaries(st.text(), st.text() | st.integers() | st.floats(), min_size=1), min_size=0, max_size=50)) @given(st.lists(st.dictionaries(st.text(), st.text() | st.integers() | st.floats(), min_size=1), min_size=0, max_size=50))
def test_fingerprint_returns_string(self, loras: list): def test_fingerprint_returns_string(self, loras: list[Dict[str, Any]]):
"""Function should always return a string.""" """Function should always return a string."""
result = calculate_recipe_fingerprint(loras) result = calculate_recipe_fingerprint(loras)
assert isinstance(result, str) assert isinstance(result, str)
@@ -178,7 +182,7 @@ class TestCalculateRecipeFingerprint:
assert result == "" assert result == ""
@given(st.lists(st.dictionaries(st.text(), st.text() | st.integers() | st.floats(), min_size=1), min_size=1, max_size=10)) @given(st.lists(st.dictionaries(st.text(), st.text() | st.integers() | st.floats(), min_size=1), min_size=1, max_size=10))
def test_fingerprint_different_inputs_produce_different_results(self, loras1: list): def test_fingerprint_different_inputs_produce_different_results(self, loras1: list[Dict[str, Any]]):
"""Different inputs should generally produce different fingerprints.""" """Different inputs should generally produce different fingerprints."""
# Create a different input by modifying the first LoRA # Create a different input by modifying the first LoRA
loras2 = loras1.copy() loras2 = loras1.copy()