From d2f955266d782f728edb56b77511263b8e84b0cb Mon Sep 17 00:00:00 2001 From: Will Miao Date: Sat, 8 Aug 2026 20:12:59 +0800 Subject: [PATCH] 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) --- tests/config/test_config_save_paths.py | 20 +- tests/conftest.py | 34 +-- tests/i18n/test_i18n.py | 14 +- tests/integration/conftest.py | 4 +- tests/integration/test_download_flow.py | 2 +- tests/integration/test_recipe_flow.py | 3 +- tests/metadata_collector/conftest.py | 14 +- .../test_metadata_collector.py | 6 +- tests/metadata_ops/test_metadata_ops.py | 11 +- tests/metadata_ops/test_readme_processor.py | 2 + tests/middleware/test_csp_middleware.py | 2 +- tests/nodes/test_nunchaku_lora.py | 4 +- tests/nodes/test_prompt_text_wildcards.py | 6 +- tests/nodes/test_save_image.py | 36 ++- tests/nodes/test_utils.py | 2 +- tests/performance/test_cache_performance.py | 30 ++- tests/routes/test_api_snapshots.py | 60 +++-- tests/routes/test_base_model_routes_smoke.py | 22 +- tests/routes/test_embedding_routes.py | 2 +- ...example_images_route_registrar_handlers.py | 23 +- tests/routes/test_example_images_routes.py | 127 ++++++---- tests/routes/test_lora_manager_lifecycle.py | 3 +- tests/routes/test_lora_routes.py | 102 +-------- tests/routes/test_misc_routes.py | 216 +++++++++--------- tests/routes/test_model_page_view.py | 12 +- tests/routes/test_model_query_handler.py | 25 +- tests/routes/test_model_update_handler.py | 57 +++-- tests/routes/test_randomizer_endpoints.py | 2 +- tests/routes/test_recipe_query_handler.py | 8 +- tests/routes/test_recipe_route_scaffolding.py | 8 +- tests/routes/test_recipe_routes.py | 14 +- tests/routes/test_route_integration.py | 6 +- tests/routes/test_settings_handler.py | 34 ++- tests/routes/test_stats_routes.py | 2 +- tests/routes/test_tag_logic_param_parsing.py | 4 +- tests/routes/test_update_routes.py | 4 +- tests/routes/test_wildcard_routes.py | 14 +- tests/services/test_aria2_downloader.py | 14 +- .../services/test_autov3_backfill_service.py | 31 +-- tests/services/test_backup_service.py | 4 +- tests/services/test_base_model_service.py | 33 +-- tests/services/test_batch_import_service.py | 14 +- tests/services/test_cache_entry_validator.py | 21 +- tests/services/test_check_pending_models.py | 7 +- tests/services/test_checkpoint_lazy_hash.py | 6 +- tests/services/test_checkpoint_scanner.py | 6 +- tests/services/test_civarchive_client.py | 8 +- .../test_civitai_base_model_service.py | 4 +- tests/services/test_civitai_client.py | 8 +- tests/services/test_civitai_image_parser.py | 4 +- tests/services/test_download_manager_basic.py | 2 + .../test_download_manager_concurrent.py | 8 +- tests/services/test_download_manager_error.py | 28 +-- tests/services/test_downloader.py | 26 ++- .../test_example_images_cleanup_service.py | 2 +- ...t_example_images_download_manager_async.py | 13 +- tests/services/test_issue_760_repro.py | 7 +- tests/services/test_license_filters.py | 10 +- .../test_license_filters_integration.py | 10 +- tests/services/test_llm_service.py | 3 + tests/services/test_metadata_service.py | 5 +- tests/services/test_metadata_sync_service.py | 14 +- .../services/test_model_lifecycle_service.py | 63 +++-- .../services/test_model_metadata_provider.py | 43 +++- tests/services/test_model_query_sub_type.py | 4 +- tests/services/test_model_scanner.py | 16 +- .../test_model_scanner_base_models.py | 10 +- tests/services/test_model_update_service.py | 1 + tests/services/test_no_tags_filter.py | 2 +- tests/services/test_persistent_model_cache.py | 3 +- tests/services/test_preview_asset_service.py | 10 +- tests/services/test_recipe_format_parser.py | 4 +- tests/services/test_recipe_repair.py | 3 +- tests/services/test_recipe_scanner.py | 15 +- tests/services/test_recipe_services.py | 14 +- tests/services/test_root_folder_recursive.py | 23 +- tests/services/test_route_support_services.py | 8 +- tests/services/test_service_registry.py | 2 +- tests/services/test_settings_manager.py | 4 +- .../services/test_sui_image_params_parser.py | 4 + tests/services/test_use_cases.py | 35 ++- tests/standalone/test_standalone_server.py | 14 +- tests/test_auto_tag_service.py | 8 +- tests/test_persistent_recipe_cache.py | 12 +- tests/test_recipe_fts_index_validation.py | 4 +- tests/test_standalone_settings.py | 5 +- tests/utils/test_civitai_utils_rewrite.py | 6 + ...st_example_images_download_manager_unit.py | 8 +- .../utils/test_example_images_file_manager.py | 27 ++- tests/utils/test_example_images_metadata.py | 4 +- .../test_example_images_processor_unit.py | 9 +- tests/utils/test_exif_utils.py | 15 +- tests/utils/test_models_sub_type.py | 8 +- tests/utils/test_preview_selection.py | 2 + tests/utils/test_utils_hypothesis.py | 10 +- 95 files changed, 953 insertions(+), 666 deletions(-) diff --git a/tests/config/test_config_save_paths.py b/tests/config/test_config_save_paths.py index 59ffe481..5e4db62c 100644 --- a/tests/config/test_config_save_paths.py +++ b/tests/config/test_config_save_paths.py @@ -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) diff --git a/tests/conftest.py b/tests/conftest.py index 71554cb4..6193dbf5 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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() diff --git a/tests/i18n/test_i18n.py b/tests/i18n/test_i18n.py index 1dc002d3..1d9becb4 100644 --- a/tests/i18n/test_i18n.py +++ b/tests/i18n/test_i18n.py @@ -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]) diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 1e16b1c0..a746c62d 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -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 diff --git a/tests/integration/test_download_flow.py b/tests/integration/test_download_flow.py index ac9735c6..89a7e985 100644 --- a/tests/integration/test_download_flow.py +++ b/tests/integration/test_download_flow.py @@ -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 diff --git a/tests/integration/test_recipe_flow.py b/tests/integration/test_recipe_flow.py index 9e3d33cd..60929e61 100644 --- a/tests/integration/test_recipe_flow.py +++ b/tests/integration/test_recipe_flow.py @@ -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: diff --git a/tests/metadata_collector/conftest.py b/tests/metadata_collector/conftest.py index c2bcab87..86498619 100644 --- a/tests/metadata_collector/conftest.py +++ b/tests/metadata_collector/conftest.py @@ -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": {}}, diff --git a/tests/metadata_collector/test_metadata_collector.py b/tests/metadata_collector/test_metadata_collector.py index 22b14788..57631654 100644 --- a/tests/metadata_collector/test_metadata_collector.py +++ b/tests/metadata_collector/test_metadata_collector.py @@ -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) diff --git a/tests/metadata_ops/test_metadata_ops.py b/tests/metadata_ops/test_metadata_ops.py index 45131103..74b0f96d 100644 --- a/tests/metadata_ops/test_metadata_ops.py +++ b/tests/metadata_ops/test_metadata_ops.py @@ -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("") == "" diff --git a/tests/metadata_ops/test_readme_processor.py b/tests/metadata_ops/test_readme_processor.py index a34bbbc5..e695c9aa 100644 --- a/tests/metadata_ops/test_readme_processor.py +++ b/tests/metadata_ops/test_readme_processor.py @@ -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 diff --git a/tests/middleware/test_csp_middleware.py b/tests/middleware/test_csp_middleware.py index 1ad1c56d..8125a2c3 100644 --- a/tests/middleware/test_csp_middleware.py +++ b/tests/middleware/test_csp_middleware.py @@ -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 diff --git a/tests/nodes/test_nunchaku_lora.py b/tests/nodes/test_nunchaku_lora.py index 6b1ac783..11416145 100644 --- a/tests/nodes/test_nunchaku_lora.py +++ b/tests/nodes/test_nunchaku_lora.py @@ -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() diff --git a/tests/nodes/test_prompt_text_wildcards.py b/tests/nodes/test_prompt_text_wildcards.py index 0ab8e91b..65a3717f 100644 --- a/tests/nodes/test_prompt_text_wildcards.py +++ b/tests/nodes/test_prompt_text_wildcards.py @@ -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(): diff --git a/tests/nodes/test_save_image.py b/tests/nodes/test_save_image.py index f4aa6b12..8c455b90 100644 --- a/tests/nodes/test_save_image.py +++ b/tests/nodes/test_save_image.py @@ -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): diff --git a/tests/nodes/test_utils.py b/tests/nodes/test_utils.py index 35151b8c..6d31ed7f 100644 --- a/tests/nodes/test_utils.py +++ b/tests/nodes/test_utils.py @@ -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() diff --git a/tests/performance/test_cache_performance.py b/tests/performance/test_cache_performance.py index 54d8c7c6..65a74d7b 100644 --- a/tests/performance/test_cache_performance.py +++ b/tests/performance/test_cache_performance.py @@ -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): diff --git a/tests/routes/test_api_snapshots.py b/tests/routes/test_api_snapshots.py index d3884df0..1e400d04 100644 --- a/tests/routes/test_api_snapshots.py +++ b/tests/routes/test_api_snapshots.py @@ -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 diff --git a/tests/routes/test_base_model_routes_smoke.py b/tests/routes/test_base_model_routes_smoke.py index f580d855..06cc8a3c 100644 --- a/tests/routes/test_base_model_routes_smoke.py +++ b/tests/routes/test_base_model_routes_smoke.py @@ -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) diff --git a/tests/routes/test_embedding_routes.py b/tests/routes/test_embedding_routes.py index fc1782a0..4ab968f8 100644 --- a/tests/routes/test_embedding_routes.py +++ b/tests/routes/test_embedding_routes.py @@ -24,7 +24,7 @@ class StubEmbeddingService: @pytest.fixture def routes(): handler = EmbeddingRoutes() - handler.service = StubEmbeddingService() + handler.service = StubEmbeddingService() # pyright: ignore[reportAttributeAccessIssue] return handler diff --git a/tests/routes/test_example_images_route_registrar_handlers.py b/tests/routes/test_example_images_route_registrar_handlers.py index a214dbbc..d560aa68 100644 --- a/tests/routes/test_example_images_route_registrar_handlers.py +++ b/tests/routes/test_example_images_route_registrar_handlers.py @@ -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 diff --git a/tests/routes/test_example_images_routes.py b/tests/routes/test_example_images_routes.py index bc926e65..9b167678 100644 --- a/tests/routes/test_example_images_routes.py +++ b/tests/routes/test_example_images_routes.py @@ -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, diff --git a/tests/routes/test_lora_manager_lifecycle.py b/tests/routes/test_lora_manager_lifecycle.py index 1bb0f1eb..f06530b9 100644 --- a/tests/routes/test_lora_manager_lifecycle.py +++ b/tests/routes/test_lora_manager_lifecycle.py @@ -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) diff --git a/tests/routes/test_lora_routes.py b/tests/routes/test_lora_routes.py index b2e15f01..aea5ce4c 100644 --- a/tests/routes/test_lora_routes.py +++ b/tests/routes/test_lora_routes.py @@ -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) diff --git a/tests/routes/test_misc_routes.py b/tests/routes/test_misc_routes.py index a735a4ff..eed69553 100644 --- a/tests/routes/test_misc_routes.py +++ b/tests/routes/test_misc_routes.py @@ -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 diff --git a/tests/routes/test_model_page_view.py b/tests/routes/test_model_page_view.py index a9275c35..55e0d773 100644 --- a/tests/routes/test_model_page_view.py +++ b/tests/routes/test_model_page_view.py @@ -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" diff --git a/tests/routes/test_model_query_handler.py b/tests/routes/test_model_query_handler.py index f39cb357..ea88b90f 100644 --- a/tests/routes/test_model_query_handler.py +++ b/tests/routes/test_model_query_handler.py @@ -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 diff --git a/tests/routes/test_model_update_handler.py b/tests/routes/test_model_update_handler.py index afafdc78..8a5fd4ef 100644 --- a/tests/routes/test_model_update_handler.py +++ b/tests/routes/test_model_update_handler.py @@ -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]] diff --git a/tests/routes/test_randomizer_endpoints.py b/tests/routes/test_randomizer_endpoints.py index 48374698..6a5e777a 100644 --- a/tests/routes/test_randomizer_endpoints.py +++ b/tests/routes/test_randomizer_endpoints.py @@ -31,7 +31,7 @@ class StubLoraService: @pytest.fixture def routes(): handler = LoraRoutes() - handler.service = StubLoraService() + handler.service = StubLoraService() # pyright: ignore[reportAttributeAccessIssue] return handler diff --git a/tests/routes/test_recipe_query_handler.py b/tests/routes/test_recipe_query_handler.py index 774c1033..a897e9c4 100644 --- a/tests/routes/test_recipe_query_handler.py +++ b/tests/routes/test_recipe_query_handler.py @@ -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"] == [ diff --git a/tests/routes/test_recipe_route_scaffolding.py b/tests/routes/test_recipe_route_scaffolding.py index 0ff95da9..20f3506a 100644 --- a/tests/routes/test_recipe_route_scaffolding.py +++ b/tests/routes/test_recipe_route_scaffolding.py @@ -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 } diff --git a/tests/routes/test_recipe_routes.py b/tests/routes/test_recipe_routes.py index 78bbad27..47686895 100644 --- a/tests/routes/test_recipe_routes.py +++ b/tests/routes/test_recipe_routes.py @@ -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 [""] 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() diff --git a/tests/routes/test_route_integration.py b/tests/routes/test_route_integration.py index dc0a1261..72424011 100644 --- a/tests/routes/test_route_integration.py +++ b/tests/routes/test_route_integration.py @@ -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) diff --git a/tests/routes/test_settings_handler.py b/tests/routes/test_settings_handler.py index 79f7452e..f7369770 100644 --- a/tests/routes/test_settings_handler.py +++ b/tests/routes/test_settings_handler.py @@ -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 diff --git a/tests/routes/test_stats_routes.py b/tests/routes/test_stats_routes.py index d27f2c16..45a2609f 100644 --- a/tests/routes/test_stats_routes.py +++ b/tests/routes/test_stats_routes.py @@ -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 diff --git a/tests/routes/test_tag_logic_param_parsing.py b/tests/routes/test_tag_logic_param_parsing.py index d10006a0..9648d2b8 100644 --- a/tests/routes/test_tag_logic_param_parsing.py +++ b/tests/routes/test_tag_logic_param_parsing.py @@ -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 diff --git a/tests/routes/test_update_routes.py b/tests/routes/test_update_routes.py index 6254d049..8d621131 100644 --- a/tests/routes/test_update_routes.py +++ b/tests/routes/test_update_routes.py @@ -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() diff --git a/tests/routes/test_wildcard_routes.py b/tests/routes/test_wildcard_routes.py index 7d86298e..b13da562 100644 --- a/tests/routes/test_wildcard_routes.py +++ b/tests/routes/test_wildcard_routes.py @@ -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" diff --git a/tests/services/test_aria2_downloader.py b/tests/services/test_aria2_downloader.py index 236785b9..bf6763e2 100644 --- a/tests/services/test_aria2_downloader.py +++ b/tests/services/test_aria2_downloader.py @@ -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( diff --git a/tests/services/test_autov3_backfill_service.py b/tests/services/test_autov3_backfill_service.py index 9b1325a3..aa1b12ea 100644 --- a/tests/services/test_autov3_backfill_service.py +++ b/tests/services/test_autov3_backfill_service.py @@ -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, '')] diff --git a/tests/services/test_backup_service.py b/tests/services/test_backup_service.py index 355fb534..bccd2550 100644 --- a/tests/services/test_backup_service.py +++ b/tests/services/test_backup_service.py @@ -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()) diff --git a/tests/services/test_base_model_service.py b/tests/services/test_base_model_service.py index 9309170e..fc54b835 100644 --- a/tests/services/test_base_model_service.py +++ b/tests/services/test_base_model_service.py @@ -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"] diff --git a/tests/services/test_batch_import_service.py b/tests/services/test_batch_import_service.py index 49b9193e..89e41f8e 100644 --- a/tests/services/test_batch_import_service.py +++ b/tests/services/test_batch_import_service.py @@ -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, ) diff --git a/tests/services/test_cache_entry_validator.py b/tests/services/test_cache_entry_validator.py index a1b7580e..8575f0e5 100644 --- a/tests/services/test_cache_entry_validator.py +++ b/tests/services/test_cache_entry_validator.py @@ -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 diff --git a/tests/services/test_check_pending_models.py b/tests/services/test_check_pending_models.py index a5c4e790..684e09a1 100644 --- a/tests/services/test_check_pending_models.py +++ b/tests/services/test_check_pending_models.py @@ -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) diff --git a/tests/services/test_checkpoint_lazy_hash.py b/tests/services/test_checkpoint_lazy_hash.py index ad2f00c6..646f31dc 100644 --- a/tests/services/test_checkpoint_lazy_hash.py +++ b/tests/services/test_checkpoint_lazy_hash.py @@ -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) diff --git a/tests/services/test_checkpoint_scanner.py b/tests/services/test_checkpoint_scanner.py index 469897f8..0d9970a0 100644 --- a/tests/services/test_checkpoint_scanner.py +++ b/tests/services/test_checkpoint_scanner.py @@ -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) diff --git a/tests/services/test_civarchive_client.py b/tests/services/test_civarchive_client.py index 8c4bf092..507e1094 100644 --- a/tests/services/test_civarchive_client.py +++ b/tests/services/test_civarchive_client.py @@ -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) diff --git a/tests/services/test_civitai_base_model_service.py b/tests/services/test_civitai_base_model_service.py index 80c2bdbe..eb007131 100644 --- a/tests/services/test_civitai_base_model_service.py +++ b/tests/services/test_civitai_base_model_service.py @@ -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.""" diff --git a/tests/services/test_civitai_client.py b/tests/services/test_civitai_client.py index cd56d15d..63f2ab93 100644 --- a/tests/services/test_civitai_client.py +++ b/tests/services/test_civitai_client.py @@ -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" diff --git a/tests/services/test_civitai_image_parser.py b/tests/services/test_civitai_image_parser.py index 1c8ec19d..a4b8498a 100644 --- a/tests/services/test_civitai_image_parser.py +++ b/tests/services/test_civitai_image_parser.py @@ -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) diff --git a/tests/services/test_download_manager_basic.py b/tests/services/test_download_manager_basic.py index a0a6d861..24035b13 100644 --- a/tests/services/test_download_manager_basic.py +++ b/tests/services/test_download_manager_basic.py @@ -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"] == [ { diff --git a/tests/services/test_download_manager_concurrent.py b/tests/services/test_download_manager_concurrent.py index 7328de7c..68a42ef0 100644 --- a/tests/services/test_download_manager_concurrent.py +++ b/tests/services/test_download_manager_concurrent.py @@ -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 diff --git a/tests/services/test_download_manager_error.py b/tests/services/test_download_manager_error.py index ba302125..6f0b3e0f 100644 --- a/tests/services/test_download_manager_error.py +++ b/tests/services/test_download_manager_error.py @@ -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(): diff --git a/tests/services/test_downloader.py b/tests/services/test_downloader.py index 84bd858b..ab276857 100644 --- a/tests/services/test_downloader.py +++ b/tests/services/test_downloader.py @@ -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-" diff --git a/tests/services/test_example_images_cleanup_service.py b/tests/services/test_example_images_cleanup_service.py index ca75ba0e..1a88d0f6 100644 --- a/tests/services/test_example_images_cleanup_service.py +++ b/tests/services/test_example_images_cleanup_service.py @@ -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() diff --git a/tests/services/test_example_images_download_manager_async.py b/tests/services/test_example_images_download_manager_async.py index 9ba9247c..df1047cd 100644 --- a/tests/services/test_example_images_download_manager_async.py +++ b/tests/services/test_example_images_download_manager_async.py @@ -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), diff --git a/tests/services/test_issue_760_repro.py b/tests/services/test_issue_760_repro.py index be8fa4ed..874be1a0 100644 --- a/tests/services/test_issue_760_repro.py +++ b/tests/services/test_issue_760_repro.py @@ -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: diff --git a/tests/services/test_license_filters.py b/tests/services/test_license_filters.py index 83b11f3d..a3e8b85a 100644 --- a/tests/services/test_license_filters.py +++ b/tests/services/test_license_filters.py @@ -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 diff --git a/tests/services/test_license_filters_integration.py b/tests/services/test_license_filters_integration.py index 4fd2ab7f..e0deb419 100644 --- a/tests/services/test_license_filters_integration.py +++ b/tests/services/test_license_filters_integration.py @@ -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 diff --git a/tests/services/test_llm_service.py b/tests/services/test_llm_service.py index c4d38fe8..c6dc7d1f 100644 --- a/tests/services/test_llm_service.py +++ b/tests/services/test_llm_service.py @@ -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 diff --git a/tests/services/test_metadata_service.py b/tests/services/test_metadata_service.py index dbb66165..95f8bfdf 100644 --- a/tests/services/test_metadata_service.py +++ b/tests/services/test_metadata_service.py @@ -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 diff --git a/tests/services/test_metadata_sync_service.py b/tests/services/test_metadata_sync_service.py index 8fb87618..0a360238 100644 --- a/tests/services/test_metadata_sync_service.py +++ b/tests/services/test_metadata_sync_service.py @@ -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, diff --git a/tests/services/test_model_lifecycle_service.py b/tests/services/test_model_lifecycle_service.py index f421fee8..df76ae47 100644 --- a/tests/services/test_model_lifecycle_service.py +++ b/tests/services/test_model_lifecycle_service.py @@ -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"}} diff --git a/tests/services/test_model_metadata_provider.py b/tests/services/test_model_metadata_provider.py index 3ad59294..c8877ccf 100644 --- a/tests/services/test_model_metadata_provider.py +++ b/tests/services/test_model_metadata_provider.py @@ -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): diff --git a/tests/services/test_model_query_sub_type.py b/tests/services/test_model_query_sub_type.py index 6282b078..306a2753 100644 --- a/tests/services/test_model_query_sub_type.py +++ b/tests/services/test_model_query_sub_type.py @@ -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: diff --git a/tests/services/test_model_scanner.py b/tests/services/test_model_scanner.py index 90298c33..ee27c004 100644 --- a/tests/services/test_model_scanner.py +++ b/tests/services/test_model_scanner.py @@ -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", diff --git a/tests/services/test_model_scanner_base_models.py b/tests/services/test_model_scanner_base_models.py index 4d1b3f3a..fc00d69f 100644 --- a/tests/services/test_model_scanner_base_models.py +++ b/tests/services/test_model_scanner_base_models.py @@ -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 diff --git a/tests/services/test_model_update_service.py b/tests/services/test_model_update_service.py index f1f2b6bc..dcb6cde1 100644 --- a/tests/services/test_model_update_service.py +++ b/tests/services/test_model_update_service.py @@ -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 diff --git a/tests/services/test_no_tags_filter.py b/tests/services/test_no_tags_filter.py index 473750f2..66da633b 100644 --- a/tests/services/test_no_tags_filter.py +++ b/tests/services/test_no_tags_filter.py @@ -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 diff --git a/tests/services/test_persistent_model_cache.py b/tests/services/test_persistent_model_cache.py index edded38f..b48b2af0 100644 --- a/tests/services/test_persistent_model_cache.py +++ b/tests/services/test_persistent_model_cache.py @@ -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, diff --git a/tests/services/test_preview_asset_service.py b/tests/services/test_preview_asset_service.py index c67d78df..68325ec8 100644 --- a/tests/services/test_preview_asset_service.py +++ b/tests/services/test_preview_asset_service.py @@ -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", diff --git a/tests/services/test_recipe_format_parser.py b/tests/services/test_recipe_format_parser.py index d99a02a8..141c49c7 100644 --- a/tests/services/test_recipe_format_parser.py +++ b/tests/services/test_recipe_format_parser.py @@ -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, diff --git a/tests/services/test_recipe_repair.py b/tests/services/test_recipe_repair.py index f9fe4759..0dae772b 100644 --- a/tests/services/test_recipe_repair.py +++ b/tests/services/test_recipe_repair.py @@ -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": { diff --git a/tests/services/test_recipe_scanner.py b/tests/services/test_recipe_scanner.py index 27978f82..2c222707 100644 --- a/tests/services/test_recipe_scanner.py +++ b/tests/services/test_recipe_scanner.py @@ -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", diff --git a/tests/services/test_recipe_services.py b/tests/services/test_recipe_services.py index 5353d2ee..82cc0fac 100644 --- a/tests/services/test_recipe_services.py +++ b/tests/services/test_recipe_services.py @@ -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) diff --git a/tests/services/test_root_folder_recursive.py b/tests/services/test_root_folder_recursive.py index e4b5abf5..51873676 100644 --- a/tests/services/test_root_folder_recursive.py +++ b/tests/services/test_root_folder_recursive.py @@ -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" diff --git a/tests/services/test_route_support_services.py b/tests/services/test_route_support_services.py index 39929832..2027de1f 100644 --- a/tests/services/test_route_support_services.py +++ b/tests/services/test_route_support_services.py @@ -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), ) diff --git a/tests/services/test_service_registry.py b/tests/services/test_service_registry.py index 711be196..9f1b1646 100644 --- a/tests/services/test_service_registry.py +++ b/tests/services/test_service_registry.py @@ -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() diff --git a/tests/services/test_settings_manager.py b/tests/services/test_settings_manager.py index 816e9a7c..3597b189 100644 --- a/tests/services/test_settings_manager.py +++ b/tests/services/test_settings_manager.py @@ -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() diff --git a/tests/services/test_sui_image_params_parser.py b/tests/services/test_sui_image_params_parser.py index 9bcaf04d..5834df25 100644 --- a/tests/services/test_sui_image_params_parser.py +++ b/tests/services/test_sui_image_params_parser.py @@ -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' diff --git a/tests/services/test_use_cases.py b/tests/services/test_use_cases.py index 3888c4bd..273ba11a 100644 --- a/tests/services/test_use_cases.py +++ b/tests/services/test_use_cases.py @@ -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) \ No newline at end of file + await use_case.execute(request) # pyright: ignore[reportArgumentType] \ No newline at end of file diff --git a/tests/standalone/test_standalone_server.py b/tests/standalone/test_standalone_server.py index 999b7bf4..7691b9ec 100644 --- a/tests/standalone/test_standalone_server.py +++ b/tests/standalone/test_standalone_server.py @@ -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", diff --git a/tests/test_auto_tag_service.py b/tests/test_auto_tag_service.py index 84df141d..dcbfe42e 100644 --- a/tests/test_auto_tag_service.py +++ b/tests/test_auto_tag_service.py @@ -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 diff --git a/tests/test_persistent_recipe_cache.py b/tests/test_persistent_recipe_cache.py index 3fd8f928..94ddc003 100644 --- a/tests/test_persistent_recipe_cache.py +++ b/tests/test_persistent_recipe_cache.py @@ -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"} diff --git a/tests/test_recipe_fts_index_validation.py b/tests/test_recipe_fts_index_validation.py index 5add49ba..9bffe073 100644 --- a/tests/test_recipe_fts_index_validation.py +++ b/tests/test_recipe_fts_index_validation.py @@ -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 [ { diff --git a/tests/test_standalone_settings.py b/tests/test_standalone_settings.py index 8e403b66..a14c423b 100644 --- a/tests/test_standalone_settings.py +++ b/tests/test_standalone_settings.py @@ -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) diff --git a/tests/utils/test_civitai_utils_rewrite.py b/tests/utils/test_civitai_utils_rewrite.py index 8d6a3108..c65598d1 100644 --- a/tests/utils/test_civitai_utils_rewrite.py +++ b/tests/utils/test_civitai_utils_rewrite.py @@ -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 diff --git a/tests/utils/test_example_images_download_manager_unit.py b/tests/utils/test_example_images_download_manager_unit.py index f078f7d8..7eec2436 100644 --- a/tests/utils/test_example_images_download_manager_unit.py +++ b/tests/utils/test_example_images_download_manager_unit.py @@ -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( diff --git a/tests/utils/test_example_images_file_manager.py b/tests/utils/test_example_images_file_manager.py index bc4129ef..702620ab 100644 --- a/tests/utils/test_example_images_file_manager.py +++ b/tests/utils/test_example_images_file_manager.py @@ -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 diff --git a/tests/utils/test_example_images_metadata.py b/tests/utils/test_example_images_metadata.py index 7498d679..095212f8 100644 --- a/tests/utils/test_example_images_metadata.py +++ b/tests/utils/test_example_images_metadata.py @@ -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 diff --git a/tests/utils/test_example_images_processor_unit.py b/tests/utils/test_example_images_processor_unit.py index 90e7df69..229d5072 100644 --- a/tests/utils/test_example_images_processor_unit.py +++ b/tests/utils/test_example_images_processor_unit.py @@ -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"] == [] diff --git a/tests/utils/test_exif_utils.py b/tests/utils/test_exif_utils.py index 9cbaad1b..b1d7797d 100644 --- a/tests/utils/test_exif_utils.py +++ b/tests/utils/test_exif_utils.py @@ -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): diff --git a/tests/utils/test_models_sub_type.py b/tests/utils/test_models_sub_type.py index 99311ded..ed4d8c46 100644 --- a/tests/utils/test_models_sub_type.py +++ b/tests/utils/test_models_sub_type.py @@ -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}], diff --git a/tests/utils/test_preview_selection.py b/tests/utils/test_preview_selection.py index 922f2670..0a744bfe 100644 --- a/tests/utils/test_preview_selection.py +++ b/tests/utils/test_preview_selection.py @@ -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) diff --git a/tests/utils/test_utils_hypothesis.py b/tests/utils/test_utils_hypothesis.py index cbf6a59a..d72e38d9 100644 --- a/tests/utils/test_utils_hypothesis.py +++ b/tests/utils/test_utils_hypothesis.py @@ -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()