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
from typing import Dict, Iterable, List
from typing import Any, Dict, Iterable, List
import pytest
@@ -146,6 +146,8 @@ def test_save_paths_repairs_empty_default_roots(monkeypatch: pytest.MonkeyPatch,
class FakeSettingsService:
active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self):
return {
@@ -183,6 +185,8 @@ def test_save_paths_repairs_stale_default_roots(monkeypatch: pytest.MonkeyPatch,
class FakeSettingsService:
active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self):
return {
@@ -220,6 +224,8 @@ def test_save_paths_keeps_valid_default_roots(monkeypatch: pytest.MonkeyPatch, t
class FakeSettingsService:
active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self):
return {
@@ -357,6 +363,8 @@ def test_save_paths_keeps_default_roots_in_extra_paths(monkeypatch: pytest.Monke
class FakeSettingsService:
active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self):
return {
@@ -409,6 +417,8 @@ def test_save_paths_keeps_default_roots_in_extra_paths_with_windows_slash_mismat
class FakeSettingsService:
active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self):
return {
@@ -460,6 +470,8 @@ def test_save_paths_repairs_empty_default_roots_to_extra_paths_when_primary_miss
class FakeSettingsService:
active_library = "comfyui"
name: str = ""
payload: Dict[str, Any] = {}
def get_libraries(self):
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(checkpoints_dir) in config_instance.base_models_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
@@ -609,7 +621,7 @@ def test_apply_library_settings_without_extra_paths(monkeypatch, tmp_path):
assert config_instance.extra_loras_roots == []
assert str(checkpoints_dir) in config_instance.base_models_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 == []
@@ -858,7 +870,7 @@ def test_save_paths_removes_stale_empty_default_when_comfyui_exists(
# dict order, returning "default".
self.active_library = "default"
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):
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
server_mock = MockModule("server")
server_mock.PromptServer = mock.MagicMock()
setattr(server_mock, "PromptServer", mock.MagicMock())
sys.modules['server'] = server_mock
folder_paths_mock = MockModule("folder_paths")
folder_paths_mock.get_folder_paths = mock.MagicMock(return_value=[])
folder_paths_mock.folder_names_and_paths = {}
setattr(folder_paths_mock, "get_folder_paths", mock.MagicMock(return_value=[]))
setattr(folder_paths_mock, "folder_names_and_paths", {})
sys.modules['folder_paths'] = folder_paths_mock
# Mock other ComfyUI modules that might be imported
comfy_mock = MockModule("comfy")
comfy_mock.utils = MockModule("comfy.utils")
comfy_mock.utils.load_torch_file = mock.MagicMock(return_value={})
comfy_mock.sd = MockModule("comfy.sd")
comfy_mock.sd.load_lora_for_models = mock.MagicMock(return_value=(None, None))
comfy_mock.model_management = MockModule("comfy.model_management")
comfy_mock.comfy_types = MockModule("comfy.comfy_types")
comfy_mock.comfy_types.IO = mock.MagicMock()
setattr(comfy_mock, "utils", MockModule("comfy.utils"))
setattr(comfy_mock.utils, "load_torch_file", mock.MagicMock(return_value={}))
setattr(comfy_mock, "sd", MockModule("comfy.sd"))
setattr(comfy_mock.sd, "load_lora_for_models", mock.MagicMock(return_value=(None, None)))
setattr(comfy_mock, "model_management", MockModule("comfy.model_management"))
setattr(comfy_mock, "comfy_types", MockModule("comfy.comfy_types"))
setattr(comfy_mock.comfy_types, "IO", mock.MagicMock())
sys.modules['comfy'] = comfy_mock
sys.modules['comfy.utils'] = comfy_mock.utils
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")
execution_mock = MockModule("execution")
execution_mock.PromptExecutor = mock.MagicMock()
setattr(execution_mock, "PromptExecutor", mock.MagicMock())
sys.modules['execution'] = execution_mock
# Mock ComfyUI nodes module
nodes_mock = MockModule("nodes")
nodes_mock.LoraLoader = mock.MagicMock()
nodes_mock.SaveImage = mock.MagicMock()
nodes_mock.NODE_CLASS_MAPPINGS = {}
setattr(nodes_mock, "LoraLoader", mock.MagicMock())
setattr(nodes_mock, "SaveImage", mock.MagicMock())
setattr(nodes_mock, "NODE_CLASS_MAPPINGS", {})
sys.modules['nodes'] = nodes_mock
@@ -347,7 +347,7 @@ def reset_singletons():
# Reset ServiceRegistry
ServiceRegistry._services = {}
ServiceRegistry._initialized = False
ServiceRegistry._initialized = False # pyright: ignore[reportAttributeAccessIssue]
# Reset ModelScanner instances
if hasattr(ModelScanner, '_instances'):
@@ -356,14 +356,14 @@ def reset_singletons():
# Reset SettingsManager
settings_manager = get_settings_manager()
if hasattr(settings_manager, '_reset'):
settings_manager._reset()
settings_manager._reset() # pyright: ignore[reportAttributeAccessIssue]
yield
# Cleanup after test
DownloadManager._instance = None
ServiceRegistry._services = {}
ServiceRegistry._initialized = False
ServiceRegistry._initialized = False # pyright: ignore[reportAttributeAccessIssue]
if hasattr(ModelScanner, '_instances'):
ModelScanner._instances.clear()
+7 -7
View File
@@ -12,7 +12,7 @@ from __future__ import annotations
import json
import re
from pathlib import Path
from typing import Dict, Iterable, Set
from typing import Any, Dict, Iterable, Set
import pytest
@@ -105,9 +105,9 @@ HTML_TRANSLATION_PATTERN = (
@pytest.fixture(scope="module")
def loaded_locales() -> Dict[str, dict]:
def loaded_locales() -> Dict[str, Any]:
"""Load locale JSON once per test module."""
locales: Dict[str, dict] = {}
locales: Dict[str, Any] = {}
for locale in EXPECTED_LOCALES:
path = LOCALES_DIR / f"{locale}.json"
@@ -131,7 +131,7 @@ def loaded_locales() -> Dict[str, dict]:
@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"])
@@ -140,7 +140,7 @@ def static_code_translation_keys() -> Set[str]:
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."""
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)
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."""
data = loaded_locales[locale]
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:])
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:
"""Locales must expose the same translation keys as English."""
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."""
class MockMetadataManager:
def __init__(self):
self.saved_metadata: List[tuple] = []
self.saved_metadata: List[tuple[str, Any]] = []
self.loaded_payloads: Dict[str, Dict[str, Any]] = {}
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)
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}"
yield base_url, port
+1 -1
View File
@@ -193,7 +193,7 @@ class TestDownloadRouteIntegration:
assert response.status == 400
# 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()
else:
body = response.body
+2 -1
View File
@@ -75,6 +75,7 @@ class TestRecipeFlowIntegration:
# Verify update
loaded = cache.load_cache()
assert loaded is not None
loaded_recipe = loaded.raw_data[0]
assert loaded_recipe["title"] == "Updated Recipe Title"
@@ -243,7 +244,7 @@ steps: 20
cfg: 7.0"""
# Basic parsing logic for testing
def parse_simple_metadata(text: str) -> dict:
def parse_simple_metadata(text: str) -> Dict[str, str]:
result = {}
for line in text.strip().split('\n'):
if ':' in line:
+7 -7
View File
@@ -44,13 +44,13 @@ def populated_registry(metadata_registry):
# Direct assignment to avoid scanner.py false positive
# (scanner.py matches _CLASS_MAPPINGS.update({...}) pattern)
nodes.NODE_CLASS_MAPPINGS["TSC_EfficientLoader"] = TSC_EfficientLoader
nodes.NODE_CLASS_MAPPINGS["SamplerCustomAdvanced"] = SamplerCustomAdvanced
nodes.NODE_CLASS_MAPPINGS["BasicScheduler"] = BasicScheduler
nodes.NODE_CLASS_MAPPINGS["KSamplerSelect"] = KSamplerSelect
nodes.NODE_CLASS_MAPPINGS["CFGGuider"] = CFGGuider
nodes.NODE_CLASS_MAPPINGS["CLIPTextEncode"] = CLIPTextEncode
nodes.NODE_CLASS_MAPPINGS["VAEDecode"] = VAEDecode
nodes.NODE_CLASS_MAPPINGS["TSC_EfficientLoader"] = TSC_EfficientLoader # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["SamplerCustomAdvanced"] = SamplerCustomAdvanced # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["BasicScheduler"] = BasicScheduler # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["KSamplerSelect"] = KSamplerSelect # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["CFGGuider"] = CFGGuider # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["CLIPTextEncode"] = CLIPTextEncode # pyright: ignore[reportAttributeAccessIssue]
nodes.NODE_CLASS_MAPPINGS["VAEDecode"] = VAEDecode # pyright: ignore[reportAttributeAccessIssue]
prompt_graph = {
"loader": {"class_type": "TSC_EfficientLoader", "inputs": {}},
@@ -1,6 +1,7 @@
import sys
import types
from types import SimpleNamespace
from typing import Any, Dict
from py.metadata_collector import metadata_processor
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:
FUNCTION = "run"
unique_id: str = ""
node = FakeNode()
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]
__name__ = "LoraLoaderLM"
nodes.NODE_CLASS_MAPPINGS["LoraLoaderLM"] = LoraLoaderLM
nodes.NODE_CLASS_MAPPINGS["LoraLoaderLM"] = LoraLoaderLM # pyright: ignore[reportAttributeAccessIssue]
prompt_graph = {
"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
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})
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 typing import Any
from unittest import mock
import pytest
@@ -28,14 +29,14 @@ from py.metadata_ops import (
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 []
class MockScanner:
"""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.update_single_model_cache = mock.AsyncMock(return_value=True)
@@ -599,7 +600,7 @@ class TestExtractGalleryTableImages:
"""
@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 \
extract_gallery_table_images
return extract_gallery_table_images(md, repo, existing_urls=existing)
@@ -643,7 +644,7 @@ class TestExtractGalleryTableImages:
class TestCleanReadmeForLlm:
@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 \
clean_readme_for_llm
return clean_readme_for_llm(md, max_length=max_length)
@@ -651,7 +652,7 @@ class TestCleanReadmeForLlm:
# -- basic guards --------------------------------------------------------
def test_none_returns_empty(self):
assert self._clean(None) == "" # type: ignore[arg-type]
assert self._clean(None) == ""
def test_empty_returns_empty(self):
assert self._clean("") == ""
@@ -19,6 +19,8 @@ _MODULE_PATH = Path(__file__).parents[2] / "py" / "services" / "agent" / "skills
def R():
"""Load the ``readme_processor`` module once per session."""
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)
spec.loader.exec_module(mod)
return mod
+1 -1
View File
@@ -32,7 +32,7 @@ def _parse_directives(header: str) -> dict[str, list[str]]:
async def _invoke_middleware(
path: str, response: web.Response, csp_header: str | None = DEFAULT_CSP
) -> web.Response:
) -> web.StreamResponse:
async def handler(_request: web.Request) -> web.Response:
if csp_header is not None:
response.headers["Content-Security-Policy"] = csp_header
+2 -2
View File
@@ -31,7 +31,7 @@ class _DummyModel:
return self
def test_nunchaku_load_lora_legacy_fallback(monkeypatch, caplog):
import folder_paths
import folder_paths # pyright: ignore[reportMissingImports]
import copy
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
def test_nunchaku_load_lora_new_logic(monkeypatch):
import folder_paths
import folder_paths # pyright: ignore[reportMissingImports]
import os
dummy_model = _DummyModel()
+4 -2
View File
@@ -1,5 +1,7 @@
from __future__ import annotations
from typing import Any, cast
from py.nodes.prompt import PromptLM
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_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():
@@ -58,7 +60,7 @@ def test_text_lm_input_types_expose_input_only_seed():
assert seed_type == "INT"
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():
+24 -12
View File
@@ -1,8 +1,9 @@
import json
import os
from typing import Any, cast
import numpy as np
import piexif
import piexif # pyright: ignore[reportMissingTypeStubs]
from PIL import Image
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"
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):
@@ -154,7 +156,8 @@ def test_save_image_skips_webp_metadata_when_disabled(monkeypatch, tmp_path):
image_path = tmp_path / "sample_00001_.webp"
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):
@@ -474,25 +477,34 @@ class TestParameterDefaultConsistency:
input_types = SaveImageLM.INPUT_TYPES()
optional = input_types["optional"]
assert optional["webp_method"][1]["default"] == 6
assert SaveImageLM.save_images.__defaults__[4] == 6 # positional: webp_method=6 is at index 4
assert SaveImageLM.process_image.__defaults__[6] == 6
widget_spec = cast(Any, optional["webp_method"])
assert widget_spec[1]["default"] == 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):
input_types = SaveImageLM.INPUT_TYPES()
optional = input_types["optional"]
assert optional["jpeg_subsampling"][1]["default"] == 0
assert SaveImageLM.save_images.__defaults__[5] == 0
assert SaveImageLM.process_image.__defaults__[7] == 0
widget_spec = cast(Any, optional["jpeg_subsampling"])
assert widget_spec[1]["default"] == 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):
input_types = SaveImageLM.INPUT_TYPES()
optional = input_types["optional"]
assert optional["add_loras_to_prompt"][1]["default"] is False
assert SaveImageLM.save_images.__defaults__[-1] is False
assert SaveImageLM.process_image.__defaults__[-1] is False
widget_spec = cast(Any, optional["add_loras_to_prompt"])
assert widget_spec[1]["default"] 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):
+1 -1
View File
@@ -34,7 +34,7 @@ class _DummyModel:
def test_nunchaku_load_lora_skips_missing_lora(monkeypatch, caplog):
import folder_paths
import folder_paths # pyright: ignore[reportMissingImports]
dummy_model = _DummyModel()
+20 -10
View File
@@ -8,6 +8,8 @@ from __future__ import annotations
import random
import string
from typing import Any, Dict, cast
import pytest
from py.services.model_hash_index import ModelHashIndex
@@ -22,9 +24,11 @@ class TestHashIndexPerformance:
def test_hash_index_lookup_small(self, benchmark):
"""Benchmark hash index lookup with 100 models."""
index, target_hash = self._create_hash_index_with_n_models(
100, return_target=True
index, target_hash = cast(
tuple[ModelHashIndex, str | None],
self._create_hash_index_with_n_models(100, return_target=True),
)
assert target_hash is not None
def lookup():
return index.get_path(target_hash)
@@ -34,9 +38,11 @@ class TestHashIndexPerformance:
def test_hash_index_lookup_medium(self, benchmark):
"""Benchmark hash index lookup with 1,000 models."""
index, target_hash = self._create_hash_index_with_n_models(
1000, return_target=True
index, target_hash = cast(
tuple[ModelHashIndex, str | None],
self._create_hash_index_with_n_models(1000, return_target=True),
)
assert target_hash is not None
def lookup():
return index.get_path(target_hash)
@@ -46,9 +52,11 @@ class TestHashIndexPerformance:
def test_hash_index_lookup_large(self, benchmark):
"""Benchmark hash index lookup with 10,000 models."""
index, target_hash = self._create_hash_index_with_n_models(
10000, return_target=True
index, target_hash = cast(
tuple[ModelHashIndex, str | None],
self._create_hash_index_with_n_models(10000, return_target=True),
)
assert target_hash is not None
def lookup():
return index.get_path(target_hash)
@@ -58,7 +66,7 @@ class TestHashIndexPerformance:
def test_hash_index_add_entry_small(self, benchmark):
"""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_path = "/path/to/new_model.safetensors"
@@ -69,7 +77,7 @@ class TestHashIndexPerformance:
def test_hash_index_add_entry_large(self, benchmark):
"""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_path = "/path/to/new_model.safetensors"
@@ -78,7 +86,9 @@ class TestHashIndexPerformance:
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.
Args:
@@ -170,7 +180,7 @@ class TestRecipeFingerprintPerformance:
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."""
loras = []
for i in range(n):
+43 -17
View File
@@ -8,8 +8,10 @@ response schemas.
from __future__ import annotations
import json
import pytest
from types import SimpleNamespace
from typing import Any
import pytest
from syrupy import SnapshotAssertion
from py.routes.handlers.misc_handlers import (
@@ -54,13 +56,35 @@ async def noop_async(*_args, **_kwargs):
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:
"""Fake prompt server for testing."""
sent = []
class Instance:
sockets: dict = {}
sockets: dict[str, Any] = {}
def send_sync(self, event, payload, sid=None):
FakePromptServer.sent.append((event, payload))
@@ -103,11 +127,11 @@ class TestSettingsHandlerSnapshots:
handler = SettingsHandler(
settings_service=settings_service,
metadata_provider_updater=noop_async,
downloader_factory=lambda: None,
downloader_factory=fake_downloader_factory,
)
response = await handler.get_settings(FakeRequest())
payload = json.loads(response.text)
response = await handler.get_settings(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = json_payload(response)
assert payload == snapshot
@@ -118,12 +142,12 @@ class TestSettingsHandlerSnapshots:
handler = SettingsHandler(
settings_service=settings_service,
metadata_provider_updater=noop_async,
downloader_factory=lambda: None,
downloader_factory=fake_downloader_factory,
)
request = FakeRequest(json_data={"language": "zh"})
response = await handler.update_settings(request)
payload = json.loads(response.text)
response = await handler.update_settings(request) # pyright: ignore[reportArgumentType]
payload = json_payload(response)
assert payload == snapshot
@@ -137,7 +161,7 @@ class TestNodeRegistryHandlerSnapshots:
node_registry = NodeRegistry()
handler = NodeRegistryHandler(
node_registry=node_registry,
prompt_server=FakePromptServer,
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False,
)
@@ -155,8 +179,8 @@ class TestNodeRegistryHandlerSnapshots:
}
)
response = await handler.register_nodes(request)
payload = json.loads(response.text)
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = json_payload(response)
assert payload == snapshot
@@ -166,13 +190,13 @@ class TestNodeRegistryHandlerSnapshots:
node_registry = NodeRegistry()
handler = NodeRegistryHandler(
node_registry=node_registry,
prompt_server=FakePromptServer,
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False,
)
request = FakeRequest(json_data={"nodes": [], "client_id": "test-client-1"})
response = await handler.register_nodes(request)
payload = json.loads(response.text)
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = json_payload(response)
assert payload == snapshot
@@ -249,10 +273,12 @@ class TestModelLibraryHandlerSnapshots:
get_embedding_scanner=scanner_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"}))
payload = json.loads(response.text)
response = await handler.check_model_exists(
FakeRequest(query={"modelId": "1"}) # pyright: ignore[reportArgumentType]
)
payload = json_payload(response)
assert payload == snapshot
+12 -10
View File
@@ -6,10 +6,10 @@ from pathlib import Path
import types
from dataclasses import dataclass, field
from typing import Optional
from typing import Any, Optional
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
from aiohttp import FormData, web
@@ -38,7 +38,9 @@ class DummyRoutes(BaseModelRoutes):
def __init__(self, service=None):
super().__init__(service)
self.set_model_update_service(NullModelUpdateService())
self.set_model_update_service(
NullModelUpdateService() # pyright: ignore[reportArgumentType]
)
@dataclass
@@ -110,7 +112,7 @@ class NullModelUpdateService:
return None
async def create_test_client(service) -> TestClient:
async def create_test_client(service) -> TestClient[Any, Any]:
routes = DummyRoutes(service)
app = web.Application()
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]
class FakeMetadata:
def __init__(self, payload: dict) -> None:
def __init__(self, payload: dict[str, Any]) -> None:
self._payload = payload
self._unknown_fields = {"legacy_field": "legacy"}
def to_dict(self) -> dict:
def to_dict(self) -> dict[str, Any]:
return self._payload.copy()
async def fake_load_metadata(path: str, *_args, **_kwargs):
assert path == str(model_path)
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))))
return True
@@ -477,7 +479,7 @@ def test_fetch_civitai_hydrates_metadata_before_sync(
*,
sha256: str,
file_path: str,
model_data: dict,
model_data: dict[str, Any],
update_cache_func,
):
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)
return True, None
save_calls: list[tuple[str, dict]] = []
captured: dict[str, dict] = {}
save_calls: list[tuple[str, dict[str, Any]]] = []
captured: dict[str, dict[str, Any]] = {}
monkeypatch.setattr(
MetadataManager, "load_metadata", staticmethod(fake_load_metadata)
+1 -1
View File
@@ -24,7 +24,7 @@ class StubEmbeddingService:
@pytest.fixture
def routes():
handler = EmbeddingRoutes()
handler.service = StubEmbeddingService()
handler.service = StubEmbeddingService() # pyright: ignore[reportAttributeAccessIssue]
return handler
@@ -3,9 +3,9 @@ from __future__ import annotations
import json
from contextlib import asynccontextmanager
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 py.routes.example_images_route_registrar import ExampleImagesRouteRegistrar
@@ -140,14 +140,14 @@ class StubFileManager:
@dataclass
class RegistrarHarness:
client: TestClient
client: TestClient[Any, Any]
download_use_case: StubDownloadUseCase
download_manager: StubDownloadManager
import_use_case: StubImportUseCase
@asynccontextmanager
async def registrar_app() -> RegistrarHarness:
async def registrar_app() -> AsyncGenerator[RegistrarHarness, None]:
app = web.Application()
download_use_case = StubDownloadUseCase()
@@ -158,8 +158,15 @@ async def registrar_app() -> RegistrarHarness:
file_manager = StubFileManager()
handler_set = ExampleImagesHandlerSet(
download=ExampleImagesDownloadHandler(download_use_case, download_manager),
management=ExampleImagesManagementHandler(import_use_case, processor, cleanup_service),
download=ExampleImagesDownloadHandler(
download_use_case, # pyright: ignore[reportArgumentType]
download_manager,
),
management=ExampleImagesManagementHandler(
import_use_case, # pyright: ignore[reportArgumentType]
processor,
cleanup_service,
),
files=ExampleImagesFileHandler(file_manager),
)
@@ -181,7 +188,7 @@ async def registrar_app() -> RegistrarHarness:
await client.close()
async def _json(response: web.StreamResponse) -> Dict[str, Any]:
async def _json(response: ClientResponse) -> Dict[str, Any]:
text = await response.text()
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
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")
harness.download_manager.check_pending_models = failing_check
+81 -46
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
import json
from contextlib import asynccontextmanager
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.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
class ExampleImagesHarness:
"""Container exposing the aiohttp client and stubbed collaborators."""
client: TestClient
client: TestClient[Any, Any]
download_manager: "StubDownloadManager"
processor: "StubExampleImagesProcessor"
file_manager: "StubExampleImagesFileManager"
@@ -35,27 +40,27 @@ class StubDownloadManager:
def __init__(self) -> None:
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))
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)))
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))
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))
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))
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))
return {"operation": "start_force_download", "payload": payload}
@@ -64,7 +69,7 @@ class StubExampleImagesProcessor:
def __init__(self) -> None:
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}
self.calls.append(("import_images", payload))
return {"operation": "import_images", "payload": payload}
@@ -122,7 +127,7 @@ class StubWebSocketManager:
@asynccontextmanager
async def example_images_app() -> ExampleImagesHarness:
async def example_images_app() -> AsyncGenerator[ExampleImagesHarness, None]:
"""Yield an ExampleImagesRoutes app wired with stubbed collaborators."""
download_manager = StubDownloadManager()
@@ -133,10 +138,10 @@ async def example_images_app() -> ExampleImagesHarness:
controller = ExampleImagesRoutes(
ws_manager=ws_manager,
download_manager=download_manager,
download_manager=download_manager, # pyright: ignore[reportArgumentType]
processor=processor,
file_manager=file_manager,
cleanup_service=cleanup_service,
file_manager=file_manager, # pyright: ignore[reportArgumentType]
cleanup_service=cleanup_service, # pyright: ignore[reportArgumentType]
)
app = web.Application()
@@ -323,23 +328,23 @@ async def test_download_handler_methods_delegate() -> None:
def __init__(self) -> None:
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))
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))
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))
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))
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))
return {"status": "force", "payload": payload}
@@ -347,35 +352,50 @@ async def test_download_handler_methods_delegate() -> None:
def __init__(self) -> None:
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)
return {"status": "started", "payload": payload}
class DummyRequest:
def __init__(self, payload: dict) -> None:
def __init__(self, payload: Dict[str, Any]) -> None:
self._payload = payload
self.query = {}
async def json(self) -> dict:
async def json(self) -> Dict[str, Any]:
return self._payload
recorder = Recorder()
use_case = StubDownloadUseCase()
handler = ExampleImagesDownloadHandler(use_case, recorder)
handler = ExampleImagesDownloadHandler(
use_case, # pyright: ignore[reportArgumentType]
recorder,
)
request = DummyRequest({"foo": "bar"})
download_response = await handler.download_example_images(request)
assert json.loads(download_response.text) == {"status": "started", "payload": {"foo": "bar"}}
status_response = await handler.get_example_images_status(request)
assert json.loads(status_response.text) == {"status": "ok"}
pause_response = await handler.pause_example_images(request)
assert json.loads(pause_response.text) == {"status": "paused"}
resume_response = await handler.resume_example_images(request)
assert json.loads(resume_response.text) == {"status": "running"}
stop_response = await handler.stop_example_images(request)
assert json.loads(stop_response.text) == {"status": "stopping"}
force_response = await handler.force_download_example_images(request)
assert json.loads(force_response.text) == {"status": "force", "payload": {"foo": "bar"}}
download_response = await handler.download_example_images(
request # pyright: ignore[reportArgumentType]
)
assert _json_response(download_response) == {"status": "started", "payload": {"foo": "bar"}}
status_response = await handler.get_example_images_status(
request # pyright: ignore[reportArgumentType]
)
assert _json_response(status_response) == {"status": "ok"}
pause_response = await handler.pause_example_images(
request # pyright: ignore[reportArgumentType]
)
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 recorder.calls == [
@@ -393,7 +413,7 @@ async def test_management_handler_methods_delegate() -> None:
def __init__(self) -> None:
self.requests: List[Any] = []
async def execute(self, request: Any) -> dict:
async def execute(self, request: Any) -> Dict[str, Any]:
self.requests.append(request)
return {"status": "imported"}
@@ -412,15 +432,23 @@ async def test_management_handler_methods_delegate() -> None:
recorder = Recorder()
cleanup_service = StubExampleImagesCleanupService()
use_case = StubImportUseCase()
handler = ExampleImagesManagementHandler(use_case, recorder, cleanup_service)
handler = ExampleImagesManagementHandler(
use_case, # pyright: ignore[reportArgumentType]
recorder,
cleanup_service,
)
request = object()
import_response = await handler.import_example_images(request)
assert json.loads(import_response.text) == {"status": "imported"}
assert await handler.delete_example_image(request) == "delete"
import_response = await handler.import_example_images(
request # pyright: ignore[reportArgumentType]
)
assert _json_response(import_response) == {"status": "imported"}
assert await handler.delete_example_image(request) == "delete" # pyright: ignore[reportArgumentType]
cleanup_service.result = {"success": True}
cleanup_response = await handler.cleanup_example_image_folders(request)
assert json.loads(cleanup_response.text) == {"success": True}
cleanup_response = await handler.cleanup_example_image_folders(
request # pyright: ignore[reportArgumentType]
)
assert _json_response(cleanup_response) == {"success": True}
assert use_case.requests == [request]
assert recorder.calls == [("delete_custom_image", request)]
assert len(cleanup_service.calls) == 1
@@ -448,9 +476,9 @@ async def test_file_handler_methods_delegate() -> None:
handler = ExampleImagesFileHandler(recorder)
request = object()
assert await handler.open_example_images_folder(request) == "open"
assert await handler.get_example_image_files(request) == "files"
assert await handler.has_example_images(request) == "has"
assert await handler.open_example_images_folder(request) == "open" # pyright: ignore[reportArgumentType]
assert await handler.get_example_image_files(request) == "files" # pyright: ignore[reportArgumentType]
assert await handler.has_example_images(request) == "has" # pyright: ignore[reportArgumentType]
assert recorder.calls == [
("open_folder", 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):
return {}
download = ExampleImagesDownloadHandler(DummyUseCase(), DummyManager())
download = ExampleImagesDownloadHandler(
DummyUseCase(), # pyright: ignore[reportArgumentType]
DummyManager(),
)
cleanup_service = StubExampleImagesCleanupService()
management = ExampleImagesManagementHandler(DummyUseCase(), DummyProcessor(), cleanup_service)
management = ExampleImagesManagementHandler(
DummyUseCase(), # pyright: ignore[reportArgumentType]
DummyProcessor(),
cleanup_service,
)
files = ExampleImagesFileHandler(object())
handler_set = ExampleImagesHandlerSet(
download=download,
+2 -1
View File
@@ -4,6 +4,7 @@ import asyncio
import logging
from pathlib import Path
from types import SimpleNamespace
from typing import Any
import pytest
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
scheduled_tasks: list[asyncio.Task] = []
scheduled_tasks: list[asyncio.Task[Any]] = []
def track_create_task(coro, *, name=None):
task = original_create_task(coro, name=name)
+2 -100
View File
@@ -5,7 +5,7 @@ from unittest.mock import MagicMock
import pytest
from py.routes.lora_routes import LoraRoutes
from server import PromptServer
from server import PromptServer # pyright: ignore[reportMissingImports]
class DummyRequest:
@@ -20,14 +20,8 @@ class DummyRequest:
class StubLoraService:
def __init__(self):
self.notes = {}
self.trigger_words = {}
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):
return self.trigger_words.get(name, [])
@@ -35,57 +29,14 @@ class StubLoraService:
async def get_lora_usage_tips_by_relative_path(self, 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
def routes():
handler = LoraRoutes()
handler.service = StubLoraService()
handler.service = StubLoraService() # pyright: ignore[reportAttributeAccessIssue]
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):
routes.service.trigger_words["demo"] = ["trigger"]
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
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):
send_mock = MagicMock()
PromptServer.instance = SimpleNamespace(send_sync=send_mock)
+114 -102
View File
@@ -5,6 +5,7 @@ import os
import subprocess
import zipfile
from types import SimpleNamespace
from typing import Any
from unittest.mock import patch, MagicMock
import pytest
@@ -34,6 +35,13 @@ from py.routes.misc_route_registrar import MISC_ROUTE_DEFINITIONS, MiscRouteRegi
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:
def __init__(self, *, json_data=None, query=None, method="POST"):
self._json_data = json_data or {}
@@ -129,8 +137,8 @@ async def test_get_settings_excludes_no_sync_keys():
downloader_factory=dummy_downloader_factory,
)
response = await handler.get_settings(FakeRequest())
payload = json.loads(response.text)
response = await handler.get_settings(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert payload["success"] is True
# 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"
request = FakeRequest(json_data={"example_images_path": str(missing_path)})
response = await handler.update_settings(request)
payload = json.loads(response.text)
response = await handler.update_settings(request) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert payload["success"] is False
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(
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["summary"]["status"] == "error"
@@ -209,8 +217,8 @@ async def test_doctor_handler_can_repair_cache():
scanner_factories=(("lora", "LoRAs", scanner_factory),),
)
response = await handler.repair_doctor_cache(FakeRequest())
payload = json.loads(response.text)
response = await handler.repair_doctor_cache(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert response.status == 200
assert payload["success"] is True
@@ -230,7 +238,7 @@ async def test_doctor_handler_exports_support_bundle():
)
response = await handler.export_doctor_bundle(
FakeRequest(
FakeRequest( # pyright: ignore[reportArgumentType]
json_data={
"summary": {"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 isinstance(response.body, bytes)
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
names = set(archive.namelist())
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(
FakeRequest(
FakeRequest( # pyright: ignore[reportArgumentType]
json_data={
"frontend_logs": [
{
@@ -276,6 +285,7 @@ async def test_doctor_handler_redacts_string_secrets_in_bundle():
)
assert response.status == 200
assert isinstance(response.body, bytes)
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
frontend_logs = archive.read("frontend-console.json").decode("utf-8")
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(
FakeRequest(
FakeRequest( # pyright: ignore[reportArgumentType]
json_data={
"frontend_logs": [
{
@@ -321,6 +331,7 @@ async def test_doctor_handler_redacts_json_shaped_string_secrets_in_bundle():
)
assert response.status == 200
assert isinstance(response.body, bytes)
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
frontend_logs = archive.read("frontend-console.json").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": [],
}
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 isinstance(response.body, bytes)
with zipfile.ZipFile(io.BytesIO(response.body), "r") as archive:
backend_logs = archive.read("backend-logs.txt").decode("utf-8")
backend_source = json.loads(
@@ -449,14 +461,14 @@ async def test_backup_handler_returns_status_and_exports(monkeypatch):
handler = BackupHandler(backup_service_factory=factory)
status_response = await handler.get_backup_status(FakeRequest())
status_payload = json.loads(status_response.text)
status_response = await handler.get_backup_status(FakeRequest()) # pyright: ignore[reportArgumentType]
status_payload = _json_payload(status_response)
assert status_payload["success"] is True
assert status_payload["status"]["backupDir"] == "/tmp/backups"
assert status_payload["status"]["enabled"] is True
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.body == b"zip-bytes"
@@ -476,8 +488,8 @@ async def test_backup_handler_rejects_missing_import_archive():
async def read(self):
return b""
response = await handler.import_backup(EmptyRequest())
payload = json.loads(response.text)
response = await handler.import_backup(EmptyRequest()) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert response.status == 400
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_wsl", lambda: False)
response = await handler.open_backup_location(FakeRequest())
payload = json.loads(response.text)
response = await handler.open_backup_location(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert response.status == 200
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),
)
response = await handler.open_wildcards_location(FakeRequest())
payload = json.loads(response.text)
response = await handler.open_wildcards_location(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert response.status == 200
assert payload["success"] is True
@@ -564,7 +576,7 @@ class RecordingRouter:
def test_misc_route_registrar_registers_all_routes():
app = SimpleNamespace(router=RecordingRouter())
registrar = MiscRouteRegistrar(app) # type: ignore[arg-type]
registrar = MiscRouteRegistrar(app) # pyright: ignore[reportArgumentType]
async def dummy_handler(_request):
return web.Response()
@@ -586,7 +598,7 @@ class FakePromptServer:
sent = []
class Instance:
sockets: dict = {}
sockets: dict[str, Any] = {}
def send_sync(self, event, payload, sid=None):
FakePromptServer.sent.append((event, payload))
@@ -599,7 +611,7 @@ async def test_register_nodes_requires_graph_id():
node_registry = NodeRegistry()
handler = NodeRegistryHandler(
node_registry=node_registry,
prompt_server=FakePromptServer,
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False,
)
@@ -609,8 +621,8 @@ async def test_register_nodes_requires_graph_id():
"client_id": "test-client-1",
}
)
response = await handler.register_nodes(request)
payload = json.loads(response.text)
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert response.status == 400
assert payload["success"] is False
@@ -622,7 +634,7 @@ async def test_register_nodes_stores_graph_identifier():
node_registry = NodeRegistry()
handler = NodeRegistryHandler(
node_registry=node_registry,
prompt_server=FakePromptServer,
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False,
)
@@ -641,8 +653,8 @@ async def test_register_nodes_stores_graph_identifier():
}
)
response = await handler.register_nodes(request)
payload = json.loads(response.text)
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert payload["success"] is True
@@ -659,7 +671,7 @@ async def test_register_nodes_defaults_graph_name_to_none():
node_registry = NodeRegistry()
handler = NodeRegistryHandler(
node_registry=node_registry,
prompt_server=FakePromptServer,
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False,
)
@@ -677,8 +689,8 @@ async def test_register_nodes_defaults_graph_name_to_none():
}
)
response = await handler.register_nodes(request)
payload = json.loads(response.text)
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert payload["success"] is True
@@ -692,7 +704,7 @@ async def test_register_nodes_includes_capabilities():
node_registry = NodeRegistry()
handler = NodeRegistryHandler(
node_registry=node_registry,
prompt_server=FakePromptServer,
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False,
)
@@ -714,8 +726,8 @@ async def test_register_nodes_includes_capabilities():
}
)
response = await handler.register_nodes(request)
payload = json.loads(response.text)
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert payload["success"] is True
@@ -734,7 +746,7 @@ async def test_register_nodes_accepts_compound_node_ids():
node_registry = NodeRegistry()
handler = NodeRegistryHandler(
node_registry=node_registry,
prompt_server=FakePromptServer,
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False,
)
@@ -758,8 +770,8 @@ async def test_register_nodes_accepts_compound_node_ids():
}
)
response = await handler.register_nodes(request)
payload = json.loads(response.text)
response = await handler.register_nodes(request) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert response.status == 200
assert payload["success"] is True
@@ -778,11 +790,11 @@ async def test_register_nodes_accepts_compound_node_ids():
@pytest.mark.asyncio
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 Instance:
sockets: dict = {}
sockets: dict[str, Any] = {}
def send_sync(self, event, payload, sid=None):
send_calls.append((event, payload))
@@ -791,7 +803,7 @@ async def test_update_node_widget_sends_payload():
handler = NodeRegistryHandler(
node_registry=NodeRegistry(),
prompt_server=RecordingPromptServer,
prompt_server=RecordingPromptServer, # pyright: ignore[reportArgumentType]
standalone_mode=False,
)
@@ -803,8 +815,8 @@ async def test_update_node_widget_sends_payload():
}
)
response = await handler.update_node_widget(request)
payload = json.loads(response.text)
response = await handler.update_node_widget(request) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert response.status == 200
assert payload["success"] is True
@@ -824,18 +836,18 @@ async def test_update_node_widget_sends_payload():
@pytest.mark.asyncio
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 Instance:
sockets: dict = {}
sockets: dict[str, Any] = {}
def send_sync(self, event, payload, sid=None):
send_calls.append((event, payload))
instance = Instance()
handler = LoraCodeHandler(RecordingPromptServer)
handler = LoraCodeHandler(RecordingPromptServer) # pyright: ignore[reportArgumentType]
request = FakeRequest(
json_data={
@@ -845,8 +857,8 @@ async def test_update_lora_code_includes_graph_identifier():
}
)
response = await handler.update_lora_code(request)
payload = json.loads(response.text)
response = await handler.update_lora_code(request) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert payload["success"] is True
assert payload["results"] == [
@@ -913,7 +925,7 @@ class FakeUserModelsProvider(FakeMetadataProvider):
self.next_cursor = next_cursor
self.estimated_total = estimated_total
self.received_usernames: list[str] = []
self.received_cursors: list = []
self.received_cursors: list[Any] = []
async def get_user_models(self, username, cursor=None):
self.received_usernames.append(username)
@@ -924,7 +936,7 @@ class FakeUserModelsProvider(FakeMetadataProvider):
return self.estimated_total
async def fake_metadata_provider_factory():
async def fake_metadata_provider_factory() -> Any:
return FakeMetadataProvider()
@@ -949,8 +961,8 @@ async def fake_metadata_archive_manager_factory():
class FakeDownloadHistoryService:
def __init__(self, downloaded_by_type=None):
self.downloaded_by_type = downloaded_by_type or {}
self.marked_downloaded: list[tuple] = []
self.marked_not_downloaded: list[tuple] = []
self.marked_downloaded: list[tuple[Any, ...]] = []
self.marked_not_downloaded: list[tuple[Any, ...]] = []
async def has_been_downloaded(self, model_type, version_id):
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(
settings_service=DummySettings(),
usage_stats_factory=lambda: SimpleNamespace(
usage_stats_factory=lambda: SimpleNamespace( # pyright: ignore[reportArgumentType]
process_execution=noop_async, get_stats=noop_async
),
prompt_server=FakePromptServer,
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
service_registry_adapter=service_registry_adapter,
metadata_provider_factory=fake_metadata_provider_factory,
metadata_archive_manager_factory=fake_metadata_archive_manager_factory,
metadata_provider_updater=noop_async,
downloader_factory=dummy_downloader_factory,
registrar_factory=registrar_factory,
registrar_factory=registrar_factory, # pyright: ignore[reportArgumentType]
)
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"
mapping = recorded_registrars[0].registered_mapping
@@ -1106,7 +1118,7 @@ async def test_get_civitai_user_models_marks_library_versions():
provider = FakeUserModelsProvider(models)
async def provider_factory():
async def provider_factory() -> Any:
return provider
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(
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["username"] == "pixel"
@@ -1239,7 +1251,7 @@ async def test_get_civitai_user_models_rewrites_civitai_previews():
provider = FakeUserModelsProvider(models)
async def provider_factory():
async def provider_factory() -> Any:
return provider
handler = ModelLibraryHandler(
@@ -1253,9 +1265,9 @@ async def test_get_civitai_user_models_rewrites_civitai_previews():
)
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
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():
provider = FakeUserModelsProvider([])
async def provider_factory():
async def provider_factory() -> Any:
return provider
handler = ModelLibraryHandler(
@@ -1288,8 +1300,8 @@ async def test_get_civitai_user_models_requires_username():
metadata_provider_factory=provider_factory,
)
response = await handler.get_civitai_user_models(FakeRequest())
payload = json.loads(response.text)
response = await handler.get_civitai_user_models(FakeRequest()) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert response.status == 400
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)
async def provider_factory():
async def provider_factory() -> Any:
return provider
handler = ModelLibraryHandler(
@@ -1332,9 +1344,9 @@ async def test_get_civitai_user_models_returns_pagination_fields():
)
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 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():
provider = FakeUserModelsProvider([], next_cursor=None, estimated_total=999)
async def provider_factory():
async def provider_factory() -> Any:
return provider
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(
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 payload["success"] is True
@@ -1391,10 +1403,10 @@ def test_ensure_handler_mapping_caches_result():
controller = MiscRoutes(
settings_service=DummySettings(),
usage_stats_factory=lambda: SimpleNamespace(
usage_stats_factory=lambda: SimpleNamespace( # pyright: ignore[reportArgumentType]
process_execution=noop_async, get_stats=noop_async
),
prompt_server=FakePromptServer,
prompt_server=FakePromptServer, # pyright: ignore[reportArgumentType]
service_registry_adapter=ServiceRegistryAdapter(
get_lora_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_provider_updater=noop_async,
downloader_factory=dummy_downloader_factory,
handler_set_factory=RecordingHandlerSet,
handler_set_factory=RecordingHandlerSet, # pyright: ignore[reportArgumentType]
)
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,
)
response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"}))
payload = json.loads(response.text)
response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"})) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert payload["success"] is True
assert payload["modelType"] == "lora"
@@ -1461,7 +1473,7 @@ async def test_check_model_exists_returns_local_versions():
@pytest.mark.asyncio
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")
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,
)
response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"}))
payload = json.loads(response.text)
response = await handler.check_model_exists(FakeRequest(query={"modelId": "5"})) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert payload == {
"success": True,
@@ -1503,9 +1515,9 @@ async def test_check_model_exists_returns_download_history_when_file_missing():
)
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 == {
"success": True,
@@ -1533,9 +1545,9 @@ async def test_model_version_download_status_endpoints():
)
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 == {
"success": True,
"modelType": "lora",
@@ -1544,7 +1556,7 @@ async def test_model_version_download_status_endpoints():
}
set_response = await handler.set_model_version_download_status(
FakeRequest(
FakeRequest( # pyright: ignore[reportArgumentType]
json_data={
"modelType": "checkpoint",
"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 == {
"success": True,
"modelType": "checkpoint",
@@ -1566,7 +1578,7 @@ async def test_model_version_download_status_endpoints():
]
set_get_response = await handler.set_model_version_download_status(
FakeRequest(
FakeRequest( # pyright: ignore[reportArgumentType]
method="GET",
query={
"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 == {
"success": True,
"modelType": "embedding",
@@ -1586,7 +1598,7 @@ async def test_model_version_download_status_endpoints():
def test_create_handler_set_uses_provided_dependencies():
recorded_handlers: list[dict] = []
recorded_handlers: list[dict[str, Any]] = []
class RecordingHandlerSet:
def __init__(self, **handlers):
@@ -1609,8 +1621,8 @@ def test_create_handler_set_uses_provided_dependencies():
controller = MiscRoutes(
settings_service=DummySettings(),
usage_stats_factory=lambda: FakeUsageStats(),
prompt_server=CustomPromptServer,
usage_stats_factory=lambda: FakeUsageStats(), # pyright: ignore[reportArgumentType]
prompt_server=CustomPromptServer, # pyright: ignore[reportArgumentType]
service_registry_adapter=ServiceRegistryAdapter(
get_lora_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_provider_updater=noop_async,
downloader_factory=dummy_downloader_factory,
handler_set_factory=RecordingHandlerSet,
node_registry=fake_node_registry,
handler_set_factory=RecordingHandlerSet, # pyright: ignore[reportArgumentType]
node_registry=fake_node_registry, # pyright: ignore[reportArgumentType]
standalone_mode_flag=True,
)
@@ -1665,7 +1677,7 @@ def test_is_wsl_returns_false_on_read_error():
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()):
result = _is_wsl()
assert result is False
@@ -1688,7 +1700,7 @@ def test_wsl_to_windows_path_returns_none_on_error():
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(
"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),),
)
response = await handler.get_doctor_diagnostics(FakeRequest(method="GET"))
payload = json.loads(response.text)
response = await handler.get_doctor_diagnostics(FakeRequest(method="GET")) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
diagnostic_map = {item["id"]: item for item in payload["diagnostics"]}
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),),
)
response = await handler.get_doctor_diagnostics(FakeRequest(method="GET"))
payload = json.loads(response.text)
response = await handler.get_doctor_diagnostics(FakeRequest(method="GET")) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
diagnostic_map = {item["id"]: item for item in payload["diagnostics"]}
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),),
)
response = await handler.resolve_filename_conflicts(FakeRequest(method="POST"))
payload = json.loads(response.text)
response = await handler.resolve_filename_conflicts(FakeRequest(method="POST")) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert payload["success"] is True
# 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),),
)
response = await handler.resolve_filename_conflicts(FakeRequest(method="POST"))
payload = json.loads(response.text)
response = await handler.resolve_filename_conflicts(FakeRequest(method="POST")) # pyright: ignore[reportArgumentType]
payload = _json_payload(response)
assert payload["success"] is True
assert payload["count"] == 0
+5 -7
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import logging
from types import SimpleNamespace
import jinja2
@@ -48,19 +49,16 @@ async def test_model_page_view_reads_version_per_request():
template_env=template_env,
template_name="dummy.html",
service=DummyService(),
settings_service=DummySettings(),
settings_service=DummySettings(), # pyright: ignore[reportArgumentType]
server_i18n=DummyI18n(),
logger=SimpleNamespace(
debug=lambda *_args, **_kwargs: None,
error=lambda *_args, **_kwargs: None,
),
logger=logging.getLogger("test_model_page_view"),
)
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"
second = await view.handle(SimpleNamespace())
second = await view.handle(SimpleNamespace()) # pyright: ignore[reportArgumentType]
assert first.text == "1.0.2-old"
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()
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
response = await handler.get_base_models(SimpleNamespace(query={"limit": "0"}))
payload = json.loads(response.text)
response = await handler.get_base_models(
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 service.received_limit == 0
@@ -33,7 +37,9 @@ async def test_model_query_handler_rejects_negative_limit_for_base_models():
service = DummyService()
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
@@ -58,9 +64,11 @@ async def test_model_query_handler_search_tags_passes_query_and_limit():
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
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["tags"] == [{"tag": "anime", "count": 3}]
@@ -73,7 +81,8 @@ async def test_model_query_handler_search_tags_defaults_limit_to_20():
service = DummySearchTagsService()
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
@@ -83,6 +92,8 @@ async def test_model_query_handler_search_tags_clamps_negative_limit():
service = DummySearchTagsService()
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
+38 -19
View File
@@ -2,6 +2,7 @@ import copy
import json
import logging
from types import SimpleNamespace
from typing import Any
import pytest
@@ -198,22 +199,24 @@ async def test_get_civitai_versions_degrades_when_download_history_unavailable(m
handler = ModelCivitaiHandler(
service=service,
settings_service=SimpleNamespace(get=lambda *_: False),
ws_manager=SimpleNamespace(),
settings_service=SimpleNamespace(get=lambda *_: False), # pyright: ignore[reportArgumentType]
ws_manager=SimpleNamespace(), # pyright: ignore[reportArgumentType]
logger=logging.getLogger(__name__),
metadata_provider_factory=metadata_provider_factory,
validate_model_type=lambda *_: True,
expected_model_types=lambda: "LoRA",
find_model_file=lambda *_: None,
metadata_sync=SimpleNamespace(),
metadata_refresh_use_case=SimpleNamespace(),
metadata_progress_callback=lambda *_args, **_kwargs: None,
metadata_sync=SimpleNamespace(), # pyright: ignore[reportArgumentType]
metadata_refresh_use_case=SimpleNamespace(), # pyright: ignore[reportArgumentType]
metadata_progress_callback=lambda *_args, **_kwargs: None, # pyright: ignore[reportArgumentType]
)
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 payload[0]["id"] == 7
@@ -284,10 +287,14 @@ async def test_refresh_model_updates_filters_records_without_updates():
async def json(self):
return {}
response = await handler.refresh_model_updates(DummyRequest())
response = await handler.refresh_model_updates(
DummyRequest() # pyright: ignore[reportArgumentType]
)
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 len(payload["records"]) == 1
assert payload["records"][0]["modelId"] == 1
@@ -347,7 +354,9 @@ async def test_refresh_model_updates_with_target_ids():
async def json(self):
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
call = update_service.calls[0]
@@ -399,7 +408,9 @@ async def test_refresh_model_updates_accepts_snake_case_ids():
async def json(self):
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
call = update_service.calls[0]
@@ -429,9 +440,9 @@ async def test_fetch_missing_license_data_updates_metadata(monkeypatch):
return None, 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)))
return True
@@ -479,10 +490,14 @@ async def test_fetch_missing_license_data_updates_metadata(monkeypatch):
async def json(self):
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
payload = json.loads(response.text)
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert len(payload["updated"]) == 3
assert provider_calls == [[10, 20]]
@@ -516,9 +531,9 @@ async def test_fetch_missing_license_data_filters_model_ids(monkeypatch):
return None, 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)))
return True
@@ -566,10 +581,14 @@ async def test_fetch_missing_license_data_filters_model_ids(monkeypatch):
async def json(self):
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
payload = json.loads(response.text)
text = response.text
assert text is not None
payload = json.loads(text)
assert payload["success"] is True
assert len(payload["updated"]) == 1
assert provider_calls == [[20]]
+1 -1
View File
@@ -31,7 +31,7 @@ class StubLoraService:
@pytest.fixture
def routes():
handler = LoraRoutes()
handler.service = StubLoraService()
handler.service = StubLoraService() # pyright: ignore[reportAttributeAccessIssue]
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__),
)
response = await handler.get_base_models(SimpleNamespace(query={"limit": "0"}))
payload = json.loads(response.text)
response = await handler.get_base_models(
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["base_models"] == [
@@ -124,7 +124,7 @@ def test_to_route_mapping_uses_handler_set():
super().__init__()
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
return DummyHandlerSet()
@@ -162,10 +162,12 @@ def test_recipe_route_registrar_binds_every_route():
self.router = FakeRouter()
app = FakeApp()
registrar = recipe_route_registrar.RecipeRouteRegistrar(app)
registrar = recipe_route_registrar.RecipeRouteRegistrar(
app # pyright: ignore[reportArgumentType]
)
handler_mapping = {
definition.handler_name: object()
definition.handler_name: lambda _request: None
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:
"""Container exposing the aiohttp client and stubbed collaborators."""
client: TestClient
client: TestClient[Any, Any]
scanner: "StubRecipeScanner"
analysis: "StubAnalysisService"
persistence: "StubPersistenceService"
@@ -92,6 +92,9 @@ class StubRecipeScanner:
candidate = Path(self.recipes_dir) / f"{recipe_id}.recipe.json"
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:
self.removed.append(recipe_id)
self.recipes.pop(recipe_id, None)
@@ -110,7 +113,7 @@ class StubAnalysisService:
self.remote_calls: List[Optional[str]] = []
self.local_calls: List[Optional[str]] = []
self.result = SimpleNamespace(payload={"loras": []}, status=200)
self._recipe_parser_factory = None
self._recipe_parser_factory: Any = None
StubAnalysisService.instances.append(self)
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)
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):
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 payload["items"] == []
assert harness.scanner.last_paginated_params is not None
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>"]
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")
payload = await response.json()
+3 -3
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import asyncio
from contextlib import asynccontextmanager
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.test_utils import TestClient, TestServer
@@ -34,7 +34,7 @@ class IntegrationCache:
class IntegrationScanner:
"""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._cache = IntegrationCache(list(items))
self._hash_index = SimpleNamespace(
@@ -68,7 +68,7 @@ class IntegrationScanner:
@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."""
server = TestServer(app)
+24 -10
View File
@@ -1,4 +1,5 @@
import json
from typing import Any, Optional
import pytest
@@ -17,7 +18,7 @@ class FakeRequest:
class DummySettings:
def __init__(self):
self.activated = None
self.should_raise = None
self.should_raise: Optional[Exception] = None
def activate_library(self, name):
if self.should_raise:
@@ -25,6 +26,13 @@ class DummySettings:
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:
async def refresh_session(self): # pragma: no cover - helper
return None
@@ -53,7 +61,7 @@ async def test_get_libraries_returns_registry(monkeypatch, handler):
monkeypatch.setattr(config, "get_library_registry_snapshot", lambda: registry)
response = await handler.get_libraries(FakeRequest())
payload = json.loads(response.text)
payload = json_payload(response)
assert response.status == 200
assert payload == {
@@ -71,7 +79,7 @@ async def test_get_libraries_handles_errors(monkeypatch, handler):
monkeypatch.setattr(config, "get_library_registry_snapshot", boom)
response = await handler.get_libraries(FakeRequest())
payload = json.loads(response.text)
payload = json_payload(response)
assert response.status == 500
assert payload["success"] is False
@@ -90,8 +98,10 @@ async def test_activate_library_success(monkeypatch):
registry = {"libraries": {"alpha": {"name": "Alpha"}}, "active_library": "alpha"}
monkeypatch.setattr(config, "get_library_registry_snapshot", lambda: registry)
response = await handler.activate_library(FakeRequest(json_data={"library": "alpha"}))
payload = json.loads(response.text)
response = await handler.activate_library(
FakeRequest(json_data={"library": "alpha"}) # pyright: ignore[reportArgumentType]
)
payload = json_payload(response)
assert response.status == 200
assert payload == {
@@ -105,7 +115,7 @@ async def test_activate_library_success(monkeypatch):
@pytest.mark.asyncio
async def test_activate_library_requires_name(handler):
response = await handler.activate_library(FakeRequest(json_data={}))
payload = json.loads(response.text)
payload = json_payload(response)
assert response.status == 400
assert payload["success"] is False
@@ -122,8 +132,10 @@ async def test_activate_library_unknown_returns_404(monkeypatch):
downloader_factory=dummy_downloader_factory,
)
response = await handler.activate_library(FakeRequest(json_data={"library": "ghost"}))
payload = json.loads(response.text)
response = await handler.activate_library(
FakeRequest(json_data={"library": "ghost"}) # pyright: ignore[reportArgumentType]
)
payload = json_payload(response)
assert response.status == 404
assert payload["success"] is False
@@ -140,8 +152,10 @@ async def test_activate_library_unexpected_error_returns_500(monkeypatch):
downloader_factory=dummy_downloader_factory,
)
response = await handler.activate_library(FakeRequest(json_data={"library": "broken"}))
payload = json.loads(response.text)
response = await handler.activate_library(
FakeRequest(json_data={"library": "broken"}) # pyright: ignore[reportArgumentType]
)
payload = json_payload(response)
assert response.status == 500
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.text == "rendered"
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 stats_routes.routes.template_env.filters["t"]("greeting") == "translated:greeting"
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 types
from typing import Any
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
@@ -19,6 +20,7 @@ class MockService:
def __init__(self):
self.model_type = "test-model"
self.last_call_kwargs: dict[str, Any] = {}
async def get_paginated_data(self, **kwargs):
# Store the kwargs for verification
+2 -2
View File
@@ -24,7 +24,7 @@ def _fake_request(body=None, query_params=None):
async def _json():
return body or {}
req.json = _json
req.json = _json # pyright: ignore[reportAttributeAccessIssue]
return req
@@ -131,7 +131,7 @@ async def test_perform_git_update_preserves_user_dirs(monkeypatch, tmp_path):
class FakeHeads:
def __getitem__(self, name):
class Head:
def checkout(self_inner):
def checkout(self):
calls.append(("head-checkout", (name,)))
return Head()
+10 -4
View File
@@ -32,9 +32,11 @@ async def test_search_wildcards_returns_results():
handler = WildcardsHandler(service=StubService())
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 payload == {
@@ -62,8 +64,12 @@ async def test_search_wildcards_handles_errors():
raise RuntimeError("boom")
handler = WildcardsHandler(service=StubService())
response = await handler.search_wildcards(FakeRequest(query={"search": "cat"}))
payload = json.loads(response.text)
response = await handler.search_wildcards(
FakeRequest(query={"search": "cat"}) # pyright: ignore[reportArgumentType]
)
text = response.text
assert text is not None
payload = json.loads(text)
assert response.status == 500
assert payload["error"] == "boom"
+7 -7
View File
@@ -6,7 +6,7 @@ from unittest.mock import AsyncMock
import pytest
from py.services.aria2_downloader import Aria2Downloader, Aria2Error
from py.services.aria2_downloader import Aria2Downloader, Aria2Error, Aria2Transfer
from py.services.aria2_transfer_state import Aria2TransferStateStore
from py.services import aria2_transfer_state
@@ -165,9 +165,9 @@ async def test_download_file_keeps_auth_headers_when_civitai_does_not_redirect(
@pytest.mark.asyncio
async def test_pause_resume_cancel_forward_to_rpc(monkeypatch):
downloader = Aria2Downloader()
downloader._transfers["download-1"] = type(
"Transfer", (), {"gid": "gid-1", "save_path": "/tmp/model.safetensors"}
)()
downloader._transfers["download-1"] = Aria2Transfer(
gid="gid-1", save_path="/tmp/model.safetensors"
)
calls = []
@@ -200,9 +200,9 @@ async def test_download_file_reuses_existing_transfer_without_add_uri(
downloader._rpc_secret = "secret"
save_path = tmp_path / "downloads" / "model.safetensors"
downloader._transfers["download-1"] = type(
"Transfer", (), {"gid": "gid-1", "save_path": str(save_path)}
)()
downloader._transfers["download-1"] = Aria2Transfer(
gid="gid-1", save_path=str(save_path)
)
rpc_calls = []
statuses = iter(
+17 -14
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
import json
from pathlib import Path
from typing import Any, Dict, List, Optional
from typing import Any, Dict, Iterator, List, Optional
import pytest
@@ -16,7 +16,7 @@ from py.services.persistent_model_cache import DEFAULT_LICENSE_FLAGS, Persistent
@pytest.fixture(autouse=True)
def reset_backfill_singleton() -> None:
def reset_backfill_singleton() -> Iterator[None]:
"""Reset the service singleton so every test starts from a fresh instance."""
Autov3BackfillService._instance = None
yield
@@ -66,7 +66,7 @@ class RecordingScanner:
self.model_type = model_type
self._persistent_cache = persistent_cache
self.entries: Dict[str, Dict[str, Any]] = {entry['file_path']: entry for entry in entries}
self.update_calls: List[tuple] = []
self.update_calls: List[tuple[str, str, str]] = []
async def update_autov3_for_model(self, model_type: str, file_path: str, autov3: str) -> bool:
self.update_calls.append((model_type, file_path, autov3))
@@ -114,7 +114,7 @@ async def test_backfill_updates_models_and_self_terminates(tmp_path: Path, monke
)
scanner = RecordingScanner('dummy', store, entries)
updated = await Autov3BackfillService.get_instance().backfill(scanner)
updated = await Autov3BackfillService.get_instance().backfill(scanner) # pyright: ignore[reportArgumentType]
# Non-safetensors files yield no embedded hash, so both are marked ''.
assert updated == 2
@@ -124,6 +124,7 @@ async def test_backfill_updates_models_and_self_terminates(tmp_path: Path, monke
assert store.get_models_missing_autov3('dummy') == []
persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data}
assert items[path_a]['autov3'] == ''
assert items[path_b]['autov3'] == ''
@@ -147,7 +148,7 @@ async def test_backfill_skips_missing_files_without_marking(tmp_path: Path, monk
)
scanner = RecordingScanner('dummy', store, entries)
updated = await Autov3BackfillService.get_instance().backfill(scanner)
updated = await Autov3BackfillService.get_instance().backfill(scanner) # pyright: ignore[reportArgumentType]
assert updated == 1
assert scanner.update_calls == [('dummy', existing, '')]
@@ -162,7 +163,7 @@ async def test_backfill_returns_zero_when_same_type_already_running(tmp_path: Pa
service = Autov3BackfillService.get_instance()
service._running_types = {'dummy'}
try:
assert await service.backfill(scanner) == 0
assert await service.backfill(scanner) == 0 # pyright: ignore[reportArgumentType]
finally:
service._running_types = set()
assert scanner.update_calls == []
@@ -195,7 +196,7 @@ async def test_backfill_runs_concurrently_for_different_model_types(tmp_path: Pa
try:
# The lora backfill must still run while checkpoint is in progress.
assert await service.backfill(lora_scanner) == 1
assert await service.backfill(lora_scanner) == 1 # pyright: ignore[reportArgumentType]
assert lora_scanner.update_calls == [('lora', lora_file, '')]
finally:
service._running_types = set()
@@ -213,7 +214,7 @@ async def test_backfill_never_raises_on_failure(tmp_path: Path, monkeypatch) ->
store.save_cache('dummy', entries, {'hash-boom': [existing]}, [])
scanner = RaisingScanner('dummy', store, entries)
updated = await Autov3BackfillService.get_instance().backfill(scanner)
updated = await Autov3BackfillService.get_instance().backfill(scanner) # pyright: ignore[reportArgumentType]
assert updated == 0
@@ -238,7 +239,7 @@ async def test_backfill_uses_default_cache_when_scanner_has_none(tmp_path: Path,
store.update_single_model(model_type, new_item, old_item)
return True
updated = await Autov3BackfillService.get_instance().backfill(BareScanner())
updated = await Autov3BackfillService.get_instance().backfill(BareScanner()) # pyright: ignore[reportArgumentType]
assert updated == 1
assert store.get_models_missing_autov3('dummy') == []
@@ -253,9 +254,9 @@ async def test_backfill_idempotent_second_run_is_noop(tmp_path: Path, monkeypatc
scanner = RecordingScanner('dummy', store, entries)
service = Autov3BackfillService.get_instance()
assert await service.backfill(scanner) == 1
assert await service.backfill(scanner) == 1 # pyright: ignore[reportArgumentType]
# A re-run has nothing left to do.
assert await service.backfill(scanner) == 0
assert await service.backfill(scanner) == 0 # pyright: ignore[reportArgumentType]
assert len(scanner.update_calls) == 1
@@ -275,7 +276,7 @@ async def test_backfill_end_to_end_through_scanner_lazy_import(tmp_path: Path, m
)
class RealScanner(ModelScanner):
def __init__(self) -> None:
def __init__(self) -> None: # pyright: ignore[reportMissingSuperCall]
self.model_type = 'dummy'
self._persistent_cache = store
self._cache = ModelCache(raw_data=[dict(e) for e in entries], folders=[])
@@ -285,6 +286,7 @@ async def test_backfill_end_to_end_through_scanner_lazy_import(tmp_path: Path, m
assert store.get_models_missing_autov3('dummy') == []
persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data}
assert items[path_a]['autov3'] == ''
assert items[path_b]['autov3'] == ''
@@ -314,12 +316,13 @@ async def test_backfill_prefers_civitai_autov3_from_sidecar(tmp_path: Path, monk
store.save_cache('dummy', [_entry(path, 'hash-ckpt')], {'hash-ckpt': [path]}, [])
scanner = RecordingScanner('dummy', store, [_entry(path, 'hash-ckpt')])
updated = await Autov3BackfillService.get_instance().backfill(scanner)
updated = await Autov3BackfillService.get_instance().backfill(scanner) # pyright: ignore[reportArgumentType]
assert updated == 1
assert scanner.update_calls == [('dummy', path, 'abcdef123456')]
persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data}
assert items[path]['autov3'] == 'abcdef123456'
# Self-terminating: the row is marked and the driving query empties.
@@ -344,7 +347,7 @@ async def test_backfill_falls_back_to_header_when_sidecar_has_no_match(tmp_path:
store.save_cache('dummy', [_entry(path, 'hash-plain')], {'hash-plain': [path]}, [])
scanner = RecordingScanner('dummy', store, [_entry(path, 'hash-plain')])
updated = await Autov3BackfillService.get_instance().backfill(scanner)
updated = await Autov3BackfillService.get_instance().backfill(scanner) # pyright: ignore[reportArgumentType]
assert updated == 1
assert scanner.update_calls == [('dummy', path, '')]
+2 -2
View File
@@ -209,8 +209,8 @@ async def test_model_update_service_migrates_legacy_snapshot_db(tmp_path, monkey
return str(legacy_db)
monkeypatch.setattr(
"py.services.persistent_model_cache.get_persistent_cache",
lambda *_args, **_kwargs: LegacyCache(),
"py.services.persistent_model_cache.PersistentModelCache.get_default",
lambda *args, **kwargs: LegacyCache(),
)
service = ModelUpdateService(settings_manager=DummySettingsManager())
+19 -14
View File
@@ -28,13 +28,14 @@ class DummyService(BaseModelService):
return model_data
class StubRepository:
class StubRepository(ModelCacheRepository):
def __init__(self, data):
super().__init__(scanner=object())
self._data = list(data)
self.parse_sort_calls = []
self.fetch_sorted_calls = []
def parse_sort(self, sort_by):
def parse_sort(self, sort_by): # pyright: ignore[reportIncompatibleMethodOverride]
params = ModelCacheRepository.parse_sort(sort_by)
self.parse_sort_calls.append(sort_by)
return params
@@ -44,8 +45,9 @@ class StubRepository:
return list(self._data)
class StubFilterSet:
class StubFilterSet(ModelFilterSet):
def __init__(self, result):
super().__init__(settings=StubSettings({}))
self.result = list(result)
self.calls = []
@@ -54,8 +56,9 @@ class StubFilterSet:
return list(self.result)
class StubSearchStrategy:
class StubSearchStrategy(SearchStrategy):
def __init__(self, search_result):
super().__init__()
self.search_result = list(search_result)
self.normalize_calls = []
self.apply_calls = []
@@ -67,7 +70,7 @@ class StubSearchStrategy:
normalized.update(options)
return normalized
def apply(self, data, search_term, options, fuzzy):
def apply(self, data, search_term, options, fuzzy=False):
self.apply_calls.append((list(data), search_term, options, fuzzy))
return list(self.search_result)
@@ -269,8 +272,9 @@ async def test_get_paginated_data_filters_and_searches_combination():
assert response["total_pages"] == 1
class PassThroughFilterSet:
class PassThroughFilterSet(ModelFilterSet):
def __init__(self):
super().__init__(settings=StubSettings({}))
self.calls = []
def apply(self, data, criteria):
@@ -278,8 +282,9 @@ class PassThroughFilterSet:
return list(data)
class NoSearchStrategy:
class NoSearchStrategy(SearchStrategy):
def __init__(self):
super().__init__()
self.normalize_calls = []
self.apply_called = False
@@ -355,7 +360,7 @@ async def test_get_paginated_data_filters_by_update_status():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -428,7 +433,7 @@ async def test_get_paginated_data_skips_items_when_update_check_fails():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -465,7 +470,7 @@ async def test_get_paginated_data_annotates_update_flags_with_bulk_dedup():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -561,7 +566,7 @@ async def test_version_grouping_same_base_prefers_matching_base():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -658,7 +663,7 @@ async def test_version_grouping_same_base_honors_latest_local_version():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -694,7 +699,7 @@ async def test_get_paginated_data_filters_update_available_only():
filter_set=filter_set,
search_strategy=search_strategy,
settings_provider=settings,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
response = await service.get_paginated_data(
@@ -1028,7 +1033,7 @@ def test_model_filter_set_supports_legacy_tag_arrays():
{"model_name": "AnimeOnly", "tags": ["anime"]},
]
criteria = FilterCriteria(tags=["style"])
criteria = FilterCriteria(tags=["style"]) # pyright: ignore[reportArgumentType]
result = filter_set.apply(data, criteria)
assert [item["model_name"] for item in result] == ["StyleOnly", "StyleAnime"]
+7 -7
View File
@@ -333,7 +333,7 @@ class TestBatchImportService:
)
service = BatchImportService(
analysis_service=analysis_service,
analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service,
ws_manager=ws_manager,
logger=logger,
@@ -445,8 +445,8 @@ class TestBatchImportServiceEdgeCases:
logger = logging.getLogger("test")
return BatchImportService(
analysis_service=analysis_service,
persistence_service=persistence_service,
analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service, # pyright: ignore[reportArgumentType]
ws_manager=ws_manager,
logger=logger,
)
@@ -506,8 +506,8 @@ class TestBatchImportServiceEdgeCases:
(tmp_path / "test.png").write_bytes(b"fake-image")
service = BatchImportService(
analysis_service=analysis_service,
persistence_service=persistence_service,
analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service, # pyright: ignore[reportArgumentType]
ws_manager=ws_manager,
logger=logger,
)
@@ -571,8 +571,8 @@ class TestInputValidation:
logger = logging.getLogger("test")
return BatchImportService(
analysis_service=analysis_service,
persistence_service=persistence_service,
analysis_service=analysis_service, # pyright: ignore[reportArgumentType]
persistence_service=persistence_service, # pyright: ignore[reportArgumentType]
ws_manager=ws_manager,
logger=logger,
)
+15 -6
View File
@@ -84,6 +84,7 @@ class TestCacheEntryValidator:
result = CacheEntryValidator.validate(entry, auto_repair=False)
assert result.is_valid is True
assert result.entry is not None
assert result.entry['sha256'] == ''
assert result.entry['hash_status'] == 'pending'
@@ -141,7 +142,7 @@ class TestCacheEntryValidator:
def test_validate_none_entry(self):
"""Test validation handles None entry"""
result = CacheEntryValidator.validate(None, auto_repair=False)
result = CacheEntryValidator.validate(None, auto_repair=False) # pyright: ignore[reportArgumentType]
assert result.is_valid is False
assert result.repaired is False
@@ -150,7 +151,7 @@ class TestCacheEntryValidator:
def test_validate_non_dict_entry(self):
"""Test validation handles non-dict entry"""
result = CacheEntryValidator.validate("not a dict", auto_repair=False)
result = CacheEntryValidator.validate("not a dict", auto_repair=False) # pyright: ignore[reportArgumentType]
assert result.is_valid is False
assert result.repaired is False
@@ -169,6 +170,7 @@ class TestCacheEntryValidator:
assert result.is_valid is True
assert result.repaired is True
assert result.entry is not None
assert result.entry['file_name'] == ''
assert result.entry['model_name'] == ''
assert result.entry['tags'] == []
@@ -186,6 +188,7 @@ class TestCacheEntryValidator:
assert result.is_valid is True
assert result.repaired is True
assert result.entry is not None
assert result.entry['size'] == 0 # Default value
assert result.entry['tags'] == [] # Default value
@@ -199,6 +202,7 @@ class TestCacheEntryValidator:
result = CacheEntryValidator.validate(entry, auto_repair=True)
assert result.is_valid is True
assert result.entry is not None
assert result.entry['sha256'] == 'abc123def456'
def test_validate_batch_all_valid(self):
@@ -262,8 +266,8 @@ class TestCacheEntryValidator:
def test_get_file_path_safe_not_dict(self):
"""Test safe file_path extraction from non-dict"""
assert CacheEntryValidator.get_file_path_safe(None) == ''
assert CacheEntryValidator.get_file_path_safe('string') == ''
assert CacheEntryValidator.get_file_path_safe(None) == '' # pyright: ignore[reportArgumentType]
assert CacheEntryValidator.get_file_path_safe('string') == '' # pyright: ignore[reportArgumentType]
def test_get_sha256_safe(self):
"""Test safe sha256 extraction"""
@@ -277,8 +281,8 @@ class TestCacheEntryValidator:
def test_get_sha256_safe_not_dict(self):
"""Test safe sha256 extraction from non-dict"""
assert CacheEntryValidator.get_sha256_safe(None) == ''
assert CacheEntryValidator.get_sha256_safe('string') == ''
assert CacheEntryValidator.get_sha256_safe(None) == '' # pyright: ignore[reportArgumentType]
assert CacheEntryValidator.get_sha256_safe('string') == '' # pyright: ignore[reportArgumentType]
def test_validate_with_all_optional_fields(self):
"""Test validation with all optional fields present"""
@@ -358,6 +362,7 @@ class TestAutov3Validation:
)
assert result.is_valid is True
assert result.entry is not None
assert result.entry['autov3'] == 'abcdef123456'
assert result.repaired is True
@@ -378,6 +383,7 @@ class TestAutov3Validation:
assert result.is_valid is True
assert result.repaired is False
assert result.entry is not None
assert result.entry['autov3'] is None
def test_validate_absent_autov3_is_valid_and_not_counted_as_repair(self):
@@ -386,6 +392,7 @@ class TestAutov3Validation:
assert result.is_valid is True
assert result.repaired is False
assert result.entry is not None
assert 'autov3' not in result.entry
def test_validate_short_autov3_still_valid_and_repaired_to_none(self):
@@ -396,6 +403,7 @@ class TestAutov3Validation:
)
assert result.is_valid is True
assert result.entry is not None
assert result.entry['autov3'] is None
assert result.repaired is True
@@ -407,5 +415,6 @@ class TestAutov3Validation:
)
assert result.is_valid is True
assert result.entry is not None
assert result.entry['autov3'] is None
assert result.repaired is True
+4 -3
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
import json
from types import SimpleNamespace
from typing import Any
import pytest
@@ -13,7 +14,7 @@ from py.utils import example_images_download_manager as download_module
class StubScanner:
"""Scanner double returning predetermined cache contents."""
def __init__(self, models: list[dict]) -> None:
def __init__(self, models: list[dict[str, Any]]) -> None:
self._cache = SimpleNamespace(raw_data=models)
async def get_cached_data(self):
@@ -58,9 +59,9 @@ class RecordingWebSocketManager:
"""Collects broadcast payloads for assertions."""
def __init__(self) -> None:
self.payloads: list[dict] = []
self.payloads: list[dict[str, Any]] = []
async def broadcast(self, payload: dict) -> None:
async def broadcast(self, payload: dict[str, Any]) -> None:
self.payloads.append(payload)
+3 -3
View File
@@ -4,7 +4,7 @@ import asyncio
import json
import os
from pathlib import Path
from typing import List
from typing import Any, List
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@@ -17,9 +17,9 @@ from py.utils.models import CheckpointMetadata
class RecordingWebSocketManager:
def __init__(self) -> None:
self.payloads: List[dict] = []
self.payloads: List[dict[str, Any]] = []
async def broadcast_init_progress(self, payload: dict) -> None:
async def broadcast_init_progress(self, payload: dict[str, Any]) -> None:
self.payloads.append(payload)
+3 -3
View File
@@ -1,6 +1,6 @@
import os
from pathlib import Path
from typing import List
from typing import Any, List
import pytest
@@ -12,9 +12,9 @@ from py.services.persistent_model_cache import PersistedCacheData
class RecordingWebSocketManager:
def __init__(self) -> None:
self.payloads: List[dict] = []
self.payloads: List[dict[str, Any]] = []
async def broadcast_init_progress(self, payload: dict) -> None:
async def broadcast_init_progress(self, payload: dict[str, Any]) -> None:
self.payloads.append(payload)
+6 -2
View File
@@ -1,4 +1,5 @@
import copy
from typing import Any, Dict
from unittest.mock import AsyncMock
import pytest
@@ -34,7 +35,7 @@ def downloader(monkeypatch):
return instance
def _base_civarchive_payload(version_id=1976567, *, trigger="mxpln", nsfw_level=31):
def _base_civarchive_payload(version_id=1976567, *, trigger="mxpln", nsfw_level=31) -> Dict[str, Any]:
version_name = "v2.0" if version_id != 1976567 else "v1.0"
file_sha = "e2b7a280d6539556f23f380b3f71e4e22bc4524445c4c96526e117c6005c6ad3"
return {
@@ -110,6 +111,7 @@ async def test_get_model_by_hash_transforms_payload(downloader):
result, error = await client.get_model_by_hash("abc")
assert error is None
assert result is not None
assert result["id"] == 1976567
assert result["nsfwLevel"] == 31
assert result["trainedWords"] == ["mxpln"]
@@ -131,7 +133,7 @@ async def test_get_model_versions_fetches_each_version(downloader):
base_payload = _base_civarchive_payload(version_id=2042594, trigger="mxpln-new", nsfw_level=5)
other_payload = _base_civarchive_payload()
responses = {
responses: Dict[Any, Dict[str, Any]] = {
(base_url, None): base_payload,
(base_url, (("modelVersionId", "2042594"),)): base_payload,
(base_url, (("modelVersionId", "1976567"),)): other_payload,
@@ -151,6 +153,7 @@ async def test_get_model_versions_fetches_each_version(downloader):
result = await client.get_model_versions("1746460")
assert result is not None
assert result["name"] == "Mixplin Style [Illustrious]"
assert result["type"] == "LORA"
versions = result["modelVersions"]
@@ -221,6 +224,7 @@ async def test_get_model_by_hash_uses_file_fallback(downloader, monkeypatch):
result, error = await client.get_model_by_hash("fallback")
assert error is None
assert result is not None
assert result["id"] == 1976567
assert result["model"]["name"] == "Mixplin Style [Illustrious]"
assert any("/models/1746460" in call["url"] for call in downloader.calls)
@@ -9,6 +9,8 @@ from py.services.civitai_base_model_service import CivitaiBaseModelService
class TestCivitaiBaseModelService:
"""Test suite for CivitaiBaseModelService."""
service = CivitaiBaseModelService()
@pytest.fixture(autouse=True)
def setup_service(self):
"""Create a fresh service instance for each test."""
@@ -46,7 +48,7 @@ class TestCivitaiBaseModelService:
def test_generate_abbreviation_edge_cases(self):
"""Test abbreviation generation edge cases."""
assert self.service.generate_abbreviation("") == "OTH"
assert self.service.generate_abbreviation(None) == "OTH"
assert self.service.generate_abbreviation(None) == "OTH" # pyright: ignore[reportArgumentType]
def test_cache_status_no_cache(self):
"""Test cache status when no cache exists."""
+7 -1
View File
@@ -97,6 +97,7 @@ async def test_get_model_by_hash_enriches_metadata(monkeypatch, downloader):
result, error = await client.get_model_by_hash("hash")
assert error is None
assert result is not None
assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"]
assert result["creator"] == {"username": "user"}
@@ -254,7 +255,7 @@ async def test_get_model_versions_bulk_success(monkeypatch, downloader):
client = await CivitaiClient.get_instance()
result = await client.get_model_versions_bulk([1, "2", 2])
result = await client.get_model_versions_bulk([1, "2", 2]) # pyright: ignore[reportArgumentType]
assert result == {
1: {
@@ -310,6 +311,7 @@ async def test_get_model_version_by_version_id(monkeypatch, downloader):
result = await client.get_model_version(version_id=7)
assert result is not None
assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"]
assert result["creator"] == {"username": "user"}
@@ -364,6 +366,7 @@ async def test_get_model_version_with_model_id_prefers_version_endpoint(monkeypa
result = await client.get_model_version(model_id=99, version_id=7)
assert result is not None
assert result["id"] == 7
assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"]
@@ -420,6 +423,7 @@ async def test_get_model_version_with_model_id_fallbacks_to_hash(monkeypatch, do
result = await client.get_model_version(model_id=99, version_id=7)
assert result is not None
assert result["id"] == 7
assert result["model"]["description"] == "desc"
assert result["model"]["tags"] == ["tag"]
@@ -461,6 +465,7 @@ async def test_get_model_version_with_model_id_builds_from_model_data(monkeypatc
result = await client.get_model_version(model_id=99, version_id=7)
assert result is not None
assert result["modelId"] == 99
assert result["model"]["name"] == "Model"
assert result["model"]["type"] == "LORA"
@@ -503,6 +508,7 @@ async def test_get_model_version_info_success(monkeypatch, downloader):
assert result == expected
assert error is None
assert result is not None
assert "comfy" not in result["images"][0]["meta"]
assert result["images"][0]["meta"]["other"] == "keep"
+2 -2
View File
@@ -91,7 +91,7 @@ async def test_parse_metadata_handles_nested_meta_and_lowercase_hashes(monkeypat
},
}
assert parser.is_metadata_matching(metadata)
assert parser.is_metadata_matching(metadata) # pyright: ignore[reportArgumentType]
result = await parser.parse_metadata(metadata)
@@ -272,7 +272,7 @@ async def test_parse_metadata_handles_modelVersionIds(monkeypatch):
"modelVersionIds": [2398829, 2398838],
}
assert parser.is_metadata_matching(metadata)
assert parser.is_metadata_matching(metadata) # pyright: ignore[reportArgumentType]
result = await parser.parse_metadata(metadata)
@@ -758,6 +758,7 @@ async def test_get_active_downloads_restores_orphaned_aria2_partial_as_paused(
downloads = await manager.get_active_downloads()
persisted = await manager._aria2_state_store.get("download-1")
assert persisted is not None
assert downloads["downloads"] == [
{
@@ -922,6 +923,7 @@ async def test_get_active_downloads_restores_persisted_aria2_without_initial_sav
downloads = await manager.get_active_downloads()
persisted = await manager._aria2_state_store.get("download-1")
assert persisted is not None
assert downloads["downloads"] == [
{
@@ -3,6 +3,7 @@
import os
from pathlib import Path
from types import SimpleNamespace
from typing import Optional
from unittest.mock import AsyncMock
import pytest
@@ -99,7 +100,7 @@ async def test_execute_download_uses_rewritten_civitai_preview(monkeypatch, tmp_
self.file_path = str(path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.preview_url: Optional[str] = None
self.autov3 = None
self.preview_nsfw_level = None
@@ -182,6 +183,7 @@ async def test_execute_download_uses_rewritten_civitai_preview(monkeypatch, tmp_
assert any("width=450,optimized=true" in url for url in preview_urls)
assert dummy_downloader.memory_calls == 0
assert optimize_called["value"] is False
assert metadata.preview_url is not None
assert metadata.preview_url.endswith(".jpeg")
assert metadata.preview_nsfw_level == 2
stored_preview = manager._active_downloads["dl"]["preview_path"]
@@ -204,7 +206,7 @@ async def test_execute_download_respects_blur_setting(monkeypatch, tmp_path):
self.file_path = str(path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.preview_url: Optional[str] = None
self.autov3 = None
self.preview_nsfw_level = None
@@ -326,7 +328,7 @@ async def test_execute_download_uses_auth_for_red_civitai_downloads(monkeypatch,
self.file_path = str(path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.preview_url: Optional[str] = None
self.autov3 = None
self.preview_nsfw_level = None
+14 -14
View File
@@ -75,7 +75,7 @@ async def test_execute_download_retries_urls(monkeypatch, tmp_path):
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -165,7 +165,7 @@ async def test_execute_download_uses_aria2_backend_for_model_files(monkeypatch,
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -272,7 +272,7 @@ async def test_execute_download_allows_anonymous_civitai_with_aria2(
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -350,7 +350,7 @@ async def test_execute_download_adjusts_checkpoint_sub_type(monkeypatch, tmp_pat
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
self.preview_nsfw_level = 0
self.sub_type = "checkpoint"
@@ -450,7 +450,7 @@ async def test_execute_download_extracts_zip_single_model(monkeypatch, tmp_path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -531,7 +531,7 @@ async def test_execute_download_extracts_zip_multiple_models(monkeypatch, tmp_pa
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -614,7 +614,7 @@ async def test_execute_download_extracts_zip_pt_embedding(monkeypatch, tmp_path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -1168,7 +1168,7 @@ async def test_execute_download_waits_for_paused_pre_transfer_gate(monkeypatch,
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return os.path.basename(self.file_path)
@@ -1278,7 +1278,7 @@ async def test_execute_download_reuses_existing_aria2_partial_path(monkeypatch,
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
return "renamed.safetensors"
@@ -1356,7 +1356,7 @@ async def test_execute_download_rejects_conflicting_aria2_partial_path(tmp_path)
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
raise AssertionError("should not rename")
@@ -1408,7 +1408,7 @@ async def test_execute_download_reassigns_same_aria2_partial_to_new_download_id(
self.sha256 = "sha256"
self.file_name = path.stem
self.preview_url = None
self.autov3 = None
self.autov3: Optional[str] = None
def generate_unique_filename(self, *_args, **_kwargs):
raise AssertionError("should not rename")
@@ -1460,9 +1460,9 @@ async def test_execute_download_reassigns_same_aria2_partial_to_new_download_id(
assert manager._active_downloads["new-download"]["file_path"] == str(target_path)
assert dummy_aria2.calls == [("reassign_transfer", "old-download", "new-download")]
assert await manager._aria2_state_store.get("old-download") is None
assert (await manager._aria2_state_store.get("new-download"))["save_path"] == str(
target_path
)
persisted = await manager._aria2_state_store.get("new-download")
assert persisted is not None
assert persisted["save_path"] == str(target_path)
def test_is_same_aria2_download_request_requires_version_id_match():
+17 -9
View File
@@ -9,7 +9,7 @@ from py.services.downloader import Downloader
class FakeStream:
def __init__(self, chunks: Sequence[Sequence] | Sequence[bytes]):
def __init__(self, chunks: Sequence[bytes | tuple[bytes, float]]):
self._chunks = list(chunks)
async def read(self, _chunk_size: int) -> bytes:
@@ -25,6 +25,7 @@ class FakeStream:
payload = item[0]
delay = item[1]
assert isinstance(payload, bytes)
await asyncio.sleep(delay)
return payload
@@ -84,11 +85,11 @@ def _build_downloader(responses, *, max_retries=0):
downloader.max_retries = max_retries
downloader.base_delay = 0
fake_session = FakeSession(responses)
downloader._session = fake_session
downloader._session = fake_session # pyright: ignore[reportAttributeAccessIssue]
downloader._session_created_at = datetime.now()
downloader._proxy_url = None
async def _noop_create_session():
downloader._session = fake_session
downloader._session = fake_session # pyright: ignore[reportAttributeAccessIssue]
downloader._session_created_at = datetime.now()
downloader._proxy_url = None
@@ -96,6 +97,13 @@ def _build_downloader(responses, *, max_retries=0):
return downloader
def _session(downloader: Downloader) -> FakeSession:
"""Return the injected fake session, asserting the runtime invariant."""
session = downloader._session
assert isinstance(session, FakeSession)
return session
@pytest.mark.asyncio
async def test_download_file_preserves_incomplete_part_when_size_mismatch(tmp_path):
target_path = tmp_path / "model" / "file.bin"
@@ -196,7 +204,7 @@ async def test_download_file_recovers_from_stall(tmp_path):
assert success is True
assert Path(result_path).read_bytes() == payload
assert downloader._session._get_calls == 2
assert _session(downloader)._get_calls == 2
assert not Path(str(target_path) + ".part").exists()
@@ -224,8 +232,8 @@ async def test_download_file_resumes_after_incomplete_integrity_check(tmp_path):
assert success is True
assert Path(result_path).read_bytes() == b"abcdef"
assert downloader._session._get_calls == 2
assert downloader._session.requests[1]["headers"]["Range"] == "bytes=3-"
assert _session(downloader)._get_calls == 2
assert _session(downloader).requests[1]["headers"]["Range"] == "bytes=3-"
assert not Path(str(target_path) + ".part").exists()
@@ -261,6 +269,6 @@ async def test_download_file_retries_redirected_url_when_range_not_honored(tmp_p
assert success is True
assert Path(result_path).read_bytes() == b"abcdef"
assert first_response.released is True
assert downloader._session.requests[0]["headers"]["Range"] == "bytes=3-"
assert downloader._session.requests[1]["url"] == redirected_url
assert downloader._session.requests[1]["headers"]["Range"] == "bytes=3-"
assert _session(downloader).requests[0]["headers"]["Range"] == "bytes=3-"
assert _session(downloader).requests[1]["url"] == redirected_url
assert _session(downloader).requests[1]["headers"]["Range"] == "bytes=3-"
@@ -54,7 +54,7 @@ async def test_cleanup_moves_empty_and_orphaned(tmp_path, monkeypatch):
result = await service.cleanup_example_image_folders()
deleted_bucket = Path(result['deleted_root'])
deleted_bucket = Path(str(result['deleted_root']))
assert result['success'] is True
assert result['moved_total'] == 2
assert not empty_folder.exists()
@@ -4,6 +4,7 @@ import asyncio
import json
from pathlib import Path
from types import SimpleNamespace
from typing import Any
import pytest
@@ -15,18 +16,18 @@ class RecordingWebSocketManager:
"""Collects broadcast payloads for assertions."""
def __init__(self) -> None:
self.payloads: list[dict] = []
self.payloads: list[dict[str, Any]] = []
async def broadcast(self, payload: dict) -> None:
async def broadcast(self, payload: dict[str, Any]) -> None:
self.payloads.append(payload)
class StubScanner:
"""Scanner double returning predetermined cache contents."""
def __init__(self, models: list[dict]) -> None:
def __init__(self, models: list[dict[str, Any]]) -> None:
self._cache = SimpleNamespace(raw_data=models)
self.sync_calls: list[tuple[str, dict]] = []
self.sync_calls: list[tuple[str, dict[str, Any]]] = []
async def get_cached_data(self):
return self._cache
@@ -39,7 +40,7 @@ class StubScanner:
break
return True
async def sync_cache_from_metadata(self, file_path: str, metadata: dict) -> bool:
async def sync_cache_from_metadata(self, file_path: str, metadata: dict[str, Any]) -> bool:
self.sync_calls.append((file_path, metadata))
for index, model in enumerate(self._cache.raw_data):
if model.get("file_path") == metadata.get("file_path"):
@@ -511,7 +512,7 @@ async def test_not_found_example_images_are_cleaned(
missing_url = "https://example.com/missing.png"
valid_url = "https://example.com/valid.png"
model_metadata = {
model_metadata: dict[str, Any] = {
"sha256": model_hash,
"model_name": "Missing Example",
"file_path": str(model_path),
+4 -3
View File
@@ -2,17 +2,18 @@ import asyncio
import json
import pytest
from pathlib import Path
from typing import Any
from py.services.settings_manager import get_settings_manager
from py.utils import example_images_download_manager as download_module
class RecordingWebSocketManager:
def __init__(self) -> None:
self.payloads: list[dict] = []
async def broadcast(self, payload: dict) -> None:
self.payloads: list[dict[str, Any]] = []
async def broadcast(self, payload: dict[str, Any]) -> None:
self.payloads.append(payload)
class StubScanner:
def __init__(self, models: list[dict]) -> None:
def __init__(self, models: list[dict[str, Any]]) -> None:
self.raw_data = models
async def get_cached_data(self):
class Cache:
+9 -1
View File
@@ -1,16 +1,24 @@
"""Tests for license-based filtering functionality."""
import pytest
from typing import Any
from unittest.mock import Mock, AsyncMock
from py.services.base_model_service import BaseModelService
from py.utils.civitai_utils import build_license_flags
from py.utils.models import BaseModelMetadata
class DummyModelService(BaseModelService):
"""Dummy implementation of BaseModelService for testing."""
def __init__(self):
super().__init__(
model_type="test",
scanner=Mock(),
metadata_class=BaseModelMetadata,
settings_provider=Mock(),
)
# Mock the required attributes
self.model_type = "test"
self.scanner = Mock()
@@ -28,7 +36,7 @@ class DummyModelService(BaseModelService):
self.scanner.get_cached_data = mock_get_cached_data
async def format_response(self, model_data: dict) -> dict:
async def format_response(self, model_data: dict[str, Any]) -> dict[str, Any]:
"""Required abstract method implementation."""
return model_data
@@ -1,10 +1,12 @@
"""Integration tests for license-based filtering in BaseModelService."""
import pytest
from typing import Any
from unittest.mock import Mock, AsyncMock
from py.services.base_model_service import BaseModelService
from py.utils.civitai_utils import build_license_flags
from py.utils.models import BaseModelMetadata
from py.services.model_query import ModelCacheRepository, ModelFilterSet, SearchStrategy, SortParams
@@ -12,6 +14,12 @@ class DummyModelService(BaseModelService):
"""Dummy implementation of BaseModelService for testing."""
def __init__(self):
super().__init__(
model_type="test",
scanner=Mock(),
metadata_class=BaseModelMetadata,
settings_provider=Mock(),
)
# Mock the required attributes
self.model_type = "test"
self.scanner = Mock()
@@ -33,7 +41,7 @@ class DummyModelService(BaseModelService):
self.scanner.get_cached_data = mock_get_cached_data
async def format_response(self, model_data: dict) -> dict:
async def format_response(self, model_data: dict[str, Any]) -> dict[str, Any]:
"""Required abstract method implementation."""
return model_data
+3
View File
@@ -57,6 +57,9 @@ class MockSession:
def __init__(self, response):
self._response = response
self.closed = False
self.last_url = None
self.last_json = None
self.last_headers = None
def post(self, url, json=None, headers=None):
self.last_url = url
+3 -2
View File
@@ -1,4 +1,5 @@
from types import SimpleNamespace
from typing import Optional
from unittest.mock import AsyncMock
import pytest
@@ -21,13 +22,13 @@ class DummyProvider(ModelMetadataProvider):
async def get_model_versions_bulk(self, model_ids):
return None
async def get_model_version(self, model_id: int = None, version_id: int = None):
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None):
return None
async def get_model_version_info(self, version_id: str):
return None, None
async def get_user_models(self, username: str):
async def get_user_models(self, username: str, cursor: Optional[str] = None):
return None
+7 -7
View File
@@ -10,7 +10,7 @@ from py.services.metadata_sync_service import MetadataSyncService
class DummySettings:
def __init__(self, values: dict | None = None) -> None:
def __init__(self, values: dict[str, Any] | None = None) -> None:
self._values = values or {}
def get(self, key: str, default=None):
@@ -19,7 +19,7 @@ class DummySettings:
def build_service(
*,
settings_values: dict | None = None,
settings_values: dict[str, Any] | None = None,
default_provider: SimpleNamespace | None = None,
provider_selector: AsyncMock | None = None,
):
@@ -43,7 +43,7 @@ def build_service(
service = MetadataSyncService(
metadata_manager=metadata_manager,
preview_service=preview_service,
settings=settings,
settings=settings, # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=default_provider_factory,
metadata_provider_selector=provider_selector,
)
@@ -194,7 +194,7 @@ async def test_fetch_and_update_model_success_updates_cache(tmp_path):
helpers.metadata_manager.hydrate_model_data.side_effect = hydrate
model_data = {
model_data: Dict[str, Any] = {
"model_name": "Local",
"folder": "root",
"file_path": str(model_path),
@@ -288,9 +288,9 @@ async def test_fetch_and_update_model_handles_missing_remote_metadata(tmp_path):
helpers.metadata_manager.hydrate_model_data.side_effect = hydrate
model_data = {
model_data: Dict[str, Any] = {
"model_name": "Local",
"folder": "sub",
"folder": "root",
"file_path": str(model_path),
}
@@ -659,7 +659,7 @@ async def test_fetch_and_update_model_does_not_overwrite_api_metadata_with_archi
helpers.default_provider.get_model_by_hash.return_value = (civarchive_payload, None)
model_path = tmp_path / "model.safetensors"
model_data = {
model_data: Dict[str, Any] = {
"model_name": "High Quality",
"metadata_source": "civitai_api",
"civitai": existing_civitai,
+40 -23
View File
@@ -1,6 +1,7 @@
import json
import os
from pathlib import Path
from typing import Any, Dict, cast
import pytest
@@ -9,6 +10,11 @@ from py.utils.metadata_manager import MetadataManager
from py.utils.models import LoraMetadata
async def _empty_metadata_loader(path: str) -> Dict[str, object]:
"""Default metadata loader for tests: return an empty payload."""
return {}
class ScannerWithRoots:
def __init__(self, roots):
self._roots = list(roots)
@@ -117,7 +123,7 @@ async def test_delete_model_rejects_path_outside_roots(tmp_path: Path):
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({"civitai": {"modelId": 1}}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
# Path within root should work (model file exists)
result = await service.delete_model(str(model))
@@ -133,7 +139,7 @@ async def test_delete_model_rejects_path_outside_roots(tmp_path: Path):
service2 = ModelLifecycleService(
scanner=scanner2,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="outside configured library"):
await service2.delete_model(str(outside))
@@ -148,7 +154,7 @@ async def test_rename_model_rejects_path_outside_roots(tmp_path: Path):
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
outside = tmp_path / "outside.safetensors"
outside.write_bytes(b"data")
@@ -170,7 +176,7 @@ async def test_bulk_delete_rejects_any_path_outside_roots(tmp_path: Path):
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="outside configured library"):
await service.bulk_delete_models([str(model_ok), str(outside)])
@@ -209,7 +215,7 @@ class VersionAwareScanner:
continue
candidate = civitai.get("modelId")
try:
normalized = int(candidate)
normalized = int(cast(Any, candidate))
except (TypeError, ValueError):
continue
if normalized != model_id:
@@ -297,6 +303,7 @@ async def test_rename_model_preserves_compound_extensions(tmp_path: Path):
assert expected_main.exists()
assert not model_path.exists()
assert isinstance(result["new_file_path"], str)
assert result["new_file_path"].endswith(f"{new_name}.safetensors")
assert expected_preview.exists()
assert not preview_path.exists()
@@ -343,7 +350,7 @@ async def test_delete_model_updates_update_service(tmp_path: Path):
scanner=scanner,
metadata_manager=metadata_manager,
metadata_loader=metadata_loader,
update_service=update_service,
update_service=update_service, # pyright: ignore[reportArgumentType]
)
result = await service.delete_model(model_path.as_posix())
@@ -396,6 +403,7 @@ async def test_rename_model_preserves_extension(tmp_path: Path):
assert expected_main.exists()
assert not model_path.exists()
assert isinstance(result["new_file_path"], str)
assert result["new_file_path"].endswith(f"{new_name}{old_extension}")
assert expected_preview.exists()
assert not preview_path.exists()
@@ -448,7 +456,12 @@ async def test_rename_model_with_dotted_basename(tmp_path: Path):
expected_main = tmp_path / f"{new_name}{old_extension}"
assert expected_main.exists()
assert result["new_file_path"] == expected_main.as_posix()
assert any(p.endswith(f"{new_name}{old_extension}") for p in result["renamed_files"])
renamed_files = result["renamed_files"]
assert isinstance(renamed_files, list)
assert any(
isinstance(p, str) and p.endswith(f"{new_name}{old_extension}")
for p in renamed_files
)
saved_metadata = json.loads((tmp_path / f"{new_name}.metadata.json").read_text())
assert saved_metadata["file_name"] == new_name
@@ -490,7 +503,9 @@ async def test_delete_model_removes_gguf_file(tmp_path: Path):
assert not model_path.exists()
assert not metadata_path.exists()
assert not preview_path.exists()
assert any(item.endswith("model.gguf") for item in result["deleted_files"])
deleted_files = result["deleted_files"]
assert isinstance(deleted_files, list)
assert any(isinstance(item, str) and item.endswith("model.gguf") for item in deleted_files)
# =============================================================================
@@ -531,10 +546,10 @@ async def test_exclude_model_marks_as_excluded(tmp_path: Path):
saved_metadata = []
class SavingMetadataManager:
async def save_metadata(self, path: str, metadata: dict):
async def save_metadata(self, path: str, metadata: Dict[str, Any]):
saved_metadata.append((path, metadata.copy()))
async def metadata_loader(path: str):
async def metadata_loader(path: str) -> Dict[str, Any]:
return metadata_payload.copy()
service = ModelLifecycleService(
@@ -546,6 +561,7 @@ async def test_exclude_model_marks_as_excluded(tmp_path: Path):
result = await service.exclude_model(str(model_path))
assert result["success"] is True
assert isinstance(result["message"], str)
assert "excluded" in result["message"].lower()
assert saved_metadata[0][1]["exclude"] is True
assert str(model_path) in scanner._excluded_models
@@ -581,7 +597,7 @@ async def test_exclude_model_updates_tag_counts(tmp_path: Path):
scanner = TagCountScanner(raw_data)
class DummyMetadataManagerLocal:
async def save_metadata(self, path: str, metadata: dict):
async def save_metadata(self, path: str, metadata: Dict[str, Any]):
pass
async def metadata_loader(path: str):
@@ -607,7 +623,7 @@ async def test_exclude_model_empty_path_raises_error():
service = ModelLifecycleService(
scanner=VersionAwareScanner([]),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="Model path is required"):
@@ -645,7 +661,7 @@ async def test_unexclude_model_restores_cache_entry(tmp_path: Path):
saved_metadata = []
class SavingMetadataManager:
async def save_metadata(self, path: str, metadata: dict):
async def save_metadata(self, path: str, metadata: Dict[str, Any]):
saved_metadata.append((path, metadata.copy()))
await MetadataManager.save_metadata(path, metadata)
@@ -663,6 +679,7 @@ async def test_unexclude_model_restores_cache_entry(tmp_path: Path):
result = await service.unexclude_model(str(model_path))
assert result["success"] is True
assert isinstance(result["message"], str)
assert "restored" in result["message"].lower()
assert scanner._excluded_models == []
assert saved_metadata[0][1]["exclude"] is False
@@ -700,7 +717,7 @@ async def test_bulk_delete_models_deletes_multiple_files(tmp_path: Path):
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
result = await service.bulk_delete_models(file_paths)
@@ -716,7 +733,7 @@ async def test_bulk_delete_models_empty_list_raises_error():
service = ModelLifecycleService(
scanner=VersionAwareScanner([]),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="No file paths provided"):
@@ -734,7 +751,7 @@ async def test_delete_model_empty_path_raises_error():
service = ModelLifecycleService(
scanner=VersionAwareScanner([]),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="Model path is required"):
@@ -747,7 +764,7 @@ async def test_rename_model_empty_path_raises_error():
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="required"):
@@ -763,7 +780,7 @@ async def test_rename_model_empty_name_raises_error(tmp_path: Path):
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="required"):
@@ -779,7 +796,7 @@ async def test_rename_model_invalid_characters_raises_error(tmp_path: Path):
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
invalid_names = [
@@ -817,7 +834,7 @@ async def test_rename_model_existing_file_raises_error(tmp_path: Path):
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
with pytest.raises(ValueError, match="already exists"):
@@ -837,7 +854,7 @@ async def test_extract_model_id_from_civitai_payload():
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
# Test civitai.modelId
@@ -863,7 +880,7 @@ async def test_extract_model_id_returns_none_for_invalid_payload():
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
assert service._extract_model_id_from_payload({}) is None
@@ -879,7 +896,7 @@ async def test_extract_model_id_handles_string_values():
service = ModelLifecycleService(
scanner=DummyScanner(),
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
metadata_loader=_empty_metadata_loader,
)
payload = {"civitai": {"modelId": "54321"}}
+40 -3
View File
@@ -6,11 +6,12 @@ from py.services import model_metadata_provider as provider_module
from py.services.errors import RateLimitError
from py.services.model_metadata_provider import (
FallbackMetadataProvider,
ModelMetadataProvider,
RateLimitRetryingProvider,
)
class RateLimitThenSuccessProvider:
class RateLimitThenSuccessProvider(ModelMetadataProvider):
def __init__(self) -> None:
self.calls = 0
@@ -20,8 +21,20 @@ class RateLimitThenSuccessProvider:
raise RateLimitError("limited", retry_after=1.0)
return {"id": "ok"}, None
async def get_model_versions(self, model_id: str):
return None
class AlwaysRateLimitedProvider:
async def get_model_version(self, model_id=None, version_id=None):
return None
async def get_model_version_info(self, version_id: str):
return None, None
async def get_user_models(self, username: str, cursor=None):
return None
class AlwaysRateLimitedProvider(ModelMetadataProvider):
def __init__(self) -> None:
self.calls = 0
@@ -29,8 +42,20 @@ class AlwaysRateLimitedProvider:
self.calls += 1
raise RateLimitError("limited")
async def get_model_versions(self, model_id: str):
return None
class TrackingProvider:
async def get_model_version(self, model_id=None, version_id=None):
return None
async def get_model_version_info(self, version_id: str):
return None, None
async def get_user_models(self, username: str, cursor=None):
return None
class TrackingProvider(ModelMetadataProvider):
def __init__(self) -> None:
self.calls = 0
@@ -38,6 +63,18 @@ class TrackingProvider:
self.calls += 1
return {"id": "secondary"}, None
async def get_model_versions(self, model_id: str):
return None
async def get_model_version(self, model_id=None, version_id=None):
return None
async def get_model_version_info(self, version_id: str):
return None, None
async def get_user_models(self, username: str, cursor=None):
return None
@pytest.mark.asyncio
async def test_fallback_retries_same_provider_on_rate_limit(monkeypatch):
+2 -2
View File
@@ -84,11 +84,11 @@ class TestResolveSubType:
def test_none_entry_returns_default(self):
"""None entry should return default."""
assert resolve_sub_type(None) == "LORA"
assert resolve_sub_type(None) == "LORA" # pyright: ignore[reportArgumentType]
def test_non_mapping_returns_default(self):
"""Non-mapping entry should return default."""
assert resolve_sub_type("invalid") == "LORA"
assert resolve_sub_type("invalid") == "LORA" # pyright: ignore[reportArgumentType]
class TestModelFilterSetWithSubType:
+8 -8
View File
@@ -2,7 +2,7 @@ import asyncio
import os
import sqlite3
from pathlib import Path
from typing import List
from typing import Any, Dict, List, Optional
from types import MethodType
import pytest
@@ -18,9 +18,9 @@ from py.utils.models import BaseModelMetadata
class RecordingWebSocketManager:
def __init__(self) -> None:
self.payloads: List[dict] = []
self.payloads: List[Dict[str, Any]] = []
async def broadcast_init_progress(self, payload: dict) -> None:
async def broadcast_init_progress(self, payload: Dict[str, Any]) -> None:
self.payloads.append(payload)
@@ -48,7 +48,7 @@ class DummyScanner(ModelScanner):
*,
hash_index: ModelHashIndex | None = None,
excluded_models: List[str] | None = None,
) -> dict:
) -> Optional[Dict[str, Any]]:
hash_index = hash_index or self._hash_index
excluded_models = excluded_models if excluded_models is not None else self._excluded_models
@@ -483,7 +483,7 @@ async def test_version_index_tracks_version_ids(tmp_path: Path):
assert cache.version_index[202]['file_path'] == second_path
assert await scanner.check_model_version_exists(101) is True
assert await scanner.check_model_version_exists('202') is True
assert await scanner.check_model_version_exists('202') is True # pyright: ignore[reportArgumentType]
assert await scanner.check_model_version_exists(999) is False
removed = await scanner._batch_update_cache_for_deleted_models([first_path])
@@ -530,7 +530,7 @@ async def test_reconcile_cache_applies_adjust_cached_entry(tmp_path: Path):
applied: List[str] = []
def _adjust(self, entry: dict) -> dict:
def _adjust(self, entry: Dict[str, Any]) -> Dict[str, Any]:
applied.append(entry["file_path"])
entry["custom_field"] = "adjusted"
return entry
@@ -607,7 +607,7 @@ async def test_reconcile_cache_removes_duplicate_alias_when_same_real_file_seen_
scanner = MultiRootDummyScanner([loras_root, extra_root])
await scanner._initialize_cache()
duplicate_entry = {
duplicate_entry: Dict[str, Any] = {
"file_path": _normalize_path(extra_root / "one.txt"),
"folder": "",
"sha256": "hash-one",
@@ -712,7 +712,7 @@ def test_cache_entries_differ_extra_key():
# ── sync_cache_from_metadata ─────────────────────────────────────────
def _make_cache_entry(**overrides) -> dict:
def _make_cache_entry(**overrides) -> Dict[str, Any]:
entry = {
"file_path": "/m/a.safetensors",
"model_name": "TestModel",
@@ -3,13 +3,19 @@ from types import SimpleNamespace
import pytest
from py.services.model_scanner import ModelScanner
from py.utils.models import BaseModelMetadata
class DummyScanner:
class DummyScanner(ModelScanner):
def __init__(self, raw_data):
super().__init__(
model_type="dummy",
model_class=BaseModelMetadata,
file_extensions={".safetensors"},
)
self._cache = SimpleNamespace(raw_data=raw_data)
async def get_cached_data(self):
async def get_cached_data(self, force_refresh: bool = False, rebuild_cache: bool = False):
return self._cache
@@ -391,6 +391,7 @@ async def test_update_in_library_versions_changes_update_state(tmp_path):
await service.update_in_library_versions("lora", 3, [31, 35])
record = await service.get_record("lora", 3)
assert record is not None
assert record.has_update() is False
+1 -1
View File
@@ -71,7 +71,7 @@ class StubLoraScanner:
def recipe_scanner(tmp_path, monkeypatch):
monkeypatch.setattr(config, "loras_roots", [str(tmp_path)])
stub = StubLoraScanner()
scanner = RecipeScanner(lora_scanner=stub)
scanner = RecipeScanner(lora_scanner=stub) # pyright: ignore[reportArgumentType]
return scanner
@pytest.mark.asyncio
@@ -1,4 +1,5 @@
from pathlib import Path
from typing import Any, Dict
import pytest
@@ -346,7 +347,7 @@ def test_update_single_model_update_hash(tmp_path: Path, monkeypatch):
# ── get_models_missing_autov3 ─────────────────────────────────────────
def _autov3_entry(file_path: str, sha256: str, autov3=None) -> dict:
def _autov3_entry(file_path: str, sha256: str, autov3=None) -> Dict[str, Any]:
"""Minimal model entry for the models table (autov3 tri-state preserved)."""
return {
'file_path': file_path,
+5 -5
View File
@@ -1,5 +1,5 @@
from pathlib import Path
from typing import Any
from typing import Any, Dict, List
import pytest
@@ -55,7 +55,7 @@ async def test_ensure_preview_prefers_rewritten_civitai_image(tmp_path):
exif_utils=exif_utils,
)
images = [
images: List[Dict[str, object]] = [
{
"url": "https://image.civitai.com/container/example/original=true/sample.jpeg",
"type": "image",
@@ -115,7 +115,7 @@ async def test_ensure_preview_falls_back_to_webp_when_rewrite_fails(tmp_path):
exif_utils=exif_utils,
)
images = [
images: List[Dict[str, object]] = [
{
"url": "https://image.civitai.com/container/example/original=true/sample.png",
"type": "image",
@@ -165,7 +165,7 @@ async def test_ensure_preview_rewrites_civitai_video(tmp_path):
exif_utils=RecordingExifUtils(),
)
images = [
images: List[Dict[str, object]] = [
{
"url": "https://image.civitai.com/container/example/original=true/sample.mp4",
"type": "video",
@@ -227,7 +227,7 @@ async def test_ensure_preview_respects_blur_setting(monkeypatch, tmp_path):
exif_utils=RecordingExifUtils(),
)
images = [
images: List[Dict[str, object]] = [
{
"url": "https://image.civitai.com/container/example/original=true/nsfw.jpeg",
"type": "image",
+3 -1
View File
@@ -1,4 +1,6 @@
import json
from typing import Any, Dict
import pytest
from py.recipes.parsers.recipe_format import RecipeFormatParser
@@ -83,7 +85,7 @@ async def test_recipe_format_parser_marks_lora_in_library_by_version(monkeypatch
fake_metadata_provider,
)
cached_entry = {
cached_entry: Dict[str, Any] = {
"file_path": "/loras/moriimee.safetensors",
"file_name": "MoriiMee Gothic Niji | LoRA Style",
"size": 4096,
+2 -1
View File
@@ -1,5 +1,6 @@
import pytest
import asyncio
from typing import Any, Dict
from unittest.mock import AsyncMock, MagicMock
from py.services.recipe_scanner import RecipeScanner
from types import SimpleNamespace
@@ -259,7 +260,7 @@ async def test_repair_all_recipes_strips_runtime_fields(setup_scanner):
recipe_scanner, mock_civitai_client, mock_metadata_provider = setup_scanner
# Recipe with runtime fields
recipe = {
recipe: Dict[str, Any] = {
"id": "r1",
"title": "Cleanup Test",
"checkpoint": {
+8 -7
View File
@@ -3,6 +3,7 @@ import json
import os
from pathlib import Path
from types import SimpleNamespace
from typing import Any, Dict
import pytest
@@ -24,7 +25,7 @@ class StubLoraScanner:
def __init__(self) -> None:
self._hash_index = StubHashIndex()
self._hash_meta: dict[str, dict[str, str]] = {}
self._models_by_name: dict[str, dict] = {}
self._models_by_name: dict[str, Dict[str, Any]] = {}
self._cache = SimpleNamespace(raw_data=[], version_index={})
async def get_cached_data(self):
@@ -44,7 +45,7 @@ class StubLoraScanner:
async def get_model_info_by_name(self, name: str):
return self._models_by_name.get(name)
def register_model(self, name: str, info: dict) -> None:
def register_model(self, name: str, info: Dict[str, Any]) -> None:
self._models_by_name[name] = info
hash_value = (info.get("sha256") or "").lower()
version_id = info.get("civitai", {}).get("id")
@@ -76,7 +77,7 @@ def recipe_scanner(tmp_path: Path, monkeypatch):
settings_manager_module.reset_settings_manager()
monkeypatch.setattr(config, "loras_roots", [str(tmp_path)])
stub = StubLoraScanner()
scanner = RecipeScanner(lora_scanner=stub)
scanner = RecipeScanner(lora_scanner=stub) # pyright: ignore[reportArgumentType]
async def _init():
await scanner.refresh_cache(force=True)
@@ -107,7 +108,7 @@ def test_recipes_dir_uses_custom_settings_path(tmp_path: Path, monkeypatch):
manager = settings_manager_module.get_settings_manager()
manager.set("recipes_path", str(custom_recipes))
scanner = RecipeScanner(lora_scanner=StubLoraScanner())
scanner = RecipeScanner(lora_scanner=StubLoraScanner()) # pyright: ignore[reportArgumentType]
resolved = scanner.recipes_dir
assert resolved == str((tmp_path / "custom_recipes").resolve())
@@ -123,7 +124,7 @@ def test_recipes_dir_falls_back_to_first_lora_root(tmp_path: Path, monkeypatch):
monkeypatch.setattr(config, "loras_roots", [str(tmp_path / "alpha")])
scanner = RecipeScanner(lora_scanner=StubLoraScanner())
scanner = RecipeScanner(lora_scanner=StubLoraScanner()) # pyright: ignore[reportArgumentType]
resolved = scanner.recipes_dir
assert resolved == str(tmp_path / "alpha" / "recipes")
@@ -719,7 +720,7 @@ async def test_initialize_waits_for_lora_scanner(monkeypatch):
ready_flag.set()
lora_scanner = StubLoraScanner()
scanner = RecipeScanner(lora_scanner=lora_scanner)
scanner = RecipeScanner(lora_scanner=lora_scanner) # pyright: ignore[reportArgumentType]
await scanner.initialize_in_background()
@@ -736,7 +737,7 @@ async def test_invalid_model_version_marked_deleted_and_not_retried(
recipes_dir = Path(config.loras_roots[0]) / "recipes"
recipes_dir.mkdir(parents=True, exist_ok=True)
recipe = {
recipe: Dict[str, Any] = {
"id": "invalid-version",
"file_path": str(recipes_dir / "invalid-version.webp"),
"title": "Invalid",
+10 -4
View File
@@ -4,8 +4,9 @@ import os
from io import BytesIO
from pathlib import Path
from types import SimpleNamespace
from typing import Any, Dict
import piexif
import piexif # pyright: ignore[reportMissingTypeStubs]
import pytest
from PIL import Image, PngImagePlugin
@@ -463,12 +464,17 @@ async def test_save_recipe_preserves_workflow_when_png_is_converted_to_webp(tmp_
image_path = Path(result.payload["image_path"])
exif_dict = piexif.load(str(image_path))
assert exif_dict is not None
exif_0th = exif_dict["0th"]
assert exif_0th is not None
assert (
exif_dict["0th"][piexif.ImageIFD.ImageDescription].decode("utf-8")
exif_0th[piexif.ImageIFD.ImageDescription].decode("utf-8")
== 'Workflow:{"nodes":[{"id":1}]}'
)
user_comment = exif_dict["Exif"][piexif.ExifIFD.UserComment]
exif_section = exif_dict["Exif"]
assert exif_section is not None
user_comment = exif_section[piexif.ExifIFD.UserComment]
decoded_comment = user_comment[8:].decode("utf-16be")
assert "prompt text" in decoded_comment
assert "Recipe metadata:" in decoded_comment
@@ -705,7 +711,7 @@ async def test_move_recipe_updates_paths(tmp_path):
matches = list(Path(self.recipes_dir).rglob(f"{target_id}.recipe.json"))
return str(matches[0]) if matches else None
async def update_recipe_metadata(self, target_id: str, metadata: dict):
async def update_recipe_metadata(self, target_id: str, metadata: Dict[str, Any]):
if target_id != recipe_id:
return False
self.recipe.update(metadata)
+14 -9
View File
@@ -2,6 +2,7 @@ import pytest
from py.services.model_query import ModelFilterSet, FilterCriteria
from py.services.recipe_scanner import RecipeScanner
from types import SimpleNamespace
from typing import Any, cast
# Mock settings
@@ -193,9 +194,9 @@ async def test_recipe_scanner_root_recursive_true():
async def get_cached_data(self):
return SimpleNamespace(raw_data=[])
scanner = RecipeScanner(lora_scanner=StubLoraScanner())
scanner = RecipeScanner(lora_scanner=StubLoraScanner()) # pyright: ignore[reportArgumentType]
# Manually populate cache for testing get_paginated_data logic
scanner._cache = SimpleNamespace(
scanner._cache = cast(Any, SimpleNamespace(
raw_data=[
{
"id": "r1",
@@ -234,13 +235,15 @@ async def test_recipe_scanner_root_recursive_true():
],
sorted_by_name=[],
version_index={},
)
))
result = await scanner.get_paginated_data(
page=1, page_size=10, folder="", recursive=True
)
assert len(result["items"]) == 2
items = result["items"]
assert isinstance(items, list)
assert len(items) == 2
@pytest.mark.asyncio
@@ -250,8 +253,8 @@ async def test_recipe_scanner_root_recursive_false():
async def get_cached_data(self):
return SimpleNamespace(raw_data=[])
scanner = RecipeScanner(lora_scanner=StubLoraScanner())
scanner._cache = SimpleNamespace(
scanner = RecipeScanner(lora_scanner=StubLoraScanner()) # pyright: ignore[reportArgumentType]
scanner._cache = cast(Any, SimpleNamespace(
raw_data=[
{
"id": "r1",
@@ -290,11 +293,13 @@ async def test_recipe_scanner_root_recursive_false():
],
sorted_by_name=[],
version_index={},
)
))
result = await scanner.get_paginated_data(
page=1, page_size=10, folder="", recursive=False
)
assert len(result["items"]) == 1
assert result["items"][0]["id"] == "r1"
items = result["items"]
assert isinstance(items, list)
assert len(items) == 1
assert items[0]["id"] == "r1"
@@ -51,10 +51,10 @@ class DummyProvider:
def __init__(self, payload: Dict[str, Any]) -> None:
self.payload = payload
async def get_model_by_hash(self, sha256: str):
async def get_model_by_hash(self, model_hash: str):
return self.payload, None
async def get_model_version(self, model_id: int, model_version_id: int | None):
async def get_model_version(self, model_id: Any = None, version_id: Any = None):
return self.payload
@@ -77,7 +77,7 @@ def test_metadata_sync_merges_remote_fields(tmp_path: Path) -> None:
service = MetadataSyncService(
metadata_manager=manager,
preview_service=preview,
settings=DummySettings(),
settings=DummySettings(), # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=lambda: asyncio.sleep(0, result=provider),
metadata_provider_selector=lambda _name=None: asyncio.sleep(0, result=provider),
)
@@ -112,7 +112,7 @@ def test_metadata_sync_fetch_and_update_updates_cache(tmp_path: Path) -> None:
service = MetadataSyncService(
metadata_manager=manager,
preview_service=preview,
settings=DummySettings(),
settings=DummySettings(), # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=lambda: asyncio.sleep(0, result=provider),
metadata_provider_selector=lambda _name=None: asyncio.sleep(0, result=provider),
)
+1 -1
View File
@@ -52,7 +52,7 @@ async def test_lazy_loaded_scanners(monkeypatch, method_name, module_path, class
async def test_lazy_loaded_websocket_manager(monkeypatch):
fake_manager = object()
module = types.ModuleType("py.services.websocket_manager")
module.ws_manager = fake_manager
setattr(module, "ws_manager", fake_manager)
monkeypatch.setitem(sys.modules, "py.services.websocket_manager", module)
first = await ServiceRegistry.get_websocket_manager()
+2 -2
View File
@@ -482,7 +482,7 @@ def test_model_name_display_setting_notifies_scanners(tmp_path, monkeypatch):
manager = _create_manager_with_settings(tmp_path, monkeypatch, initial)
loop = asyncio.new_event_loop()
loop._thread_id = 1
setattr(loop, "_thread_id", 1)
class DummyScanner:
def __init__(self):
@@ -530,7 +530,7 @@ def test_model_name_display_setting_notifies_scanners(tmp_path, monkeypatch):
assert dummy_scanner.calls == ["file_name"]
assert dispatched_loops == [dummy_scanner.loop]
finally:
loop._thread_id = None
setattr(loop, "_thread_id", None)
loop.close()
@@ -8,6 +8,8 @@ from py.recipes.parsers import SuiImageParamsParser
class TestSuiImageParamsParser:
"""Test cases for SuiImageParamsParser."""
parser: SuiImageParamsParser = SuiImageParamsParser()
def setup_method(self):
"""Set up test fixtures."""
self.parser = SuiImageParamsParser()
@@ -116,6 +118,7 @@ class TestSuiImageParamsParser:
result = await self.parser.parse_metadata(metadata_str)
loras = result.get('loras')
assert isinstance(loras, list)
assert len(loras) == 1
assert loras[0]['type'] == 'lora'
assert loras[0]['name'] == 'test_lora'
@@ -142,6 +145,7 @@ class TestSuiImageParamsParser:
result = await self.parser.parse_metadata(metadata_str)
loras = result.get('loras')
assert isinstance(loras, list)
assert len(loras) == 1
assert loras[0]['type'] == 'lora'
+25 -10
View File
@@ -1,11 +1,14 @@
import asyncio
import logging
from dataclasses import dataclass
from types import SimpleNamespace
from typing import Any, Dict, List, Optional
import pytest
from py.services.model_file_service import AutoOrganizeResult
from py.services.download_coordinator import DownloadCoordinator
from py.services.metadata_sync_service import MetadataSyncService
from py.services.model_file_service import AutoOrganizeResult, ModelFileService
from py.services.use_cases import (
AutoOrganizeInProgressError,
AutoOrganizeUseCase,
@@ -26,6 +29,7 @@ from py.utils.example_images_download_manager import (
)
from py.utils.example_images_processor import (
ExampleImagesImportError,
ExampleImagesProcessor,
ExampleImagesValidationError,
)
from py.utils.metadata_manager import MetadataManager
@@ -44,13 +48,13 @@ class StubLockProvider:
return self._lock
class StubFileService:
class StubFileService(ModelFileService):
def __init__(self) -> None:
super().__init__(scanner=None, model_type="lora")
self.calls: List[Dict[str, Any]] = []
async def auto_organize_models(
self,
*,
file_paths: Optional[List[str]] = None,
progress_callback=None,
exclusion_patterns=None,
@@ -65,8 +69,15 @@ class StubFileService:
return result
class StubMetadataSync:
class StubMetadataSync(MetadataSyncService):
def __init__(self) -> None:
super().__init__(
metadata_manager=object(),
preview_service=object(),
settings=StubSettings(), # pyright: ignore[reportArgumentType]
default_metadata_provider_factory=lambda: asyncio.sleep(0, result=None), # pyright: ignore[reportArgumentType]
metadata_provider_selector=lambda _name=None: asyncio.sleep(0, result=None), # pyright: ignore[reportArgumentType]
)
self.calls: List[Dict[str, Any]] = []
async def fetch_and_update_model(self, **kwargs: Any):
@@ -94,8 +105,12 @@ class ProgressCollector:
self.events.append(payload)
class StubDownloadCoordinator:
class StubDownloadCoordinator(DownloadCoordinator):
def __init__(self, *, error: Optional[str] = None) -> None:
super().__init__(
ws_manager=SimpleNamespace(generate_download_id=lambda: "abc123"),
download_manager_factory=lambda: asyncio.sleep(0, result=None),
)
self.error = error
self.payloads: List[Dict[str, Any]] = []
@@ -125,13 +140,13 @@ class StubExampleImagesDownloadManager:
return {"success": True, "message": "ok"}
class StubExampleImagesProcessor:
class StubExampleImagesProcessor(ExampleImagesProcessor):
def __init__(self) -> None:
self.calls: List[Dict[str, Any]] = []
self.error: Optional[str] = None
self.response: Dict[str, Any] = {"success": True}
async def import_images(self, model_hash: str, files: List[str]) -> Dict[str, Any]:
async def import_images(self, model_hash: str, files: List[str]) -> Dict[str, Any]: # pyright: ignore[reportIncompatibleMethodOverride]
self.calls.append({"model_hash": model_hash, "files": files})
if self.error == "validation":
raise ExampleImagesValidationError("missing")
@@ -464,7 +479,7 @@ async def test_import_example_images_use_case_delegates() -> None:
use_case = ImportExampleImagesUseCase(processor=processor)
request = DummyJsonRequest({"model_hash": "abc", "file_paths": ["/tmp/file"]})
result = await use_case.execute(request)
result = await use_case.execute(request) # pyright: ignore[reportArgumentType]
assert processor.calls == [{"model_hash": "abc", "files": ["/tmp/file"]}]
assert result == {"success": True}
@@ -477,7 +492,7 @@ async def test_import_example_images_use_case_maps_validation_error() -> None:
request = DummyJsonRequest({"model_hash": None, "file_paths": []})
with pytest.raises(ImportExampleImagesValidationError):
await use_case.execute(request)
await use_case.execute(request) # pyright: ignore[reportArgumentType]
async def test_import_example_images_use_case_propagates_generic_error() -> None:
@@ -487,4 +502,4 @@ async def test_import_example_images_use_case_propagates_generic_error() -> None
request = DummyJsonRequest({"model_hash": "abc", "file_paths": ["/tmp/file"]})
with pytest.raises(ExampleImagesImportError):
await use_case.execute(request)
await use_case.execute(request) # pyright: ignore[reportArgumentType]
+7 -7
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import json
from pathlib import Path
from types import ModuleType, SimpleNamespace
from typing import List, Tuple
from typing import Any, Generator, List, Tuple
import pytest
from aiohttp import web
@@ -13,11 +13,11 @@ from aiohttp import web
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
def standalone_module(monkeypatch) -> ModuleType:
def standalone_module(monkeypatch) -> Generator[ModuleType, None, None]:
"""Load the ``standalone`` module with a lightweight ``LoraManager`` stub."""
import importlib
@@ -41,7 +41,7 @@ def standalone_module(monkeypatch) -> ModuleType:
async def _cleanup(cls, app): # pragma: no cover - compatibility shim
return None
stub_module.LoraManager = _StubLoraManager
stub_module.LoraManager = _StubLoraManager # pyright: ignore[reportAttributeAccessIssue]
sys.modules["py.lora_manager"] = stub_module
module = importlib.import_module("standalone")
@@ -54,7 +54,7 @@ def standalone_module(monkeypatch) -> ModuleType:
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."""
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."""
app = web.Application()
route_calls: List[Tuple[str, dict]] = []
route_calls: List[Tuple[str, dict[str, Any]]] = []
app[ROUTE_CALLS_KEY] = route_calls
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 "/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 {
"/ws/fetch-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"))
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:
@@ -208,18 +208,18 @@ class TestAutoTagCategories:
re.compile(pattern, re.IGNORECASE)
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 "LOW" in MODE_TAGS
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 "T2V" in VIDEO_MODE_TAGS
assert "TI2V" in VIDEO_MODE_TAGS
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 "video" in DEFAULT_ENABLED_GROUPS
assert "speed" not in DEFAULT_ENABLED_GROUPS
+10 -2
View File
@@ -3,7 +3,7 @@
import json
import os
import tempfile
from typing import Dict, List
from typing import Any, Dict, List
import pytest
@@ -27,7 +27,7 @@ def temp_db_path():
@pytest.fixture
def sample_recipes() -> List[Dict]:
def sample_recipes() -> List[Dict[str, Any]]:
"""Create sample recipe data."""
return [
{
@@ -133,6 +133,7 @@ class TestPersistentRecipeCache:
# Load and verify
loaded = cache.load_cache()
assert loaded is not None
r1 = next(r for r in loaded.raw_data if r["id"] == "recipe-001")
assert r1["title"] == "Updated Title"
assert r1["favorite"] is False
@@ -147,6 +148,7 @@ class TestPersistentRecipeCache:
# Load and verify
loaded = cache.load_cache()
assert loaded is not None
assert len(loaded.raw_data) == 1
assert loaded.raw_data[0]["id"] == "recipe-002"
@@ -203,6 +205,7 @@ class TestPersistentRecipeCache:
cache.save_cache(recipes)
loaded = cache.load_cache()
assert loaded is not None
assert len(loaded.raw_data) == 1
assert loaded.raw_data[0]["id"] == "valid-001"
@@ -251,6 +254,7 @@ class TestPersistentRecipeCache:
cache.save_cache(recipes)
loaded = cache.load_cache()
assert loaded is not None
loras = loaded.raw_data[0]["loras"]
assert len(loras) == 2
assert loras[0]["modelVersionId"] == 12345
@@ -509,6 +513,7 @@ class TestPersistentRecipeCache:
cache.save_cache(sample_recipes)
loaded = cache.load_cache()
assert loaded is not None
assert loaded.image_id_map == {}
def test_image_id_map_survives_recipe_update(self, temp_db_path, sample_recipes):
@@ -522,6 +527,7 @@ class TestPersistentRecipeCache:
cache.update_recipe(updated)
loaded = cache.load_cache()
assert loaded is not None
assert loaded.image_id_map == {"123": "recipe-alpha"}
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"})
loaded = cache.load_cache()
assert loaded is not None
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):
@@ -542,4 +549,5 @@ class TestPersistentRecipeCache:
cache.save_image_id_map({"222": "new-only"})
loaded = cache.load_cache()
assert loaded is not None
assert loaded.image_id_map == {"222": "new-only"}
+2 -2
View File
@@ -2,7 +2,7 @@
import os
import tempfile
from typing import Dict, List
from typing import Any, Dict, List
import pytest
@@ -25,7 +25,7 @@ def temp_db_path():
@pytest.fixture
def sample_recipes() -> List[Dict]:
def sample_recipes() -> List[Dict[str, Any]]:
"""Create sample recipe data for FTS indexing."""
return [
{
+3 -2
View File
@@ -1,6 +1,7 @@
import importlib
import json
from pathlib import Path
from typing import Any
import pytest
@@ -21,12 +22,12 @@ def reset_settings(tmp_path, monkeypatch):
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:
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"
with example_path.open('r', encoding='utf-8') as handle:
return json.load(handle)
@@ -76,6 +76,7 @@ class TestRewritePreviewUrl:
for url in test_cases:
result, was_rewritten = rewrite_preview_url(url, "image")
assert was_rewritten is True
assert result is not None
assert "width=450,optimized=true" in result
def test_handles_urls_with_explicit_port(self):
@@ -83,6 +84,7 @@ class TestRewritePreviewUrl:
url = "https://image.civitai.com:443/checkpoints/original=true"
result, was_rewritten = rewrite_preview_url(url, "image")
assert was_rewritten is True
assert result is not None
assert "width=450,optimized=true" in result
# Port is preserved in the URL (this is acceptable behavior)
assert ":443" in result
@@ -104,6 +106,8 @@ class TestRewritePreviewUrl:
result2, was2 = rewrite_preview_url(url, "Video")
assert was1 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 result2
@@ -119,6 +123,7 @@ class TestRewritePreviewUrl:
url = "https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/abc123/original=true/12345.png"
result, was_rewritten = rewrite_preview_url(url, "image")
assert was_rewritten is True
assert result is not None
assert result.startswith(
"https://image.civitai.com/xG1nkqKTMzGDvpLrqFT7WA/abc123/"
)
@@ -129,6 +134,7 @@ class TestRewritePreviewUrl:
url = "https://image.civitai.com/original=true/test.png"
result, was_rewritten = rewrite_preview_url(url, None)
assert was_rewritten is True
assert result is not None
assert "transcode=true" not in result
assert "width=450,optimized=true" in result
@@ -2,7 +2,7 @@ from __future__ import annotations
import asyncio
import time
from typing import Any, Dict
from typing import Any, Dict, Generator
import pytest
@@ -19,7 +19,7 @@ class RecordingWebSocketManager:
@pytest.fixture(autouse=True)
def restore_settings() -> None:
def restore_settings() -> Generator[None, None, None]:
manager = get_settings_manager()
original = manager.settings.copy()
try:
@@ -45,7 +45,9 @@ async def test_start_download_requires_configured_path(
result = await manager.start_download({"auto_mode": 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(
+16 -11
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
import json
import os
import subprocess
from typing import Any, Dict
from typing import Any, Dict, Generator
import pytest
@@ -20,8 +20,13 @@ class JsonRequest:
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)
def restore_settings() -> None:
def restore_settings() -> Generator[None, None, None]:
manager = get_settings_manager()
original = manager.settings.copy()
try:
@@ -54,7 +59,7 @@ async def test_open_folder_requires_existing_model_directory(monkeypatch: pytest
request = JsonRequest({"model_hash": model_hash})
response = await ExampleImagesFileManager.open_folder(request)
body = json.loads(response.text)
body = _parse_json(response)
assert body["success"] is True
# 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})
response = await ExampleImagesFileManager.open_folder(request)
body = json.loads(response.text)
body = _parse_json(response)
assert response.status == 200
assert body == {
@@ -126,7 +131,7 @@ async def test_open_folder_returns_uri_mode_with_rendered_template(
request = JsonRequest({"model_hash": model_hash})
response = await ExampleImagesFileManager.open_folder(request)
body = json.loads(response.text)
body = _parse_json(response)
assert response.status == 200
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")
response = await ExampleImagesFileManager.open_folder(JsonRequest({"model_hash": model_hash}))
body = json.loads(response.text)
body = _parse_json(response)
assert response.status == 400
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})
response = await ExampleImagesFileManager.open_folder(request)
body = json.loads(response.text)
body = _parse_json(response)
assert response.status == 400
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})
response = await ExampleImagesFileManager.get_files(request)
body = json.loads(response.text)
body = _parse_json(response)
assert response.status == 200
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})
response = await ExampleImagesFileManager.has_images(request)
body = json.loads(response.text)
body = _parse_json(response)
assert body["has_images"] is True
empty_request = JsonRequest({}, {"model_hash": "missing"})
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
async def test_has_images_requires_model_hash() -> None:
response = await ExampleImagesFileManager.has_images(JsonRequest({}, {}))
body = json.loads(response.text)
body = _parse_json(response)
assert response.status == 400
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")
metadata_path = tmp_path / "preserve.metadata.json"
existing_payload = {
existing_payload: Dict[str, Any] = {
"model_name": "Example",
"file_path": str(model_file),
"civitai": {
@@ -200,7 +200,7 @@ async def test_update_metadata_from_local_examples_generates_entries(monkeypatch
model_dir = tmp_path / model_hash
model_dir.mkdir()
(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):
return True
@@ -4,7 +4,7 @@ import json
import os
from pathlib import Path
from types import SimpleNamespace
from typing import Any, Dict, Tuple
from typing import Any, Dict, Generator, Tuple
import pytest
@@ -15,7 +15,7 @@ from py.utils.example_images_paths import get_model_folder
@pytest.fixture(autouse=True)
def restore_settings() -> None:
def restore_settings() -> Generator[None, None, None]:
manager = get_settings_manager()
original = manager.settings.copy()
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))
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["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")
metadata_path = tmp_path / "keep.metadata.json"
existing_metadata = {
existing_metadata: Dict[str, Any] = {
"model_name": "Keep",
"file_path": str(model_file),
"civitai": {
@@ -313,6 +313,7 @@ async def test_delete_custom_image_preserves_existing_metadata(monkeypatch: pyte
)
assert response.status == 200
assert response.text is not None
body = json.loads(response.text)
assert body["success"] is True
assert body["custom_images"] == []
+11 -4
View File
@@ -1,6 +1,7 @@
import json
from typing import Any, Dict
import piexif
import piexif # pyright: ignore[reportMissingTypeStubs]
from PIL import Image, PngImagePlugin
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)
exif_dict = piexif.load(str(optimized_path))
assert exif_dict["0th"] is not None
assert (
exif_dict["0th"][piexif.ImageIFD.ImageDescription].decode("utf-8")
== 'Workflow:{"nodes": [{"id": 1}]}'
)
assert exif_dict["Exif"] is not None
user_comment = exif_dict["Exif"][piexif.ExifIFD.UserComment]
assert user_comment.startswith(b"UNICODE\0")
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))
assert updated_exif["0th"] is not None
assert (
updated_exif["0th"][piexif.ImageIFD.ImageDescription].decode("utf-8")
== 'Workflow:{"nodes":[{"id":1}]}'
)
assert updated_exif["Exif"] is not None
updated_comment = updated_exif["Exif"][piexif.ExifIFD.UserComment]
assert (
updated_comment[8:].decode("utf-16be")
@@ -147,10 +152,10 @@ def test_update_image_metadata_preserves_png_workflow(tmp_path):
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."""
# ISOBMFF box 1: JXL signature box (size=12, type='JXL ', signature)
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
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."""
compressed = brotli.compress(json.dumps(payload_json).encode("utf-8"))
brob_payload = b"comf" + compressed
@@ -262,6 +267,7 @@ class TestIsobmffBrotliExtraction:
path.write_bytes(data)
result = ExifUtils._load_structured_metadata(str(path))
assert result["prompt"] is not None
assert json.loads(result["prompt"]) == {"text": "hello", "negative": "bad"}
def test_extract_workflow_as_list(self, tmp_path):
@@ -272,6 +278,7 @@ class TestIsobmffBrotliExtraction:
path.write_bytes(data)
result = ExifUtils._load_structured_metadata(str(path))
assert result["workflow"] is not None
assert json.loads(result["workflow"]) == [{"id": 1}, {"id": 2}]
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."""
from typing import Any, Dict
import pytest
from py.utils.models import (
BaseModelMetadata,
@@ -44,7 +46,7 @@ class TestCheckpointMetadataSubType:
def test_checkpoint_from_civitai_info_uses_sub_type(self):
"""from_civitai_info should use sub_type from version_info."""
version_info = {
version_info: Dict[str, Any] = {
"baseModel": "SDXL",
"model": {"name": "Test", "description": "", "tags": []},
"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):
"""from_civitai_info should use sub_type from version_info."""
version_info = {
version_info: Dict[str, Any] = {
"baseModel": "SD1.5",
"model": {"name": "Test", "description": "", "tags": []},
"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):
"""from_civitai_info should extract type from civitai data."""
version_info = {
version_info: Dict[str, Any] = {
"baseModel": "SDXL",
"model": {"name": "Test", "description": "", "tags": [], "type": "Lora"},
"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)
assert selected is not None
assert selected["url"] == "nsfw"
assert level == 32
@@ -40,6 +41,7 @@ def test_select_preview_respects_configurable_threshold(threshold_name, expected
mature_threshold=NSFW_LEVELS[threshold_name],
)
assert selected is not None
assert selected["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 typing import Any, Dict
import pytest
from hypothesis import given, settings, strategies as st
@@ -80,6 +82,8 @@ class TestNormalizePath:
@given(st.text(alphabet='abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789._-/\\') | st.none())
def test_normalize_path_is_idempotent_for_ascii(self, path: str | None):
"""Normalizing an already normalized ASCII path should not change it."""
if path is None:
return
normalized = normalize_path(path)
renormalized = normalize_path(normalized)
assert normalized == renormalized
@@ -160,14 +164,14 @@ class TestCalculateRecipeFingerprint:
"""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))
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."""
fp1 = calculate_recipe_fingerprint(loras)
fp2 = calculate_recipe_fingerprint(loras)
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))
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."""
result = calculate_recipe_fingerprint(loras)
assert isinstance(result, str)
@@ -178,7 +182,7 @@ class TestCalculateRecipeFingerprint:
assert result == ""
@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."""
# Create a different input by modifying the first LoRA
loras2 = loras1.copy()