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