diff --git a/py/config.py b/py/config.py index b929049f..8e96b2e8 100644 --- a/py/config.py +++ b/py/config.py @@ -1,9 +1,13 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. import os import platform import posixpath import threading from pathlib import Path -import folder_paths # type: ignore +import folder_paths # pyright: ignore[reportMissingImports] from typing import Any, Dict, Iterable, List, Mapping, Optional, Set, Tuple import logging import json @@ -90,7 +94,7 @@ def _resolve_valid_default_root( def _normalize_folder_paths_for_comparison( - folder_paths: Mapping[str, Iterable[str]], + folder_paths: Mapping[str, Any], ) -> Dict[str, Set[str]]: """Normalize folder paths for comparison across libraries.""" @@ -482,7 +486,7 @@ class Config: import ctypes FILE_ATTRIBUTE_REPARSE_POINT = 0x400 - attrs = ctypes.windll.kernel32.GetFileAttributesW(str(path)) # type: ignore[attr-defined] + attrs = ctypes.windll.kernel32.GetFileAttributesW(str(path)) # pyright: ignore[reportAttributeAccessIssue] return attrs != -1 and (attrs & FILE_ATTRIBUTE_REPARSE_POINT) except Exception as e: logger.error(f"Error checking Windows reparse point: {e}") @@ -491,7 +495,7 @@ class Config: logger.error(f"Error checking link status for {path}: {e}") return False - def _entry_is_symlink(self, entry: os.DirEntry) -> bool: + def _entry_is_symlink(self, entry: os.DirEntry[str]) -> bool: """Check if a directory entry is a symlink, including Windows junctions.""" if entry.is_symlink(): return True @@ -500,7 +504,7 @@ class Config: import ctypes FILE_ATTRIBUTE_REPARSE_POINT = 0x400 - attrs = ctypes.windll.kernel32.GetFileAttributesW(entry.path) # type: ignore[attr-defined] + attrs = ctypes.windll.kernel32.GetFileAttributesW(entry.path) # pyright: ignore[reportAttributeAccessIssue] return attrs != -1 and (attrs & FILE_ATTRIBUTE_REPARSE_POINT) except Exception: pass @@ -1126,8 +1130,8 @@ class Config: def _apply_library_paths( self, - folder_paths: Mapping[str, Iterable[str]], - extra_folder_paths: Optional[Mapping[str, Iterable[str]]] = None, + folder_paths: Mapping[str, Any], + extra_folder_paths: Optional[Mapping[str, Any]] = None, recipes_path: str = "", ) -> None: self._path_mappings.clear() @@ -1432,12 +1436,13 @@ class Config: # ('_lm_config_cache') that is NEVER removed from sys.modules (its key does # NOT start with 'py.'), so it survives re-imports of py.* modules. _CONFIG_SENTINEL = "_lm_config_cache" +config: Config if _CONFIG_SENTINEL in _sys.modules: # Re-import: reuse the existing singleton from the sentinel. - config: Config = _sys.modules[_CONFIG_SENTINEL].config # type: ignore[valid-type] + config = _sys.modules[_CONFIG_SENTINEL].config else: - config: Config = Config() + config = Config() # Register the sentinel so re-imports of py.config find us. _sentinel_mod = _types.ModuleType(_CONFIG_SENTINEL) - _sentinel_mod.config = config + setattr(_sentinel_mod, "config", config) _sys.modules[_CONFIG_SENTINEL] = _sentinel_mod diff --git a/py/lora_manager.py b/py/lora_manager.py index f304bc77..4b02555d 100644 --- a/py/lora_manager.py +++ b/py/lora_manager.py @@ -14,7 +14,7 @@ standalone_mode = ( if not standalone_mode: setup_logging() -from server import PromptServer # type: ignore +from server import PromptServer # pyright: ignore[reportMissingImports] from .config import config from .services.model_service_factory import ( diff --git a/py/metadata_collector/__init__.py b/py/metadata_collector/__init__.py index f3dde058..c7287ba3 100644 --- a/py/metadata_collector/__init__.py +++ b/py/metadata_collector/__init__.py @@ -22,7 +22,7 @@ if not standalone_mode: logger.info("ComfyUI Metadata Collector initialized") - def get_metadata(prompt_id=None): # type: ignore[no-redef] + def get_metadata(prompt_id=None): # pyright: ignore[reportRedeclaration] """Helper function to get metadata from the registry""" registry = MetadataRegistry() return registry.get_metadata(prompt_id) @@ -31,6 +31,6 @@ else: def init(): logger.info("ComfyUI Metadata Collector disabled in standalone mode") - def get_metadata(prompt_id=None): # type: ignore[no-redef] + def get_metadata(prompt_id=None): # pyright: ignore[reportRedeclaration] """Dummy implementation for standalone mode""" return {} diff --git a/py/metadata_collector/metadata_hook.py b/py/metadata_collector/metadata_hook.py index b9b40e09..8cd41b99 100644 --- a/py/metadata_collector/metadata_hook.py +++ b/py/metadata_collector/metadata_hook.py @@ -16,7 +16,7 @@ class MetadataHook: execution = None try: # Try direct import first - import execution # type: ignore + import execution # pyright: ignore[reportMissingImports] except ImportError: # Try to locate from system modules for module_name in sys.modules: diff --git a/py/metadata_collector/metadata_registry.py b/py/metadata_collector/metadata_registry.py index 856771b1..15f4d8e4 100644 --- a/py/metadata_collector/metadata_registry.py +++ b/py/metadata_collector/metadata_registry.py @@ -1,5 +1,6 @@ import time -from nodes import NODE_CLASS_MAPPINGS # type: ignore +from typing import Any +from nodes import NODE_CLASS_MAPPINGS # pyright: ignore[reportMissingImports, reportAttributeAccessIssue] from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE @@ -9,6 +10,15 @@ class MetadataRegistry: _instance = None + current_prompt_id: Any = None + current_prompt: Any = None + metadata: dict[str, Any] = {} + prompt_metadata: dict[str, Any] = {} + executed_nodes: set[str] = set() + node_cache: dict[str, Any] = {} + max_prompt_history: int = 3 + metadata_categories: list[str] = METADATA_CATEGORIES + def __new__(cls): if cls._instance is None: cls._instance = super().__new__(cls) diff --git a/py/metadata_ops/__init__.py b/py/metadata_ops/__init__.py index d5b33efe..3ecb957c 100644 --- a/py/metadata_ops/__init__.py +++ b/py/metadata_ops/__init__.py @@ -43,7 +43,7 @@ SCANNER_GETTER_NAMES = tuple(SCANNER_TYPE_MAP.keys()) async def _find_model_entry( model_path: str, -) -> tuple[object, object, str | None] | tuple[None, None, None]: +) -> tuple[Any, object, str | None] | tuple[None, None, None]: """Iterate all scanners and return the first (scanner, entry, getter_name) that owns *model_path*. Returns ``(None, None, None)`` when no scanner claims it. @@ -73,7 +73,7 @@ async def _find_model_entry( async def _find_scanner_for_model( model_path: str, -) -> tuple[object, object] | tuple[None, None]: +) -> tuple[Any, object] | tuple[None, None]: """Find the (scanner, cache_entry) responsible for *model_path*.""" scanner, entry, _ = await _find_model_entry(model_path) return scanner, entry diff --git a/py/nodes/checkpoint_loader.py b/py/nodes/checkpoint_loader.py index 801a509f..87dc497a 100644 --- a/py/nodes/checkpoint_loader.py +++ b/py/nodes/checkpoint_loader.py @@ -1,7 +1,7 @@ import logging -from typing import List, Tuple -import comfy.sd # type: ignore -import folder_paths # type: ignore +from typing import Any, List, Tuple +import comfy.sd # pyright: ignore[reportMissingImports] +import folder_paths # pyright: ignore[reportMissingImports] from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui logger = logging.getLogger(__name__) @@ -18,9 +18,9 @@ class CheckpointLoaderLM: CATEGORY = "Lora Manager/loaders" @classmethod - def INPUT_TYPES(s): + def INPUT_TYPES(cls): # Get list of checkpoint names from scanner (includes extra folder paths) - checkpoint_names = s._get_checkpoint_names() + checkpoint_names = cls._get_checkpoint_names() return { "required": { "ckpt_name": ( @@ -89,7 +89,7 @@ class CheckpointLoaderLM: logger.error(f"Error getting checkpoint names: {e}") return [] - def load_checkpoint(self, ckpt_name: str) -> Tuple: + def load_checkpoint(self, ckpt_name: str) -> Tuple[Any, Any, Any]: """Load a checkpoint by name, supporting extra folder paths Args: diff --git a/py/nodes/create_hook_lora.py b/py/nodes/create_hook_lora.py index a5f44fb8..d192cde5 100644 --- a/py/nodes/create_hook_lora.py +++ b/py/nodes/create_hook_lora.py @@ -57,8 +57,8 @@ class CreateHookLoraLM: del text # used by the frontend widget only # Lazy imports: comfy is not available in CI/test environment at module level - import comfy.hooks # type: ignore # noqa: C0415 - import comfy.utils # type: ignore # noqa: C0415 + import comfy.hooks # pyright: ignore[reportMissingImports] # noqa: C0415 + import comfy.utils # pyright: ignore[reportMissingImports] # noqa: C0415 prev_hooks: comfy.hooks.HookGroup | None = kwargs.get("prev_hooks") diff --git a/py/nodes/lora_loader.py b/py/nodes/lora_loader.py index 1ac52d27..74da5804 100644 --- a/py/nodes/lora_loader.py +++ b/py/nodes/lora_loader.py @@ -1,8 +1,8 @@ import importlib import logging -import comfy.sd # type: ignore -import comfy.utils # type: ignore +import comfy.sd # pyright: ignore[reportMissingImports] +import comfy.utils # pyright: ignore[reportMissingImports] from ..utils.utils import get_lora_info_absolute from .utils import ( diff --git a/py/nodes/lora_stack_combiner.py b/py/nodes/lora_stack_combiner.py index 4c5d8e5d..9f3412da 100644 --- a/py/nodes/lora_stack_combiner.py +++ b/py/nodes/lora_stack_combiner.py @@ -73,7 +73,7 @@ class LoraStackCombinerLM: stack = inspect.stack() if len(stack) > 2 and stack[2].function == "get_input_info": - optional_inputs = _LoraStackOptionalInputs(optional_inputs) # type: ignore[assignment] + optional_inputs = _LoraStackOptionalInputs(optional_inputs) # pyright: ignore[reportAssignmentType] return { "required": {}, diff --git a/py/nodes/nunchaku_qwen.py b/py/nodes/nunchaku_qwen.py index 70c2e7d5..196affda 100644 --- a/py/nodes/nunchaku_qwen.py +++ b/py/nodes/nunchaku_qwen.py @@ -15,15 +15,15 @@ import os import re from collections import defaultdict from pathlib import Path -from typing import Dict, List, Optional, Tuple, Union +from typing import Any, Dict, List, Optional, Tuple, Union, cast -import comfy.utils # type: ignore -import folder_paths # type: ignore +import comfy.utils # pyright: ignore[reportMissingImports] +import folder_paths # pyright: ignore[reportMissingImports] import torch import torch.nn as nn from safetensors import safe_open -from nunchaku.lora.flux.nunchaku_converter import ( +from nunchaku.lora.flux.nunchaku_converter import ( # pyright: ignore[reportMissingTypeStubs] pack_lowrank_weight, unpack_lowrank_weight, ) @@ -87,10 +87,6 @@ def _rename_layer_underscore_layer_name(old_name: str) -> str: return new_name -def _is_indexable_module(module): - return isinstance(module, (nn.ModuleList, nn.Sequential, list, tuple)) - - def _get_module_by_name(model: nn.Module, name: str) -> Optional[nn.Module]: if not name: return model @@ -100,7 +96,7 @@ def _get_module_by_name(model: nn.Module, name: str) -> Optional[nn.Module]: continue if hasattr(module, part): module = getattr(module, part) - elif part.isdigit() and _is_indexable_module(module): + elif part.isdigit() and isinstance(module, (nn.ModuleList, nn.Sequential, list, tuple)): try: module = module[int(part)] except (IndexError, TypeError): @@ -267,7 +263,9 @@ def _handle_proj_out_split(lora_dict: Dict[str, Dict[str, torch.Tensor]], base_k return result, consumed -def _apply_lora_to_module(module: nn.Module, a_tensor: torch.Tensor, b_tensor: torch.Tensor, module_name: str, model: nn.Module) -> None: +def _apply_lora_to_module(module: Any, a_tensor: torch.Tensor, b_tensor: torch.Tensor, module_name: str, model: Any) -> None: + # These modules are dynamic torch containers; monkey-patched attributes + # below are set at runtime, so the module/model types are deliberately Any. if not hasattr(module, "in_features") or not hasattr(module, "out_features"): raise ValueError(f"{module_name}: unsupported module without in/out features") if a_tensor.shape[1] != module.in_features or b_tensor.shape[0] != module.out_features: @@ -336,7 +334,7 @@ def _apply_lora_to_module(module: nn.Module, a_tensor: torch.Tensor, b_tensor: t raise ValueError(f"{module_name}: unsupported module type {type(module)}") -def reset_lora_v2(model: nn.Module) -> None: +def reset_lora_v2(model: Any) -> None: slots = getattr(model, "_lora_slots", None) if not slots: return @@ -344,6 +342,7 @@ def reset_lora_v2(model: nn.Module) -> None: module = _get_module_by_name(model, name) if module is None: continue + module = cast(Any, module) module_type = info.get("type", "nunchaku") if module_type == "nunchaku": base_rank = info["base_rank"] @@ -371,7 +370,7 @@ def reset_lora_v2(model: nn.Module) -> None: def compose_loras_v2(model: nn.Module, lora_configs: List[Tuple[Union[str, Path, Dict[str, torch.Tensor]], float]], apply_awq_mod: bool = True) -> bool: del apply_awq_mod # retained for interface compatibility reset_lora_v2(model) - aggregated_weights: Dict[str, List[Dict[str, object]]] = defaultdict(list) + aggregated_weights: Dict[str, List[Dict[str, Any]]] = defaultdict(list) saw_supported_format = False unresolved_targets = 0 @@ -471,7 +470,7 @@ def compose_loras_v2(model: nn.Module, lora_configs: List[Tuple[Union[str, Path, class ComfyQwenImageWrapperLM(nn.Module): def __init__(self, model: nn.Module, config=None, apply_awq_mod: bool = True): super().__init__() - self.model = model + self.model: Any = model self.config = {} if config is None else config self.dtype = next(model.parameters()).dtype self.loras: List[Tuple[Union[str, Path, Dict[str, torch.Tensor]], float]] = [] diff --git a/py/nodes/prompt.py b/py/nodes/prompt.py index 0c649099..0bebfb26 100644 --- a/py/nodes/prompt.py +++ b/py/nodes/prompt.py @@ -67,7 +67,7 @@ class PromptLM: stack = inspect.stack() if len(stack) > 2 and stack[2].function == "get_input_info": - optional_inputs = _PromptOptionalInputs(optional_inputs) # type: ignore[assignment] + optional_inputs = _PromptOptionalInputs(optional_inputs) # pyright: ignore[reportAssignmentType] return { "required": { @@ -126,7 +126,7 @@ class PromptLM: else: prompt = expanded_text - from nodes import CLIPTextEncode # type: ignore + from nodes import CLIPTextEncode # pyright: ignore[reportMissingImports, reportAttributeAccessIssue] conditioning = CLIPTextEncode().encode(clip, prompt)[0] return (conditioning, prompt) diff --git a/py/nodes/save_image.py b/py/nodes/save_image.py index ddc8d2c2..82e03124 100644 --- a/py/nodes/save_image.py +++ b/py/nodes/save_image.py @@ -5,7 +5,7 @@ import time import uuid from typing import Any, Dict, Optional import numpy as np -import folder_paths # type: ignore +import folder_paths # pyright: ignore[reportMissingImports] from ..services.service_registry import ServiceRegistry from ..metadata_collector.metadata_processor import MetadataProcessor from ..metadata_collector import get_metadata @@ -13,7 +13,7 @@ from ..utils.constants import CARD_PREVIEW_WIDTH from ..utils.exif_utils import ExifUtils from ..utils.utils import calculate_recipe_fingerprint, sanitize_folder_name from PIL import Image, PngImagePlugin -import piexif +import piexif # pyright: ignore[reportMissingTypeStubs] import logging # Civitai-compatible sampler name mapping: ComfyUI internal → A1111 display name @@ -355,7 +355,7 @@ class SaveImageLM: type_lower = model_type.lower() if model_type else "other" return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}" - def format_metadata(self, metadata_dict: dict, add_loras_to_prompt: bool = False) -> str: + def format_metadata(self, metadata_dict: dict[str, Any], add_loras_to_prompt: bool = False) -> str: """Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources.""" if not metadata_dict: return "" @@ -396,7 +396,7 @@ class SaveImageLM: ckpt_display_name = os.path.splitext(os.path.basename(checkpoint))[0] # Resolve LoRA hash and Civitai data from local cache - loras_data: list[dict] = [] + loras_data: list[dict[str, Any]] = [] for lora_name, strength in lora_entries: lora_hash, lora_civitai, lora_base_model = self._resolve_model_cache_entry( "lora_scanner", lora_name @@ -418,9 +418,9 @@ class SaveImageLM: hashes[f"LORA:{lora['name']}"] = lora["hash"][:10].upper() # Build Civitai resources JSON array - civitai_resources: list[dict] = [] + civitai_resources: list[dict[str, Any]] = [] if ckpt_civitai.get("id", 0) > 0: - ckpt_resource: dict = {} + ckpt_resource: dict[str, Any] = {} ckpt_type = (ckpt_civitai.get("model") or {}).get("type", "Checkpoint") model_id = ckpt_civitai.get("modelId", 0) version_id = ckpt_civitai.get("id", 0) @@ -439,7 +439,7 @@ class SaveImageLM: lora_civitai = lora["civitai"] if not lora_civitai or lora_civitai.get("id", 0) <= 0: continue - lora_resource: dict = {"weight": lora["strength"]} + lora_resource: dict[str, Any] = {"weight": lora["strength"]} lora_type = (lora_civitai.get("model") or {}).get("type", "LORA") model_id = lora_civitai.get("modelId", 0) version_id = lora_civitai.get("id", 0) diff --git a/py/nodes/unet_loader.py b/py/nodes/unet_loader.py index d8d909c2..0ab324b5 100644 --- a/py/nodes/unet_loader.py +++ b/py/nodes/unet_loader.py @@ -1,7 +1,7 @@ import logging import os -from typing import List, Tuple -import comfy.sd # type: ignore +from typing import Any, List, Tuple +import comfy.sd # pyright: ignore[reportMissingImports] from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui logger = logging.getLogger(__name__) @@ -34,9 +34,9 @@ class UNETLoaderLM: CATEGORY = "Lora Manager/loaders" @classmethod - def INPUT_TYPES(s): + def INPUT_TYPES(cls): # Get list of unet names from scanner (includes extra folder paths) - unet_names = s._get_unet_names() + unet_names = cls._get_unet_names() return { "required": { "unet_name": ( @@ -105,7 +105,7 @@ class UNETLoaderLM: logger.error(f"Error getting unet names: {e}") return [] - def load_unet(self, unet_name: str, weight_dtype: str) -> Tuple: + def load_unet(self, unet_name: str, weight_dtype: str) -> Tuple[Any, ...]: """Load a diffusion model by name, supporting extra folder paths Args: @@ -148,7 +148,7 @@ class UNETLoaderLM: def _load_gguf_unet( self, unet_path: str, unet_name: str, weight_dtype: str - ) -> Tuple: + ) -> Tuple[Any, ...]: """Load a GGUF format diffusion model Args: diff --git a/py/nodes/utils.py b/py/nodes/utils.py index 12f2fc1e..44a33bc7 100644 --- a/py/nodes/utils.py +++ b/py/nodes/utils.py @@ -1,3 +1,6 @@ +from typing import Any + + class AnyType(str): """A special class that is always equal in not equal comparisons. Credit to pythongosssss""" @@ -6,7 +9,7 @@ class AnyType(str): # Credit to Regis Gaughan, III (rgthree) -class FlexibleOptionalInputType(dict): +class FlexibleOptionalInputType(dict[str, Any]): """A special class to make flexible nodes that pass data to our python handlers. Enables both flexible/dynamic input types (like for Any Switch) or a dynamic number of inputs @@ -23,6 +26,7 @@ class FlexibleOptionalInputType(dict): """ def __init__(self, type): + super().__init__() self.type = type def __getitem__(self, key): @@ -40,7 +44,7 @@ import re import logging import copy import sys -import folder_paths # type: ignore +import folder_paths # pyright: ignore[reportMissingImports] logger = logging.getLogger(__name__) @@ -70,7 +74,7 @@ def extract_lora_name(lora_path): return apply_lora_syntax_format(name_no_ext) -def parse_lora_syntax(text: str) -> list[dict]: +def parse_lora_syntax(text: str) -> list[dict[str, Any]]: """Parse syntax from text input into a list of dicts. Each entry contains: name, model_strength, clip_strength. diff --git a/py/recipes/base.py b/py/recipes/base.py index 036b9dd0..5bff6628 100644 --- a/py/recipes/base.py +++ b/py/recipes/base.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. """Base classes for recipe parsers.""" import json @@ -38,7 +42,7 @@ class RecipeMetadataParser(ABC): pass @staticmethod - async def populate_lora_from_civitai(lora_entry: Dict[str, Any], civitai_info_tuple: Tuple[Dict[str, Any], Optional[str]], + async def populate_lora_from_civitai(lora_entry: Dict[str, Any], civitai_info_tuple: Tuple[Dict[str, Any] | None, str | None] | Dict[str, Any], recipe_scanner=None, base_model_counts=None, hash_value=None) -> Optional[Dict[str, Any]]: """ Populate a lora entry with information from Civitai API response @@ -194,7 +198,7 @@ class RecipeMetadataParser(ABC): return lora_entry @staticmethod - async def populate_checkpoint_from_civitai(checkpoint: Dict[str, Any], civitai_info: Dict[str, Any]) -> Dict[str, Any]: + async def populate_checkpoint_from_civitai(checkpoint: Dict[str, Any], civitai_info: Dict[str, Any] | Tuple[Dict[str, Any] | None, str | None] | None) -> Dict[str, Any]: """ Populate checkpoint information from Civitai API response diff --git a/py/recipes/enrichment.py b/py/recipes/enrichment.py index db1475af..045d20be 100644 --- a/py/recipes/enrichment.py +++ b/py/recipes/enrichment.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. import logging import json import os diff --git a/py/recipes/factory.py b/py/recipes/factory.py index 963ea710..3427d8b7 100644 --- a/py/recipes/factory.py +++ b/py/recipes/factory.py @@ -1,6 +1,7 @@ """Factory for creating recipe metadata parsers.""" import logging +from typing import Any from .parsers import ( RecipeFormatParser, ComfyMetadataParser, @@ -31,7 +32,8 @@ class RecipeParserFactory: # First, try CivitaiApiMetadataParser for dict input if isinstance(metadata, dict): try: - if CivitaiApiMetadataParser().is_metadata_matching(metadata): + user_comment: Any = metadata + if CivitaiApiMetadataParser().is_metadata_matching(user_comment): return CivitaiApiMetadataParser() except Exception as e: logger.debug(f"CivitaiApiMetadataParser check failed: {e}") diff --git a/py/recipes/parsers/automatic.py b/py/recipes/parsers/automatic.py index 19368214..4029f2d5 100644 --- a/py/recipes/parsers/automatic.py +++ b/py/recipes/parsers/automatic.py @@ -52,7 +52,7 @@ class AutomaticMetadataParser(RecipeMetadataParser): negative_and_params = "" # Initialize metadata - metadata = { + metadata: Dict[str, Any] = { "prompt": prompt, "loras": [] } diff --git a/py/recipes/parsers/civitai_image.py b/py/recipes/parsers/civitai_image.py index 1c0efc96..ad9167c1 100644 --- a/py/recipes/parsers/civitai_image.py +++ b/py/recipes/parsers/civitai_image.py @@ -14,15 +14,16 @@ logger = logging.getLogger(__name__) class CivitaiApiMetadataParser(RecipeMetadataParser): """Parser for Civitai image metadata format""" - def is_metadata_matching(self, metadata) -> bool: + def is_metadata_matching(self, user_comment) -> bool: """Check if the metadata matches the Civitai image metadata format Args: - metadata: The metadata from the image (dict) + user_comment: The metadata from the image (dict) Returns: bool: True if this parser can handle the metadata """ + metadata = user_comment if not metadata or not isinstance(metadata, dict): return False @@ -73,7 +74,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): return False - async def parse_metadata( # type: ignore[override] + async def parse_metadata( # pyright: ignore[reportIncompatibleMethodOverride] self, user_comment, recipe_scanner=None, civitai_client=None, local_cache: dict[str, Any] | None = None, ) -> Dict[str, Any]: @@ -89,8 +90,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): Returns: Dict containing parsed recipe data """ - metadata: Dict[str, Any] = user_comment # type: ignore[assignment] - metadata = user_comment + metadata: Dict[str, Any] = user_comment try: # Get metadata provider instead of using civitai_client directly metadata_provider = await get_default_metadata_provider() @@ -116,7 +116,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): metadata = inner_meta # Initialize result structure - result = { + result: Dict[str, Any] = { "base_model": None, "loras": [], "model": None, @@ -125,10 +125,10 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): } # Track already added LoRAs to prevent duplicates - added_loras = {} # key: model_version_id or hash, value: index in result["loras"] + added_loras: Dict[str, Any] = {} # key: model_version_id or hash, value: index in result["loras"] # Extract hash information from hashes field for LoRA matching - lora_hashes = {} + lora_hashes: Dict[str, Any] = {} if "hashes" in metadata and isinstance(metadata["hashes"], dict): for key, hash_value in metadata["hashes"].items(): key_str = str(key) @@ -184,7 +184,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): if model_info: result["base_model"] = model_info.get("baseModel", "") - base_model_counts = {} + base_model_counts: Dict[str, int] = {} # Process standard resources array if "resources" in metadata and isinstance(metadata["resources"], list): @@ -196,7 +196,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser): # identification because it has an explicit type field and hash, # unlike modelVersionIds which is a flat list with no type info. if resource_type == "model": - checkpoint_entry = { + checkpoint_entry: Dict[str, Any] = { "id": 0, "modelId": 0, "name": resource.get("name", "Unknown Model"), diff --git a/py/recipes/parsers/meta_format.py b/py/recipes/parsers/meta_format.py index 2c512103..caae7323 100644 --- a/py/recipes/parsers/meta_format.py +++ b/py/recipes/parsers/meta_format.py @@ -30,7 +30,7 @@ class MetaFormatParser(RecipeMetadataParser): prompt = parts[0].strip() # Initialize metadata - metadata = {"prompt": prompt, "loras": []} + metadata: Dict[str, Any] = {"prompt": prompt, "loras": []} # Extract negative prompt and parameters if available if len(parts) > 1: diff --git a/py/recipes/parsers/recipe_format.py b/py/recipes/parsers/recipe_format.py index ab681ca4..b97eaa60 100644 --- a/py/recipes/parsers/recipe_format.py +++ b/py/recipes/parsers/recipe_format.py @@ -148,7 +148,7 @@ class RecipeFormatParser(RecipeMetadataParser): checkpoint_data = recipe_metadata.get('checkpoint') or {} if isinstance(checkpoint_data, dict) and checkpoint_data: version_id = checkpoint_data.get('modelVersionId') or checkpoint_data.get('id') - checkpoint_entry = { + checkpoint_entry: Dict[str, Any] = { 'id': version_id or 0, 'modelId': checkpoint_data.get('modelId', 0), 'name': checkpoint_data.get('name', 'Unknown Checkpoint'), diff --git a/py/routes/base_model_routes.py b/py/routes/base_model_routes.py index ca2f3d6a..2ebed727 100644 --- a/py/routes/base_model_routes.py +++ b/py/routes/base_model_routes.py @@ -2,7 +2,7 @@ from __future__ import annotations import logging from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Callable, Dict, Mapping +from typing import TYPE_CHECKING, Awaitable, Callable, Dict, Mapping import jinja2 from aiohttp import web @@ -84,7 +84,7 @@ class BaseModelRoutes(ABC): self.metadata_progress_callback = WebSocketBroadcastCallback() self._handler_set: ModelHandlerSet | None = None - self._handler_mapping: Dict[str, Callable[[web.Request], web.StreamResponse]] | None = None + self._handler_mapping: Dict[str, Callable[[web.Request], Awaitable[web.Response]]] | None = None self._preview_service = PreviewAssetService( metadata_manager=MetadataManager, @@ -131,7 +131,7 @@ class BaseModelRoutes(ABC): self._handler_set = None self._handler_mapping = None - def _ensure_handler_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]: + def _ensure_handler_mapping(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]: if self._handler_mapping is None: handler_set = self._create_handler_set() self._handler_set = handler_set @@ -220,7 +220,7 @@ class BaseModelRoutes(ABC): ) @property - def route_handlers(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]: + def route_handlers(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]: return self._ensure_handler_mapping() def setup_routes(self, app: web.Application, prefix: str) -> None: @@ -237,7 +237,7 @@ class BaseModelRoutes(ABC): """Setup model-specific routes.""" raise NotImplementedError - def _parse_specific_params(self, request: web.Request) -> Dict: + def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]: """Parse model-specific parameters - to be overridden by subclasses.""" return {} @@ -253,7 +253,7 @@ class BaseModelRoutes(ABC): """Find the appropriate model file from the files list - can be overridden by subclasses.""" return next((file for file in files if file.get("type") in ("Model", "Diffusion Model") and file.get("primary") is True), None) - def get_handler(self, name: str) -> Callable[[web.Request], web.StreamResponse]: + def get_handler(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]: """Expose handlers for subclasses or tests.""" return self._ensure_handler_mapping()[name] @@ -285,7 +285,7 @@ class BaseModelRoutes(ABC): ) return self.model_lifecycle_service - def _make_handler_proxy(self, name: str) -> Callable[[web.Request], web.StreamResponse]: + def _make_handler_proxy(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]: async def proxy(request: web.Request) -> web.StreamResponse: try: handler = self.get_handler(name) diff --git a/py/routes/base_recipe_routes.py b/py/routes/base_recipe_routes.py index 1b7eaf12..59b7c51d 100644 --- a/py/routes/base_recipe_routes.py +++ b/py/routes/base_recipe_routes.py @@ -4,7 +4,7 @@ from __future__ import annotations import logging import os -from typing import Callable, Mapping +from typing import Awaitable, Callable, Mapping import jinja2 from aiohttp import web @@ -61,7 +61,9 @@ class BaseRecipeRoutes: self._i18n_registered = False self._startup_hooks_registered = False self._handler_set: RecipeHandlerSet | None = None - self._handler_mapping: dict[str, Callable] | None = None + self._handler_mapping: Mapping[ + str, Callable[[web.Request], Awaitable[web.StreamResponse]] + ] | None = None async def attach_dependencies(self, app: web.Application | None = None) -> None: """Resolve shared services from the registry.""" @@ -84,7 +86,9 @@ class BaseRecipeRoutes: app.on_startup.append(self.attach_dependencies) self._startup_hooks_registered = True - def to_route_mapping(self) -> Mapping[str, Callable]: + def to_route_mapping( + self, + ) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]: """Return a mapping of handler name to coroutine for registrar binding.""" if self._handler_mapping is None: @@ -124,17 +128,17 @@ class BaseRecipeRoutes: or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0" ) if not standalone_mode: - from ..metadata_collector import get_metadata # type: ignore[import-not-found] - from ..metadata_collector.metadata_processor import ( # type: ignore[import-not-found] + from ..metadata_collector import get_metadata # pyright: ignore[reportMissingImports] + from ..metadata_collector.metadata_processor import ( # pyright: ignore[reportMissingImports] MetadataProcessor, ) - from ..metadata_collector.metadata_registry import ( # type: ignore[import-not-found] + from ..metadata_collector.metadata_registry import ( # pyright: ignore[reportMissingImports] MetadataRegistry, ) else: # pragma: no cover - optional dependency path - get_metadata = None # type: ignore[assignment] - MetadataProcessor = None # type: ignore[assignment] - MetadataRegistry = None # type: ignore[assignment] + get_metadata = None # pyright: ignore[reportAssignmentType] + MetadataProcessor = None # pyright: ignore[reportAssignmentType] + MetadataRegistry = None # pyright: ignore[reportAssignmentType] analysis_service = RecipeAnalysisService( exif_utils=ExifUtils, diff --git a/py/routes/checkpoint_routes.py b/py/routes/checkpoint_routes.py index 49b8ad5d..28587696 100644 --- a/py/routes/checkpoint_routes.py +++ b/py/routes/checkpoint_routes.py @@ -1,5 +1,5 @@ import logging -from typing import Dict, List, Set +from typing import Any, Dict, List, Set from aiohttp import web from .base_model_routes import BaseModelRoutes @@ -28,13 +28,13 @@ class CheckpointRoutes(BaseModelRoutes): # Attach service dependencies self.attach_service(self.service) - def setup_routes(self, app: web.Application): + def setup_routes(self, app: web.Application, prefix: str = "checkpoints"): """Setup Checkpoint routes""" # Schedule service initialization on app startup app.on_startup.append(lambda _: self.initialize_services()) - + # Setup common routes with 'checkpoints' prefix (includes page route) - super().setup_routes(app, 'checkpoints') + super().setup_routes(app, prefix) def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str): """Setup Checkpoint-specific routes""" @@ -53,9 +53,9 @@ class CheckpointRoutes(BaseModelRoutes): """Get expected model types string for error messages""" return "Checkpoint" - def _parse_specific_params(self, request: web.Request) -> Dict: + def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]: """Parse Checkpoint-specific parameters""" - params: Dict = {} + params: Dict[str, Any] = {} if 'checkpoint_hash' in request.query: params['hash_filters'] = {'single_hash': request.query['checkpoint_hash'].lower()} @@ -70,7 +70,7 @@ class CheckpointRoutes(BaseModelRoutes): """Get detailed information for a specific checkpoint by name""" try: name = request.match_info.get('name', '') - checkpoint_info = await self.service.get_model_info_by_name(name) + checkpoint_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue] if checkpoint_info: return web.json_response(checkpoint_info) @@ -89,7 +89,7 @@ class CheckpointRoutes(BaseModelRoutes): roots.extend(config.checkpoints_roots or []) roots.extend(config.extra_checkpoints_roots or []) # Remove duplicates while preserving order - seen: set = set() + seen: set[str] = set() unique_roots: List[str] = [] for root in roots: if root and root not in seen: @@ -114,7 +114,7 @@ class CheckpointRoutes(BaseModelRoutes): roots.extend(config.unet_roots or []) roots.extend(config.extra_unet_roots or []) # Remove duplicates while preserving order - seen: set = set() + seen: set[str] = set() unique_roots: List[str] = [] for root in roots: if root and root not in seen: diff --git a/py/routes/embedding_routes.py b/py/routes/embedding_routes.py index 5268dc4a..ac00374e 100644 --- a/py/routes/embedding_routes.py +++ b/py/routes/embedding_routes.py @@ -26,13 +26,13 @@ class EmbeddingRoutes(BaseModelRoutes): # Attach service dependencies self.attach_service(self.service) - def setup_routes(self, app: web.Application): + def setup_routes(self, app: web.Application, prefix: str = "embeddings"): """Setup Embedding routes""" # Schedule service initialization on app startup app.on_startup.append(lambda _: self.initialize_services()) - + # Setup common routes with 'embeddings' prefix (includes page route) - super().setup_routes(app, 'embeddings') + super().setup_routes(app, prefix) def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str): """Setup Embedding-specific routes""" @@ -51,7 +51,7 @@ class EmbeddingRoutes(BaseModelRoutes): """Get detailed information for a specific embedding by name""" try: name = request.match_info.get('name', '') - embedding_info = await self.service.get_model_info_by_name(name) + embedding_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue] if embedding_info: return web.json_response(embedding_info) diff --git a/py/routes/example_images_routes.py b/py/routes/example_images_routes.py index aed4e0fd..3ce68e33 100644 --- a/py/routes/example_images_routes.py +++ b/py/routes/example_images_routes.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from typing import Callable, Mapping +from typing import Any, Awaitable, Callable, Mapping from aiohttp import web @@ -35,7 +35,7 @@ class ExampleImagesRoutes: *, ws_manager, download_manager: DownloadManager | None = None, - processor=ExampleImagesProcessor, + processor: Any = ExampleImagesProcessor, file_manager=ExampleImagesFileManager, cleanup_service: ExampleImagesCleanupService | None = None, ) -> None: @@ -46,7 +46,9 @@ class ExampleImagesRoutes: self._file_manager = file_manager self._cleanup_service = cleanup_service or ExampleImagesCleanupService() self._handler_set: ExampleImagesHandlerSet | None = None - self._handler_mapping: Mapping[str, Callable[[web.Request], web.StreamResponse]] | None = None + self._handler_mapping: Mapping[ + str, Callable[[web.Request], Awaitable[web.StreamResponse]] + ] | None = None @classmethod def setup_routes(cls, app: web.Application, *, ws_manager) -> None: @@ -61,7 +63,9 @@ class ExampleImagesRoutes: registrar = ExampleImagesRouteRegistrar(app) registrar.register_routes(self.to_route_mapping()) - def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]: + def to_route_mapping( + self, + ) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]: """Return the registrar-compatible mapping of handler names to callables.""" if self._handler_mapping is None: diff --git a/py/routes/handlers/example_images_handlers.py b/py/routes/handlers/example_images_handlers.py index 3100fee2..a9a78f96 100644 --- a/py/routes/handlers/example_images_handlers.py +++ b/py/routes/handlers/example_images_handlers.py @@ -3,7 +3,7 @@ from __future__ import annotations import logging from dataclasses import dataclass -from typing import Callable, Mapping +from typing import Awaitable, Callable, Mapping from aiohttp import web @@ -170,7 +170,7 @@ class ExampleImagesHandlerSet: management: ExampleImagesManagementHandler files: ExampleImagesFileHandler - def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]: + def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]: """Flatten handler methods into the registrar mapping.""" return { diff --git a/py/routes/handlers/misc_handlers.py b/py/routes/handlers/misc_handlers.py index 6cc2da27..91526ec4 100644 --- a/py/routes/handlers/misc_handlers.py +++ b/py/routes/handlers/misc_handlers.py @@ -276,7 +276,7 @@ def _collect_comfyui_session_logs( ) -> dict[str, Any]: if log_entries is None: try: - import app.logger as comfy_logger + import app.logger as comfy_logger # pyright: ignore[reportMissingImports] log_entries = list(comfy_logger.get_logs() or []) except Exception as exc: # pragma: no cover - environment dependent @@ -422,10 +422,10 @@ class PromptServerProtocol(Protocol): """Subset of PromptServer used by the handlers.""" instance: "PromptServerProtocol" - sockets: dict # maps clientId (sid) → WebSocketResponse + sockets: dict[str, Any] # maps clientId (sid) → WebSocketResponse def send_sync( - self, event: str, payload: dict | None = None, sid: str | None = None + self, event: str, payload: dict[str, Any] | None = None, sid: str | None = None ) -> None: # pragma: no cover - protocol ... @@ -443,7 +443,12 @@ class UsageStatsFactory(Protocol): class MetadataProviderProtocol(Protocol): async def get_model_versions( self, model_id: int - ) -> dict | None: # pragma: no cover - protocol + ) -> dict[str, Any] | None: # pragma: no cover - protocol + ... + + async def get_user_models( + self, username: str, cursor: str | None = None + ) -> Any: # pragma: no cover - protocol ... @@ -466,16 +471,16 @@ class MetadataArchiveManagerProtocol(Protocol): class BackupServiceProtocol(Protocol): async def create_snapshot( self, *, snapshot_type: str = "manual", persist: bool = False - ) -> dict: # pragma: no cover - protocol + ) -> dict[str, Any]: # pragma: no cover - protocol ... - async def restore_snapshot(self, archive_path: str) -> dict: # pragma: no cover - protocol + async def restore_snapshot(self, archive_path: str) -> dict[str, Any]: # pragma: no cover - protocol ... - def get_status(self) -> dict: # pragma: no cover - protocol + def get_status(self) -> dict[str, Any]: # pragma: no cover - protocol ... - def get_available_snapshots(self) -> list[dict]: # pragma: no cover - protocol + def get_available_snapshots(self) -> list[dict[str, Any]]: # pragma: no cover - protocol ... @@ -491,7 +496,7 @@ class NodeRegistry: def __init__(self) -> None: self._lock = asyncio.Lock() # sid → {unique_id → node_info} - self._tab_nodes: Dict[str, Dict[str, dict]] = {} + self._tab_nodes: Dict[str, Dict[str, dict[str, Any]]] = {} self._ready = asyncio.Event() self._waiting_clients: set[str] = set() @@ -504,7 +509,7 @@ class NodeRegistry: # Helpers to build one node dict (extracted so it's reused for each tab) # ------------------------------------------------------------------ @staticmethod - def _build_node_dict(node: dict) -> dict: + def _build_node_dict(node: dict[str, Any]) -> dict[str, Any]: node_id = node["node_id"] graph_id = str(node["graph_id"]) unique_id = f"{graph_id}:{node_id}" @@ -513,11 +518,11 @@ class NodeRegistry: bgcolor = node.get("bgcolor") or DEFAULT_NODE_COLOR raw_capabilities = node.get("capabilities") - capabilities: dict = {} + capabilities: dict[str, Any] = {} if isinstance(raw_capabilities, dict): capabilities = dict(raw_capabilities) - raw_widget_names: list | None = node.get("widget_names") + raw_widget_names: list[Any] | None = node.get("widget_names") if not isinstance(raw_widget_names, list): capability_widget_names = capabilities.get("widget_names") raw_widget_names = ( @@ -565,9 +570,9 @@ class NodeRegistry: # ------------------------------------------------------------------ # Public API # ------------------------------------------------------------------ - async def register_nodes(self, sid: str, nodes: list[dict]) -> None: + async def register_nodes(self, sid: str, nodes: list[dict[str, Any]]) -> None: """Register/replace the node list for a single ComfyUI tab (identified by *sid*).""" - tab_nodes: dict[str, dict] = {} + tab_nodes: dict[str, dict[str, Any]] = {} for node in nodes: nd = self._build_node_dict(node) tab_nodes[nd["unique_id"]] = nd @@ -602,7 +607,7 @@ class NodeRegistry: except asyncio.TimeoutError: return False - async def get_merged_registry(self, active_sids: set[str] | None = None) -> dict: + async def get_merged_registry(self, active_sids: set[str] | None = None) -> dict[str, Any]: """Return the union of all known tab nodes, pruning any tab that is no longer connected.""" async with self._lock: @@ -619,8 +624,8 @@ class NodeRegistry: len(stale_sids), stale_sids, ) - merged: dict[str, dict] = {} - tab_info: dict[str, dict] = {} + merged: dict[str, dict[str, Any]] = {} + tab_info: dict[str, dict[str, Any]] = {} for sid, nodes in self._tab_nodes.items(): tab_info[sid] = { "node_count": len(nodes), @@ -653,7 +658,7 @@ class SupportersHandler: def __init__(self, logger: logging.Logger | None = None) -> None: self._logger = logger or logging.getLogger(__name__) - def _load_supporters(self) -> dict: + def _load_supporters(self) -> dict[str, Any]: """Load supporters data from JSON file.""" try: current_file = os.path.abspath(__file__) @@ -1229,10 +1234,8 @@ class DoctorHandler: settings_snapshot = _sanitize_sensitive_data( getattr(self._settings, "settings", {}) or {} ) - startup_messages_getter = getattr(self._settings, "get_startup_messages", None) - startup_messages = ( - list(startup_messages_getter()) if callable(startup_messages_getter) else [] - ) + startup_messages_getter: Any = getattr(self._settings, "get_startup_messages", None) + startup_messages = list(startup_messages_getter()) if startup_messages_getter else [] environment = { "app_version": app_version, @@ -1439,7 +1442,7 @@ class SettingsHandler: *, settings_service=None, metadata_provider_updater: Callable[ - [], Awaitable[None] + [], Awaitable[Any] ] = update_metadata_providers, downloader_factory: Callable[ [], Awaitable[DownloaderProtocol] @@ -1484,8 +1487,8 @@ class SettingsHandler: settings_file = getattr(self._settings, "settings_file", None) if settings_file: response_data["settings_file"] = settings_file - messages_getter = getattr(self._settings, "get_startup_messages", None) - messages = list(messages_getter()) if callable(messages_getter) else [] + messages_getter: Any = getattr(self._settings, "get_startup_messages", None) + messages = list(messages_getter()) if messages_getter else [] return web.json_response( { "success": True, @@ -2005,11 +2008,11 @@ async def _noop_backup_service() -> None: @dataclass class ServiceRegistryAdapter: - get_lora_scanner: Callable[[], Awaitable] - get_checkpoint_scanner: Callable[[], Awaitable] - get_embedding_scanner: Callable[[], Awaitable] - get_downloaded_version_history_service: Callable[[], Awaitable] - get_backup_service: Callable[[], Awaitable] = _noop_backup_service + get_lora_scanner: Callable[[], Awaitable[Any]] + get_checkpoint_scanner: Callable[[], Awaitable[Any]] + get_embedding_scanner: Callable[[], Awaitable[Any]] + get_downloaded_version_history_service: Callable[[], Awaitable[Any]] + get_backup_service: Callable[[], Awaitable[Any]] = _noop_backup_service class ModelLibraryHandler: @@ -2050,8 +2053,8 @@ class ModelLibraryHandler: return await self._service_registry.get_downloaded_version_history_service() @staticmethod - def _with_downloaded_flag(versions: list[dict]) -> list[dict]: - enriched: list[dict] = [] + def _with_downloaded_flag(versions: list[dict[str, Any]]) -> list[dict[str, Any]]: + enriched: list[dict[str, Any]] = [] for version in versions: entry = dict(version) entry.setdefault("hasBeenDownloaded", True) @@ -2244,7 +2247,7 @@ class ModelLibraryHandler: checkpoint_scanner = await self._service_registry.get_checkpoint_scanner() embedding_scanner = await self._service_registry.get_embedding_scanner() - results: list[dict] = [] + results: list[dict[str, Any]] = [] for model_id in model_ids: lora_versions = await lora_scanner.get_model_versions_by_id(model_id) if lora_versions: @@ -2353,7 +2356,7 @@ class ModelLibraryHandler: ) try: - model_version_id = int(data.get("modelVersionId")) + model_version_id = int(data.get("modelVersionId")) # pyright: ignore[reportArgumentType] except (TypeError, ValueError): return web.json_response( {"success": False, "error": "Parameter modelVersionId must be an integer"}, @@ -2465,10 +2468,10 @@ class ModelLibraryHandler: "checkpoint": checkpoint_scanner, "embedding": embedding_scanner, } - scanner = scanner_map.get(found_type) + scanner = scanner_map.get(found_type or "") if scanner: - persist = getattr(scanner, "_persist_current_cache", None) - if callable(persist): + persist: Any = getattr(scanner, "_persist_current_cache", None) + if persist: await persist() history_service = await self._get_download_history_service() @@ -2649,13 +2652,13 @@ class ModelLibraryHandler: } lora_type_aliases = {model_type.lower() for model_type in VALID_LORA_TYPES} - type_scanner_map: Dict[str, object | None] = { + type_scanner_map: Dict[str, Any] = { **{alias: lora_scanner for alias in lora_type_aliases}, "checkpoint": checkpoint_scanner, "textualinversion": embedding_scanner, } - versions: list[dict] = [] + versions: list[dict[str, Any]] = [] history_service = await self._get_download_history_service() model_ids: list[int] = [] model_count = 0 @@ -2707,6 +2710,8 @@ class ModelLibraryHandler: tags_value = model.get("tags") tags = tags_value if isinstance(tags_value, list) else [] model_id = model.get("id") + if model_id is None: + continue try: model_id_int = int(model_id) except (TypeError, ValueError): @@ -2722,6 +2727,8 @@ class ModelLibraryHandler: continue version_id = version.get("id") + if version_id is None: + continue try: version_id_int = int(version_id) except (TypeError, ValueError): @@ -2783,7 +2790,7 @@ class MetadataArchiveHandler: ] = get_metadata_archive_manager, settings_service=None, metadata_provider_updater: Callable[ - [], Awaitable[None] + [], Awaitable[Any] ] = update_metadata_providers, ) -> None: self._metadata_archive_manager_factory = metadata_archive_manager_factory @@ -2930,7 +2937,7 @@ class BackupHandler: if request.content_type.startswith("multipart/"): reader = await request.multipart() - field = await reader.next() + field: Any = await reader.next() uploaded = False while field is not None: if getattr(field, "filename", None): @@ -3549,7 +3556,7 @@ class NodeRegistryHandler: except (TypeError, ValueError): parsed_node_id = node_identifier - payload: dict = { + payload: dict[str, Any] = { "id": parsed_node_id, "value": value, "mode": mode, @@ -3673,7 +3680,7 @@ class NodeRegistryHandler: except (TypeError, ValueError): parsed_node_id = node_identifier - payload: dict = { + payload: dict[str, Any] = { "id": parsed_node_id, "value": value, "mode": mode, @@ -3740,8 +3747,8 @@ class MiscHandlerSet: doctor: DoctorHandler, example_workflows: ExampleWorkflowsHandler, base_model: BaseModelHandlerSet, - hf_handler: HfHandler | None = None, - agent_handler: AgentHandler | None = None, + hf_handler: Any = None, + agent_handler: Any = None, ) -> None: self.health = health self.settings = settings diff --git a/py/routes/handlers/model_handlers.py b/py/routes/handlers/model_handlers.py index 585916a2..385ba993 100644 --- a/py/routes/handlers/model_handlers.py +++ b/py/routes/handlers/model_handlers.py @@ -71,7 +71,7 @@ class ModelPageView: self._server_i18n = server_i18n self._logger = logger - def _load_supporters(self) -> dict: + def _load_supporters(self) -> dict[str, Any]: """Load supporters data from JSON file.""" try: current_file = os.path.abspath(__file__) @@ -152,7 +152,7 @@ class ModelPageView: self._template_env.filters["t"] = ( self._server_i18n.create_template_filter() ) - self._template_env._i18n_filter_added = True # type: ignore[attr-defined] + self._template_env._i18n_filter_added = True # pyright: ignore[reportAttributeAccessIssue] from ...services.llm_service import PROVIDER_PRESETS @@ -199,7 +199,7 @@ class ModelListingHandler: self, *, service, - parse_specific_params: Callable[[web.Request], Dict], + parse_specific_params: Callable[[web.Request], Dict[str, Any]], logger: logging.Logger, ) -> None: self._service = service @@ -287,7 +287,7 @@ class ModelListingHandler: ) return web.json_response({"error": str(exc)}, status=500) - def _parse_common_params(self, request: web.Request) -> Dict: + def _parse_common_params(self, request: web.Request) -> Dict[str, Any]: page = int(request.query.get("page", "1")) page_size = min(int(request.query.get("page_size", "20")), 100) sort_by = request.query.get("sort_by", "name") @@ -658,7 +658,7 @@ class ModelManagementHandler: try: reader = await request.multipart() - field = await reader.next() + field: Any = await reader.next() if field is None or field.name != "preview_file": raise ValueError("Expected 'preview_file' field") content_type = field.headers.get("Content-Type", "image/png") @@ -700,7 +700,7 @@ class ModelManagementHandler: { "success": True, "preview_url": config.get_preview_static_url( - result["preview_path"] + str(result["preview_path"]) ), "preview_nsfw_level": result["preview_nsfw_level"], } @@ -781,7 +781,7 @@ class ModelManagementHandler: result = await self._preview_service.replace_preview( model_path=model_path, - preview_data=preview_data, + preview_data=preview_bytes, content_type=content_type, original_filename=original_filename, nsfw_level=nsfw_level, @@ -793,7 +793,7 @@ class ModelManagementHandler: { "success": True, "preview_url": config.get_preview_static_url( - result["preview_path"] + str(result["preview_path"]) ), "preview_nsfw_level": result["preview_nsfw_level"], } @@ -2060,7 +2060,7 @@ class ModelCivitaiHandler: settings_service: SettingsManager, ws_manager: WebSocketManager, logger: logging.Logger, - metadata_provider_factory: Callable[[], Awaitable], + metadata_provider_factory: Callable[[], Awaitable[Any]], validate_model_type: Callable[[str], bool], expected_model_types: Callable[[], str], find_model_file: Callable[ @@ -2125,7 +2125,7 @@ class ModelCivitaiHandler: downloaded_version_ids = set( await history_service.get_downloaded_version_ids( self._service.model_type, - model_id, + int(model_id), ) ) except Exception as exc: # pragma: no cover - defensive logging @@ -2402,8 +2402,8 @@ class ModelUpdateHandler: self._logger.error("Failed to fetch license info: %s", exc, exc_info=True) return web.json_response({"success": False, "error": str(exc)}, status=500) - updated: List[Dict[str, str]] = [] - errors: List[Dict[str, str]] = [] + updated: List[Dict[str, Any]] = [] + errors: List[Dict[str, Any]] = [] for model_id in model_ids: license_payload = license_map.get(model_id) if not license_payload: @@ -2416,6 +2416,7 @@ class ModelUpdateHandler: model_section = civitai_section.get("model") if not isinstance(model_section, Mapping): model_section = {} + model_section = dict(model_section) model_section.update(resolved_payload) civitai_section["model"] = model_section metadata_payload["civitai"] = civitai_section @@ -2431,7 +2432,7 @@ class ModelUpdateHandler: ) errors.append({"filePath": metadata_path, "error": str(exc)}) - response_payload = {"success": True, "updated": updated} + response_payload: Dict[str, Any] = {"success": True, "updated": updated} missing_model_ids = [mid for mid in model_ids if mid not in license_map] if missing_model_ids: response_payload["missingModelIds"] = missing_model_ids @@ -2780,6 +2781,7 @@ class ModelUpdateHandler: civitai_payload = metadata_payload.get("civitai") if not isinstance(civitai_payload, Mapping): civitai_payload = {} + civitai_payload = dict(civitai_payload) model_payload = civitai_payload.get("model") if not isinstance(model_payload, Mapping): @@ -2824,7 +2826,7 @@ class ModelUpdateHandler: return aggregated - def _extract_target_model_ids(self, payload: Dict) -> Optional[List[int]]: + def _extract_target_model_ids(self, payload: Dict[str, Any]) -> Optional[List[int]]: if not isinstance(payload, Mapping): return None @@ -2852,7 +2854,7 @@ class ModelUpdateHandler: return {} to_dict = getattr(metadata, "to_dict", None) - if callable(to_dict): + if to_dict: try: return to_dict() except Exception: @@ -2863,7 +2865,7 @@ class ModelUpdateHandler: return {} - async def _read_json(self, request: web.Request) -> Dict: + async def _read_json(self, request: web.Request) -> Dict[str, Any]: if not request.can_read_body: return {} try: @@ -2895,7 +2897,7 @@ class ModelUpdateHandler: record, *, version_context: Optional[Dict[int, Dict[str, Any]]] = None, - ) -> Dict: + ) -> Dict[str, Any]: context = version_context or {} # Check user setting for hiding early access versions hide_early_access = False @@ -2924,7 +2926,7 @@ class ModelUpdateHandler: @staticmethod def _serialize_version( version, context: Optional[Dict[str, Any]] - ) -> Dict: + ) -> Dict[str, Any]: context = context or {} preview_override = context.get("preview_override") preview_url = ( diff --git a/py/routes/handlers/recipe_handlers.py b/py/routes/handlers/recipe_handlers.py index 367baa7f..40e2212f 100644 --- a/py/routes/handlers/recipe_handlers.py +++ b/py/routes/handlers/recipe_handlers.py @@ -1082,10 +1082,10 @@ class RecipeManagementHandler: *, image_url: str, name: str, - lora_entries: list, - checkpoint_entry: dict, - gen_params_request: dict, - tags: list, + lora_entries: list[Any], + checkpoint_entry: Dict[str, Any] | None, + gen_params_request: Dict[str, Any] | None, + tags: list[Any], base_model: str, source_path: str, ) -> web.Response: @@ -1678,7 +1678,7 @@ class RecipeManagementHandler: if not provider: return "" - version_info = await provider.get_model_version_info(version_id) + version_info = await provider.get_model_version_info(str(version_id)) if isinstance(version_info, tuple): version_info = version_info[0] @@ -2391,7 +2391,7 @@ class RecipeAnalysisHandler: content_type = request.headers.get("Content-Type", "") if "multipart/form-data" in content_type: reader = await request.multipart() - field = await reader.next() + field: Any = await reader.next() if field is None or field.name != "image": raise RecipeValidationError("No image field found") image_chunks = bytearray() diff --git a/py/routes/lora_routes.py b/py/routes/lora_routes.py index a2df45c7..36ea5d90 100644 --- a/py/routes/lora_routes.py +++ b/py/routes/lora_routes.py @@ -1,8 +1,8 @@ import asyncio import logging from aiohttp import web -from typing import Dict -from server import PromptServer # type: ignore +from typing import Any, Dict +from server import PromptServer # pyright: ignore[reportMissingImports] from .base_model_routes import BaseModelRoutes from .model_route_registrar import ModelRouteRegistrar @@ -31,13 +31,13 @@ class LoraRoutes(BaseModelRoutes): # Attach service dependencies self.attach_service(self.service) - def setup_routes(self, app: web.Application): + def setup_routes(self, app: web.Application, prefix: str = "loras"): """Setup LoRA routes""" # Schedule service initialization on app startup app.on_startup.append(lambda _: self.initialize_services()) # Setup common routes with 'loras' prefix (includes page route) - super().setup_routes(app, "loras") + super().setup_routes(app, prefix) def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str): """Setup LoRA-specific routes""" @@ -73,7 +73,7 @@ class LoraRoutes(BaseModelRoutes): "POST", "/api/lm/{prefix}/get_trigger_words", prefix, self.get_trigger_words ) - def _parse_specific_params(self, request: web.Request) -> Dict: + def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]: """Parse LoRA-specific parameters""" params = {} @@ -119,25 +119,6 @@ class LoraRoutes(BaseModelRoutes): logger.error(f"Error getting letter counts: {e}") return web.json_response({"success": False, "error": str(e)}, status=500) - async def get_lora_notes(self, request: web.Request) -> web.Response: - """Get notes for a specific LoRA file""" - try: - lora_name = request.query.get("name") - if not lora_name: - return web.Response(text="Lora file name is required", status=400) - - notes = await self.service.get_lora_notes(lora_name) - if notes is not None: - return web.json_response({"success": True, "notes": notes}) - else: - return web.json_response( - {"success": False, "error": "LoRA not found in cache"}, status=404 - ) - - except Exception as e: - logger.error(f"Error getting lora notes: {e}", exc_info=True) - return web.json_response({"success": False, "error": str(e)}, status=500) - async def get_lora_trigger_words(self, request: web.Request) -> web.Response: """Get trigger words for a specific LoRA file""" try: @@ -168,52 +149,6 @@ class LoraRoutes(BaseModelRoutes): logger.error(f"Error getting lora usage tips by path: {e}", exc_info=True) return web.json_response({"success": False, "error": str(e)}, status=500) - async def get_lora_preview_url(self, request: web.Request) -> web.Response: - """Get the static preview URL for a LoRA file""" - try: - lora_name = request.query.get("name") - if not lora_name: - return web.Response(text="Lora file name is required", status=400) - - preview_url = await self.service.get_lora_preview_url(lora_name) - if preview_url: - return web.json_response({"success": True, "preview_url": preview_url}) - else: - return web.json_response( - { - "success": False, - "error": "No preview URL found for the specified lora", - }, - status=404, - ) - - except Exception as e: - logger.error(f"Error getting lora preview URL: {e}", exc_info=True) - return web.json_response({"success": False, "error": str(e)}, status=500) - - async def get_lora_civitai_url(self, request: web.Request) -> web.Response: - """Get the Civitai URL for a LoRA file""" - try: - lora_name = request.query.get("name") - if not lora_name: - return web.Response(text="Lora file name is required", status=400) - - result = await self.service.get_lora_civitai_url(lora_name) - if result["civitai_url"]: - return web.json_response({"success": True, **result}) - else: - return web.json_response( - { - "success": False, - "error": "No Civitai data found for the specified lora", - }, - status=404, - ) - - except Exception as e: - logger.error(f"Error getting lora Civitai URL: {e}", exc_info=True) - return web.json_response({"success": False, "error": str(e)}, status=500) - async def get_random_loras(self, request: web.Request) -> web.Response: """Get random LoRAs based on filters and strength ranges""" try: @@ -337,7 +272,7 @@ class LoraRoutes(BaseModelRoutes): graph_identifier = entry.get("graph_id") try: - parsed_node_id = int(node_identifier) + parsed_node_id = int(node_identifier) # pyright: ignore[reportArgumentType] except (TypeError, ValueError): parsed_node_id = node_identifier diff --git a/py/routes/misc_route_registrar.py b/py/routes/misc_route_registrar.py index 5ba051e1..558c2129 100644 --- a/py/routes/misc_route_registrar.py +++ b/py/routes/misc_route_registrar.py @@ -5,7 +5,7 @@ miscellaneous endpoints share a consistent registration flow. """ from dataclasses import dataclass -from typing import Callable, Iterable, Mapping +from typing import Any, Callable, Iterable, Mapping from aiohttp import web @@ -147,7 +147,7 @@ class MiscRouteRegistrar: handler_lookup[definition.handler_name], ) - def _bind(self, method: str, path: str, handler: Callable) -> None: + def _bind(self, method: str, path: str, handler: Callable[..., Any]) -> None: add_method_name = self._METHOD_MAP[method.upper()] add_method = getattr(self._app.router, add_method_name) add_method(path, handler) diff --git a/py/routes/misc_routes.py b/py/routes/misc_routes.py index e0e77013..f572a731 100644 --- a/py/routes/misc_routes.py +++ b/py/routes/misc_routes.py @@ -7,7 +7,7 @@ import os from typing import Awaitable, Callable, Mapping from aiohttp import web -from server import PromptServer # type: ignore +from server import PromptServer # pyright: ignore[reportMissingImports] from ..services.metadata_service import ( get_metadata_archive_manager, diff --git a/py/routes/model_route_registrar.py b/py/routes/model_route_registrar.py index 7e6e130f..935eb4c8 100644 --- a/py/routes/model_route_registrar.py +++ b/py/routes/model_route_registrar.py @@ -3,7 +3,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Callable, Iterable, Mapping +from typing import Any, Callable, Iterable, Mapping from aiohttp import web @@ -174,15 +174,15 @@ class ModelRouteRegistrar: handler_lookup[definition.handler_name], ) - def add_route(self, method: str, path: str, handler: Callable) -> None: + def add_route(self, method: str, path: str, handler: Callable[..., Any]) -> None: self._bind_route(method, path, handler) def add_prefixed_route( - self, method: str, path_template: str, prefix: str, handler: Callable + self, method: str, path_template: str, prefix: str, handler: Callable[..., Any] ) -> None: self._bind_route(method, path_template.replace("{prefix}", prefix), handler) - def _bind_route(self, method: str, path: str, handler: Callable) -> None: + def _bind_route(self, method: str, path: str, handler: Callable[..., Any]) -> None: add_method_name = self._METHOD_MAP[method.upper()] add_method = getattr(self._app.router, add_method_name) add_method(path, handler) diff --git a/py/routes/recipe_route_registrar.py b/py/routes/recipe_route_registrar.py index 375d4180..76dca31c 100644 --- a/py/routes/recipe_route_registrar.py +++ b/py/routes/recipe_route_registrar.py @@ -3,7 +3,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Callable, Mapping +from typing import Any, Callable, Mapping from aiohttp import web @@ -105,7 +105,7 @@ class RecipeRouteRegistrar: handler = handler_lookup[definition.handler_name] self._bind_route(definition.method, definition.path, handler) - def _bind_route(self, method: str, path: str, handler: Callable) -> None: + def _bind_route(self, method: str, path: str, handler: Callable[..., Any]) -> None: add_method_name = self._METHOD_MAP[method.upper()] add_method = getattr(self._app.router, add_method_name) add_method(path, handler) diff --git a/py/routes/stats_routes.py b/py/routes/stats_routes.py index 22efe863..da520ab9 100644 --- a/py/routes/stats_routes.py +++ b/py/routes/stats_routes.py @@ -40,10 +40,11 @@ class StatsRoutes: """Route handlers for Statistics page and API endpoints""" def __init__(self): - self.lora_scanner = None - self.checkpoint_scanner = None - self.embedding_scanner = None - self.usage_stats = None + self.lora_scanner: Any = None + self.checkpoint_scanner: Any = None + self.embedding_scanner: Any = None + self.usage_stats: Any = None + self._i18n_filter_added = False self.template_env = jinja2.Environment( loader=jinja2.FileSystemLoader(config.templates_path), autoescape=True @@ -95,9 +96,9 @@ class StatsRoutes: server_i18n.set_locale(user_language) # 为模板环境添加i18n过滤器 - if not hasattr(self.template_env, '_i18n_filter_added'): + if not self._i18n_filter_added: self.template_env.filters['t'] = server_i18n.create_template_filter() - self.template_env._i18n_filter_added = True + self._i18n_filter_added = True template = self.template_env.get_template('statistics.html') rendered = template.render( @@ -549,7 +550,7 @@ class StatsRoutes: 'error': str(e) }, status=500) - def _count_unused_models(self, models: List[Dict], usage_data: Dict) -> int: + def _count_unused_models(self, models: List[Dict[str, Any]], usage_data: Dict[str, Any]) -> int: """Count models that have never been used""" used_hashes = set(usage_data.keys()) unused_count = 0 @@ -560,7 +561,7 @@ class StatsRoutes: return unused_count - def _get_top_used_models(self, usage_data: Dict, model_map: Dict, limit: int) -> List[Dict]: + def _get_top_used_models(self, usage_data: Dict[str, Any], model_map: Dict[str, Any], limit: int) -> List[Dict[str, Any]]: """Get top used models with their metadata""" sorted_usage = sorted(usage_data.items(), key=lambda x: x[1].get('total', 0), reverse=True) @@ -578,7 +579,7 @@ class StatsRoutes: return top_models - def _get_usage_timeline(self, usage_data: Dict, days: int) -> List[Dict]: + def _get_usage_timeline(self, usage_data: Dict[str, Any], days: int) -> List[Dict[str, Any]]: """Get usage timeline for the past N days""" timeline = [] today = datetime.now() @@ -614,7 +615,7 @@ class StatsRoutes: return list(reversed(timeline)) # Oldest to newest - def _format_size(self, size_bytes: int) -> str: + def _format_size(self, size_bytes: float) -> str: """Format file size in human readable format""" for unit in ['B', 'KB', 'MB', 'GB', 'TB']: if size_bytes < 1024.0: diff --git a/py/routes/update_routes.py b/py/routes/update_routes.py index 163f40bb..84853172 100644 --- a/py/routes/update_routes.py +++ b/py/routes/update_routes.py @@ -6,7 +6,7 @@ import shutil import tempfile import asyncio from aiohttp import web, ClientError -from typing import Dict, List +from typing import Any, Dict, List, cast from ..utils.settings_paths import ensure_settings_file from ..services.downloader import get_downloader @@ -467,9 +467,10 @@ class UpdateRoutes: if not success: logger.error(f"Failed to fetch release info: {data}") return False, "" - - zip_url = data.get("zipball_url") - version = data.get("tag_name", "unknown") + + release_payload = cast(dict[str, Any], data) + zip_url = release_payload.get("zipball_url", "") + version = release_payload.get("tag_name", "unknown") # Download ZIP to temporary file with tempfile.NamedTemporaryFile(delete=False, suffix=".zip") as tmp_zip: @@ -580,9 +581,10 @@ class UpdateRoutes: logger.warning("Failed to fetch GitHub commit: %s", data) return "main", [], 0, "" - commit_sha = data.get('sha', '')[:7] - commit_message = data.get('commit', {}).get('message', '') - commit_date = data.get('commit', {}).get('committer', {}).get('date', '')[:10] + commit_payload = cast(dict[str, Any], data) + commit_sha = commit_payload.get('sha', '')[:7] + commit_message = commit_payload.get('commit', {}).get('message', '') + commit_date = commit_payload.get('commit', {}).get('committer', {}).get('date', '')[:10] version = f"main-{commit_sha}" changelog = [commit_message] if commit_message else [] @@ -598,10 +600,11 @@ class UpdateRoutes: custom_headers={'Accept': 'application/vnd.github+json'} ) if c_ok: - if c_data.get('status') in ('ahead', 'diverged'): - behind_by = c_data.get('ahead_by', 0) + compare_payload = cast(dict[str, Any], c_data) + if compare_payload.get('status') in ('ahead', 'diverged'): + behind_by = compare_payload.get('ahead_by', 0) else: - behind_by = c_data.get('behind_by', 0) + behind_by = compare_payload.get('behind_by', 0) return version, changelog, behind_by, commit_date @@ -706,7 +709,7 @@ class UpdateRoutes: logger.info(f"Successfully updated to {new_version}") return True, new_version - except git.exc.GitError as e: + except git.exc.GitError as e: # pyright: ignore[reportAttributeAccessIssue] logger.error(f"Git error during update: {e}") return False, "" except Exception as e: @@ -767,7 +770,7 @@ class UpdateRoutes: return git_info @staticmethod - async def _get_remote_version() -> tuple[str, List[str], List[Dict]]: + async def _get_remote_version() -> tuple[str, List[str], List[Dict[str, Any]]]: """ Fetch remote version from GitHub Returns: @@ -789,7 +792,7 @@ class UpdateRoutes: # Parse releases releases = [] - for i, release in enumerate(data): + for i, release in enumerate(cast(list[dict[str, Any]], data)): version = release.get('tag_name', '') if not version.startswith('v'): version = f"v{version}" diff --git a/py/services/agent/agent_service.py b/py/services/agent/agent_service.py index ee8dfbbf..c7d39e4f 100644 --- a/py/services/agent/agent_service.py +++ b/py/services/agent/agent_service.py @@ -117,7 +117,7 @@ def _render_prompt(template: str, variables: Dict[str, Any]) -> str: Uses simple regex substitution — no Jinja2 dependency needed. """ - def replace(match: re.Match) -> str: + def replace(match: re.Match[str]) -> str: key = match.group(1).strip() value = variables.get(key, "") if isinstance(value, (dict, list)): diff --git a/py/services/agent/post_processor.py b/py/services/agent/post_processor.py index a13c8975..c863c5ff 100644 --- a/py/services/agent/post_processor.py +++ b/py/services/agent/post_processor.py @@ -295,7 +295,7 @@ class PostProcessor: normalises every tag to lowercase for case-insensitive dedup. """ merged: List[str] = [] - seen: set = set() + seen: set[str] = set() for tag in list(existing) + list(new): t = tag.strip().lower() if t and t not in seen: diff --git a/py/services/agent/skill_registry.py b/py/services/agent/skill_registry.py index d3f2c4ce..8ef099bd 100644 --- a/py/services/agent/skill_registry.py +++ b/py/services/agent/skill_registry.py @@ -49,7 +49,7 @@ _FRONTMATTER_RE = re.compile( ) -def _parse_skill_file(path: Path) -> tuple[dict, str]: +def _parse_skill_file(path: Path) -> tuple[dict[str, Any], str]: """Read a prompt definition file (``prompt.md`` or legacy ``SKILL.md``) and return (frontmatter_dict, body_text). diff --git a/py/services/agent/skills/enrich_hf_metadata/readme_processor.py b/py/services/agent/skills/enrich_hf_metadata/readme_processor.py index b5e8d829..92953800 100644 --- a/py/services/agent/skills/enrich_hf_metadata/readme_processor.py +++ b/py/services/agent/skills/enrich_hf_metadata/readme_processor.py @@ -9,7 +9,7 @@ from __future__ import annotations import html as html_module import re -from typing import List, Tuple +from typing import Any, List, Tuple _REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)") @@ -18,10 +18,10 @@ _REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)") def extract_simple_markdown_images( markdown_text: str, repo: str, - existing_urls: set | None = None, + existing_urls: set[str] | None = None, default_width: int = 512, default_height: int = 512, -) -> list[dict]: +) -> list[dict[str, Any]]: """Extract standalone markdown images from the README body. Matches ``![alt](url)`` on lines that are NOT part of a markdown table @@ -36,8 +36,8 @@ def extract_simple_markdown_images( return [] base_url = f"https://huggingface.co/{repo}/resolve/main" - images: list[dict] = [] - seen_urls: set = set(existing_urls) if existing_urls else set() + images: list[dict[str, Any]] = [] + seen_urls: set[str] = set(existing_urls) if existing_urls else set() # Collect lines that are NOT inside fenced code blocks lines = markdown_text.split("\n") @@ -86,10 +86,10 @@ def extract_simple_markdown_images( def extract_html_img_tags( markdown_text: str, repo: str, - existing_urls: set | None = None, + existing_urls: set[str] | None = None, default_width: int = 512, default_height: int = 512, -) -> list[dict]: +) -> list[dict[str, Any]]: """Extract image URLs from HTML ```` tags in the README. Many HF collection repos (e.g. ``deadman44/Z-Image_LoRA``) use raw HTML @@ -103,8 +103,8 @@ def extract_html_img_tags( return [] base_url = f"https://huggingface.co/{repo}/resolve/main" - images: list[dict] = [] - seen_urls: set = set(existing_urls) if existing_urls else set() + images: list[dict[str, Any]] = [] + seen_urls: set[str] = set(existing_urls) if existing_urls else set() for m in re.finditer( r']*src=\"([^\"]+)\"', @@ -175,7 +175,7 @@ def extract_gallery_images( repo: str, default_width: int = 512, default_height: int = 512, -) -> List[dict]: +) -> List[dict[str, Any]]: """Extract widget/gallery images from the YAML frontmatter of a HF README. Args: @@ -196,7 +196,7 @@ def extract_gallery_images( if not frontmatter: return [] - images: List[dict] = [] + images: List[dict[str, Any]] = [] base_url = f"https://huggingface.co/{repo}/resolve/main" w = default_width or 512 h = default_height or 512 @@ -258,7 +258,7 @@ def extract_gallery_images( text = raw_text if url: - image: dict = { + image: dict[str, Any] = { "url": url, "type": "image", "nsfwLevel": 0, @@ -276,10 +276,10 @@ def extract_gallery_images( def extract_gallery_table_images( markdown_text: str, repo: str, - existing_urls: set | None = None, + existing_urls: set[str] | None = None, default_width: int = 512, default_height: int = 512, -) -> list[dict]: +) -> list[dict[str, Any]]: """Extract images from ``| Preview | Prompt |`` markdown gallery tables. Many HF READMEs include a sample-gallery table in the body (outside @@ -295,8 +295,8 @@ def extract_gallery_table_images( return [] base_url = f"https://huggingface.co/{repo}/resolve/main" - images: list[dict] = [] - seen_urls: set = set(existing_urls) if existing_urls else set() + images: list[dict[str, Any]] = [] + seen_urls: set[str] = set(existing_urls) if existing_urls else set() lines = markdown_text.split("\n") n = len(lines) i = 0 @@ -514,7 +514,7 @@ def _strip_standalone_images(text: str) -> str: URL was stripped entirely, making it impossible for the LLM to return a ``preview_url`` for repos that use HTML ```` tags exclusively. """ - def _img_to_md(match: re.Match) -> str: + def _img_to_md(match: re.Match[str]) -> str: """Convert an ```` tag to markdown image syntax ``![alt](src)``.""" tag = match.group(0) src_m = re.search(r'src="([^"]+)"', tag) or re.search(r"src='([^']+)'", tag) @@ -942,7 +942,7 @@ def _strip_badge_images(text: str) -> str: "twitter", "colab", "gradio", "space", ) - def _should_remove(m: re.Match) -> str: + def _should_remove(m: re.Match[str]) -> str: alt = (m.group(1) or "").lower() for kw in badge_keywords: if kw in alt: diff --git a/py/services/aria2_downloader.py b/py/services/aria2_downloader.py index 5b3aae63..9da16b43 100644 --- a/py/services/aria2_downloader.py +++ b/py/services/aria2_downloader.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. from __future__ import annotations import asyncio @@ -23,7 +27,7 @@ logger = logging.getLogger(__name__) def _try_certifi_ca_path() -> str | None: """Return the certifi CA bundle path if available, else None.""" try: - import certifi # type: ignore[import-untyped] + import certifi # pyright: ignore[reportMissingTypeStubs] path = certifi.where() if os.path.isfile(path): @@ -84,7 +88,7 @@ class Aria2Downloader: self._transfers: Dict[str, Aria2Transfer] = {} self._poll_interval = 0.5 self._state_store = Aria2TransferStateStore() - self._stderr_reader_task: Optional[asyncio.Task] = None + self._stderr_reader_task: Optional[asyncio.Task[Any]] = None @property def is_running(self) -> bool: @@ -190,7 +194,7 @@ class Aria2Downloader: download_id, ) - options: Dict[str, str] = { + options: Dict[str, Any] = { "dir": save_dir, "out": out_name, "continue": "true", diff --git a/py/services/auto_tag_service.py b/py/services/auto_tag_service.py index 89c0f966..1b076dcf 100644 --- a/py/services/auto_tag_service.py +++ b/py/services/auto_tag_service.py @@ -8,7 +8,7 @@ from filename, base_model, and CivitAI version name — no manual tagging requir from __future__ import annotations import re -from typing import Dict, List, Set +from typing import Any, Dict, List, Set # ── Tag category definitions ────────────────────────────────────────── # Each category maps a display label to a regex pattern. @@ -52,7 +52,7 @@ AUTO_TAG_GROUPS = { DEFAULT_ENABLED_GROUPS = {"mode", "video"} -def _collect_sources(model_data: Dict) -> List[str]: +def _collect_sources(model_data: Dict[str, Any]) -> List[str]: """Collect all text sources from model data for tag matching.""" sources: List[str] = [] @@ -73,7 +73,7 @@ def _collect_sources(model_data: Dict) -> List[str]: return sources -def extract_auto_tags(model_data: Dict) -> List[str]: +def extract_auto_tags(model_data: Dict[str, Any]) -> List[str]: """Extract auto-detected tags from model metadata. Uses a two-layer approach: diff --git a/py/services/autov3_backfill_service.py b/py/services/autov3_backfill_service.py index 07b3d98f..60ea965d 100644 --- a/py/services/autov3_backfill_service.py +++ b/py/services/autov3_backfill_service.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. """Backfill the AutoV3 checked state for models loaded from a persisted snapshot. The SQLite persistent cache predates the AutoV3 feature, so entries hydrated @@ -66,7 +70,7 @@ class Autov3BackfillService: # initialize concurrently (lora_manager.py), so a global guard would # silently skip every type but the first to start. Each model type # runs its own backfill; a duplicate trigger for the same type no-ops. - self._running_types: set = set() + self._running_types: set[str] = set() @classmethod def get_instance(cls) -> "Autov3BackfillService": diff --git a/py/services/backup_service.py b/py/services/backup_service.py index c852f7ca..ea568b3a 100644 --- a/py/services/backup_service.py +++ b/py/services/backup_service.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. from __future__ import annotations import asyncio diff --git a/py/services/base_model_service.py b/py/services/base_model_service.py index 524e43b6..f22868b3 100644 --- a/py/services/base_model_service.py +++ b/py/services/base_model_service.py @@ -2,7 +2,7 @@ from abc import ABC, abstractmethod import asyncio import re import random -from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING +from typing import Any, Awaitable, Dict, List, Optional, Type, Union, TYPE_CHECKING, cast import logging import os import time @@ -70,24 +70,24 @@ class BaseModelService(ABC): page: int, page_size: int, sort_by: str = "name", - folder: str = None, - folder_include: list = None, - folder_exclude: list = None, - search: str = None, + folder: str | None = None, + folder_include: list[str] | None = None, + folder_exclude: list[str] | None = None, + search: str | None = None, fuzzy_search: bool = False, - base_models: list = None, - model_types: list = None, + base_models: list[str] | None = None, + model_types: list[str] | None = None, tags: Optional[Dict[str, str]] = None, auto_tags: Optional[Dict[str, str]] = None, - search_options: dict = None, - hash_filters: dict = None, + search_options: dict[str, Any] | None = None, + hash_filters: dict[str, Any] | None = None, favorites_only: bool = False, update_available_only: bool = False, credit_required: Optional[bool] = None, allow_selling_generated_content: Optional[bool] = None, tag_logic: str = "any", **kwargs, - ) -> Dict: + ) -> Dict[str, Any]: """Get paginated and filtered model data""" overall_start = time.perf_counter() @@ -178,8 +178,8 @@ class BaseModelService(ABC): ufs = self.settings.get("version_grouping", "same_base") group_by_base = ufs == "same_base" - model_groups: Dict[Any, List[Dict]] = {} - ungrouped_standalone: List[Dict] = [] + model_groups: Dict[Any, List[Dict[str, Any]]] = {} + ungrouped_standalone: List[Dict[str, Any]] = [] for item in sorted_data: mid = self._extract_group_key(item) if mid is None: @@ -249,7 +249,7 @@ class BaseModelService(ABC): filter_duration = time.perf_counter() - t1 post_filter_count = len(filtered_data) - annotated_for_filter: Optional[List[Dict]] = None + annotated_for_filter: Optional[List[Dict[str, Any]]] = None t2 = time.perf_counter() if update_available_only: annotated_for_filter = await self._annotate_update_flags(filtered_data) @@ -296,11 +296,11 @@ class BaseModelService(ABC): page: int, page_size: int, sort_by: str = "name", - search: str = None, + search: str | None = None, fuzzy_search: bool = False, - search_options: dict = None, + search_options: dict[str, Any] | None = None, **kwargs, - ) -> Dict: + ) -> Dict[str, Any]: """Get paginated excluded model data.""" excluded_paths = list(self.scanner.get_excluded_models()) excluded_entries: List[Dict[str, Any]] = [] @@ -326,7 +326,7 @@ class BaseModelService(ABC): ] persist_current_cache = getattr(self.scanner, "_persist_current_cache", None) if callable(persist_current_cache): - await persist_current_cache() + await cast(Awaitable[Any], persist_current_cache()) excluded_entries = self._sort_entries(excluded_entries, sort_by) @@ -444,11 +444,11 @@ class BaseModelService(ABC): return entry async def _apply_hash_filters( - self, data: List[Dict], hash_filters: Dict - ) -> List[Dict]: + self, data: List[Dict[str, Any]], hash_filters: Dict[str, Any] + ) -> List[Dict[str, Any]]: """Apply hash-based filtering (SHA256 and AutoV3).""" - def matches_hash_set(item: Dict, hash_set: set) -> bool: + def matches_hash_set(item: Dict[str, Any], hash_set: set[str]) -> bool: """Check whether an item matches any hash in the set. Compares the item's ``sha256`` field and its non-empty ``autov3`` @@ -476,18 +476,18 @@ class BaseModelService(ABC): async def _apply_common_filters( self, - data: List[Dict], - folder: str = None, - folder_include: list = None, - folder_exclude: list = None, - base_models: list = None, - model_types: list = None, + data: List[Dict[str, Any]], + folder: str | None = None, + folder_include: list[str] | None = None, + folder_exclude: list[str] | None = None, + base_models: list[str] | None = None, + model_types: list[str] | None = None, tags: Optional[Dict[str, str]] = None, auto_tags: Optional[Dict[str, str]] = None, favorites_only: bool = False, - search_options: dict = None, + search_options: dict[str, Any] | None = None, tag_logic: str = "any", - ) -> List[Dict]: + ) -> List[Dict[str, Any]]: """Apply common filters that work across all model types""" normalized_options = self.search_strategy.normalize_options(search_options) criteria = FilterCriteria( @@ -506,24 +506,24 @@ class BaseModelService(ABC): async def _apply_search_filters( self, - data: List[Dict], + data: List[Dict[str, Any]], search: str, fuzzy_search: bool, - search_options: dict, - ) -> List[Dict]: + search_options: dict[str, Any] | None, + ) -> List[Dict[str, Any]]: """Apply search filtering""" normalized_options = self.search_strategy.normalize_options(search_options) return self.search_strategy.apply( data, search, normalized_options, fuzzy_search ) - async def _apply_specific_filters(self, data: List[Dict], **kwargs) -> List[Dict]: + async def _apply_specific_filters(self, data: List[Dict[str, Any]], **kwargs) -> List[Dict[str, Any]]: """Apply model-specific filters - to be overridden by subclasses if needed""" return data async def _apply_credit_required_filter( - self, data: List[Dict], credit_required: bool - ) -> List[Dict]: + self, data: List[Dict[str, Any]], credit_required: bool + ) -> List[Dict[str, Any]]: """Apply credit required filtering based on license_flags. Args: @@ -553,8 +553,8 @@ class BaseModelService(ABC): return filtered_data async def _apply_allow_selling_filter( - self, data: List[Dict], allow_selling: bool - ) -> List[Dict]: + self, data: List[Dict[str, Any]], allow_selling: bool + ) -> List[Dict[str, Any]]: """Apply allow selling generated content filtering based on license_flags. Args: @@ -586,8 +586,8 @@ class BaseModelService(ABC): async def _annotate_update_flags( self, - items: List[Dict], - ) -> List[Dict]: + items: List[Dict[str, Any]], + ) -> List[Dict[str, Any]]: """Attach an update_available flag to each response item. Items without a civitai model id default to False. @@ -602,7 +602,7 @@ class BaseModelService(ABC): item["update_available"] = False return annotated - id_to_items: Dict[int, List[Dict]] = {} + id_to_items: Dict[int, List[Dict[str, Any]]] = {} ordered_ids: List[int] = [] for item in annotated: model_id = self._extract_model_id(item) @@ -639,7 +639,7 @@ class BaseModelService(ABC): record_method = getattr(self.update_service, "get_records_bulk", None) if callable(record_method): try: - records = await record_method(self.model_type, ordered_ids) + records = await cast(Awaitable[Any], record_method(self.model_type, ordered_ids)) resolved = { model_id: record.has_update(hide_early_access=hide_early_access) for model_id, record in records.items() @@ -659,11 +659,11 @@ class BaseModelService(ABC): bulk_method = getattr(self.update_service, "has_updates_bulk", None) if callable(bulk_method): try: - resolved = await bulk_method( + resolved = await cast(Awaitable[Any], bulk_method( self.model_type, ordered_ids, hide_early_access=hide_early_access, - ) + )) except Exception as exc: logger.error( "Failed to resolve update status in bulk for %s models (%s): %s", @@ -725,7 +725,7 @@ class BaseModelService(ABC): return annotated @staticmethod - def _extract_hf_group_key(item: Dict) -> Optional[str]: + def _extract_hf_group_key(item: Dict[str, Any]) -> Optional[str]: """Extract `hf:{owner}/{repo}` from item's ``hf_url``, or None.""" hf_url = item.get("hf_url") if isinstance(item, dict) else None if not hf_url or not isinstance(hf_url, str): @@ -738,7 +738,7 @@ class BaseModelService(ABC): return f"hf:{m.group(1)}" @staticmethod - def _extract_group_key(item: Dict) -> Union[int, str, None]: + def _extract_group_key(item: Dict[str, Any]) -> Union[int, str, None]: """Return the group identity key: CivitAI modelId (int) or HF repo (str). Preference order: @@ -752,7 +752,7 @@ class BaseModelService(ABC): return BaseModelService._extract_hf_group_key(item) @staticmethod - def _extract_model_id(item: Dict) -> Optional[int]: + def _extract_model_id(item: Dict[str, Any]) -> Optional[int]: civitai = item.get("civitai") if isinstance(item, dict) else None if not isinstance(civitai, dict): return None @@ -765,7 +765,7 @@ class BaseModelService(ABC): return None @staticmethod - def _extract_version_id(item: Dict) -> Optional[int]: + def _extract_version_id(item: Dict[str, Any]) -> Optional[int]: civitai = item.get("civitai") if isinstance(item, dict) else None if not isinstance(civitai, dict): return None @@ -778,7 +778,7 @@ class BaseModelService(ABC): return None @staticmethod - def _extract_base_model(item: Dict) -> Optional[str]: + def _extract_base_model(item: Dict[str, Any]) -> Optional[str]: value = item.get("base_model") if value is None: return None @@ -830,7 +830,7 @@ class BaseModelService(ABC): return highest_by_base - def _paginate(self, data: List[Dict], page: int, page_size: int) -> Dict: + def _paginate(self, data: List[Dict[str, Any]], page: int, page_size: int) -> Dict[str, Any]: """Apply pagination to filtered data""" total_items = len(data) start_idx = (page - 1) * page_size @@ -845,7 +845,7 @@ class BaseModelService(ABC): } @abstractmethod - async def format_response(self, model_data: Dict) -> Optional[Dict]: + async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]: """Format model data for API response - must be implemented by subclasses. Subclasses should return None for corrupted entries so the handler @@ -854,17 +854,17 @@ class BaseModelService(ABC): pass # Common service methods that delegate to scanner - async def get_top_tags(self, limit: int = 20) -> List[Dict]: + async def get_top_tags(self, limit: int = 20) -> List[Dict[str, Any]]: """Get top tags sorted by frequency""" return await self.scanner.get_top_tags(limit) async def search_tags( self, query: str, limit: int = 50 - ) -> List[Dict]: + ) -> List[Dict[str, Any]]: """Search tags by substring, sorted by frequency""" return await self.scanner.search_tags(query, limit) - async def get_base_models(self, limit: int = 20) -> List[Dict]: + async def get_base_models(self, limit: int = 20) -> List[Dict[str, Any]]: """Get base models sorted by frequency""" return await self.scanner.get_base_models(limit) @@ -931,7 +931,7 @@ class BaseModelService(ABC): """Get model root directories""" return self.scanner.get_model_roots() - def filter_civitai_data(self, data: Dict, minimal: bool = False) -> Dict: + def filter_civitai_data(self, data: Dict[str, Any], minimal: bool = False) -> Dict[str, Any]: """Filter relevant fields from CivitAI data""" if not data: return {} @@ -957,7 +957,7 @@ class BaseModelService(ABC): ) return {k: data[k] for k in fields if k in data} - async def get_folder_tree(self, model_root: str) -> Dict: + async def get_folder_tree(self, model_root: str) -> Dict[str, Any]: """Get hierarchical folder tree for a specific model root""" cache = await self.scanner.get_cached_data() @@ -986,7 +986,7 @@ class BaseModelService(ABC): return tree - async def get_unified_folder_tree(self) -> Dict: + async def get_unified_folder_tree(self) -> Dict[str, Any]: """Get unified folder tree across all model roots""" cache = await self.scanner.get_cached_data() @@ -1015,7 +1015,7 @@ class BaseModelService(ABC): return unified_tree - async def get_model_notes(self, model_name: str) -> Optional[dict]: + async def get_model_notes(self, model_name: str) -> Optional[dict[str, Any]]: """Get notes and file_path for a specific model file. Supports both simple names (``OWSMianne_ANIMA_V1``) and full-path @@ -1147,7 +1147,7 @@ class BaseModelService(ABC): return {"civitai_url": None, "model_id": None, "version_id": None} - async def get_model_metadata(self, file_path: str) -> Optional[Dict]: + async def get_model_metadata(self, file_path: str) -> Optional[Dict[str, Any]]: """Load full metadata for a single model. Listing/search endpoints return lightweight cache entries; this method performs @@ -1243,7 +1243,7 @@ class BaseModelService(ABC): return True @staticmethod - def _relative_path_sort_key(relative_path: str, include_terms: List[str]) -> tuple: + def _relative_path_sort_key(relative_path: str, include_terms: List[str]) -> tuple[int, int, int, str]: """Sort paths by how well they satisfy the include tokens. Sorts based on path without extension for consistent ordering. @@ -1276,12 +1276,12 @@ class BaseModelService(ABC): offset: int = 0, *, folder: Optional[str] = None, - folder_include: Optional[list] = None, - folder_exclude: Optional[list] = None, - base_models: Optional[list] = None, - model_types: Optional[list] = None, - tags: Optional[dict] = None, - auto_tags: Optional[dict] = None, + folder_include: Optional[list[str]] = None, + folder_exclude: Optional[list[str]] = None, + base_models: Optional[list[str]] = None, + model_types: Optional[list[str]] = None, + tags: Optional[dict[str, str]] = None, + auto_tags: Optional[dict[str, str]] = None, tag_logic: str = "any", credit_required: Optional[bool] = None, allow_selling_generated_content: Optional[bool] = None, diff --git a/py/services/checkpoint_scanner.py b/py/services/checkpoint_scanner.py index ae4dc8e9..4d1e2900 100644 --- a/py/services/checkpoint_scanner.py +++ b/py/services/checkpoint_scanner.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. import asyncio import json import logging @@ -427,7 +431,7 @@ class CheckpointScanner(ModelScanner): roots.extend(config.extra_checkpoints_roots or []) roots.extend(config.extra_unet_roots or []) # Remove duplicates while preserving order - seen: set = set() + seen: set[str] = set() unique_roots: List[str] = [] for root in roots: if root not in seen: diff --git a/py/services/checkpoint_service.py b/py/services/checkpoint_service.py index 0fd3f257..1c6922e5 100644 --- a/py/services/checkpoint_service.py +++ b/py/services/checkpoint_service.py @@ -1,6 +1,6 @@ import os import logging -from typing import Dict, Optional +from typing import Any, Dict, Optional from .base_model_service import BaseModelService from .auto_tag_service import extract_auto_tags @@ -21,58 +21,58 @@ class CheckpointService(BaseModelService): """ super().__init__("checkpoint", scanner, CheckpointMetadata, update_service=update_service) - async def format_response(self, checkpoint_data: Dict) -> Optional[Dict]: + async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]: """Format Checkpoint data for API response. Returns None when the entry is missing critical fields (corrupted cache row), so the handler layer can filter it out. See issue #730. """ # Guard against corrupted cache entries missing critical fields - file_path = checkpoint_data.get("file_path") + file_path = model_data.get("file_path") if not file_path or not isinstance(file_path, str): logger.warning( "Skipping corrupted checkpoint entry (missing file_path): %s", - checkpoint_data.get("file_name", ""), + model_data.get("file_name", ""), ) return None # Get sub_type from cache entry (new canonical field) - sub_type = checkpoint_data.get("sub_type", "checkpoint") + sub_type = model_data.get("sub_type", "checkpoint") - file_name = checkpoint_data.get("file_name") or "" - model_name = checkpoint_data.get("model_name") or file_name - folder = checkpoint_data.get("folder") or "" + file_name = model_data.get("file_name") or "" + model_name = model_data.get("model_name") or file_name + folder = model_data.get("folder") or "" return { "model_name": model_name, "file_name": file_name, - "preview_url": config.get_preview_static_url(checkpoint_data.get("preview_url", "")), - "preview_nsfw_level": checkpoint_data.get("preview_nsfw_level", 0), - "base_model": checkpoint_data.get("base_model", ""), + "preview_url": config.get_preview_static_url(model_data.get("preview_url", "")), + "preview_nsfw_level": model_data.get("preview_nsfw_level", 0), + "base_model": model_data.get("base_model", ""), "folder": folder, - "sha256": checkpoint_data.get("sha256", ""), + "sha256": model_data.get("sha256", ""), "file_path": file_path.replace(os.sep, "/"), - "file_size": checkpoint_data.get("size", 0), - "modified": checkpoint_data.get("modified", ""), - "tags": checkpoint_data.get("tags", []), - "from_civitai": checkpoint_data.get("from_civitai", True), - "usage_count": checkpoint_data.get("usage_count", 0), - "notes": checkpoint_data.get("notes", ""), + "file_size": model_data.get("size", 0), + "modified": model_data.get("modified", ""), + "tags": model_data.get("tags", []), + "from_civitai": model_data.get("from_civitai", True), + "usage_count": model_data.get("usage_count", 0), + "notes": model_data.get("notes", ""), "sub_type": sub_type, - "favorite": checkpoint_data.get("favorite", False), - "exclude": bool(checkpoint_data.get("exclude", False)), - "update_available": bool(checkpoint_data.get("update_available", False)), - "skip_metadata_refresh": bool(checkpoint_data.get("skip_metadata_refresh", False)), - "civitai": self.filter_civitai_data(checkpoint_data.get("civitai", {}), minimal=True), - "auto_tags": checkpoint_data.get("auto_tags") or extract_auto_tags(checkpoint_data), - "version_count": checkpoint_data.get("version_count"), - "hf_url": checkpoint_data.get("hf_url", ""), + "favorite": model_data.get("favorite", False), + "exclude": bool(model_data.get("exclude", False)), + "update_available": bool(model_data.get("update_available", False)), + "skip_metadata_refresh": bool(model_data.get("skip_metadata_refresh", False)), + "civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True), + "auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data), + "version_count": model_data.get("version_count"), + "hf_url": model_data.get("hf_url", ""), } - def find_duplicate_hashes(self) -> Dict: + def find_duplicate_hashes(self) -> Dict[str, Any]: """Find Checkpoints with duplicate SHA256 hashes""" return self.scanner._hash_index.get_duplicate_hashes() - def find_duplicate_filenames(self) -> Dict: + def find_duplicate_filenames(self) -> Dict[str, Any]: """Find Checkpoints with conflicting filenames""" return self.scanner._hash_index.get_duplicate_filenames() diff --git a/py/services/civarchive_client.py b/py/services/civarchive_client.py index 954633a2..0c300ea7 100644 --- a/py/services/civarchive_client.py +++ b/py/services/civarchive_client.py @@ -1,8 +1,12 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. import json import logging import asyncio from copy import deepcopy -from typing import Optional, Dict, Tuple, List +from typing import Any, Optional, Dict, Tuple, List, cast from .model_metadata_provider import CivArchiveModelMetadataProvider, ModelMetadataProviderManager from .downloader import get_downloader from .errors import RateLimitError @@ -37,8 +41,8 @@ class CivArchiveClient: async def _request_json( self, path: str, - params: Optional[Dict[str, str]] = None - ) -> Tuple[Optional[Dict], Optional[str]]: + params: Optional[Dict[str, Any]] = None + ) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: """Call CivArchive API and return JSON payload""" success, payload = await self._make_request(path, params=params) if not success: @@ -52,12 +56,12 @@ class CivArchiveClient: self, path: str, *, - params: Optional[Dict[str, str]] = None, - ) -> Tuple[bool, Dict | str]: + params: Optional[Dict[str, Any]] = None, + ) -> Tuple[bool, Dict[str, Any] | str]: """Wrapper around downloader.make_request that surfaces rate limits.""" downloader = await get_downloader() - kwargs: Dict[str, Dict[str, str]] = {} + kwargs: Dict[str, Dict[str, Any]] = {} if params: safe_params = {str(key): str(value) for key, value in params.items() if value is not None} if safe_params: @@ -73,10 +77,11 @@ class CivArchiveClient: if payload.provider is None: payload.provider = "civarchive_api" raise payload - return success, payload + # RateLimitError is always raised above, so the returned payload is a dict or str. + return success, cast(Dict[str, Any] | str, payload) @staticmethod - def _normalize_payload(payload: Dict) -> Dict: + def _normalize_payload(payload: Dict[str, Any]) -> Dict[str, Any]: """Unwrap CivArchive responses that wrap content under a data key""" if not isinstance(payload, dict): return {} @@ -86,12 +91,12 @@ class CivArchiveClient: return payload @staticmethod - def _split_context(payload: Dict) -> Tuple[Dict, Dict, List[Dict]]: + def _split_context(payload: Dict[str, Any]) -> Tuple[Dict[str, Any], Dict[str, Any], List[Dict[str, Any]]]: """Separate version payload from surrounding model context""" data = CivArchiveClient._normalize_payload(payload) - context: Dict = {} - fallback_files: List[Dict] = [] - version: Dict = {} + context: Dict[str, Any] = {} + fallback_files: List[Dict[str, Any]] = [] + version: Dict[str, Any] = {} for key, value in data.items(): if key in {"version", "model"}: @@ -115,7 +120,7 @@ class CivArchiveClient: return context, version, fallback_files @staticmethod - def _ensure_list(value) -> List: + def _ensure_list(value: Any) -> List[Any]: if isinstance(value, list): return value if value is None: @@ -123,7 +128,7 @@ class CivArchiveClient: return [value] @staticmethod - def _build_model_info(context: Dict) -> Dict: + def _build_model_info(context: Dict[str, Any]) -> Dict[str, Any]: tags = context.get("tags") if not isinstance(tags, list): tags = list(tags) if isinstance(tags, (set, tuple)) else ([] if tags is None else [tags]) @@ -136,7 +141,7 @@ class CivArchiveClient: } @staticmethod - def _build_creator_info(context: Dict) -> Dict: + def _build_creator_info(context: Dict[str, Any]) -> Dict[str, Any]: username = context.get("creator_username") or context.get("username") or "" image = context.get("creator_image") or context.get("creator_avatar") or "" creator: Dict[str, Optional[str]] = { @@ -150,7 +155,7 @@ class CivArchiveClient: return creator @staticmethod - def _transform_file_entry(file_data: Dict) -> Dict: + def _transform_file_entry(file_data: Dict[str, Any]) -> Dict[str, Any]: mirrors = file_data.get("mirrors") or [] if not isinstance(mirrors, list): mirrors = [mirrors] @@ -165,7 +170,7 @@ class CivArchiveClient: if not name and available_mirror: name = available_mirror.get("filename") - transformed: Dict = { + transformed: Dict[str, Any] = { "id": file_data.get("id"), "sizeKB": file_data.get("sizeKB"), "name": name, @@ -216,23 +221,23 @@ class CivArchiveClient: def _transform_files( self, - files: Optional[List[Dict]], - fallback_files: Optional[List[Dict]] = None - ) -> List[Dict]: - candidates: List[Dict] = [] + files: Optional[List[Dict[str, Any]]], + fallback_files: Optional[List[Dict[str, Any]]] = None + ) -> List[Dict[str, Any]]: + candidates: List[Dict[str, Any]] = [] if isinstance(files, list) and files: candidates = files elif isinstance(fallback_files, list): candidates = fallback_files - transformed_files: List[Dict] = [] + transformed_files: List[Dict[str, Any]] = [] for file_data in candidates: if isinstance(file_data, dict): transformed_files.append(self._transform_file_entry(file_data)) # Sort: .safetensors first, .ckpt second, others last # so the backend fallback (no file_params) prefers safetensors - def _sort_key(f: Dict) -> int: + def _sort_key(f: Dict[str, Any]) -> int: fname = f.get("name") or "" if isinstance(fname, str): lower = fname.lower() @@ -247,10 +252,10 @@ class CivArchiveClient: def _transform_version( self, - context: Dict, - version: Dict, - fallback_files: Optional[List[Dict]] = None - ) -> Optional[Dict]: + context: Dict[str, Any], + version: Dict[str, Any], + fallback_files: Optional[List[Dict[str, Any]]] = None + ) -> Optional[Dict[str, Any]]: if not version: return None @@ -291,7 +296,7 @@ class CivArchiveClient: return version_copy - async def _resolve_version_from_files(self, payload: Dict) -> Optional[Dict]: + async def _resolve_version_from_files(self, payload: Dict[str, Any]) -> Optional[Dict[str, Any]]: """Fallback to fetch version data when only file metadata is available""" data = self._normalize_payload(payload) files = data.get("files") or payload.get("files") or [] @@ -323,7 +328,7 @@ class CivArchiveClient: return resolved return None - async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: """Find model by SHA256 hash value using CivArchive API""" try: payload, error = await self._request_json(f"/sha256/{model_hash.lower()}") @@ -332,12 +337,12 @@ class CivArchiveClient: return None, "Model not found" return None, error - context, version_data, fallback_files = self._split_context(payload) + context, version_data, fallback_files = self._split_context(cast(Dict[str, Any], payload)) transformed = self._transform_version(context, version_data, fallback_files) if transformed: return transformed, None - resolved = await self._resolve_version_from_files(payload) + resolved = await self._resolve_version_from_files(cast(Dict[str, Any], payload)) if resolved: return resolved, None @@ -350,7 +355,7 @@ class CivArchiveClient: logger.error(f"Error fetching CivArchive model by hash {model_hash[:10]}: {e}") return None, str(e) - async def get_model_versions(self, model_id: str) -> Optional[Dict]: + async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]: """Get all versions of a model using CivArchive API""" try: payload, error = await self._request_json(f"/models/{model_id}") @@ -364,7 +369,7 @@ class CivArchiveClient: context, version_data, fallback_files = self._split_context(payload) versions_meta = data.get("versions") or [] - transformed_versions: List[Dict] = [] + transformed_versions: List[Dict[str, Any]] = [] for meta in versions_meta: if not isinstance(meta, dict): continue @@ -381,7 +386,7 @@ class CivArchiveClient: if primary_version: transformed_versions.insert(0, primary_version) - ordered_versions: List[Dict] = [] + ordered_versions: List[Dict[str, Any]] = [] seen_ids = set() for version in transformed_versions: version_id = version.get("id") @@ -402,7 +407,7 @@ class CivArchiveClient: logger.error(f"Error fetching CivArchive model versions for {model_id}: {e}") return None - async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]: + async def get_model_version(self, model_id: int | str | None = None, version_id: int | str | None = None) -> Optional[Dict[str, Any]]: """Get specific model version using CivArchive API Args: @@ -459,7 +464,7 @@ class CivArchiveClient: logger.error(f"Error fetching CivArchive model version via API {model_id}/{version_id}: {e}") return None - async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: """ Fetch model version metadata using a known bogus model lookup CivArchive lacks a direct version lookup API, this uses a workaround (which we handle in the main model request now) diff --git a/py/services/civitai_base_model_service.py b/py/services/civitai_base_model_service.py index 1cd61e50..9b9db658 100644 --- a/py/services/civitai_base_model_service.py +++ b/py/services/civitai_base_model_service.py @@ -283,7 +283,7 @@ class CivitaiBaseModelService: return None if isinstance(result, str): - data = json.loads(result) + data: Any = json.loads(result) else: data = result diff --git a/py/services/civitai_client.py b/py/services/civitai_client.py index 27f581ee..d084ff07 100644 --- a/py/services/civitai_client.py +++ b/py/services/civitai_client.py @@ -1,10 +1,14 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. import asyncio import copy import logging import os import time from collections import OrderedDict -from typing import Any, Optional, Dict, Tuple, List, Sequence +from typing import Any, Optional, Dict, Tuple, List, Sequence, cast from .connectivity_guard import ( OFFLINE_FRIENDLY_MESSAGE, is_expected_offline_error, @@ -58,7 +62,7 @@ class CivitaiClient: # Uses OrderedDict with LRU eviction at MAX_CACHE_ENTRIES to prevent # unbounded growth in long-running server processes. self._version_info_cache: OrderedDict[ - str, Tuple[Optional[Dict], Optional[str]] + str, Tuple[Optional[Dict[str, Any]], Optional[str]] ] = OrderedDict() self._MAX_CACHE_ENTRIES = 500 @@ -72,7 +76,7 @@ class CivitaiClient: *, use_auth: bool = False, **kwargs, - ) -> Tuple[bool, Dict | str]: + ) -> Tuple[bool, Dict[str, Any] | str]: """Wrapper around downloader.make_request that surfaces rate limits, with retry for transient server errors (5xx, Cloudflare 524, network flakiness).""" @@ -86,7 +90,8 @@ class CivitaiClient: **kwargs, ) if success: - return True, result + # RateLimitError is raised below; a successful result is dict or str. + return True, cast(Dict[str, Any] | str, result) if isinstance(result, RateLimitError): if result.provider is None: @@ -126,7 +131,7 @@ class CivitaiClient: return False, "Unexpected error in _make_request" @staticmethod - def _remove_comfy_metadata(model_version: Optional[Dict]) -> None: + def _remove_comfy_metadata(model_version: Optional[Dict[str, Any]]) -> None: """Remove Comfy-specific metadata from model version images.""" if not isinstance(model_version, dict): return @@ -173,7 +178,7 @@ class CivitaiClient: async def get_model_by_hash( self, model_hash: str - ) -> Tuple[Optional[Dict], Optional[str]]: + ) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: try: success, version = await self._make_request( "GET", @@ -220,7 +225,7 @@ class CivitaiClient: # Ensure directory exists os.makedirs(os.path.dirname(save_path), exist_ok=True) with open(save_path, "wb") as f: - f.write(content) + f.write(content if isinstance(content, bytes) else content.encode("utf-8")) return True return False except Exception as e: @@ -275,7 +280,7 @@ class CivitaiClient: return True return False - async def get_model_versions(self, model_id: str) -> Optional[Dict]: + async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]: """Get all versions of a model with local availability info""" try: success, result = await self._make_request( @@ -283,7 +288,7 @@ class CivitaiClient: f"{self.base_url}/models/{model_id}", use_auth=True, ) - if success: + if success and isinstance(result, dict): # Also return model type along with versions return { "modelVersions": result.get("modelVersions", []), @@ -317,7 +322,7 @@ class CivitaiClient: async def get_model_versions_bulk( self, model_ids: Sequence[int] - ) -> Optional[Dict[int, Dict]]: + ) -> Optional[Dict[int, Dict[str, Any]]]: """Fetch model metadata for multiple ids using the batch API.""" deduped: Dict[int, None] = {} @@ -347,13 +352,13 @@ class CivitaiClient: if not isinstance(items, list): return {} - payload: Dict[int, Dict] = {} + payload: Dict[int, Dict[str, Any]] = {} for item in items: if not isinstance(item, dict): continue model_id = item.get("id") try: - normalized_id = int(model_id) + normalized_id = int(cast(Any, model_id)) except (TypeError, ValueError): continue payload[normalized_id] = { @@ -373,8 +378,8 @@ class CivitaiClient: return None async def get_model_version( - self, model_id: int = None, version_id: int = None - ) -> Optional[Dict]: + self, model_id: int | None = None, version_id: int | None = None + ) -> Optional[Dict[str, Any]]: """Get specific model version with additional metadata.""" try: if model_id is None and version_id is not None: @@ -392,7 +397,7 @@ class CivitaiClient: logger.error(f"Error fetching model version: {e}") return None - async def _get_version_by_id_only(self, version_id: int) -> Optional[Dict]: + async def _get_version_by_id_only(self, version_id: int) -> Optional[Dict[str, Any]]: version = await self._fetch_version_by_id(version_id) if version is None: return None @@ -411,7 +416,7 @@ class CivitaiClient: async def _get_version_with_model_id( self, model_id: int, version_id: Optional[int] - ) -> Optional[Dict]: + ) -> Optional[Dict[str, Any]]: model_data = await self._fetch_model_data(model_id) if not model_data: return None @@ -464,20 +469,20 @@ class CivitaiClient: self._remove_comfy_metadata(version) return version - async def _fetch_model_data(self, model_id: int) -> Optional[Dict]: + async def _fetch_model_data(self, model_id: int) -> Optional[Dict[str, Any]]: success, data = await self._make_request( "GET", f"{self.base_url}/models/{model_id}", use_auth=True, ) - if success: + if success and isinstance(data, dict): return data if is_expected_offline_error(data): return None logger.warning(f"Failed to fetch model data for model {model_id}") return None - async def _fetch_version_by_id(self, version_id: Optional[int]) -> Optional[Dict]: + async def _fetch_version_by_id(self, version_id: Optional[int]) -> Optional[Dict[str, Any]]: if version_id is None: return None @@ -486,7 +491,7 @@ class CivitaiClient: f"{self.base_url}/model-versions/{version_id}", use_auth=True, ) - if success: + if success and isinstance(version, dict): return version if is_expected_offline_error(version): return None @@ -494,7 +499,7 @@ class CivitaiClient: logger.warning(f"Failed to fetch version by id {version_id}") return None - async def _fetch_version_by_hash(self, model_hash: Optional[str]) -> Optional[Dict]: + async def _fetch_version_by_hash(self, model_hash: Optional[str]) -> Optional[Dict[str, Any]]: if not model_hash: return None @@ -503,7 +508,7 @@ class CivitaiClient: f"{self.base_url}/model-versions/by-hash/{model_hash}", use_auth=True, ) - if success: + if success and isinstance(version, dict): return version if is_expected_offline_error(version): return None @@ -512,8 +517,8 @@ class CivitaiClient: return None def _select_target_version( - self, model_data: Dict, model_id: int, version_id: Optional[int] - ) -> Optional[Dict]: + self, model_data: Dict[str, Any], model_id: int, version_id: Optional[int] + ) -> Optional[Dict[str, Any]]: model_versions = model_data.get("modelVersions", []) if not model_versions: logger.warning(f"No model versions found for model {model_id}") @@ -532,7 +537,7 @@ class CivitaiClient: return model_versions[0] - def _extract_primary_model_hash(self, version_entry: Dict) -> Optional[str]: + def _extract_primary_model_hash(self, version_entry: Dict[str, Any]) -> Optional[str]: for file_info in version_entry.get("files", []): if file_info.get("type") == "Model" and file_info.get("primary"): hashes = file_info.get("hashes", {}) @@ -542,8 +547,8 @@ class CivitaiClient: return None def _build_version_from_model_data( - self, version_entry: Dict, model_id: int, model_data: Dict - ) -> Dict: + self, version_entry: Dict[str, Any], model_id: int, model_data: Dict[str, Any] + ) -> Dict[str, Any]: version = copy.deepcopy(version_entry) version.pop("index", None) version["modelId"] = model_id @@ -555,7 +560,7 @@ class CivitaiClient: } return version - def _enrich_version_with_model_data(self, version: Dict, model_data: Dict) -> None: + def _enrich_version_with_model_data(self, version: Dict[str, Any], model_data: Dict[str, Any]) -> None: model_info = version.get("model") if not isinstance(model_info, dict): model_info = {} @@ -571,7 +576,7 @@ class CivitaiClient: async def get_model_version_info( self, version_id: str - ) -> Tuple[Optional[Dict], Optional[str]]: + ) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: """Fetch model version metadata from Civitai Args: @@ -596,7 +601,7 @@ class CivitaiClient: logger.debug("Resolving Civitai model version info: %s", url) success, result = await self._make_request("GET", url, use_auth=True) - if success: + if success and isinstance(result, dict): logger.debug("Successfully fetched model version info for: %s", version_id) self._remove_comfy_metadata(result) self._version_info_cache[version_id] = (result, None) @@ -626,7 +631,7 @@ class CivitaiClient: async def get_image_info( self, image_id: str, source_url: str | None = None - ) -> Optional[Dict]: + ) -> Optional[Dict[str, Any]]: """Fetch image information from Civitai API Args: @@ -659,7 +664,7 @@ class CivitaiClient: ) return None - if result and "items" in result and isinstance(result["items"], list): + if isinstance(result, dict) and "items" in result and isinstance(result["items"], list): items = result["items"] for item in items: @@ -699,7 +704,7 @@ class CivitaiClient: async def get_model_versions_by_hashes( self, hashes: List[str] - ) -> Optional[List[Dict]]: + ) -> Optional[List[Dict[str, Any]]]: """Fetch full version details for up to 100 SHA256 hashes via the batch endpoint. Uses POST /api/v1/model-versions/by-hash which returns full version @@ -716,7 +721,7 @@ class CivitaiClient: return [] BATCH_SIZE = 100 - all_versions: List[Dict] = [] + all_versions: List[Dict[str, Any]] = [] for start in range(0, len(hashes), BATCH_SIZE): batch = hashes[start : start + BATCH_SIZE] @@ -736,7 +741,7 @@ class CivitaiClient: continue if isinstance(result, list): - all_versions.extend(result) + all_versions.extend(cast(Any, result)) else: logger.debug( "Unexpected by-hash response type: %s", type(result) diff --git a/py/services/download_coordinator.py b/py/services/download_coordinator.py index ddfc859b..dbbe2e17 100644 --- a/py/services/download_coordinator.py +++ b/py/services/download_coordinator.py @@ -18,7 +18,7 @@ class DownloadCoordinator: self, *, ws_manager, - download_manager_factory: Callable[[], Awaitable], + download_manager_factory: Callable[[], Awaitable[Any]], ) -> None: self._ws_manager = ws_manager self._download_manager_factory = download_manager_factory diff --git a/py/services/download_manager.py b/py/services/download_manager.py index 628555bd..e9b40a9e 100644 --- a/py/services/download_manager.py +++ b/py/services/download_manager.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. import copy import logging import os @@ -8,7 +12,7 @@ import zipfile from concurrent.futures import ThreadPoolExecutor from collections import OrderedDict import uuid -from typing import Dict, List, Optional, Set, Tuple +from typing import Any, Dict, List, Optional, Set, Tuple, cast from urllib.parse import urlparse from ..utils.models import LoraMetadata, CheckpointMetadata, EmbeddingMetadata from ..utils.constants import ( @@ -121,7 +125,7 @@ class DownloadManager: "delay": 0, } ) - except DownloadInProgressError: + except DownloadInProgressError: # pyright: ignore[reportPossiblyUnboundVariable] logger.info( "Skipping automatic example images download for %s; another example images download is already running", model_hash, @@ -170,7 +174,7 @@ class DownloadManager: logger.error("aria2 download failed for %s: %s", download_url, exc) return False, str(exc) - download_kwargs = { + download_kwargs: Dict[str, Any] = { "progress_callback": progress_callback, "use_auth": use_auth, } @@ -204,16 +208,16 @@ class DownloadManager: async def download_from_civitai( self, - model_id: int = None, - model_version_id: int = None, - save_dir: str = None, + model_id: int | None = None, + model_version_id: int | None = None, + save_dir: str | None = None, relative_path: str = "", progress_callback=None, use_default_paths: bool = False, - download_id: str = None, - source: str = None, - file_params: Dict = None, - ) -> Dict: + download_id: str | None = None, + source: str | None = None, + file_params: Dict[str, Any] | None = None, + ) -> Dict[str, Any]: """Download model from Civitai with task tracking and concurrency control Args: @@ -309,14 +313,14 @@ class DownloadManager: async def _download_with_semaphore( self, task_id: str, - model_id: int, - model_version_id: int, - save_dir: str, + model_id: int | None, + model_version_id: int | None, + save_dir: str | None, relative_path: str, progress_callback=None, use_default_paths: bool = False, - source: str = None, - file_params: Dict = None, + source: str | None = None, + file_params: Dict[str, Any] | None = None, ): """Execute download with semaphore to limit concurrency""" # Update status to waiting @@ -380,7 +384,8 @@ class DownloadManager: # Use original download implementation try: # Check for cancellation before starting - if asyncio.current_task().cancelled(): + current_task = asyncio.current_task() + if current_task is not None and current_task.cancelled(): raise asyncio.CancelledError() result = await self._execute_original_download( @@ -484,11 +489,11 @@ class DownloadManager: # Schedule cleanup of download record after delay asyncio.create_task(self._cleanup_download_record(task_id)) - def _start_background_download_task(self, download_id: str, coroutine) -> asyncio.Task: + def _start_background_download_task(self, download_id: str, coroutine) -> asyncio.Task[Any]: task = asyncio.create_task(coroutine) self._download_tasks[download_id] = task - def _cleanup_done_task(done_task: asyncio.Task) -> None: + def _cleanup_done_task(done_task: asyncio.Task[Any]) -> None: current_task = self._download_tasks.get(download_id) if current_task is done_task: self._download_tasks.pop(download_id, None) @@ -530,7 +535,7 @@ class DownloadManager: async def _cleanup_cancelled_download_files( self, download_id: str, - download_info: Optional[Dict], + download_info: Optional[Dict[str, Any]], ) -> None: target_files = set() persisted = await self._aria2_state_store.get(download_id) @@ -603,13 +608,13 @@ class DownloadManager: self, download_id: str, *, - extra: Optional[Dict] = None, + extra: Optional[Dict[str, Any]] = None, ) -> None: info = self._active_downloads.get(download_id) if not info: return - payload = { + payload: Dict[str, Any] = { "download_id": download_id, "model_id": info.get("model_id"), "model_version_id": info.get("model_version_id"), @@ -631,7 +636,7 @@ class DownloadManager: await self._aria2_state_store.upsert(download_id, payload) - def _build_restored_download_info(self, record: Dict, save_path: str) -> Dict: + def _build_restored_download_info(self, record: Dict[str, Any], save_path: str) -> Dict[str, Any]: return { "model_id": record.get("model_id"), "model_version_id": record.get("model_version_id"), @@ -653,8 +658,8 @@ class DownloadManager: def _is_same_aria2_download_request( self, - current_info: Optional[Dict], - persisted_record: Dict, + current_info: Optional[Dict[str, Any]], + persisted_record: Dict[str, Any], ) -> bool: if not isinstance(current_info, dict): return False @@ -666,13 +671,15 @@ class DownloadManager: return current_version_id == persisted_version_id - def _build_download_urls_from_file_info(self, file_info: Dict, source: str = None) -> List[str]: + def _build_download_urls_from_file_info(self, file_info: Dict[str, Any], source: str | None = None) -> List[str]: mirrors = file_info.get("mirrors") or [] download_urls: List[str] = [] if mirrors: for mirror in mirrors: if mirror.get("deletedAt") is None and mirror.get("url"): - download_urls.append(normalize_civitai_download_url(mirror["url"])) + normalized_url = normalize_civitai_download_url(mirror["url"]) + if normalized_url: + download_urls.append(normalized_url) if source == "civarchive" and len(download_urls) > 1: civitai_urls = [ @@ -688,7 +695,9 @@ class DownloadManager: if not download_urls: download_url = file_info.get("downloadUrl") if download_url: - download_urls.append(normalize_civitai_download_url(download_url)) + normalized_url = normalize_civitai_download_url(download_url) + if normalized_url: + download_urls.append(normalized_url) return download_urls @@ -696,8 +705,8 @@ class DownloadManager: self, *, model_type: str, - version_info: Dict, - file_info: Dict, + version_info: Dict[str, Any], + file_info: Dict[str, Any], save_path: str, ): if model_type == "checkpoint": @@ -706,7 +715,7 @@ class DownloadManager: return EmbeddingMetadata.from_civitai_info(version_info, file_info, save_path) return LoraMetadata.from_civitai_info(version_info, file_info, save_path) - def _resolve_save_path_from_persisted_record(self, record: Dict) -> Optional[str]: + def _resolve_save_path_from_persisted_record(self, record: Dict[str, Any]) -> Optional[str]: save_path = record.get("save_path") or record.get("file_path") if isinstance(save_path, str) and save_path: return os.path.abspath(save_path) @@ -728,7 +737,7 @@ class DownloadManager: return os.path.abspath(os.path.join(save_dir, file_name)) - async def _resume_restored_aria2_download(self, download_id: str, record: Dict) -> Dict: + async def _resume_restored_aria2_download(self, download_id: str, record: Dict[str, Any]) -> Dict[str, Any]: try: if download_id in self._active_downloads: self._active_downloads[download_id]["status"] = "downloading" @@ -842,7 +851,7 @@ class DownloadManager: self, previous_download_id: str, new_download_id: str, - persisted_record: Dict, + persisted_record: Dict[str, Any], save_path: str, ) -> None: aria2_downloader = await get_aria2_downloader() @@ -938,7 +947,7 @@ class DownloadManager: except Exception: status_payload = None - if status_payload is not None: + if status_payload is not None and isinstance(gid, str): remote_status = status_payload.get("status", "") if remote_status in {"active", "waiting", "paused"}: await aria2_downloader.restore_transfer(download_id, gid, save_path) @@ -1115,17 +1124,17 @@ class DownloadManager: async def _execute_original_download( self, - model_id, - model_version_id, - save_dir, - relative_path, + model_id: int | None, + model_version_id: int | None, + save_dir: str | None, + relative_path: str, progress_callback, - use_default_paths, - download_id=None, - transfer_backend="python", - source=None, - file_params=None, - ): + use_default_paths: bool, + download_id: str | None = None, + transfer_backend: str = "python", + source: str | None = None, + file_params: Dict[str, Any] | None = None, + ) -> Dict[str, Any]: """Wrapper for original download_from_civitai implementation""" try: # Check if model version already exists in library @@ -1172,7 +1181,7 @@ class DownloadManager: # Get version info based on the provided identifier version_info = await metadata_provider.get_model_version( - model_id, model_version_id + cast(int, model_id), cast(int, model_version_id) ) if not version_info: @@ -1183,7 +1192,7 @@ class DownloadManager: ) metadata_provider = await get_default_metadata_provider() version_info = await metadata_provider.get_model_version( - model_id, model_version_id + cast(int, model_id), cast(int, model_version_id) ) if not version_info: @@ -1388,6 +1397,8 @@ class DownloadManager: relative_path = self._calculate_relative_path(version_info, model_type) # Update save directory with relative path if provided + if not save_dir: + return {"success": False, "error": "No save directory specified"} if relative_path: base_save_dir = save_dir save_dir = os.path.join(save_dir, relative_path) @@ -1561,6 +1572,11 @@ class DownloadManager: version_info, file_info, save_path ) logger.info(f"Creating EmbeddingMetadata for {file_name}") + else: + return { + "success": False, + "error": f'Unsupported model type "{model_type}"', + } # 6. Start download process if transfer_backend == "aria2" and download_id: @@ -1580,7 +1596,7 @@ class DownloadManager: }, ) - execute_kwargs = { + execute_kwargs: Dict[str, Any] = { "download_urls": download_urls, "save_dir": save_dir, "metadata": metadata, @@ -1627,7 +1643,8 @@ class DownloadManager: ) # If early_access_msg exists and download failed, replace error message - if "early_access_msg" in locals() and not result.get("success", False): + early_access_msg = locals().get("early_access_msg") + if early_access_msg and not result.get("success", False): result["error"] = early_access_msg return result @@ -1652,7 +1669,7 @@ class DownloadManager: self, model_type: str, model_id_value, - version_info: Dict, + version_info: Dict[str, Any], fallback_version_id=None, file_path: str | None = None, ) -> None: @@ -1683,8 +1700,8 @@ class DownloadManager: try: await history_service.mark_downloaded( model_type, - int(version_id), - model_id=int(resolved_model_id) if resolved_model_id is not None else None, + int(cast(Any, version_id)), + model_id=int(cast(Any, resolved_model_id)) if resolved_model_id is not None else None, source="download", file_path=file_path, ) @@ -1701,7 +1718,7 @@ class DownloadManager: self, model_type: str, model_id_value, - version_info: Dict, + version_info: Dict[str, Any], fallback_version_id=None, ) -> None: """Ensure update tracking reflects a newly downloaded version.""" @@ -1725,7 +1742,7 @@ class DownloadManager: if isinstance(model_info, dict): resolved_model_id = model_info.get("id") try: - resolved_model_id = int(resolved_model_id) + resolved_model_id = int(cast(Any, resolved_model_id)) except (TypeError, ValueError): logger.debug( "Skipping update sync; invalid model id: %s", resolved_model_id @@ -1736,7 +1753,7 @@ class DownloadManager: if version_id is None: version_id = fallback_version_id try: - version_id = int(version_id) + version_id = int(cast(Any, version_id)) except (TypeError, ValueError): logger.debug( "Skipping update sync; invalid version id for model %s: %s", @@ -1773,7 +1790,7 @@ class DownloadManager: for entry in local_versions or []: vid = entry.get("versionId") try: - version_ids.add(int(vid)) + version_ids.add(int(cast(Any, vid))) except (TypeError, ValueError): continue @@ -1795,7 +1812,7 @@ class DownloadManager: ) def _calculate_relative_path( - self, version_info: Dict, model_type: str = "lora" + self, version_info: Dict[str, Any], model_type: str = "lora" ) -> str: """Calculate relative path using template from settings @@ -1871,21 +1888,22 @@ class DownloadManager: download_urls: List[str], save_dir: str, metadata, - version_info: Dict, + version_info: Dict[str, Any], relative_path: str, progress_callback=None, model_type: str = "lora", - download_id: str = None, + download_id: str | None = None, transfer_backend: Optional[str] = None, - ) -> Dict: + ) -> Dict[str, Any]: """Execute the actual download process including preview images and model files""" - metadata_entries: List = [] + metadata_entries: List[Any] = [] metadata_files_for_cleanup: List[str] = [] extracted_paths: List[str] = [] metadata_path = "" preview_targets: List[str] = [] preview_path: str | None = None preview_nsfw_level = 0 + save_path: str | None = None transfer_backend = (transfer_backend or self._get_model_download_backend()).lower() try: resolved, save_path = await self._resolve_download_target_path( @@ -1933,9 +1951,9 @@ class DownloadManager: mature_threshold=mature_threshold, ) - preview_url = selected_image.get("url") if selected_image else None + preview_url = cast(Optional[str], selected_image.get("url")) if selected_image else None media_type = ( - (selected_image.get("type") or "").lower() if selected_image else "" + cast(str, selected_image.get("type") or "").lower() if selected_image else "" ) def _extension_from_url(url: str, fallback: str) -> str: @@ -1959,9 +1977,10 @@ class DownloadManager: preview_url, media_type="video" ) attempt_urls: List[str] = [] - if rewritten: + if rewritten and rewritten_url: attempt_urls.append(rewritten_url) - attempt_urls.append(preview_url) + if preview_url: + attempt_urls.append(preview_url) seen_attempts = set() for attempt in attempt_urls: @@ -1978,7 +1997,7 @@ class DownloadManager: rewritten_url, rewritten = rewrite_preview_url( preview_url, media_type="image" ) - if rewritten: + if rewritten and rewritten_url: preview_ext = _extension_from_url(preview_url, ".png") preview_path = os.path.splitext(save_path)[0] + preview_ext success, _ = await downloader.download_file( @@ -2004,7 +2023,9 @@ class DownloadManager: ) if success: with open(temp_path, "wb") as temp_file_handle: - temp_file_handle.write(content) + temp_file_handle.write( + content if isinstance(content, bytes) else content.encode("utf-8") + ) preview_path = ( os.path.splitext(save_path)[0] + ".webp" ) @@ -2056,6 +2077,8 @@ class DownloadManager: last_error = None for download_url in download_urls: download_url = normalize_civitai_download_url(download_url) + if download_url is None: + continue use_auth = download_url.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES) if transfer_backend == "aria2" and download_id: await self._persist_aria2_state( @@ -2239,7 +2262,7 @@ class DownloadManager: entry, normalized_file_path, adjust_root ) if adjusted_entry is not None: - entry = adjusted_entry + entry = cast(Any, adjusted_entry) metadata_entries[index] = entry metadata_file_path = ( @@ -2359,11 +2382,11 @@ class DownloadManager: async def _build_metadata_entries( self, base_metadata, file_paths: List[str] - ) -> List: + ) -> List[Any]: if not file_paths: return [] - entries: List = [] + entries: List[Any] = [] for index, file_path in enumerate(file_paths): entry = base_metadata if index == 0 else copy.deepcopy(base_metadata) # Update file paths without modifying size and modified timestamps @@ -2406,7 +2429,7 @@ class DownloadManager: return destination def _distribute_preview_to_entries( - self, preview_path: str, entries: List + self, preview_path: str, entries: List[Any] ) -> List[str]: if not preview_path or not entries: return [] @@ -2465,7 +2488,7 @@ class DownloadManager: progress_callback, normalized_snapshot, rounded_progress ) - async def cancel_download(self, download_id: str) -> Dict: + async def cancel_download(self, download_id: str) -> Dict[str, Any]: """Cancel an active download by download_id Args: @@ -2547,7 +2570,7 @@ class DownloadManager: self._download_tasks.pop(download_id, None) await self._aria2_state_store.remove(download_id) - async def skip_download(self, download_id: str) -> Dict: + async def skip_download(self, download_id: str) -> Dict[str, Any]: """Skip a download while preserving all partial files on disk. Removes all in-memory tracking (asyncio task, semaphore, active/pause @@ -2630,7 +2653,7 @@ class DownloadManager: # Preserve aria2 state store entry so the partial download # info survives restarts and can be resumed later - async def pause_download(self, download_id: str) -> Dict: + async def pause_download(self, download_id: str) -> Dict[str, Any]: """Pause an active download without losing progress.""" await self._restore_persisted_downloads() @@ -2677,7 +2700,7 @@ class DownloadManager: return {"success": True, "message": "Download paused successfully"} - async def resume_download(self, download_id: str) -> Dict: + async def resume_download(self, download_id: str) -> Dict[str, Any]: """Resume a previously paused download.""" await self._restore_persisted_downloads() @@ -2694,7 +2717,7 @@ class DownloadManager: self._pause_events[download_id] = pause_control self._active_downloads[download_id] = self._build_restored_download_info( persisted, - os.path.abspath(save_path), + os.path.abspath(cast(str, save_path)), ) if pause_control.is_set(): @@ -2821,7 +2844,7 @@ class DownloadManager: elif asyncio.iscoroutine(result): await result - async def get_active_downloads(self) -> Dict: + async def get_active_downloads(self) -> Dict[str, Any]: """Get information about all active downloads Returns: diff --git a/py/services/downloaded_version_history_service.py b/py/services/downloaded_version_history_service.py index 083d8fb9..80cc26b2 100644 --- a/py/services/downloaded_version_history_service.py +++ b/py/services/downloaded_version_history_service.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. from __future__ import annotations import asyncio diff --git a/py/services/downloader.py b/py/services/downloader.py index e1823869..3925d093 100644 --- a/py/services/downloader.py +++ b/py/services/downloader.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. """ Unified download manager for all HTTP/HTTPS downloads in the application. @@ -20,7 +24,7 @@ from dataclasses import dataclass from datetime import datetime, timedelta from email.utils import parsedate_to_datetime from urllib.parse import urlparse -from typing import Optional, Dict, Tuple, Callable, Union, Awaitable +from typing import Optional, Dict, Tuple, Callable, Union, Awaitable, Any, cast from ..services.settings_manager import get_settings_manager from .connectivity_guard import ( OFFLINE_COOLDOWN_ERROR, @@ -204,6 +208,7 @@ class Downloader: # Double check after acquiring lock if self._session is None or self._should_refresh_session(): await self._create_session() + assert self._session is not None return self._session @property @@ -231,7 +236,7 @@ class Downloader: ) try: - timeout_value = float(raw_value) + timeout_value = float(cast(Any, raw_value)) except (TypeError, ValueError): timeout_value = default_timeout @@ -243,7 +248,7 @@ class Downloader: raw_value = os.environ.get("COMFYUI_DOWNLOAD_MAX_RETRIES") try: - retries = int(raw_value) + retries = int(cast(Any, raw_value)) except (TypeError, ValueError): retries = default_retries @@ -320,7 +325,7 @@ class Downloader: # CA coverage across different Python environments (especially # embedded/compatibility Python builds). try: - import certifi # type: ignore[import-untyped] + import certifi # pyright: ignore[reportMissingTypeStubs] ca_path = certifi.where() ssl_context = ssl.create_default_context(cafile=ca_path) @@ -330,7 +335,7 @@ class Downloader: logger.debug("SSL: certifi unavailable; using system default CA bundle") # Optimize TCP connection parameters - connector_kwargs = dict( + connector_kwargs: Dict[str, Any] = dict( ssl=ssl_context, limit=8, # Concurrent connections ttl_dns_cache=300, # DNS cache timeout @@ -890,7 +895,7 @@ class Downloader: use_auth: bool = False, custom_headers: Optional[Dict[str, str]] = None, return_headers: bool = False, - ) -> Tuple[bool, Union[bytes, str], Optional[Dict]]: + ) -> Tuple[bool, Union[bytes, str], Optional[Dict[str, Any]]]: """ Download a file to memory (for small files like preview images) @@ -976,7 +981,7 @@ class Downloader: url: str, use_auth: bool = False, custom_headers: Optional[Dict[str, str]] = None, - ) -> Tuple[bool, Union[Dict, str]]: + ) -> Tuple[bool, Union[Dict[str, Any], str]]: """ Get response headers without downloading the full content @@ -1036,7 +1041,7 @@ class Downloader: use_auth: bool = False, custom_headers: Optional[Dict[str, str]] = None, **kwargs, - ) -> Tuple[bool, Union[Dict, str]]: + ) -> Tuple[bool, Union[Dict[str, Any], str, RateLimitError]]: """ Make a generic HTTP request and return JSON response diff --git a/py/services/embedding_scanner.py b/py/services/embedding_scanner.py index cfb4ef2d..3341e1ec 100644 --- a/py/services/embedding_scanner.py +++ b/py/services/embedding_scanner.py @@ -27,7 +27,7 @@ class EmbeddingScanner(ModelScanner): roots.extend(config.embeddings_roots or []) roots.extend(config.extra_embeddings_roots or []) # Remove duplicates while preserving order - seen: set = set() + seen: set[str] = set() unique_roots: List[str] = [] for root in roots: if root and root not in seen: diff --git a/py/services/embedding_service.py b/py/services/embedding_service.py index f668ec8b..7913f581 100644 --- a/py/services/embedding_service.py +++ b/py/services/embedding_service.py @@ -1,6 +1,6 @@ import os import logging -from typing import Dict, Optional +from typing import Any, Dict, Optional from .base_model_service import BaseModelService from .auto_tag_service import extract_auto_tags @@ -21,58 +21,58 @@ class EmbeddingService(BaseModelService): """ super().__init__("embedding", scanner, EmbeddingMetadata, update_service=update_service) - async def format_response(self, embedding_data: Dict) -> Optional[Dict]: + async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]: """Format Embedding data for API response. Returns None when the entry is missing critical fields (corrupted cache row), so the handler layer can filter it out. See issue #730. """ # Guard against corrupted cache entries missing critical fields - file_path = embedding_data.get("file_path") + file_path = model_data.get("file_path") if not file_path or not isinstance(file_path, str): logger.warning( "Skipping corrupted embedding entry (missing file_path): %s", - embedding_data.get("file_name", ""), + model_data.get("file_name", ""), ) return None # Get sub_type from cache entry (new canonical field) - sub_type = embedding_data.get("sub_type", "embedding") + sub_type = model_data.get("sub_type", "embedding") - file_name = embedding_data.get("file_name") or "" - model_name = embedding_data.get("model_name") or file_name - folder = embedding_data.get("folder") or "" + file_name = model_data.get("file_name") or "" + model_name = model_data.get("model_name") or file_name + folder = model_data.get("folder") or "" return { "model_name": model_name, "file_name": file_name, - "preview_url": config.get_preview_static_url(embedding_data.get("preview_url", "")), - "preview_nsfw_level": embedding_data.get("preview_nsfw_level", 0), - "base_model": embedding_data.get("base_model", ""), + "preview_url": config.get_preview_static_url(model_data.get("preview_url", "")), + "preview_nsfw_level": model_data.get("preview_nsfw_level", 0), + "base_model": model_data.get("base_model", ""), "folder": folder, - "sha256": embedding_data.get("sha256", ""), + "sha256": model_data.get("sha256", ""), "file_path": file_path.replace(os.sep, "/"), - "file_size": embedding_data.get("size", 0), - "modified": embedding_data.get("modified", ""), - "tags": embedding_data.get("tags", []), - "from_civitai": embedding_data.get("from_civitai", True), - # "usage_count": embedding_data.get("usage_count", 0), # TODO: Enable when embedding usage tracking is implemented - "notes": embedding_data.get("notes", ""), + "file_size": model_data.get("size", 0), + "modified": model_data.get("modified", ""), + "tags": model_data.get("tags", []), + "from_civitai": model_data.get("from_civitai", True), + # "usage_count": model_data.get("usage_count", 0), # TODO: Enable when embedding usage tracking is implemented + "notes": model_data.get("notes", ""), "sub_type": sub_type, - "favorite": embedding_data.get("favorite", False), - "exclude": bool(embedding_data.get("exclude", False)), - "update_available": bool(embedding_data.get("update_available", False)), - "skip_metadata_refresh": bool(embedding_data.get("skip_metadata_refresh", False)), - "civitai": self.filter_civitai_data(embedding_data.get("civitai", {}), minimal=True), - "auto_tags": embedding_data.get("auto_tags") or extract_auto_tags(embedding_data), - "version_count": embedding_data.get("version_count"), - "hf_url": embedding_data.get("hf_url", ""), + "favorite": model_data.get("favorite", False), + "exclude": bool(model_data.get("exclude", False)), + "update_available": bool(model_data.get("update_available", False)), + "skip_metadata_refresh": bool(model_data.get("skip_metadata_refresh", False)), + "civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True), + "auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data), + "version_count": model_data.get("version_count"), + "hf_url": model_data.get("hf_url", ""), } - def find_duplicate_hashes(self) -> Dict: + def find_duplicate_hashes(self) -> Dict[str, Any]: """Find Embeddings with duplicate SHA256 hashes""" return self.scanner._hash_index.get_duplicate_hashes() - def find_duplicate_filenames(self) -> Dict: + def find_duplicate_filenames(self) -> Dict[str, Any]: """Find Embeddings with conflicting filenames""" return self.scanner._hash_index.get_duplicate_filenames() diff --git a/py/services/example_images_cleanup_service.py b/py/services/example_images_cleanup_service.py index 4157a8e9..ee5f0374 100644 --- a/py/services/example_images_cleanup_service.py +++ b/py/services/example_images_cleanup_service.py @@ -35,7 +35,7 @@ class CleanupResult: def to_dict(self) -> Dict[str, object]: """Convert the dataclass to a serialisable dictionary.""" - data = { + data: Dict[str, object] = { "success": self.success, "checked_folders": self.checked_folders, "moved_empty_folders": self.moved_empty_folders, diff --git a/py/services/lora_scanner.py b/py/services/lora_scanner.py index cdf4fcea..cf111a3f 100644 --- a/py/services/lora_scanner.py +++ b/py/services/lora_scanner.py @@ -1,10 +1,12 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. import logging from typing import List from ..utils.models import LoraMetadata -from ..config import config from .model_scanner import ModelScanner -from .model_hash_index import ModelHashIndex # Changed from LoraHashIndex to ModelHashIndex import sys logger = logging.getLogger(__name__) @@ -15,8 +17,10 @@ class LoraScanner(ModelScanner): def __init__(self): # Define supported file extensions file_extensions = {'.safetensors'} - + # Initialize parent class with ModelHashIndex + from .model_hash_index import ModelHashIndex + super().__init__( model_type="lora", model_class=LoraMetadata, @@ -26,11 +30,13 @@ class LoraScanner(ModelScanner): def get_model_roots(self) -> List[str]: """Get lora root directories (including extra paths)""" + from ..config import config + roots: List[str] = [] roots.extend(config.loras_roots or []) roots.extend(config.extra_loras_roots or []) # Remove duplicates while preserving order - seen: set = set() + seen: set[str] = set() unique_roots: List[str] = [] for root in roots: if root and root not in seen: @@ -68,8 +74,12 @@ class LoraScanner(ModelScanner): test_hash = next(iter(self._hash_index._hash_to_path.keys())) test_path = self._hash_index.get_path(test_hash) logger.debug(f"\nTest lookup by hash: {test_hash[:8]}... -> {test_path}") + if test_path is None: + return # Also test reverse lookup test_hash_result = self._hash_index.get_hash(test_path) + if test_hash_result is None: + return logger.debug(f"Test reverse lookup: {test_path} -> {test_hash_result[:8]}...\n\n") diff --git a/py/services/lora_service.py b/py/services/lora_service.py index 0f65813b..fa4863b7 100644 --- a/py/services/lora_service.py +++ b/py/services/lora_service.py @@ -1,7 +1,7 @@ import logging import json import os -from typing import Dict, List, Optional +from typing import Any, Dict, List, Optional from .base_model_service import BaseModelService from .model_query import resolve_sub_type @@ -24,7 +24,7 @@ class LoraService(BaseModelService): """ super().__init__("lora", scanner, LoraMetadata, update_service=update_service) - async def format_response(self, lora_data: Dict) -> Optional[Dict]: + async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]: """Format LoRA data for API response. Returns None when the entry is missing critical fields (corrupted cache @@ -32,56 +32,56 @@ class LoraService(BaseModelService): whole listing request. See issue #730. """ # Guard against corrupted cache entries missing critical fields - file_path = lora_data.get("file_path") + file_path = model_data.get("file_path") if not file_path or not isinstance(file_path, str): logger.warning( "Skipping corrupted LoRA entry (missing file_path): %s", - lora_data.get("file_name", ""), + model_data.get("file_name", ""), ) return None # Resolve sub_type using priority: sub_type > model_type > civitai.model.type > default # Normalize to lowercase for consistent API responses - sub_type = resolve_sub_type(lora_data).lower() + sub_type = resolve_sub_type(model_data).lower() - file_name = lora_data.get("file_name") or "" - model_name = lora_data.get("model_name") or file_name - folder = lora_data.get("folder") or "" + file_name = model_data.get("file_name") or "" + model_name = model_data.get("model_name") or file_name + folder = model_data.get("folder") or "" return { "model_name": model_name, "file_name": file_name, "preview_url": config.get_preview_static_url( - lora_data.get("preview_url", "") + model_data.get("preview_url", "") ), - "preview_nsfw_level": lora_data.get("preview_nsfw_level", 0), - "base_model": lora_data.get("base_model", ""), + "preview_nsfw_level": model_data.get("preview_nsfw_level", 0), + "base_model": model_data.get("base_model", ""), "folder": folder, - "sha256": lora_data.get("sha256", ""), + "sha256": model_data.get("sha256", ""), "file_path": file_path.replace(os.sep, "/"), - "file_size": lora_data.get("size", 0), - "modified": lora_data.get("modified", ""), - "tags": lora_data.get("tags", []), - "from_civitai": lora_data.get("from_civitai", True), - "usage_count": lora_data.get("usage_count", 0), - "usage_tips": lora_data.get("usage_tips", ""), - "notes": lora_data.get("notes", ""), - "favorite": lora_data.get("favorite", False), - "exclude": bool(lora_data.get("exclude", False)), - "update_available": bool(lora_data.get("update_available", False)), + "file_size": model_data.get("size", 0), + "modified": model_data.get("modified", ""), + "tags": model_data.get("tags", []), + "from_civitai": model_data.get("from_civitai", True), + "usage_count": model_data.get("usage_count", 0), + "usage_tips": model_data.get("usage_tips", ""), + "notes": model_data.get("notes", ""), + "favorite": model_data.get("favorite", False), + "exclude": bool(model_data.get("exclude", False)), + "update_available": bool(model_data.get("update_available", False)), "skip_metadata_refresh": bool( - lora_data.get("skip_metadata_refresh", False) + model_data.get("skip_metadata_refresh", False) ), "sub_type": sub_type, "civitai": self.filter_civitai_data( - lora_data.get("civitai", {}), minimal=True + model_data.get("civitai", {}), minimal=True ), - "auto_tags": lora_data.get("auto_tags") or extract_auto_tags(lora_data), - "version_count": lora_data.get("version_count"), - "hf_url": lora_data.get("hf_url", ""), + "auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data), + "version_count": model_data.get("version_count"), + "hf_url": model_data.get("hf_url", ""), } - async def _apply_specific_filters(self, data: List[Dict], **kwargs) -> List[Dict]: + async def _apply_specific_filters(self, data: List[Dict[str, Any]], **kwargs) -> List[Dict[str, Any]]: """Apply LoRA-specific filters""" # Handle first_letter filter for LoRAs first_letter = kwargs.get("first_letter") @@ -152,7 +152,7 @@ class LoraService(BaseModelService): return data - def _filter_by_first_letter(self, data: List[Dict], letter: str) -> List[Dict]: + def _filter_by_first_letter(self, data: List[Dict[str, Any]], letter: str) -> List[Dict[str, Any]]: """Filter data by first letter of model name Special handling: @@ -307,7 +307,7 @@ class LoraService(BaseModelService): return None @staticmethod - def get_recommended_strength_from_lora_data(lora_data: Dict) -> Optional[float]: + def get_recommended_strength_from_lora_data(lora_data: Dict[str, Any]) -> Optional[float]: """Parse usage_tips JSON and extract recommended model strength.""" try: usage_tips = lora_data.get("usage_tips", "") @@ -320,7 +320,7 @@ class LoraService(BaseModelService): @staticmethod def get_recommended_clip_strength_from_lora_data( - lora_data: Dict, + lora_data: Dict[str, Any], ) -> Optional[float]: """Parse usage_tips JSON and extract recommended clip strength.""" try: @@ -332,7 +332,7 @@ class LoraService(BaseModelService): except (json.JSONDecodeError, TypeError, AttributeError): return None - async def get_lora_metadata_by_filename(self, filename: str) -> Optional[Dict]: + async def get_lora_metadata_by_filename(self, filename: str) -> Optional[Dict[str, Any]]: """Return cached raw metadata for a LoRA matching the given filename.""" cache = await self.scanner.get_cached_data(force_refresh=False) @@ -357,11 +357,11 @@ class LoraService(BaseModelService): return None - def find_duplicate_hashes(self) -> Dict: + def find_duplicate_hashes(self) -> Dict[str, Any]: """Find LoRAs with duplicate SHA256 hashes""" return self.scanner._hash_index.get_duplicate_hashes() - def find_duplicate_filenames(self) -> Dict: + def find_duplicate_filenames(self) -> Dict[str, Any]: """Find LoRAs with conflicting filenames""" return self.scanner._hash_index.get_duplicate_filenames() @@ -373,8 +373,8 @@ class LoraService(BaseModelService): use_same_clip_strength: bool = True, clip_strength_min: float = 0.0, clip_strength_max: float = 1.0, - locked_loras: Optional[List[Dict]] = None, - pool_config: Optional[Dict] = None, + locked_loras: Optional[List[Dict[str, Any]]] = None, + pool_config: Optional[Dict[str, Any]] = None, count_mode: str = "fixed", count_min: int = 3, count_max: int = 7, @@ -382,7 +382,7 @@ class LoraService(BaseModelService): recommended_strength_scale_min: float = 0.5, recommended_strength_scale_max: float = 1.0, seed: Optional[int] = None, - ) -> List[Dict]: + ) -> List[Dict[str, Any]]: """ Get random LoRAs with specified strength ranges. @@ -513,8 +513,8 @@ class LoraService(BaseModelService): return result_loras async def _apply_pool_filters( - self, available_loras: List[Dict], pool_config: Dict - ) -> List[Dict]: + self, available_loras: List[Dict[str, Any]], pool_config: Dict[str, Any] + ) -> List[Dict[str, Any]]: """ Apply pool_config filters to available LoRAs. @@ -671,8 +671,8 @@ class LoraService(BaseModelService): return available_loras async def get_cycler_list( - self, pool_config: Optional[Dict] = None, sort_by: str = "filename" - ) -> List[Dict]: + self, pool_config: Optional[Dict[str, Any]] = None, sort_by: str = "filename" + ) -> List[Dict[str, Any]]: """ Get filtered and sorted LoRA list for cycling. diff --git a/py/services/metadata_service.py b/py/services/metadata_service.py index c40694f1..9024a77e 100644 --- a/py/services/metadata_service.py +++ b/py/services/metadata_service.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. import os import logging from .model_metadata_provider import ( @@ -170,7 +174,7 @@ def _wrap_provider_with_rate_limit(provider_name: str | None, provider: ModelMet return RateLimitRetryingProvider(provider, label=provider_name) -async def get_metadata_provider(provider_name: str = None): +async def get_metadata_provider(provider_name: str | None = None): """Get a specific metadata provider or default provider with rate-limit handling.""" provider_manager = await ModelMetadataProviderManager.get_instance() diff --git a/py/services/metadata_sync_service.py b/py/services/metadata_sync_service.py index d5b4af36..13e8c671 100644 --- a/py/services/metadata_sync_service.py +++ b/py/services/metadata_sync_service.py @@ -6,7 +6,7 @@ import json import logging import os from datetime import datetime -from typing import Any, Awaitable, Callable, Dict, Iterable, Optional +from typing import Any, Awaitable, Callable, Dict, Iterable, Optional, Protocol from ..services.settings_manager import SettingsManager from ..utils.civitai_utils import resolve_license_payload @@ -18,14 +18,14 @@ from .errors import RateLimitError logger = logging.getLogger(__name__) -class MetadataProviderProtocol: +class MetadataProviderProtocol(Protocol): """Subset of metadata provider interface consumed by the sync service.""" - async def get_model_by_hash(self, sha256: str) -> tuple[Optional[Dict[str, Any]], Optional[str]]: + async def get_model_by_hash(self, model_hash: str) -> tuple[Optional[Dict[str, Any]], Optional[str]]: ... async def get_model_version( - self, model_id: int, model_version_id: Optional[int] + self, model_id: Any = None, version_id: Any = None ) -> Optional[Dict[str, Any]]: ... @@ -39,8 +39,8 @@ class MetadataSyncService: metadata_manager, preview_service, settings: SettingsManager, - default_metadata_provider_factory: Callable[[], Awaitable[MetadataProviderProtocol]], - metadata_provider_selector: Callable[[str], Awaitable[MetadataProviderProtocol]], + default_metadata_provider_factory: Callable[..., Awaitable[MetadataProviderProtocol]], + metadata_provider_selector: Callable[..., Awaitable[MetadataProviderProtocol]], ) -> None: self._metadata_manager = metadata_manager self._preview_service = preview_service @@ -492,7 +492,7 @@ class MetadataSyncService: if not file_paths: raise ValueError("No file paths provided for verification") - results = { + results: Dict[str, Any] = { "verified_as_duplicates": True, "mismatched_files": [], "new_hash_map": {}, diff --git a/py/services/model_cache.py b/py/services/model_cache.py index deb5cc6e..597c3fd8 100644 --- a/py/services/model_cache.py +++ b/py/services/model_cache.py @@ -31,17 +31,22 @@ DISPLAY_NAME_MODES = {"model_name", "file_name"} class ModelCache: """Cache structure for model data with extensible sorting.""" - raw_data: List[Dict] + raw_data: List[Dict[str, Any]] folders: List[str] - version_index: Dict[int, Dict] = field(default_factory=dict) + version_index: Dict[int, Dict[str, Any]] = field(default_factory=dict) model_id_index: Dict[int, List[Dict[str, Any]]] = field(default_factory=dict) name_display_mode: str = "model_name" + _lock: Any = field(init=False, repr=False, default=None) + # Cache for last sort: (sort_key, order, seed) -> sorted list + _last_sort: Tuple[Optional[str], str, Optional[str]] = field( + init=False, repr=False, default=(None, "asc", None) + ) + _last_sorted_data: List[Dict[str, Any]] = field( + init=False, repr=False, default_factory=list + ) def __post_init__(self): self._lock = asyncio.Lock() - # Cache for last sort: (sort_key, order, seed) -> sorted list - self._last_sort: Tuple[Optional[str], str, Optional[str]] = (None, "asc", None) - self._last_sorted_data: List[Dict] = [] self._normalize_raw_data() self.name_display_mode = self._normalize_display_mode(self.name_display_mode) # Default sort on init @@ -64,7 +69,7 @@ class ModelCache: return "" return str(value) - def _normalize_item(self, item: Dict) -> None: + def _normalize_item(self, item: Dict[str, Any]) -> None: """Ensure core metadata fields are present and string typed.""" if not isinstance(item, dict): @@ -80,7 +85,7 @@ class ModelCache: for item in self.raw_data: self._normalize_item(item) - def _get_display_name(self, item: Dict) -> str: + def _get_display_name(self, item: Dict[str, Any]) -> str: """Return the value used for name-based sorting based on display settings.""" if self.name_display_mode == "file_name": @@ -114,7 +119,7 @@ class ModelCache: for item in self.raw_data: self.add_to_version_index(item) - def add_to_version_index(self, item: Dict) -> None: + def add_to_version_index(self, item: Dict[str, Any]) -> None: """Register a cache item in the version/model indexes if possible.""" civitai_data = item.get('civitai') if isinstance(item, dict) else None @@ -143,7 +148,7 @@ class ModelCache: else: versions.append(descriptor) - def remove_from_version_index(self, item: Dict) -> None: + def remove_from_version_index(self, item: Dict[str, Any]) -> None: """Remove a cache item from the version/model indexes if present.""" civitai_data = item.get('civitai') if isinstance(item, dict) else None @@ -177,7 +182,7 @@ class ModelCache: def _build_version_descriptor( self, - item: Dict, + item: Dict[str, Any], civitai_data: Dict[str, Any], version_id: int, ) -> Optional[Dict[str, Any]]: @@ -204,8 +209,8 @@ class ModelCache: async def resort(self): """Resort cached data according to last sort mode if set""" async with self._lock: - if self._last_sort[0] is not None: - sort_key, order, seed = self._last_sort + sort_key, order, seed = self._last_sort + if sort_key is not None: sorted_data = self._sort_data(self.raw_data, sort_key, order, seed) self._last_sorted_data = sorted_data # Update folder list @@ -219,7 +224,7 @@ class ModelCache: self.folders = sorted(list(all_folders), key=lambda x: x.lower()) self.rebuild_version_index() - def _sort_data(self, data: List[Dict], sort_key: str, order: str, seed: Optional[str] = None) -> List[Dict]: + def _sort_data(self, data: List[Dict[str, Any]], sort_key: str, order: str, seed: Optional[str] = None) -> List[Dict[str, Any]]: """Sort data by sort_key and order""" start_time = time.perf_counter() reverse = (order == 'desc') @@ -293,7 +298,7 @@ class ModelCache: logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration) return result - async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc', seed: Optional[str] = None) -> List[Dict]: + async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc', seed: Optional[str] = None) -> List[Dict[str, Any]]: """Get sorted data by sort_key and order, using cache if possible""" async with self._lock: cache_key = (sort_key, order, seed) @@ -321,8 +326,8 @@ class ModelCache: self.name_display_mode = normalized - if self._last_sort[0] == 'name': - sort_key, order, seed = self._last_sort + sort_key, order, seed = self._last_sort + if sort_key == 'name': self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order, seed) async def update_preview_url(self, file_path: str, preview_url: str, preview_nsfw_level: int) -> bool: diff --git a/py/services/model_file_service.py b/py/services/model_file_service.py index 2537e2b8..9a7a5848 100644 --- a/py/services/model_file_service.py +++ b/py/services/model_file_service.py @@ -41,7 +41,7 @@ class AutoOrganizeResult: def to_dict(self) -> Dict[str, Any]: """Convert result to dictionary""" - result = { + result: Dict[str, Any] = { 'success': self.status != 'error', 'status': self.status, 'message': f'Auto-organize {self.operation_type} completed: {self.success_count} moved, {self.skipped_count} skipped, {self.failure_count} failed out of {self.total} total', @@ -418,6 +418,8 @@ class ModelFileService: """Calculate the target directory for a model""" if is_flat_structure: file_path = model.get('file_path') + if not isinstance(file_path, str): + return None current_dir = os.path.dirname(file_path) # Check if already in root directory diff --git a/py/services/model_hash_index.py b/py/services/model_hash_index.py index fb796fa1..103f39ae 100644 --- a/py/services/model_hash_index.py +++ b/py/services/model_hash_index.py @@ -35,6 +35,7 @@ class ModelHashIndex: # Track duplicates by filename - FIXED LOGIC is_re_registration = False + existing_hash: Optional[str] = None if filename in self._filename_to_hash: existing_hash = self._filename_to_hash[filename] existing_path = self._hash_to_path.get(existing_hash) @@ -101,7 +102,7 @@ class ModelHashIndex: """Extract filename without extension from path""" return os.path.splitext(os.path.basename(file_path))[0] - def remove_by_path(self, file_path: str, hash_val: str = None) -> None: + def remove_by_path(self, file_path: str, hash_val: Optional[str] = None) -> None: """Remove entry by file path""" filename = self._get_filename_from_path(file_path) diff --git a/py/services/model_lifecycle_service.py b/py/services/model_lifecycle_service.py index 2aba5e3a..54189bed 100644 --- a/py/services/model_lifecycle_service.py +++ b/py/services/model_lifecycle_service.py @@ -4,7 +4,7 @@ from __future__ import annotations import logging import os -from typing import Any, Awaitable, Callable, Dict, Iterable, List, Mapping, Optional, TYPE_CHECKING +from typing import Any, Awaitable, Callable, Dict, Iterable, List, Mapping, Optional, TYPE_CHECKING, cast from ..services.service_registry import ServiceRegistry from ..utils.constants import PREVIEW_EXTENSIONS @@ -87,8 +87,8 @@ class ModelLifecycleService: scanner, metadata_manager, metadata_loader: Callable[[str], Awaitable[Dict[str, object]]], - recipe_scanner_factory: Callable[[], Awaitable] | None = None, - update_service: "ModelUpdateService" | None = None, + recipe_scanner_factory: Callable[[], Awaitable[Any]] | None = None, + update_service: Optional["ModelUpdateService"] = None, ) -> None: self._scanner = scanner self._metadata_manager = metadata_manager @@ -146,7 +146,7 @@ class ModelLifecycleService: persist_current_cache = getattr(self._scanner, "_persist_current_cache", None) if callable(persist_current_cache): - await persist_current_cache() + await cast(Awaitable[Any], persist_current_cache()) return {"success": True, "deleted_files": deleted_files} @@ -252,7 +252,7 @@ class ModelLifecycleService: persist_current_cache = getattr(self._scanner, "_persist_current_cache", None) if callable(persist_current_cache): - await persist_current_cache() + await cast(Awaitable[Any], persist_current_cache()) message = f"Model {os.path.basename(file_path)} excluded" return {"success": True, "message": message} @@ -357,7 +357,8 @@ class ModelLifecycleService: if os.path.exists(metadata_path): metadata = await self._metadata_loader(metadata_path) - hash_value = metadata.get("sha256") if isinstance(metadata, dict) else None + raw_hash = metadata.get("sha256") if isinstance(metadata, dict) else None + hash_value = raw_hash if isinstance(raw_hash, str) else None renamed_files: List[str] = [] new_metadata_path: Optional[str] = None diff --git a/py/services/model_metadata_provider.py b/py/services/model_metadata_provider.py index 999c7936..9625f738 100644 --- a/py/services/model_metadata_provider.py +++ b/py/services/model_metadata_provider.py @@ -10,7 +10,7 @@ from .errors import RateLimitError, ResourceNotFoundError try: from bs4 import BeautifulSoup except ImportError as exc: - BeautifulSoup = None # type: ignore[assignment] + BeautifulSoup = None # pyright: ignore[reportAssignmentType] _BS4_IMPORT_ERROR = exc else: _BS4_IMPORT_ERROR = None @@ -18,7 +18,7 @@ else: try: import aiosqlite except ImportError as exc: - aiosqlite = None # type: ignore[assignment] + aiosqlite = None # pyright: ignore[reportAssignmentType] _AIOSQLITE_IMPORT_ERROR = exc else: _AIOSQLITE_IMPORT_ERROR = None @@ -105,24 +105,24 @@ class ModelMetadataProvider(ABC): """Base abstract class for all model metadata providers""" @abstractmethod - async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: """Find model by hash value""" pass @abstractmethod - async def get_model_versions(self, model_id: str) -> Optional[Dict]: + async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]: """Get all versions of a model with their details""" pass async def get_model_versions_bulk( self, model_ids: Sequence[int] - ) -> Optional[Dict[int, Dict]]: + ) -> Optional[Dict[int, Dict[str, Any]]]: """Fetch model versions for multiple model ids when supported.""" raise NotImplementedError async def get_model_versions_by_hashes( self, hashes: List[str] - ) -> Optional[List[Dict]]: + ) -> Optional[List[Dict[str, Any]]]: """Fetch full version details for multiple SHA256 hashes. Used specifically to retrieve ``usageControl`` which is only @@ -133,17 +133,17 @@ class ModelMetadataProvider(ABC): raise NotImplementedError @abstractmethod - async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]: + async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]: """Get specific model version with additional metadata""" pass @abstractmethod - async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: """Fetch model version metadata""" pass @abstractmethod - async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]: + async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]: """Fetch one page of models owned by the specified user. Returns ``{"items": [...], "nextCursor": }`` on success, @@ -161,29 +161,29 @@ class CivitaiModelMetadataProvider(ModelMetadataProvider): def __init__(self, civitai_client): self.client = civitai_client - async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: return await self.client.get_model_by_hash(model_hash) - async def get_model_versions(self, model_id: str) -> Optional[Dict]: + async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]: return await self.client.get_model_versions(model_id) async def get_model_versions_bulk( self, model_ids: Sequence[int] - ) -> Optional[Dict[int, Dict]]: + ) -> Optional[Dict[int, Dict[str, Any]]]: return await self.client.get_model_versions_bulk(model_ids) async def get_model_versions_by_hashes( self, hashes: List[str] - ) -> Optional[List[Dict]]: + ) -> Optional[List[Dict[str, Any]]]: return await self.client.get_model_versions_by_hashes(hashes) - async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]: + async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]: return await self.client.get_model_version(model_id, version_id) - async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: return await self.client.get_model_version_info(version_id) - async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]: + async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]: return await self.client.get_user_models(username, cursor) async def get_creator_model_count(self, username: str) -> Optional[int]: @@ -195,19 +195,19 @@ class CivArchiveModelMetadataProvider(ModelMetadataProvider): def __init__(self, civarchive_client): self.client = civarchive_client - async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: return await self.client.get_model_by_hash(model_hash) - async def get_model_versions(self, model_id: str) -> Optional[Dict]: + async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]: return await self.client.get_model_versions(model_id) - async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]: + async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]: return await self.client.get_model_version(model_id, version_id) - async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: return await self.client.get_model_version_info(version_id) - async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]: + async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]: """Not supported by CivArchive provider""" return None @@ -218,7 +218,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider): self.db_path = db_path self._aiosqlite = _require_aiosqlite() - async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: """Find model by hash value from SQLite database""" async with self._aiosqlite.connect(self.db_path) as db: # Look up in model_files table to get model_id and version_id @@ -243,7 +243,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider): result = await self._get_version_with_model_data(db, model_id, version_id) return result, None if result else "Error retrieving model data" - async def get_model_versions(self, model_id: str) -> Optional[Dict]: + async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]: """Get all versions of a model from SQLite database""" async with self._aiosqlite.connect(self.db_path) as db: db.row_factory = self._aiosqlite.Row @@ -299,7 +299,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider): 'name': model_name } - async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]: + async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]: """Get specific model version with additional metadata from SQLite database""" if not model_id and not version_id: return None @@ -339,7 +339,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider): # Now we have both model_id and version_id, get the full data return await self._get_version_with_model_data(db, model_id, version_id) - async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: """Fetch model version metadata from SQLite database""" async with self._aiosqlite.connect(self.db_path) as db: db.row_factory = self._aiosqlite.Row @@ -358,11 +358,11 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider): version_data = await self._get_version_with_model_data(db, model_id, version_id) return version_data, None - async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]: + async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]: """Listing models by username is not supported for archive database""" return None - async def _get_version_with_model_data(self, db, model_id, version_id) -> Optional[Dict]: + async def _get_version_with_model_data(self, db, model_id, version_id) -> Optional[Dict[str, Any]]: """Helper to build version data with model information""" # Get version details version_query = "SELECT name, base_model, data FROM model_versions WHERE id = ? AND model_id = ?" @@ -485,7 +485,7 @@ class FallbackMetadataProvider(ModelMetadataProvider): jitter_ratio=self._rate_limit_jitter_ratio, ) - async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: for provider, label in self._iter_providers(): try: result, error = await self._call_with_rate_limit( @@ -507,7 +507,7 @@ class FallbackMetadataProvider(ModelMetadataProvider): continue return None, "Model not found" - async def get_model_versions(self, model_id: str) -> Optional[Dict]: + async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]: not_found_confirmed = False for provider, label in self._iter_providers(): try: @@ -538,7 +538,7 @@ class FallbackMetadataProvider(ModelMetadataProvider): continue return None - async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]: + async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]: for provider, label in self._iter_providers(): try: result = await self._call_with_rate_limit( @@ -561,7 +561,7 @@ class FallbackMetadataProvider(ModelMetadataProvider): continue return None - async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: for provider, label in self._iter_providers(): try: result, error = await self._call_with_rate_limit( @@ -585,7 +585,7 @@ class FallbackMetadataProvider(ModelMetadataProvider): async def get_model_versions_by_hashes( self, hashes: List[str] - ) -> Optional[List[Dict]]: + ) -> Optional[List[Dict[str, Any]]]: for provider, label in self._iter_providers(): try: result = await self._call_with_rate_limit( @@ -613,7 +613,7 @@ class FallbackMetadataProvider(ModelMetadataProvider): continue return None - async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]: + async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]: for provider, label in self._iter_providers(): try: result = await self._call_with_rate_limit( @@ -681,14 +681,14 @@ class RateLimitRetryingProvider(ModelMetadataProvider): def __getattr__(self, item): return getattr(self._provider, item) - async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: return await self._rate_limit_helper.run( self._label, self._provider.get_model_by_hash, model_hash, ) - async def get_model_versions(self, model_id: str) -> Optional[Dict]: + async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]: return await self._rate_limit_helper.run( self._label, self._provider.get_model_versions, @@ -698,7 +698,7 @@ class RateLimitRetryingProvider(ModelMetadataProvider): async def get_model_versions_bulk( self, model_ids: Sequence[int], - ) -> Optional[Dict[int, Dict]]: + ) -> Optional[Dict[int, Dict[str, Any]]]: return await self._rate_limit_helper.run( self._label, self._provider.get_model_versions_bulk, @@ -707,14 +707,14 @@ class RateLimitRetryingProvider(ModelMetadataProvider): async def get_model_versions_by_hashes( self, hashes: List[str] - ) -> Optional[List[Dict]]: + ) -> Optional[List[Dict[str, Any]]]: return await self._rate_limit_helper.run( self._label, self._provider.get_model_versions_by_hashes, hashes, ) - async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]: + async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]: return await self._rate_limit_helper.run( self._label, self._provider.get_model_version, @@ -722,14 +722,14 @@ class RateLimitRetryingProvider(ModelMetadataProvider): version_id, ) - async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: return await self._rate_limit_helper.run( self._label, self._provider.get_model_version_info, version_id, ) - async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]: + async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]: return await self._rate_limit_helper.run( self._label, self._provider.get_user_models, @@ -762,12 +762,12 @@ class ModelMetadataProviderManager: if is_default or self.default_provider is None: self.default_provider = name - async def get_model_by_hash(self, model_hash: str, provider_name: str = None) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_by_hash(self, model_hash: str, provider_name: Optional[str] = None) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: """Find model by hash using specified or default provider""" provider = self._get_provider(provider_name) return await provider.get_model_by_hash(model_hash) - async def get_model_versions(self, model_id: str, provider_name: str = None) -> Optional[Dict]: + async def get_model_versions(self, model_id: str, provider_name: Optional[str] = None) -> Optional[Dict[str, Any]]: """Get model versions using specified or default provider""" provider = self._get_provider(provider_name) return await provider.get_model_versions(model_id) @@ -775,8 +775,8 @@ class ModelMetadataProviderManager: async def get_model_versions_bulk( self, model_ids: Sequence[int], - provider_name: str = None, - ) -> Optional[Dict[int, Dict]]: + provider_name: Optional[str] = None, + ) -> Optional[Dict[int, Dict[str, Any]]]: """Fetch model versions for multiple model ids when supported by provider.""" provider = self._get_provider(provider_name) try: @@ -784,12 +784,12 @@ class ModelMetadataProviderManager: except NotImplementedError: return None - async def get_model_version(self, model_id: int = None, version_id: int = None, provider_name: str = None) -> Optional[Dict]: + async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None, provider_name: Optional[str] = None) -> Optional[Dict[str, Any]]: """Get specific model version using specified or default provider""" provider = self._get_provider(provider_name) return await provider.get_model_version(model_id, version_id) - async def get_model_version_info(self, version_id: str, provider_name: str = None) -> Tuple[Optional[Dict], Optional[str]]: + async def get_model_version_info(self, version_id: str, provider_name: Optional[str] = None) -> Tuple[Optional[Dict[str, Any]], Optional[str]]: """Fetch model version info using specified or default provider""" provider = self._get_provider(provider_name) return await provider.get_model_version_info(version_id) @@ -797,8 +797,8 @@ class ModelMetadataProviderManager: async def get_model_versions_by_hashes( self, hashes: List[str], - provider_name: str = None, - ) -> Optional[List[Dict]]: + provider_name: Optional[str] = None, + ) -> Optional[List[Dict[str, Any]]]: provider = self._get_provider(provider_name) try: return await provider.get_model_versions_by_hashes(hashes) @@ -808,19 +808,19 @@ class ModelMetadataProviderManager: async def get_user_models( self, username: str, - provider_name: str = None, + provider_name: Optional[str] = None, cursor: Optional[str] = None, - ) -> Optional[Dict]: + ) -> Optional[Dict[str, Any]]: """Fetch one page of models owned by the specified user""" provider = self._get_provider(provider_name) return await provider.get_user_models(username, cursor) - async def get_creator_model_count(self, username: str, provider_name: str = None) -> Optional[int]: + async def get_creator_model_count(self, username: str, provider_name: Optional[str] = None) -> Optional[int]: """Best-effort published model count for the specified user""" provider = self._get_provider(provider_name) return await provider.get_creator_model_count(username) - def _get_provider(self, provider_name: str = None) -> ModelMetadataProvider: + def _get_provider(self, provider_name: Optional[str] = None) -> ModelMetadataProvider: """Get provider by name or default provider""" if provider_name: if provider_name not in self.providers: diff --git a/py/services/model_query.py b/py/services/model_query.py index dabe88e0..b5ba2daa 100644 --- a/py/services/model_query.py +++ b/py/services/model_query.py @@ -12,6 +12,7 @@ from typing import ( Tuple, Protocol, Callable, + cast, ) from ..utils.constants import NSFW_LEVELS @@ -309,7 +310,7 @@ class ModelFilterSet: else: include_tags.add(normalized) else: - include_tags = {tag.strip().lower() for tag in tag_filters if tag} + include_tags = {tag.strip().lower() for tag in cast(Iterable[Any], tag_filters) if tag} if include_tags: tag_logic = criteria.tag_logic.lower() if criteria.tag_logic else "any" diff --git a/py/services/model_scanner.py b/py/services/model_scanner.py index b06c1027..c9add7df 100644 --- a/py/services/model_scanner.py +++ b/py/services/model_scanner.py @@ -5,7 +5,7 @@ import asyncio import time import shutil from dataclasses import dataclass -from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional, Set, Type, Union +from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional, Sequence, Set, Type, Union, cast from ..utils.models import BaseModelMetadata, autov3_from_civitai_files from ..config import config @@ -29,7 +29,7 @@ logger = logging.getLogger(__name__) class CacheBuildResult: """Represents the outcome of scanning model files for cache building.""" - raw_data: List[Dict] + raw_data: List[Dict[str, Any]] hash_index: ModelHashIndex tags_count: Dict[str, int] excluded_models: List[str] @@ -59,7 +59,7 @@ class ModelScanner: lock = cls._get_lock() async with lock: if cls not in cls._instances: - cls._instances[cls] = cls() + cls._instances[cls] = cls() # pyright: ignore[reportCallIssue] return cls._instances[cls] def __init__(self, model_type: str, model_class: Type[BaseModelMetadata], file_extensions: Set[str], hash_index: Optional[ModelHashIndex] = None): @@ -78,7 +78,7 @@ class ModelScanner: self.model_type = model_type self.model_class = model_class self.file_extensions = file_extensions - self._cache = None + self._cache: Any = None self._hash_index = hash_index or ModelHashIndex() self._tags_count = {} # Dictionary to store tag counts self._is_initializing = False # Flag to track initialization state @@ -183,7 +183,7 @@ class ModelScanner: is_mapping = isinstance(source, Mapping) def get_value(key: str, default: Any = None) -> Any: - if is_mapping: + if isinstance(source, Mapping): return source.get(key, default) sentinel = object() @@ -772,7 +772,7 @@ class ModelScanner: else: await self._reconcile_cache() - return self._cache + return cast(ModelCache, self._cache) async def _initialize_cache(self) -> None: """Initialize or refresh the cache""" @@ -932,6 +932,8 @@ class ModelScanner: ) continue model_data = validation_result.entry + if model_data is None: + continue self._ensure_license_flags(model_data) # Add to cache @@ -992,8 +994,8 @@ class ModelScanner: self._cache.raw_data = [item for item in self._cache.raw_data if item['file_path'] not in missing_files] dedup_removed = 0 - seen_paths: set = set() - deduped: list = [] + seen_paths: set[str] = set() + deduped: list[Dict[str, Any]] = [] for item in reversed(self._cache.raw_data): path = item.get('file_path', '') if path not in seen_paths: @@ -1108,7 +1110,7 @@ class ModelScanner: *, hash_index: Optional[ModelHashIndex] = None, excluded_models: Optional[List[str]] = None - ) -> Dict: + ) -> Optional[Dict[str, Any]]: """Process a single model file and return its metadata""" hash_index = hash_index or self._hash_index excluded_models = excluded_models if excluded_models is not None else self._excluded_models @@ -1132,7 +1134,7 @@ class ModelScanner: file_name = os.path.splitext(os.path.basename(file_path))[0] file_info['name'] = file_name - metadata = self.model_class.from_civitai_info(version_info, file_info, file_path) + metadata = cast(Any, self.model_class).from_civitai_info(version_info, file_info, file_path) metadata.preview_url = find_preview_file(file_name, os.path.dirname(file_path)) await MetadataManager.save_metadata(file_path, metadata) logger.info(f"Created metadata from .civitai.info for {file_path} (Reason: .civitai.info was found but .metadata.json was missing)") @@ -1169,6 +1171,8 @@ class ModelScanner: if metadata is None: metadata = await self._create_default_metadata(file_path) + assert metadata is not None + # Hook: allow subclasses to adjust metadata metadata = self.adjust_metadata(metadata, file_path, root_path) @@ -1296,7 +1300,7 @@ class ModelScanner: async def _sync_download_history( self, - raw_data: List[Mapping[str, Any]], + raw_data: Sequence[Mapping[str, Any]], *, source: str, ) -> None: @@ -1345,7 +1349,7 @@ class ModelScanner: ) -> CacheBuildResult: """Collect metadata for all model files.""" - raw_data: List[Dict] = [] + raw_data: List[Dict[str, Any]] = [] hash_index = ModelHashIndex() tags_count: Dict[str, int] = {} excluded_models: List[str] = [] @@ -1409,6 +1413,8 @@ class ModelScanner: ) continue result = validation_result.entry + if result is None: + continue self._ensure_license_flags(result) raw_data.append(result) @@ -1448,7 +1454,7 @@ class ModelScanner: excluded_models=excluded_models ) - async def add_model_to_cache(self, metadata_dict: Dict, folder: str = '') -> bool: + async def add_model_to_cache(self, metadata_dict: Dict[str, Any], folder: str = '') -> bool: """Add a model to the cache Args: @@ -1461,7 +1467,8 @@ class ModelScanner: try: if self._cache is None: await self.get_cached_data() - + assert self._cache is not None + # Update folder in metadata metadata_dict['folder'] = folder @@ -1496,7 +1503,7 @@ class ModelScanner: logger.error(f"Error adding model to cache: {e}") return False - async def move_model(self, source_path: str, target_path: str) -> Optional[str]: + async def move_model(self, source_path: str, target_path: str) -> Optional[Dict[str, Any]]: """Move a model and its associated files to a new location Args: @@ -1530,7 +1537,7 @@ class ModelScanner: # Check for filename conflicts and auto-rename if necessary from ..utils.models import BaseModelMetadata final_filename = BaseModelMetadata.generate_unique_filename( - target_path, base_name, file_ext, get_source_hash + target_path, base_name, file_ext, lambda: get_source_hash() or "" ) target_file = os.path.join(target_path, final_filename).replace(os.sep, '/') @@ -1578,7 +1585,7 @@ class ModelScanner: logger.error(f"Error moving associated file {source_file}: {e}") # Handle metadata file specially to update paths - if source_metadata and os.path.exists(source_metadata): + if source_metadata and moved_metadata_path and os.path.exists(source_metadata): try: shutil.move(source_metadata, moved_metadata_path) metadata = await self._update_metadata_paths(moved_metadata_path, target_file) @@ -1596,7 +1603,7 @@ class ModelScanner: logger.error(f"Error moving model: {e}", exc_info=True) return None - async def _update_metadata_paths(self, metadata_path: str, model_path: str) -> Dict: + async def _update_metadata_paths(self, metadata_path: str, model_path: str) -> Optional[Dict[str, Any]]: """Update file paths in metadata file""" try: with open(metadata_path, 'r', encoding='utf-8') as f: @@ -1622,7 +1629,7 @@ class ModelScanner: logger.error(f"Error updating metadata paths: {e}", exc_info=True) return None - async def update_single_model_cache(self, original_path: str, new_path: str, metadata: Dict, recalculate_type: bool = False) -> Union[bool, Dict]: + async def update_single_model_cache(self, original_path: str, new_path: str, metadata: Optional[Dict[str, Any]], recalculate_type: bool = False) -> Union[bool, Dict[str, Any]]: """Update cache after a model has been moved or modified""" cache = await self.get_cached_data() @@ -1645,6 +1652,7 @@ class ModelScanner: ] cache_modified = bool(existing_item) or bool(metadata) + cache_entry: Optional[Dict[str, Any]] = None if metadata: normalized_new_path = new_path.replace(os.sep, '/') @@ -1695,7 +1703,9 @@ class ModelScanner: if cache_modified: await self._persist_current_cache() - return cache_entry if metadata else True + if metadata and cache_entry is not None: + return cache_entry + return True async def sync_cache_from_metadata( self, file_path: str, metadata_dict: Dict[str, Any] @@ -1820,8 +1830,8 @@ class ModelScanner: existing_entry.update(desired_entry) # ---- Incremental tag count update ---- - new_tags: set = set(desired_entry.get("tags") or []) - old_tag_set: set = set(old_tags) + new_tags: set[str] = set(desired_entry.get("tags") or []) + old_tag_set: set[str] = set(old_tags) for tag in old_tag_set - new_tags: current = self._tags_count.get(tag, 0) if current <= 1: @@ -2020,7 +2030,7 @@ class ModelScanner: return None - async def get_top_tags(self, limit: int = 20) -> List[Dict[str, any]]: + async def get_top_tags(self, limit: int = 20) -> List[Dict[str, Any]]: """Get top tags sorted by count. If limit is 0, return all tags.""" await self.get_cached_data() @@ -2036,7 +2046,7 @@ class ModelScanner: async def search_tags( self, query: str, limit: int = 50 - ) -> List[Dict[str, any]]: + ) -> List[Dict[str, Any]]: """Search tags by case-insensitive substring match, sorted by count. If query is empty, behaves like get_top_tags (returns top ``limit`` @@ -2059,7 +2069,7 @@ class ModelScanner: return matched return matched[:limit] - async def get_base_models(self, limit: int = 20) -> List[Dict[str, any]]: + async def get_base_models(self, limit: int = 20) -> List[Dict[str, Any]]: """Get base models sorted by count. If limit is 0, return all.""" cache = await self.get_cached_data() @@ -2140,7 +2150,7 @@ class ModelScanner: await self._persist_current_cache() return updated - async def bulk_delete_models(self, file_paths: List[str]) -> Dict: + async def bulk_delete_models(self, file_paths: List[str]) -> Dict[str, Any]: """Delete multiple models and update cache in a batch operation Args: @@ -2338,7 +2348,7 @@ class ModelScanner: logger.error(f"Error checking model version existence: {e}") return False - async def get_model_versions_by_id(self, model_id: int) -> List[Dict]: + async def get_model_versions_by_id(self, model_id: int) -> List[Dict[str, Any]]: """Get all versions of a model by its ID Args: diff --git a/py/services/model_service_factory.py b/py/services/model_service_factory.py index 0e1da069..c38faf8a 100644 --- a/py/services/model_service_factory.py +++ b/py/services/model_service_factory.py @@ -6,13 +6,13 @@ logger = logging.getLogger(__name__) class ModelServiceFactory: """Factory for managing model services and routes""" - _services: Dict[str, Type] = {} - _routes: Dict[str, Type] = {} + _services: Dict[str, Type[Any]] = {} + _routes: Dict[str, Type[Any]] = {} _initialized_services: Dict[str, Any] = {} _initialized_routes: Dict[str, Any] = {} @classmethod - def register_model_type(cls, model_type: str, service_class: Type, route_class: Type): + def register_model_type(cls, model_type: str, service_class: Type[Any], route_class: Type[Any]): """Register a new model type with its service and route classes Args: @@ -24,7 +24,7 @@ class ModelServiceFactory: cls._routes[model_type] = route_class @classmethod - def get_service_class(cls, model_type: str) -> Type: + def get_service_class(cls, model_type: str) -> Type[Any]: """Get service class for a model type Args: @@ -41,7 +41,7 @@ class ModelServiceFactory: return cls._services[model_type] @classmethod - def get_route_class(cls, model_type: str) -> Type: + def get_route_class(cls, model_type: str) -> Type[Any]: """Get route class for a model type Args: @@ -87,7 +87,7 @@ class ModelServiceFactory: logger.error(f"Failed to setup routes for {model_type}: {e}", exc_info=True) @classmethod - def get_registered_types(cls) -> list: + def get_registered_types(cls) -> list[str]: """Get list of all registered model types Returns: diff --git a/py/services/model_update_service.py b/py/services/model_update_service.py index d079da4c..fe30f473 100644 --- a/py/services/model_update_service.py +++ b/py/services/model_update_service.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. """Service for tracking remote model version updates.""" from __future__ import annotations @@ -336,9 +340,9 @@ class ModelUpdateService: return try: - from .persistent_model_cache import get_persistent_cache + from .persistent_model_cache import PersistentModelCache - legacy_path = get_persistent_cache(self._library_name).get_database_path() + legacy_path = PersistentModelCache.get_default(self._library_name).get_database_path() except Exception: return @@ -735,7 +739,7 @@ class ModelUpdateService: ) results: Dict[int, ModelUpdateRecord] = {} - prefetched: Dict[int, Mapping] = {} + prefetched: Dict[int, Mapping[Any, Any]] = {} fetch_targets: List[int] = [] if metadata_provider and local_versions: @@ -834,7 +838,7 @@ class ModelUpdateService: model_id: int, version_ids: Sequence[int], *, - version_info: Optional[Mapping] = None, + version_info: Optional[Mapping[str, Any]] = None, ) -> ModelUpdateRecord: """Persist a new set of in-library version identifiers.""" @@ -954,7 +958,11 @@ class ModelUpdateService: records = self._get_records_bulk(model_type, normalized_ids) return { - model_id: records.get(model_id).has_update(hide_early_access=hide_early_access) if records.get(model_id) else False + model_id: ( + records[model_id].has_update(hide_early_access=hide_early_access) + if model_id in records + else False + ) for model_id in normalized_ids } @@ -980,7 +988,7 @@ class ModelUpdateService: metadata_provider, *, force_refresh: bool = False, - prefetched_response: Optional[Mapping] = None, + prefetched_response: Optional[Mapping[str, Any]] = None, all_local_version_ids: Optional[Sequence[int]] = None, ) -> Optional[ModelUpdateRecord]: normalized_local = self._normalize_sequence(local_versions) @@ -1010,7 +1018,7 @@ class ModelUpdateService: fallback_attempted = False fallback_error_message: Optional[str] = None mark_model_as_ignored = False - response: Optional[Mapping] = None + response: Optional[Mapping[str, Any]] = None if metadata_provider and should_fetch: response = prefetched_response if response is None: @@ -1122,7 +1130,7 @@ class ModelUpdateService: async def _enrich_version_entries( self, metadata_provider, - responses_by_model_id: Dict[int, Mapping], + responses_by_model_id: Dict[int, Mapping[Any, Any]], ) -> None: """Enrich version entries with ``usageControl`` via batch hash endpoint. @@ -1151,7 +1159,7 @@ class ModelUpdateService: all_hashes = list(version_ids_by_hash.keys()) BATCH_SIZE = 100 - enrichment: Dict[int, Dict] = {} + enrichment: Dict[int, Dict[str, Any]] = {} try: for start in range(0, len(all_hashes), BATCH_SIZE): batch = all_hashes[start : start + BATCH_SIZE] @@ -1208,7 +1216,7 @@ class ModelUpdateService: version["earlyAccessEndsAt"] = extra["earlyAccessEndsAt"] @staticmethod - def _collect_hashes_from_response(response: Mapping) -> Dict[int, str]: + def _collect_hashes_from_response(response: Mapping[str, Any]) -> Dict[int, str]: """Extract ``{version_id: sha256}`` from a model-level API response. Returns an empty dict if the response structure is unexpected. @@ -1229,7 +1237,7 @@ class ModelUpdateService: return result @staticmethod - def _extract_sha256_from_version_entry(entry: Mapping) -> Optional[str]: + def _extract_sha256_from_version_entry(entry: Mapping[str, Any]) -> Optional[str]: """Return the SHA256 hash from the primary model file of a version entry.""" files = entry.get("files") if not isinstance(files, list): @@ -1253,22 +1261,19 @@ class ModelUpdateService: self, metadata_provider, model_ids: Sequence[int], - ) -> Dict[int, Mapping]: + ) -> Dict[int, Mapping[Any, Any]]: """Fetch model metadata in batches of up to 100 ids.""" BATCH_SIZE = 100 normalized = self._normalize_sequence(model_ids) - if not normalized: + provider = metadata_provider + if not normalized or provider is None: return {} - aggregated: Dict[int, Mapping] = {} + aggregated: Dict[int, Mapping[Any, Any]] = {} total_ids = len(normalized) total_batches = (total_ids + BATCH_SIZE - 1) // BATCH_SIZE - provider_name = ( - metadata_provider.__class__.__name__ - if metadata_provider is not None - else "unknown" - ) + provider_name = provider.__class__.__name__ for batch_index, start in enumerate(range(0, total_ids, BATCH_SIZE), start=1): chunk = normalized[start : start + BATCH_SIZE] logger.info( @@ -1279,7 +1284,7 @@ class ModelUpdateService: provider_name, ) try: - response = await metadata_provider.get_model_versions_bulk(chunk) + response = await provider.get_model_versions_bulk(chunk) except RateLimitError: raise if response is None: @@ -1356,7 +1361,7 @@ class ModelUpdateService: model_type: Optional[str] = None, model_id: Optional[int] = None, last_checked_at: Optional[float] = None, - version_info: Optional[Mapping] = None, + version_info: Optional[Mapping[str, Any]] = None, ) -> ModelUpdateRecord: local_set = set(normalized_local) # When folder-filtering, also consider versions in other folders @@ -1578,7 +1583,7 @@ class ModelUpdateService: if not isinstance(files, Iterable): return None - def parse_size(entry: Mapping) -> Optional[int]: + def parse_size(entry: Mapping[str, Any]) -> Optional[int]: size_kb = entry.get("sizeKB") if size_kb is None: return None @@ -1664,8 +1669,8 @@ class ModelUpdateService: return {} ids = list(model_ids) - status_rows: list = [] - version_rows: list = [] + status_rows: list[sqlite3.Row] = [] + version_rows: list[sqlite3.Row] = [] with self._connect() as conn: for start in range(0, len(ids), self._SQLITE_MAX_VARIABLES): diff --git a/py/services/persistent_model_cache.py b/py/services/persistent_model_cache.py index af6032aa..d3bb057f 100644 --- a/py/services/persistent_model_cache.py +++ b/py/services/persistent_model_cache.py @@ -4,7 +4,7 @@ import os import sqlite3 import threading from dataclasses import dataclass, field -from typing import Dict, List, Mapping, Optional, Sequence, Tuple +from typing import Any, Dict, List, Mapping, Optional, Sequence, Tuple from ..utils.cache_paths import CacheType, resolve_cache_path_with_migration @@ -15,7 +15,7 @@ logger = logging.getLogger(__name__) class PersistedCacheData: """Lightweight structure returned by the persistent cache.""" - raw_data: List[Dict] + raw_data: List[Dict[str, Any]] hash_rows: List[Tuple[str, str]] excluded_models: List[str] autov3_hash_rows: List[Tuple[str, str]] = field(default_factory=list) @@ -70,8 +70,8 @@ class PersistentModelCache: self._db_path = db_path or self._resolve_default_path(self._library_name) self._db_lock = threading.Lock() self._schema_initialized = False + directory = os.path.dirname(self._db_path) try: - directory = os.path.dirname(self._db_path) if directory: os.makedirs(directory, exist_ok=True) except Exception as exc: # pragma: no cover - defensive guard @@ -134,7 +134,7 @@ class PersistentModelCache: logger.warning("Failed to load persisted cache for %s: %s", model_type, exc) return None - raw_data: List[Dict] = [] + raw_data: List[Dict[str, Any]] = [] for row in rows: file_path: str = row["file_path"] trained_words = [] @@ -145,7 +145,7 @@ class PersistentModelCache: trained_words = [] creator_username = row["civitai_creator_username"] - civitai: Optional[Dict] = None + civitai: Optional[Dict[str, Any]] = None civitai_has_data = any( row[col] is not None for col in ("civitai_id", "civitai_model_id", "civitai_model_type", "civitai_name") @@ -223,7 +223,7 @@ class PersistentModelCache: autov3_hash_rows=autov3_pairs, ) - def save_cache(self, model_type: str, raw_data: Sequence[Dict], hash_index: Dict[str, List[str]], excluded_models: Sequence[str], autov3_hash_index: Optional[Dict[str, List[str]]] = None) -> None: + def save_cache(self, model_type: str, raw_data: Sequence[Dict[str, Any]], hash_index: Dict[str, List[str]], excluded_models: Sequence[str], autov3_hash_index: Optional[Dict[str, List[str]]] = None) -> None: if not self.is_enabled(): return if not self._schema_initialized: @@ -238,7 +238,7 @@ class PersistentModelCache: conn.execute("BEGIN") model_rows = [self._prepare_model_row(model_type, item) for item in raw_data] - model_map: Dict[str, Tuple] = { + model_map: Dict[str, Tuple[Any, ...]] = { row[1]: row for row in model_rows if row[1] # row[1] is file_path } @@ -279,8 +279,8 @@ class PersistentModelCache: to_remove_models, ) - insert_rows: List[Tuple] = [] - update_rows: List[Tuple] = [] + insert_rows: List[Tuple[Any, ...]] = [] + update_rows: List[Tuple[Any, ...]] = [] for file_path, row in model_map.items(): existing = existing_model_map.get(file_path) @@ -312,11 +312,11 @@ class PersistentModelCache: "SELECT file_path, tag FROM model_tags WHERE model_type = ?", (model_type,), ).fetchall() - existing_tags: Dict[str, set] = {} + existing_tags: Dict[str, set[str]] = {} for row in existing_tags_rows: existing_tags.setdefault(row["file_path"], set()).add(row["tag"]) - new_tags: Dict[str, set] = {} + new_tags: Dict[str, set[str]] = {} for item in raw_data: file_path = item.get("file_path") if not file_path: @@ -355,14 +355,14 @@ class PersistentModelCache: "SELECT sha256, file_path FROM hash_index WHERE model_type = ?", (model_type,), ).fetchall() - existing_hash_map: Dict[str, set] = {} + existing_hash_map: Dict[str, set[str]] = {} for row in existing_hash_rows: sha_value = (row["sha256"] or "").lower() if not sha_value: continue existing_hash_map.setdefault(sha_value, set()).add(row["file_path"]) - new_hash_map: Dict[str, set] = {} + new_hash_map: Dict[str, set[str]] = {} for sha_value, paths in hash_index.items(): normalized_sha = (sha_value or "").lower() if not normalized_sha: @@ -401,14 +401,14 @@ class PersistentModelCache: "SELECT autov3, file_path FROM autov3_index WHERE model_type = ?", (model_type,), ).fetchall() - existing_autov3_map: Dict[str, set] = {} + existing_autov3_map: Dict[str, set[str]] = {} for row in existing_autov3_rows: autov3_value = (row["autov3"] or "").lower() if not autov3_value: continue existing_autov3_map.setdefault(autov3_value, set()).add(row["file_path"]) - new_autov3_map: Dict[str, set] = {} + new_autov3_map: Dict[str, set[str]] = {} for autov3_value, paths in autov3_hash_index.items(): normalized_autov3 = (autov3_value or "").lower() if not normalized_autov3: @@ -600,7 +600,7 @@ class PersistentModelCache: conn.row_factory = sqlite3.Row return conn - def _prepare_model_row(self, model_type: str, item: Dict) -> Tuple: + def _prepare_model_row(self, model_type: str, item: Dict[str, Any]) -> Tuple[Any, ...]: civitai = item.get("civitai") or {} trained_words = civitai.get("trainedWords") if isinstance(trained_words, str): @@ -675,8 +675,8 @@ class PersistentModelCache: def update_single_model( self, model_type: str, - new_item: Dict, - old_item: Optional[Dict] = None, + new_item: Dict[str, Any], + old_item: Optional[Dict[str, Any]] = None, ) -> None: """Update a single model row in the persistent cache. @@ -715,8 +715,8 @@ class PersistentModelCache: conn.execute(self._insert_model_sql(), row) # --- tags --- - new_tags: set = set(new_item.get("tags") or []) - old_tags: set = set(old_item.get("tags") or []) if old_item else set() + new_tags: set[str] = set(new_item.get("tags") or []) + old_tags: set[str] = set(old_item.get("tags") or []) if old_item else set() tags_to_delete = old_tags - new_tags tags_to_insert = new_tags - old_tags diff --git a/py/services/persistent_recipe_cache.py b/py/services/persistent_recipe_cache.py index 952b5418..6c4af1b2 100644 --- a/py/services/persistent_recipe_cache.py +++ b/py/services/persistent_recipe_cache.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. """SQLite-based persistent cache for recipe metadata. This module provides fast recipe cache persistence using SQLite, enabling @@ -13,7 +17,7 @@ import os import sqlite3 import threading from dataclasses import dataclass, field -from typing import Dict, List, Optional, Set, Tuple +from typing import Any, Dict, List, Optional, Set, Tuple from ..utils.cache_paths import CacheType, resolve_cache_path_with_migration @@ -24,7 +28,7 @@ logger = logging.getLogger(__name__) class PersistedRecipeData: """Lightweight structure returned by the persistent recipe cache.""" - raw_data: List[Dict] + raw_data: List[Dict[str, Any]] file_stats: Dict[str, Tuple[float, int]] # json_path -> (mtime, size) image_id_map: Dict[str, str] = field(default_factory=dict) """Precomputed mapping of civitai image_id → recipe_id.""" @@ -63,8 +67,8 @@ class PersistentRecipeCache: self._db_path = db_path or self._resolve_default_path(self._library_name) self._db_lock = threading.Lock() self._schema_initialized = False + directory = os.path.dirname(self._db_path) try: - directory = os.path.dirname(self._db_path) if directory: os.makedirs(directory, exist_ok=True) except Exception as exc: @@ -140,7 +144,7 @@ class PersistentRecipeCache: logger.warning("Failed to load persisted recipe cache: %s", exc) return None - raw_data: List[Dict] = [] + raw_data: List[Dict[str, Any]] = [] file_stats: Dict[str, Tuple[float, int]] = {} for row in rows: @@ -162,7 +166,7 @@ class PersistentRecipeCache: def save_cache( self, - recipes: List[Dict], + recipes: List[Dict[str, Any]], json_paths: Optional[Dict[str, str]] = None, image_id_map: Optional[Dict[str, str]] = None, ) -> None: @@ -251,7 +255,7 @@ class PersistentRecipeCache: except Exception: return {} - def update_recipe(self, recipe: Dict, json_path: Optional[str] = None) -> None: + def update_recipe(self, recipe: Dict[str, Any], json_path: Optional[str] = None) -> None: """Update or insert a single recipe in the cache. Args: @@ -439,7 +443,7 @@ class PersistentRecipeCache: conn.row_factory = sqlite3.Row return conn - def _prepare_recipe_row(self, recipe: Dict, json_path: str) -> Tuple: + def _prepare_recipe_row(self, recipe: Dict[str, Any], json_path: str) -> Tuple[Any, ...]: """Convert a recipe dict to a row tuple for SQLite insertion.""" loras = recipe.get("loras") loras_json = json.dumps(loras) if loras else None @@ -486,7 +490,7 @@ class PersistentRecipeCache: tags_json, ) - def _row_to_recipe(self, row: sqlite3.Row) -> Dict: + def _row_to_recipe(self, row: sqlite3.Row) -> Dict[str, Any]: """Convert a SQLite row to a recipe dictionary.""" loras = [] if row["loras_json"]: diff --git a/py/services/preview_asset_service.py b/py/services/preview_asset_service.py index 8268553b..ac84ec5e 100644 --- a/py/services/preview_asset_service.py +++ b/py/services/preview_asset_service.py @@ -22,7 +22,7 @@ class PreviewAssetService: self, *, metadata_manager, - downloader_factory: Callable[[], Awaitable], + downloader_factory: Callable[[], Awaitable[Any]], exif_utils, ) -> None: self._metadata_manager = metadata_manager @@ -69,6 +69,8 @@ class PreviewAssetService: if not preview_url: return + preview_url = str(preview_url) + def extension_from_url(url: str, fallback: str) -> str: try: parsed = urlparse(url) diff --git a/py/services/recipe_cache.py b/py/services/recipe_cache.py index 048e8ce8..61c5542d 100644 --- a/py/services/recipe_cache.py +++ b/py/services/recipe_cache.py @@ -1,5 +1,5 @@ import asyncio -from typing import Iterable, List, Dict, Optional +from typing import Any, Iterable, List, Dict, Optional from dataclasses import dataclass, field from natsort import natsorted @@ -8,12 +8,13 @@ from natsort import natsorted class RecipeCache: """Cache structure for Recipe data""" - raw_data: List[Dict] - sorted_by_name: List[Dict] - sorted_by_date: List[Dict] + raw_data: List[Dict[str, Any]] + sorted_by_name: List[Dict[str, Any]] + sorted_by_date: List[Dict[str, Any]] folders: List[str] | None = None - folder_tree: Dict | None = None + folder_tree: Dict[str, Any] | None = None image_id_map: Dict[str, str] = field(default_factory=dict) + _lock: Any = field(init=False, repr=False, default=None) """Mapping of civitai image_id → recipe_id, precomputed at cache build time. Built once during cache initialization (O(n)) so that @@ -40,7 +41,7 @@ class RecipeCache: ) async def update_recipe_metadata( - self, recipe_id: str, metadata: Dict, *, resort: bool = True + self, recipe_id: str, metadata: Dict[str, Any], *, resort: bool = True ) -> bool: """Update metadata for a specific recipe in all cached data @@ -60,7 +61,7 @@ class RecipeCache: return True return False # Recipe not found - async def add_recipe(self, recipe_data: Dict, *, resort: bool = False) -> None: + async def add_recipe(self, recipe_data: Dict[str, Any], *, resort: bool = False) -> None: """Add a new recipe to the cache.""" async with self._lock: @@ -70,7 +71,7 @@ class RecipeCache: async def remove_recipe( self, recipe_id: str, *, resort: bool = False - ) -> Optional[Dict]: + ) -> Optional[Dict[str, Any]]: """Remove a recipe from the cache by ID. Args: @@ -91,7 +92,7 @@ class RecipeCache: async def bulk_remove( self, recipe_ids: Iterable[str], *, resort: bool = False - ) -> List[Dict]: + ) -> List[Dict[str, Any]]: """Remove multiple recipes from the cache.""" id_set = {str(recipe_id) for recipe_id in recipe_ids} @@ -111,7 +112,7 @@ class RecipeCache: return removed async def replace_recipe( - self, recipe_id: str, new_data: Dict, *, resort: bool = False + self, recipe_id: str, new_data: Dict[str, Any], *, resort: bool = False ) -> bool: """Replace cached data for a recipe.""" @@ -124,7 +125,7 @@ class RecipeCache: return True return False - async def get_recipe(self, recipe_id: str) -> Optional[Dict]: + async def get_recipe(self, recipe_id: str) -> Optional[Dict[str, Any]]: """Return a shallow copy of a cached recipe.""" async with self._lock: @@ -133,7 +134,7 @@ class RecipeCache: return dict(recipe) return None - async def snapshot(self) -> List[Dict]: + async def snapshot(self) -> List[Dict[str, Any]]: """Return a copy of all cached recipes.""" async with self._lock: diff --git a/py/services/recipe_fts_index.py b/py/services/recipe_fts_index.py index 29357ee1..0d30ab32 100644 --- a/py/services/recipe_fts_index.py +++ b/py/services/recipe_fts_index.py @@ -58,8 +58,8 @@ class RecipeFTSIndex: self._warned_not_ready = False # Ensure directory exists + directory = os.path.dirname(self._db_path) try: - directory = os.path.dirname(self._db_path) if directory: os.makedirs(directory, exist_ok=True) except Exception as exc: @@ -509,7 +509,7 @@ class RecipeFTSIndex: (recipe_id,) ) - def _prepare_fts_row(self, recipe: Dict[str, Any]) -> tuple: + def _prepare_fts_row(self, recipe: Dict[str, Any]) -> tuple[str, str, str, str, str, str, str]: """Prepare a row tuple for FTS insertion.""" recipe_id = str(recipe.get('id', '')) title = str(recipe.get('title', '')) diff --git a/py/services/recipe_scanner.py b/py/services/recipe_scanner.py index a4545a4f..bc178a79 100644 --- a/py/services/recipe_scanner.py +++ b/py/services/recipe_scanner.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. from __future__ import annotations import asyncio @@ -5,28 +9,20 @@ import json import logging import os import time -from typing import Any, Callable, Dict, Iterable, List, Optional, Set, Tuple +from typing import Any, Callable, Dict, Iterable, List, Optional, Set, Tuple, Union, cast from ..config import config from .recipe_cache import RecipeCache -from .recipe_fts_index import RecipeFTSIndex -from .persistent_recipe_cache import ( - PersistentRecipeCache, - get_persistent_recipe_cache, - PersistedRecipeData, -) -from .service_registry import ServiceRegistry -from .lora_scanner import LoraScanner -from .metadata_service import get_default_metadata_provider -from .checkpoint_scanner import CheckpointScanner -from .settings_manager import get_settings_manager from .recipes.errors import RecipeNotFoundError -from ..utils.civitai_utils import extract_civitai_image_id -from ..utils.utils import calculate_recipe_fingerprint from natsort import natsorted import sys import re -from ..recipes.merger import GenParamsMerger -from ..recipes.enrichment import RecipeEnricher +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from .lora_scanner import LoraScanner + from .checkpoint_scanner import CheckpointScanner + from .recipe_fts_index import RecipeFTSIndex + from .persistent_recipe_cache import PersistentRecipeCache, PersistedRecipeData logger = logging.getLogger(__name__) @@ -48,8 +44,12 @@ class RecipeScanner: if cls._instance is None: if not lora_scanner: # Get lora scanner from service registry if not provided + from .service_registry import ServiceRegistry + lora_scanner = await ServiceRegistry.get_lora_scanner() if not checkpoint_scanner: + from .service_registry import ServiceRegistry + checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner() cls._instance = cls(lora_scanner, checkpoint_scanner) return cls._instance @@ -77,17 +77,18 @@ class RecipeScanner: if not hasattr(self, "_initialized"): self._cache: Optional[RecipeCache] = None self._initialization_lock = asyncio.Lock() - self._initialization_task: Optional[asyncio.Task] = None + self._initialization_task: Optional[asyncio.Task[Any]] = None self._is_initializing = False self._mutation_lock = asyncio.Lock() - self._post_scan_task: Optional[asyncio.Task] = None - self._resort_tasks: Set[asyncio.Task] = set() + self._post_scan_task: Optional[asyncio.Task[Any]] = None + self._resort_tasks: Set[asyncio.Task[Any]] = set() self._cancel_requested = False # FTS index for fast search self._fts_index: Optional[RecipeFTSIndex] = None - self._fts_index_task: Optional[asyncio.Task] = None + self._fts_index_task: Optional[asyncio.Task[Any]] = None # Persistent cache for fast startup self._persistent_cache: Optional[PersistentRecipeCache] = None + self._civitai_client: Any = None # Lazily initialized from registry self._json_path_map: Dict[str, str] = {} # recipe_id -> json_path if lora_scanner: self._lora_scanner = lora_scanner @@ -123,6 +124,8 @@ class RecipeScanner: # Reset persistent cache instance for new library self._persistent_cache = None self._json_path_map = {} + from .persistent_recipe_cache import PersistentRecipeCache + PersistentRecipeCache.clear_instances() self._cache = None @@ -140,6 +143,8 @@ class RecipeScanner: async def _get_civitai_client(self): """Lazily initialize CivitaiClient from registry""" if self._civitai_client is None: + from .service_registry import ServiceRegistry + self._civitai_client = await ServiceRegistry.get_civitai_client() return self._civitai_client @@ -157,7 +162,7 @@ class RecipeScanner: return self._cancel_requested async def repair_all_recipes( - self, progress_callback: Optional[Callable[[Dict], Any]] = None + self, progress_callback: Optional[Callable[[Dict[str, Any]], Any]] = None ) -> Dict[str, Any]: """Repair all recipes by enrichment with Civitai and embedded metadata. @@ -339,6 +344,8 @@ class RecipeScanner: # 3. Use Enricher to repair/enrich try: + from ..recipes.enrichment import RecipeEnricher + updated = await RecipeEnricher.enrich_recipe(recipe, civitai_client) except Exception as e: logger.error(f"Error enriching recipe {recipe.get('id')}: {e}") @@ -490,6 +497,7 @@ class RecipeScanner: 3. Fall back to full directory scan if cache miss or reconciliation fails 4. Persist results for next startup """ + loop = None try: # Ensure cache exists to avoid None reference errors if self._cache is None: @@ -507,6 +515,8 @@ class RecipeScanner: # Initialize persistent cache if self._persistent_cache is None: + from .persistent_recipe_cache import get_persistent_recipe_cache + self._persistent_cache = get_persistent_recipe_cache() recipes_dir = self.recipes_dir @@ -592,13 +602,14 @@ class RecipeScanner: return self._cache if hasattr(self, "_cache") else None finally: # Clean up the event loop - loop.close() + if loop is not None: + loop.close() def _reconcile_recipe_cache( self, persisted: PersistedRecipeData, recipes_dir: str, - ) -> Tuple[List[Dict], bool, Dict[str, str]]: + ) -> Tuple[List[Dict[str, Any]], bool, Dict[str, str]]: """Reconcile persisted cache with current filesystem state. Args: @@ -608,7 +619,7 @@ class RecipeScanner: Returns: Tuple of (recipes list, changed flag, json_paths dict). """ - recipes: List[Dict] = [] + recipes: List[Dict[str, Any]] = [] json_paths: Dict[str, str] = {} changed = False @@ -625,12 +636,12 @@ class RecipeScanner: continue # Build recipe_id -> recipe lookup (O(n) instead of O(n²)) - recipe_by_id: Dict[str, Dict] = { + recipe_by_id: Dict[str, Dict[str, Any]] = { str(r.get("id", "")): r for r in persisted.raw_data if r.get("id") } # Build json_path -> recipe lookup from file_stats (O(m)) - persisted_by_path: Dict[str, Dict] = {} + persisted_by_path: Dict[str, Dict[str, Any]] = {} for json_path in persisted.file_stats.keys(): basename = os.path.basename(json_path) if basename.lower().endswith(".recipe.json"): @@ -696,7 +707,7 @@ class RecipeScanner: def _backfill_source_path_if_needed( self, - recipes: List[Dict], + recipes: List[Dict[str, Any]], json_paths: Dict[str, str], ) -> bool: """Backfill source_path from recipe JSON files if missing from cache. @@ -724,7 +735,7 @@ class RecipeScanner: def _full_directory_scan_sync( self, recipes_dir: str - ) -> Tuple[List[Dict], Dict[str, str]]: + ) -> Tuple[List[Dict[str, Any]], Dict[str, str]]: """Perform a full synchronous directory scan for recipes. Args: @@ -733,7 +744,7 @@ class RecipeScanner: Returns: Tuple of (recipes list, json_paths dict). """ - recipes: List[Dict] = [] + recipes: List[Dict[str, Any]] = [] json_paths: Dict[str, str] = {} # Get all recipe JSON files @@ -756,7 +767,7 @@ class RecipeScanner: return recipes, json_paths - def _load_recipe_file_sync(self, recipe_path: str) -> Optional[Dict]: + def _load_recipe_file_sync(self, recipe_path: str) -> Optional[Dict[str, Any]]: """Load a single recipe file synchronously. Args: @@ -835,6 +846,8 @@ class RecipeScanner: def _sort_cache_sync(self) -> None: """Sort cache data synchronously.""" + if self._cache is None: + return try: # Sort by name self._cache.sorted_by_name = natsorted( @@ -868,6 +881,8 @@ class RecipeScanner: source = recipe.get("source_path") if not source: continue + from ..utils.civitai_utils import extract_civitai_image_id + image_id = extract_civitai_image_id(source) if image_id and image_id not in mapping: recipe_id = recipe.get("id") @@ -950,6 +965,8 @@ class RecipeScanner: return try: + from .recipe_fts_index import RecipeFTSIndex + self._fts_index = RecipeFTSIndex() # Check if existing index is valid @@ -987,7 +1004,7 @@ class RecipeScanner: _build_fts(), name="recipe_fts_index_build" ) - def _search_with_fts(self, search: str, search_options: Dict) -> Optional[Set[str]]: + def _search_with_fts(self, search: str, search_options: Dict[str, Any]) -> Optional[Set[str]]: """Search recipes using FTS index if available. Args: @@ -1002,7 +1019,7 @@ class RecipeScanner: return None # Build the set of fields to search based on search_options - fields: Set[str] = set() + fields: Optional[Set[str]] = set() if search_options.get("title", True): fields.add("title") if search_options.get("tags", True): @@ -1033,12 +1050,12 @@ class RecipeScanner: return None def _update_fts_index_for_recipe( - self, recipe: Dict[str, Any], operation: str = "add" + self, recipe: Union[Dict[str, Any], str], operation: str = "add" ) -> None: """Update FTS index for a single recipe (add, update, or remove). Args: - recipe: The recipe dictionary. + recipe: The recipe dictionary, or a recipe ID string for removal. operation: One of 'add', 'update', or 'remove'. """ if not self._fts_index or not self._fts_index.is_ready(): @@ -1053,7 +1070,7 @@ class RecipeScanner: ) self._fts_index.remove_recipe(recipe_id) elif operation in ("add", "update"): - self._fts_index.update_recipe(recipe) + self._fts_index.update_recipe(cast(Dict[str, Any], recipe)) except Exception as exc: logger.debug("Failed to update FTS index for recipe: %s", exc) @@ -1071,6 +1088,8 @@ class RecipeScanner: if value in (None, ""): continue + from ..recipes.merger import GenParamsMerger + normalized_key = GenParamsMerger.NORMALIZATION_MAPPING.get(key, key) if normalized_key not in GenParamsMerger.ALLOWED_KEYS: continue @@ -1130,7 +1149,8 @@ class RecipeScanner: def _schedule_resort(self, *, name_only: bool = False) -> None: """Schedule a background resort of the recipe cache.""" - if not self._cache: + cache = self._cache + if not cache: return # Keep folder metadata up to date alongside sort order @@ -1138,7 +1158,7 @@ class RecipeScanner: async def _resort_wrapper() -> None: try: - await self._cache.resort(name_only=name_only) + await cache.resort(name_only=name_only) except Exception as exc: # pragma: no cover - defensive logging logger.error( "Recipe Scanner: error resorting cache: %s", exc, exc_info=True @@ -1164,10 +1184,10 @@ class RecipeScanner: except Exception: return "" - def _build_folder_tree(self, folders: list[str]) -> dict: + def _build_folder_tree(self, folders: list[str]) -> Dict[str, Any]: """Build a nested folder tree structure from relative folder paths.""" - tree: dict[str, dict] = {} + tree: dict[str, Dict[str, Any]] = {} for folder in folders: if not folder: continue @@ -1208,18 +1228,20 @@ class RecipeScanner: cache = await self.get_cached_data() self._update_folder_metadata(cache) - return cache.folders + return cache.folders or [] - async def get_folder_tree(self) -> dict: + async def get_folder_tree(self) -> Dict[str, Any]: """Return a hierarchical tree of recipe folders for sidebar navigation.""" cache = await self.get_cached_data() self._update_folder_metadata(cache) - return cache.folder_tree + return cache.folder_tree or {} @property def recipes_dir(self) -> str: """Get path to recipes directory""" + from .settings_manager import get_settings_manager + custom_recipes_dir = get_settings_manager().get("recipes_path", "") if isinstance(custom_recipes_dir, str) and custom_recipes_dir.strip(): recipes_dir = os.path.abspath( @@ -1242,7 +1264,7 @@ class RecipeScanner: # If cache is already initialized and no refresh is needed, return it immediately if self._cache is not None and not force_refresh: self._update_folder_metadata() - return self._cache + return cast(RecipeCache, self._cache) # If another initialization is already in progress, wait for it to complete if self._is_initializing and not force_refresh: @@ -1293,7 +1315,7 @@ class RecipeScanner: self._schedule_post_scan_enrichment() self._schedule_fts_index_build() - return self._cache + return cast(RecipeCache, self._cache) except Exception as e: logger.error( @@ -1344,6 +1366,8 @@ class RecipeScanner: source = recipe_data.get("source_path") if source: + from ..utils.civitai_utils import extract_civitai_image_id + image_id = extract_civitai_image_id(source) if image_id: recipe_id_value = recipe_data.get("id") @@ -1410,7 +1434,7 @@ class RecipeScanner: self._persistent_cache.save_image_id_map(cache.image_id_map) return len(removed) - async def scan_all_recipes(self) -> List[Dict]: + async def scan_all_recipes(self) -> List[Dict[str, Any]]: """Scan all recipe JSON files and return metadata""" recipes = [] recipes_dir = self.recipes_dir @@ -1436,7 +1460,7 @@ class RecipeScanner: return recipes - async def _load_recipe_file(self, recipe_path: str) -> Optional[Dict]: + async def _load_recipe_file(self, recipe_path: str) -> Optional[Dict[str, Any]]: """Load recipe data from a JSON file""" try: with open(recipe_path, "r", encoding="utf-8") as f: @@ -1517,6 +1541,8 @@ class RecipeScanner: # Calculate and update fingerprint if missing if "loras" in recipe_data and "fingerprint" not in recipe_data: + from ..utils.utils import calculate_recipe_fingerprint + fingerprint = calculate_recipe_fingerprint(recipe_data["loras"]) recipe_data["fingerprint"] = fingerprint @@ -1548,7 +1574,7 @@ class RecipeScanner: with open(recipe_path, "w", encoding="utf-8") as file_obj: json.dump(recipe_data, file_obj, indent=4, ensure_ascii=False) - async def _update_lora_information(self, recipe_data: Dict) -> bool: + async def _update_lora_information(self, recipe_data: Dict[str, Any]) -> bool: """Update LoRA information with hash and file_name Returns: @@ -1575,14 +1601,14 @@ class RecipeScanner: if isinstance(model_version_id, int) and model_version_id > 0: # Try to find in lora cache first hash_from_cache = await self._find_hash_in_lora_cache( - model_version_id + str(model_version_id) ) if hash_from_cache: lora["hash"] = hash_from_cache metadata_updated = True else: # If not in cache, fetch from Civitai - result = await self._get_hash_from_civitai(model_version_id) + result = await self._get_hash_from_civitai(str(model_version_id)) if isinstance(result, tuple): hash_from_civitai, is_deleted = result if hash_from_civitai: @@ -1645,14 +1671,16 @@ class RecipeScanner: logger.error(f"Error finding hash in lora cache: {e}") return None - async def _get_hash_from_civitai(self, model_version_id: str) -> Optional[str]: + async def _get_hash_from_civitai(self, model_version_id: str) -> Tuple[Optional[str], bool]: """Get hash from Civitai API""" try: # Get metadata provider instead of civitai client directly + from .metadata_service import get_default_metadata_provider + metadata_provider = await get_default_metadata_provider() if not metadata_provider: logger.error("Failed to get metadata provider") - return None + return None, False version_info, error_msg = await metadata_provider.get_model_version_info( model_version_id @@ -1733,7 +1761,7 @@ class RecipeScanner: return version_index.get(normalized_id) - async def _determine_base_model(self, loras: List[Dict]) -> Optional[str]: + async def _determine_base_model(self, loras: List[Dict[str, Any]]) -> Optional[str]: """Determine the most common base model among LoRAs""" base_models = {} @@ -1956,11 +1984,11 @@ class RecipeScanner: page: int, page_size: int, sort_by: str = "date", - search: str = None, - filters: dict = None, - search_options: dict = None, - lora_hash: str = None, - checkpoint_hash: str = None, + search: Optional[str] = None, + filters: Optional[Dict[str, Any]] = None, + search_options: Optional[Dict[str, Any]] = None, + lora_hash: Optional[str] = None, + checkpoint_hash: Optional[str] = None, bypass_filters: bool = True, folder: str | None = None, recursive: bool = True, @@ -2220,7 +2248,7 @@ class RecipeScanner: return result - async def get_recipe_by_id(self, recipe_id: str) -> dict: + async def get_recipe_by_id(self, recipe_id: str) -> Optional[Dict[str, Any]]: """Get a single recipe by ID with all metadata and formatted URLs Args: @@ -2312,7 +2340,7 @@ class RecipeScanner: return self._normalize_recipe_gen_params(recipe_data) - def _format_file_url(self, file_path: str) -> str: + def _format_file_url(self, file_path: Optional[str]) -> str: """Format file path as URL for serving in web UI""" if not file_path: return "/loras_static/images/no-preview.png" @@ -2360,7 +2388,7 @@ class RecipeScanner: return None - async def update_recipe_metadata(self, recipe_id: str, metadata: dict) -> bool: + async def update_recipe_metadata(self, recipe_id: str, metadata: Dict[str, Any]) -> bool: """Update recipe metadata (like title and tags) in both file system and cache Args: @@ -2465,6 +2493,8 @@ class RecipeScanner: lora_entry["modelVersionName"] = civitai_info.get("name", "") lora_entry["modelVersionId"] = civitai_info.get("id") + from ..utils.utils import calculate_recipe_fingerprint + recipe_data["fingerprint"] = calculate_recipe_fingerprint( recipe_data.get("loras", []) ) @@ -2696,7 +2726,7 @@ class RecipeScanner: return file_updated_count, cache_updated_count - async def find_recipes_by_fingerprint(self, fingerprint: str) -> list: + async def find_recipes_by_fingerprint(self, fingerprint: str) -> List[Dict[str, Any]]: """Find recipes with a matching fingerprint Args: @@ -2727,7 +2757,7 @@ class RecipeScanner: return matching_recipes - async def find_all_duplicate_recipes(self) -> dict: + async def find_all_duplicate_recipes(self) -> Dict[str, List[Any]]: """Find all recipe duplicates based on fingerprints Returns: @@ -2753,7 +2783,7 @@ class RecipeScanner: return duplicate_groups - async def find_duplicate_recipes_by_source(self) -> dict: + async def find_duplicate_recipes_by_source(self) -> Dict[str, List[Any]]: """Find all recipe duplicates based on source_path (Civitai image URLs) Returns: diff --git a/py/services/recipes/analysis_service.py b/py/services/recipes/analysis_service.py index 5f5da302..70fcb37f 100644 --- a/py/services/recipes/analysis_service.py +++ b/py/services/recipes/analysis_service.py @@ -101,6 +101,7 @@ class RecipeAnalysisService: temp_path = None metadata: Optional[dict[str, Any]] = None + image_info: Optional[dict[str, Any]] = None is_video = False extension = ".jpg" # Default @@ -413,7 +414,7 @@ class RecipeAnalysisService: error_msg = "This image does not contain any generation metadata (prompt, models, or parameters)" else: error_msg = "No parser found for this image" - payload = {"error": error_msg, "loras": []} + payload: dict[str, Any] = {"error": error_msg, "loras": []} if include_image_base64 and image_path: payload["image_base64"] = self._encode_file(image_path) payload["is_video"] = is_video @@ -494,7 +495,7 @@ class RecipeAnalysisService: getattr(tensor_image, "dtype", None), ) - import torch # type: ignore[import-not-found] + import torch # pyright: ignore[reportMissingImports] if isinstance(tensor_image, torch.Tensor): image_np = tensor_image.cpu().numpy() diff --git a/py/services/recipes/persistence_service.py b/py/services/recipes/persistence_service.py index 6852faf5..10381422 100644 --- a/py/services/recipes/persistence_service.py +++ b/py/services/recipes/persistence_service.py @@ -9,7 +9,7 @@ import shutil import time import uuid from dataclasses import dataclass -from typing import Any, Dict, Iterable, Optional +from typing import Any, Awaitable, Dict, Iterable, Optional, cast from ...config import config from ...recipes.constants import GEN_PARAM_KEYS @@ -72,6 +72,8 @@ class RecipePersistenceService: f"Missing required fields: {', '.join(missing_fields)}" ) + assert metadata is not None + resolved_image_bytes = self._resolve_image_bytes(image_bytes, image_base64) recipes_dir = target_dir or recipe_scanner.recipes_dir os.makedirs(recipes_dir, exist_ok=True) @@ -650,7 +652,9 @@ class RecipePersistenceService: for candidate in candidates: try: - checkpoint_info = await lookup(candidate) + checkpoint_info = await cast( + Awaitable[Any], lookup(candidate) + ) except Exception as exc: self._logger.debug( "Failed to lookup checkpoint %s while saving widget recipe: %s", diff --git a/py/services/server_i18n.py b/py/services/server_i18n.py index f038fbd4..3006b8a3 100644 --- a/py/services/server_i18n.py +++ b/py/services/server_i18n.py @@ -55,7 +55,7 @@ class ServerI18nManager: logger.warning(f"Locale {locale} not found, using 'en'") self.current_locale = 'en' - def get_translation(self, key: str, params: Dict[str, Any] = None, **kwargs) -> str: + def get_translation(self, key: str, params: Dict[str, Any] | None = None, **kwargs) -> str: """Get translation for a key with optional parameters (supports both dict and keyword args)""" # Merge kwargs into params for convenience if params is None: @@ -100,7 +100,7 @@ class ServerI18nManager: return value - def get_available_locales(self) -> list: + def get_available_locales(self) -> list[str]: """Get list of available locales""" return list(self.translations.keys()) diff --git a/py/services/service_registry.py b/py/services/service_registry.py index 162579c7..e4ff3ec3 100644 --- a/py/services/service_registry.py +++ b/py/services/service_registry.py @@ -1,3 +1,7 @@ +# pyright: reportImportCycles=false +# Lazy (function-local) imports still count as static edges in basedpyright's +# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms +# import cycles. Breaking them would require an architectural refactor. import asyncio import logging from typing import Optional, Dict, Any, TypeVar, Type diff --git a/py/services/settings_manager.py b/py/services/settings_manager.py index b0323ebc..86278fe8 100644 --- a/py/services/settings_manager.py +++ b/py/services/settings_manager.py @@ -12,6 +12,7 @@ from threading import Lock from typing import ( Any, Awaitable, + Coroutine, Dict, Iterable, List, @@ -411,7 +412,12 @@ class SettingsManager: needs_library_bootstrap = not isinstance(libraries, dict) or not libraries - if not needs_library_bootstrap and top_level_has_paths and len(libraries) == 1: + if ( + not needs_library_bootstrap + and top_level_has_paths + and isinstance(libraries, Mapping) + and len(libraries) == 1 + ): only_library_payload = next(iter(libraries.values())) if isinstance(only_library_payload, Mapping): folder_payload = only_library_payload.get("folder_paths") @@ -455,6 +461,9 @@ class SettingsManager: ): seed_library_name = target_name + if not isinstance(libraries, dict) or not libraries: + return + sanitized_libraries: Dict[str, Dict[str, Any]] = {} changed = False for name, data in libraries.items(): @@ -594,7 +603,7 @@ class SettingsManager: return payload def _normalize_folder_paths( - self, folder_paths: Mapping[str, Iterable[str]] + self, folder_paths: Mapping[str, Any] ) -> Dict[str, List[str]]: normalized: Dict[str, List[str]] = {} for key, values in folder_paths.items(): @@ -623,7 +632,7 @@ class SettingsManager: candidate_values = [values] else: try: - candidate_values = list(values) # type: ignore[arg-type] + candidate_values = list(values) # pyright: ignore[reportArgumentType] except TypeError: continue @@ -656,7 +665,7 @@ class SettingsManager: def _validate_folder_paths( self, library_name: str, - folder_paths: Mapping[str, Iterable[str]], + folder_paths: Mapping[str, Any], ) -> None: """Ensure folder paths do not overlap with other libraries. @@ -1119,7 +1128,7 @@ class SettingsManager: return [] if isinstance(value, str): - candidates: Iterable[str] = ( + candidates: Iterable[Any] = ( value.replace("\n", ",").replace(";", ",").split(",") ) elif isinstance(value, Sequence) and not isinstance( @@ -1167,7 +1176,7 @@ class SettingsManager: return [] if isinstance(value, str): - candidates: Iterable[str] = ( + candidates: Iterable[Any] = ( value.replace("\n", ",").replace(";", ",").split(",") ) elif isinstance(value, Sequence) and not isinstance( @@ -1207,7 +1216,7 @@ class SettingsManager: return [] if isinstance(value, str): - candidates: Iterable[str] = ( + candidates: Iterable[Any] = ( value.replace("\n", ",").replace(";", ",").split(",") ) elif isinstance(value, Sequence) and not isinstance( @@ -1595,11 +1604,11 @@ class SettingsManager: if key == "folder_paths" and isinstance(value, Mapping): active_name = self.get_active_library_name() self._validate_folder_paths(active_name, value) - self._update_active_library_entry(folder_paths=value) # type: ignore[arg-type] + self._update_active_library_entry(folder_paths=value) # pyright: ignore[reportArgumentType] elif key == "extra_folder_paths" and isinstance(value, Mapping): active_name = self.get_active_library_name() self._validate_folder_paths(active_name, value) - self._update_active_library_entry(extra_folder_paths=value) # type: ignore[arg-type] + self._update_active_library_entry(extra_folder_paths=value) # pyright: ignore[reportArgumentType] elif key == "default_lora_root": self._update_active_library_entry(default_lora_root=str(value)) elif key == "default_checkpoint_root": @@ -1752,12 +1761,12 @@ class SettingsManager: """Trigger cache resorting when the model name display preference updates.""" try: - from .service_registry import ServiceRegistry # type: ignore + from .service_registry import ServiceRegistry # pyright: ignore[reportImportCycles] except Exception: # pragma: no cover - registry optional in some contexts return display_mode = value if isinstance(value, str) else "model_name" - pending: List[Tuple[Optional[asyncio.AbstractEventLoop], Awaitable[Any]]] = [] + pending: List[Tuple[Optional[asyncio.AbstractEventLoop], Coroutine[Any, Any, Any]]] = [] def _resolve_service_loop(service: Any) -> Optional[asyncio.AbstractEventLoop]: loop = getattr(service, "loop", None) @@ -2118,7 +2127,7 @@ class SettingsManager: logger.debug("Failed to apply library settings to config: %s", exc) try: - from .service_registry import ServiceRegistry # type: ignore + from .service_registry import ServiceRegistry # pyright: ignore[reportImportCycles] for service_name in ( "lora_scanner", diff --git a/py/services/tag_fts_index.py b/py/services/tag_fts_index.py index e49261ad..a5f9cf1f 100644 --- a/py/services/tag_fts_index.py +++ b/py/services/tag_fts_index.py @@ -18,7 +18,7 @@ import sqlite3 import threading import time from pathlib import Path -from typing import Dict, List, Optional, Set +from typing import Any, Dict, List, Optional, Set from ..utils.cache_paths import CacheType, resolve_cache_path_with_migration @@ -87,10 +87,11 @@ class TagFTSIndex: self._indexing_in_progress = False self._schema_initialized = False self._warned_not_ready = False + self._needs_rebuild = False # Ensure directory exists + directory = os.path.dirname(self._db_path) try: - directory = os.path.dirname(self._db_path) if directory: os.makedirs(directory, exist_ok=True) except Exception as exc: @@ -358,7 +359,7 @@ class TagFTSIndex: finally: self._indexing_in_progress = False - def _insert_batch(self, conn: sqlite3.Connection, rows: List[tuple]) -> None: + def _insert_batch(self, conn: sqlite3.Connection, rows: List[tuple[str, int, int, str]]) -> None: """Insert a batch of rows into the database. Each row is a tuple of (tag_name, category, post_count, aliases). @@ -443,7 +444,7 @@ class TagFTSIndex: categories: Optional[List[int]] = None, limit: int = 20, offset: int = 0, - ) -> List[Dict]: + ) -> List[Dict[str, Any]]: """Search tags using FTS5 with prefix matching. Supports alias search: if the query matches an alias rather than @@ -530,7 +531,7 @@ class TagFTSIndex: categories: Optional[List[int]], limit: int, offset: int, - ) -> tuple[str, list[object]]: + ) -> tuple[str, list[int | str]]: """Build the SQL statement and params for a tag search.""" # Escape special LIKE characters and add wildcard query_escaped = ( diff --git a/py/services/tag_update_service.py b/py/services/tag_update_service.py index d2a9a7c1..c1081384 100644 --- a/py/services/tag_update_service.py +++ b/py/services/tag_update_service.py @@ -28,7 +28,8 @@ class TagUpdateService: metadata_path = f"{base}.metadata.json" metadata = await metadata_loader(metadata_path) - existing_tags = list(metadata.get("tags", [])) + raw_tags = metadata.get("tags", []) + existing_tags = list(raw_tags) if isinstance(raw_tags, list) else [] existing_lower = [tag.lower() for tag in existing_tags] tags_added: List[str] = [] diff --git a/py/services/use_cases/auto_organize_use_case.py b/py/services/use_cases/auto_organize_use_case.py index 06306bf0..b10d7c96 100644 --- a/py/services/use_cases/auto_organize_use_case.py +++ b/py/services/use_cases/auto_organize_use_case.py @@ -13,9 +13,11 @@ class AutoOrganizeLockProvider(Protocol): def is_auto_organize_running(self) -> bool: """Return ``True`` when an auto-organize operation is in-flight.""" + ... async def get_auto_organize_lock(self) -> asyncio.Lock: """Return the asyncio lock guarding auto-organize operations.""" + ... class AutoOrganizeInProgressError(RuntimeError): diff --git a/py/services/use_cases/bulk_metadata_refresh_use_case.py b/py/services/use_cases/bulk_metadata_refresh_use_case.py index 98136afd..3a650471 100644 --- a/py/services/use_cases/bulk_metadata_refresh_use_case.py +++ b/py/services/use_cases/bulk_metadata_refresh_use_case.py @@ -81,7 +81,7 @@ class BulkMetadataRefreshUseCase: async def emit(status: str, **extra: Any) -> None: if progress_callback is None: return - payload = { + payload: Dict[str, Any] = { "status": status, "total": total_models, "processed": processed, diff --git a/py/services/use_cases/example_images/import_example_images_use_case.py b/py/services/use_cases/example_images/import_example_images_use_case.py index 547b2f4e..e7d51614 100644 --- a/py/services/use_cases/example_images/import_example_images_use_case.py +++ b/py/services/use_cases/example_images/import_example_images_use_case.py @@ -5,9 +5,10 @@ from __future__ import annotations import os import tempfile from contextlib import suppress -from typing import Any, Dict, List +from typing import Any, Dict, List, cast from aiohttp import web +from aiohttp.multipart import BodyPartReader from ....utils.example_images_processor import ( ExampleImagesImportError, @@ -35,7 +36,8 @@ class ImportExampleImagesUseCase: if request.content_type and "multipart/form-data" in request.content_type: reader = await request.multipart() - first_field = await reader.next() + first_field_raw = await reader.next() + first_field = cast(BodyPartReader, first_field_raw) if first_field_raw is not None else None if first_field and first_field.name == "model_hash": model_hash = await first_field.text() else: @@ -43,7 +45,8 @@ class ImportExampleImagesUseCase: if first_field is not None: await self._collect_upload_file(first_field, files_to_import, temp_files) - async for field in reader: + async for raw_field in reader: + field = cast(BodyPartReader, raw_field) if field.name == "model_hash" and not model_hash: model_hash = await field.text() elif field.name == "files": @@ -53,6 +56,8 @@ class ImportExampleImagesUseCase: model_hash = data.get("model_hash") files_to_import = list(data.get("file_paths", [])) + if not model_hash: + raise ImportExampleImagesValidationError("Missing model_hash parameter") result = await self._processor.import_images(model_hash, files_to_import) return result except ExampleImagesValidationError as exc: diff --git a/py/services/websocket_manager.py b/py/services/websocket_manager.py index 2428cc19..9ffd89fd 100644 --- a/py/services/websocket_manager.py +++ b/py/services/websocket_manager.py @@ -1,6 +1,6 @@ import logging from aiohttp import web -from typing import Set, Dict, Optional +from typing import Set, Dict, Optional, Any from uuid import uuid4 import asyncio from datetime import datetime, timedelta @@ -15,13 +15,13 @@ class WebSocketManager: self._init_websockets: Set[web.WebSocketResponse] = set() # New set for initialization progress clients self._download_websockets: Dict[str, web.WebSocketResponse] = {} # New dict for download-specific clients # Add progress tracking dictionary - self._download_progress: Dict[str, Dict] = {} + self._download_progress: Dict[str, Dict[str, Any]] = {} # Cache last initialization progress payloads - self._last_init_progress: Dict[str, Dict] = {} + self._last_init_progress: Dict[str, Dict[str, Any]] = {} # Add auto-organize progress tracking - self._auto_organize_progress: Optional[Dict] = None + self._auto_organize_progress: Optional[Dict[str, Any]] = None # Add recipe repair progress tracking - self._recipe_repair_progress: Optional[Dict] = None + self._recipe_repair_progress: Optional[Dict[str, Any]] = None self._auto_organize_lock = asyncio.Lock() async def handle_connection(self, request: web.Request) -> web.WebSocketResponse: @@ -95,7 +95,7 @@ class WebSocketManager: self.cleanup_download_progress(download_id) logger.debug(f"Delayed cleanup completed for download {download_id}") - async def broadcast(self, data: Dict): + async def broadcast(self, data: Dict[str, Any]): """Broadcast message to all connected clients""" if not self._websockets: return @@ -106,7 +106,7 @@ class WebSocketManager: except Exception as e: logger.error(f"Error sending progress: {e}") - async def broadcast_init_progress(self, data: Dict): + async def broadcast_init_progress(self, data: Dict[str, Any]): """Broadcast initialization progress to connected clients""" payload = dict(data) if data else {} @@ -145,7 +145,7 @@ class WebSocketManager: except Exception as e: logger.debug(f'Error sending cached initialization progress: {e}') - def _get_init_progress_key(self, data: Dict) -> str: + def _get_init_progress_key(self, data: Dict[str, Any]) -> str: """Return a stable key for caching initialization progress payloads""" page_type = data.get('pageType') if page_type: @@ -155,7 +155,7 @@ class WebSocketManager: return f'scanner:{scanner_type}' return 'global' - async def broadcast_download_progress(self, download_id: str, data: Dict): + async def broadcast_download_progress(self, download_id: str, data: Dict[str, Any]): """Send progress update to specific download client""" progress_entry = { 'progress': data.get('progress', 0), @@ -183,7 +183,7 @@ class WebSocketManager: except Exception as e: logger.error(f"Error sending download progress: {e}") - async def broadcast_auto_organize_progress(self, data: Dict): + async def broadcast_auto_organize_progress(self, data: Dict[str, Any]): """Broadcast auto-organize progress to connected clients""" # Store progress data in memory self._auto_organize_progress = data @@ -191,7 +191,7 @@ class WebSocketManager: # Broadcast via WebSocket await self.broadcast(data) - async def broadcast_recipe_repair_progress(self, data: Dict): + async def broadcast_recipe_repair_progress(self, data: Dict[str, Any]): """Broadcast recipe repair progress to connected clients""" # Store progress data in memory self._recipe_repair_progress = data @@ -199,7 +199,7 @@ class WebSocketManager: # Broadcast via WebSocket await self.broadcast(data) - def get_auto_organize_progress(self) -> Optional[Dict]: + def get_auto_organize_progress(self) -> Optional[Dict[str, Any]]: """Get current auto-organize progress""" return self._auto_organize_progress @@ -207,7 +207,7 @@ class WebSocketManager: """Clear auto-organize progress data""" self._auto_organize_progress = None - def get_recipe_repair_progress(self) -> Optional[Dict]: + def get_recipe_repair_progress(self) -> Optional[Dict[str, Any]]: """Get current recipe repair progress""" return self._recipe_repair_progress @@ -234,7 +234,7 @@ class WebSocketManager: """Get the auto-organize lock""" return self._auto_organize_lock - def get_download_progress(self, download_id: str) -> Optional[Dict]: + def get_download_progress(self, download_id: str) -> Optional[Dict[str, Any]]: """Get progress information for a specific download""" return self._download_progress.get(download_id) @@ -255,7 +255,7 @@ class WebSocketManager: self._download_progress.pop(download_id, None) logger.debug(f"Cleaned up old download progress for {download_id}") - async def broadcast_cache_health_warning(self, report: 'HealthReport', page_type: str = None): + async def broadcast_cache_health_warning(self, report: 'HealthReport', page_type: Optional[str] = None): """ Broadcast cache health warning to frontend. diff --git a/py/services/websocket_progress_callback.py b/py/services/websocket_progress_callback.py index 21423044..ba496ef6 100644 --- a/py/services/websocket_progress_callback.py +++ b/py/services/websocket_progress_callback.py @@ -24,6 +24,6 @@ class WebSocketProgressCallback(ProgressCallback): class WebSocketBroadcastCallback: """Generic WebSocket progress callback broadcasting to all clients.""" - async def on_progress(self, progress_data: Dict[str, Any]) -> None: + async def on_progress(self, payload: Dict[str, Any]) -> None: """Send the provided payload to all connected clients.""" - await ws_manager.broadcast(progress_data) + await ws_manager.broadcast(payload) diff --git a/py/utils/civitai_utils.py b/py/utils/civitai_utils.py index 194f573d..9466d0dd 100644 --- a/py/utils/civitai_utils.py +++ b/py/utils/civitai_utils.py @@ -224,7 +224,7 @@ def _normalize_commercial_values(value: Any) -> Sequence[str]: if result: return result try: - if len(value) == 0: # type: ignore[arg-type] + if len(value) == 0: # pyright: ignore[reportArgumentType] return [] except TypeError: pass diff --git a/py/utils/example_images_download_manager.py b/py/utils/example_images_download_manager.py index cd5f8552..80a754ff 100644 --- a/py/utils/example_images_download_manager.py +++ b/py/utils/example_images_download_manager.py @@ -35,7 +35,7 @@ class ExampleImagesDownloadError(RuntimeError): class DownloadInProgressError(ExampleImagesDownloadError): """Raised when a download is already running.""" - def __init__(self, progress_snapshot: dict) -> None: + def __init__(self, progress_snapshot: Dict[str, Any]) -> None: super().__init__("Download already in progress") self.progress_snapshot = progress_snapshot @@ -54,7 +54,7 @@ class DownloadConfigurationError(ExampleImagesDownloadError): logger = logging.getLogger(__name__) -class _DownloadProgress(dict): +class _DownloadProgress(dict[str, Any]): """Mutable mapping maintaining download progress with set-aware serialisation.""" def __init__(self) -> None: @@ -80,7 +80,7 @@ class _DownloadProgress(dict): rate_limited_models=set(), ) - def snapshot(self) -> dict: + def snapshot(self) -> Dict[str, Any]: """Return a JSON-serialisable snapshot of the current progress.""" snapshot = dict(self) @@ -149,7 +149,7 @@ class DownloadManager: """Manages downloading example images for models.""" def __init__(self, *, ws_manager, state_lock: asyncio.Lock | None = None) -> None: - self._download_task: asyncio.Task | None = None + self._download_task: asyncio.Task[Any] | None = None self._is_downloading = False self._progress = _DownloadProgress() self._ws_manager = ws_manager @@ -162,7 +162,7 @@ class DownloadManager: return "" return ensure_library_root_exists(library_name) - async def start_download(self, options: dict): + async def start_download(self, options: Dict[str, Any]): """Start downloading example images for models.""" # Step 1: Parse options (fast, non-blocking) @@ -269,7 +269,7 @@ class DownloadManager: return {"success": True, "message": "Download started", "status": snapshot} - def _handle_download_task_done(self, task: asyncio.Task, output_dir: str) -> None: + def _handle_download_task_done(self, task: asyncio.Task[Any], output_dir: str) -> None: """Handle download task completion, including saving progress on error.""" try: # This will re-raise any exception from the task @@ -282,7 +282,7 @@ class DownloadManager: except Exception as save_error: logger.error(f"Failed to save progress after task failure: {save_error}") - async def get_status(self, request) -> dict: + async def get_status(self, request) -> Dict[str, Any]: """Get the current status of example images download.""" return { @@ -291,7 +291,7 @@ class DownloadManager: "status": self._progress.snapshot(), } - async def _load_progress_file(self, output_dir: str) -> tuple[str, set, set, set]: + async def _load_progress_file(self, output_dir: str) -> tuple[str, set[str], set[str], set[str]]: """Load progress file from disk. Returns (progress_file_path, processed_models, failed_models, rate_limited_models). This is a separate async method to allow running in executor to avoid blocking event loop. @@ -301,7 +301,7 @@ class DownloadManager: None, self._load_progress_file_sync, output_dir ) - def _load_progress_file_sync(self, output_dir: str) -> tuple[str, set, set, set]: + def _load_progress_file_sync(self, output_dir: str) -> tuple[str, set[str], set[str], set[str]]: """Synchronous implementation of progress file loading. Returns: @@ -356,7 +356,7 @@ class DownloadManager: return progress_file, processed_models, failed_models, rate_limited_models - def _load_progress_sets_sync(self, progress_file: str) -> tuple[set, set]: + def _load_progress_sets_sync(self, progress_file: str) -> tuple[set[str], set[str]]: """Load only the processed and failed model sets from progress file. This is a lighter version for quick checks without legacy migration. @@ -377,7 +377,7 @@ class DownloadManager: return processed_models, failed_models - async def check_pending_models(self, model_types: list[str]) -> dict: + async def check_pending_models(self, model_types: list[str]) -> Dict[str, Any]: """Quickly check how many models need example images downloaded. This is a lightweight check that avoids the overhead of starting @@ -1000,7 +1000,7 @@ class DownloadManager: except Exception as e: logger.error(f"Failed to save progress file: {e}") - async def start_force_download(self, options: dict): + async def start_force_download(self, options: Dict[str, Any]): """Force download example images for specific models.""" async with self._state_lock: diff --git a/py/utils/example_images_metadata.py b/py/utils/example_images_metadata.py index 42c1529f..2fed61d1 100644 --- a/py/utils/example_images_metadata.py +++ b/py/utils/example_images_metadata.py @@ -93,7 +93,7 @@ class MetadataUpdater: """Handles updating model metadata related to example images""" @staticmethod - async def refresh_model_metadata(model_hash, model_name, scanner_type, scanner, progress: dict | None = None): + async def refresh_model_metadata(model_hash, model_name, scanner_type, scanner, progress: dict[str, Any] | None = None): """Refresh model metadata from CivitAI Args: @@ -263,8 +263,9 @@ class MetadataUpdater: model_copy: Optional[Dict[str, Any]] = None try: model_copy = model.copy() - model_copy.pop('folder', None) - await MetadataManager.save_metadata(file_path, model_copy) + if model_copy is not None: + model_copy.pop('folder', None) + await MetadataManager.save_metadata(file_path, model_copy) logger.info(f"Saved metadata for {model.get('model_name')}") except Exception as e: logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}") @@ -371,8 +372,9 @@ class MetadataUpdater: if file_path: try: model_copy = model_data.copy() - model_copy.pop('folder', None) - await MetadataManager.save_metadata(file_path, model_copy) + if model_copy is not None: + model_copy.pop('folder', None) + await MetadataManager.save_metadata(file_path, model_copy) logger.info(f"Saved metadata for {model_data.get('model_name')}") except Exception as e: logger.error(f"Failed to save metadata: {str(e)}") @@ -553,7 +555,7 @@ class MetadataUpdater: images = civitai.get("images") if isinstance(images, list) and images: - stale: list[int] = [] + stale_images: list[int] = [] for idx, img in enumerate(images): if img.get("url", ""): @@ -563,15 +565,15 @@ class MetadataUpdater: prefix = f"image_{idx}." if not any(f.startswith(prefix) for f in dir_entries): - stale.append(idx) + stale_images.append(idx) - if stale: - for idx in reversed(stale): + if stale_images: + for idx in reversed(stale_images): images.pop(idx) has_changes = True logger.info( "Pruned %d stale image entry(ies) for %s", - len(stale), + len(stale_images), getattr(metadata, "model_name", model_hash), ) diff --git a/py/utils/example_images_migration.py b/py/utils/example_images_migration.py index 99786f1e..abf606a4 100644 --- a/py/utils/example_images_migration.py +++ b/py/utils/example_images_migration.py @@ -371,7 +371,7 @@ class ExampleImagesMigration: found = True break - if not found: + if not found or old_path is None: logger.warning(f"Could not find file for index {index} in {model_hash}, skipping") continue diff --git a/py/utils/example_images_processor.py b/py/utils/example_images_processor.py index 2853c955..ed6fed54 100644 --- a/py/utils/example_images_processor.py +++ b/py/utils/example_images_processor.py @@ -4,6 +4,7 @@ import os import re import random import string +from typing import Any from aiohttp import web from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS from ..services.service_registry import ServiceRegistry @@ -259,7 +260,7 @@ class ExampleImagesProcessor: logger.debug("File already exists, skipping download for %s", image_url) continue - async def _attempt_download() -> tuple: + async def _attempt_download() -> tuple[bool, Any, Any]: logger.debug("Downloading media file %s for %s", i, model_name) return await downloader.download_to_memory( image_url, diff --git a/py/utils/exif_utils.py b/py/utils/exif_utils.py index be04a227..169fbb06 100644 --- a/py/utils/exif_utils.py +++ b/py/utils/exif_utils.py @@ -4,13 +4,13 @@ import logging import os import struct from io import BytesIO -from typing import Any, Optional, Tuple +from typing import Any, Optional, Tuple, cast -import piexif +import piexif # pyright: ignore[reportMissingTypeStubs] from PIL import Image, PngImagePlugin try: - import brotli + import brotli # pyright: ignore[reportMissingTypeStubs] _BROTLI_AVAILABLE = True except ImportError: brotli = None @@ -38,7 +38,7 @@ class ExifUtils: """Utility functions for working with EXIF data in images""" @staticmethod - def _parse_isobmff_boxes(data: bytes, offset: int = 0) -> list[dict]: + def _parse_isobmff_boxes(data: bytes, offset: int = 0) -> list[dict[str, Any]]: boxes = [] while offset + 8 <= len(data): size = struct.unpack('>I', data[offset:offset + 4])[0] @@ -78,7 +78,7 @@ class ExifUtils: _BROTLI_MAX_DECOMPRESSED = 2 * 1024 * 1024 @staticmethod - def _extract_isobmff_brotli(image_path: str) -> Optional[dict]: + def _extract_isobmff_brotli(image_path: str) -> Optional[dict[str, Any]]: try: with open(image_path, 'rb') as f: data = f.read() @@ -107,7 +107,7 @@ class ExifUtils: if _BROTLI_AVAILABLE: try: - decompressed = brotli.decompress(compressed) + decompressed = brotli.decompress(compressed) # pyright: ignore[reportOptionalMemberAccess] if len(decompressed) > ExifUtils._BROTLI_MAX_DECOMPRESSED: logger.warning( "Brotli metadata too large (%d bytes, max %d), ignoring", @@ -126,7 +126,9 @@ class ExifUtils: except Exception: return None - result = {"parameters": None, "prompt": None, "workflow": None, "comment": None} + result: dict[str, Optional[str]] = { + "parameters": None, "prompt": None, "workflow": None, "comment": None + } if isinstance(meta.get("prompt"), (dict, list)): result["prompt"] = json.dumps(meta["prompt"]) elif isinstance(meta.get("prompt"), str): @@ -161,7 +163,7 @@ class ExifUtils: @staticmethod def _load_structured_metadata(image_path: str) -> dict[str, Optional[str]]: - metadata = { + metadata: dict[str, Optional[str]] = { "parameters": None, "prompt": None, "workflow": None, @@ -197,13 +199,14 @@ class ExifUtils: logger.debug(f"Error loading EXIF data: {e}") exif_dict = {} - if piexif.ExifIFD.UserComment in exif_dict.get("Exif", {}): + exif_ifd = exif_dict.get("Exif") + if exif_ifd and piexif.ExifIFD.UserComment in exif_ifd: metadata["comment"] = ExifUtils._decode_user_comment( - exif_dict["Exif"][piexif.ExifIFD.UserComment] + exif_ifd[piexif.ExifIFD.UserComment] ) image_description = ExifUtils._decode_exif_text( - exif_dict.get("0th", {}).get(piexif.ImageIFD.ImageDescription) + (exif_dict.get("0th") or {}).get(piexif.ImageIFD.ImageDescription) ) if image_description: if image_description.startswith("Workflow:"): @@ -253,19 +256,26 @@ class ExifUtils: workflow = metadata_fields.get("workflow") prompt = metadata_fields.get("prompt") + # Work on local references, then write the (possibly new) IFD dicts back. + exif_ifd = exif_dict.get("Exif") or {} + exif_0th = exif_dict.get("0th") or {} + if parameters: - exif_dict["Exif"][piexif.ExifIFD.UserComment] = ( + exif_ifd[piexif.ExifIFD.UserComment] = ( b"UNICODE\0" + parameters.encode("utf-16be") ) else: - exif_dict["Exif"].pop(piexif.ExifIFD.UserComment, None) + exif_ifd.pop(piexif.ExifIFD.UserComment, None) if workflow: - exif_dict["0th"][piexif.ImageIFD.ImageDescription] = f"Workflow:{workflow}" + exif_0th[piexif.ImageIFD.ImageDescription] = f"Workflow:{workflow}" elif prompt: - exif_dict["0th"][piexif.ImageIFD.ImageDescription] = prompt + exif_0th[piexif.ImageIFD.ImageDescription] = prompt else: - exif_dict["0th"].pop(piexif.ImageIFD.ImageDescription, None) + exif_0th.pop(piexif.ImageIFD.ImageDescription, None) + + exif_dict["Exif"] = exif_ifd + exif_dict["0th"] = exif_0th return piexif.dump(exif_dict) @@ -326,7 +336,7 @@ class ExifUtils: exif_bytes = ExifUtils._build_exif_bytes( metadata_fields, img.info.get("exif") ) - save_kwargs = {"exif": exif_bytes} + save_kwargs: dict[str, Any] = {"exif": exif_bytes} if img_format == "WEBP": save_kwargs["quality"] = 85 @@ -499,12 +509,12 @@ class ExifUtils: else: # It's binary data - validate data try: - with BytesIO(image_data) as temp_buf: + with BytesIO(cast(bytes, image_data)) as temp_buf: test_img = Image.open(temp_buf) # Verify the image can be fully loaded width, height = test_img.size # If successful, reopen for processing - img = Image.open(BytesIO(image_data)) + img = Image.open(BytesIO(cast(bytes, image_data))) except Exception as e: logger.error(f"Invalid binary image data: {e}") raise ValueError(f"Cannot process corrupt image data: {e}") @@ -521,7 +531,7 @@ class ExifUtils: import tempfile with tempfile.NamedTemporaryFile(suffix='.jpg', delete=False) as temp_file: temp_path = temp_file.name - temp_file.write(image_data) + temp_file.write(cast(bytes, image_data)) try: metadata_fields = ExifUtils._load_structured_metadata(temp_path) except Exception as e: @@ -542,7 +552,7 @@ class ExifUtils: # Resize the image with error handling try: - resized_img = img.resize((target_width, new_height), Image.LANCZOS) + resized_img = img.resize((target_width, new_height), Image.Resampling.LANCZOS) except Exception as e: logger.error(f"Failed to resize image: {e}") # Return original image if resize fails diff --git a/py/utils/lora_metadata.py b/py/utils/lora_metadata.py index 87d72516..dc376481 100644 --- a/py/utils/lora_metadata.py +++ b/py/utils/lora_metadata.py @@ -1,5 +1,5 @@ from safetensors import safe_open -from typing import Dict, List, Tuple +from typing import Dict, List, Optional, Tuple from .model_utils import determine_base_model import os import logging @@ -7,7 +7,7 @@ import json logger = logging.getLogger(__name__) -async def extract_lora_metadata(file_path: str) -> Dict: +async def extract_lora_metadata(file_path: str) -> Dict[str, str]: """Extract essential metadata from safetensors file""" try: with safe_open(file_path, framework="pt", device="cpu") as f: @@ -20,7 +20,7 @@ async def extract_lora_metadata(file_path: str) -> Dict: logger.error(f"Error reading metadata from {file_path}: {str(e)}") return {"base_model": "Unknown"} -async def extract_checkpoint_metadata(file_path: str) -> dict: +async def extract_checkpoint_metadata(file_path: str) -> dict[str, str]: """Extract metadata from a checkpoint file to determine model type and base model""" try: # Analyze filename for clues about the model @@ -83,7 +83,7 @@ async def extract_checkpoint_metadata(file_path: str) -> dict: # Return default values return {'base_model': 'Unknown', 'model_type': 'checkpoint'} -async def extract_trained_words(file_path: str) -> Tuple[List[Tuple[str, int]], str]: +async def extract_trained_words(file_path: str) -> Tuple[List[Tuple[str, int]], Optional[str]]: """Extract trained words from a safetensors file and sort by frequency Args: diff --git a/py/utils/metadata_manager.py b/py/utils/metadata_manager.py index 7dbe1b84..baac249e 100644 --- a/py/utils/metadata_manager.py +++ b/py/utils/metadata_manager.py @@ -3,9 +3,9 @@ import os import json import logging import time -from typing import Any, Dict, Optional, Type, Union +from typing import Any, Dict, Optional, Type, Union, cast -from .models import BaseModelMetadata, LoraMetadata +from .models import BaseModelMetadata, CheckpointMetadata, EmbeddingMetadata, LoraMetadata from .file_utils import normalize_path, find_preview_file, calculate_sha256, calculate_autov3 from .lora_metadata import extract_lora_metadata, extract_checkpoint_metadata @@ -56,13 +56,13 @@ class MetadataManager: return None, True # should_skip = True @staticmethod - async def load_metadata_payload(file_path: str) -> Dict: + async def load_metadata_payload(file_path: str) -> Dict[str, Any]: """ Load metadata and return it as a dictionary, including any unknown fields. Falls back to reading the raw JSON file if parsing into a model class fails. """ - payload: Dict = {} + payload: Dict[str, Any] = {} metadata_obj, should_skip = await MetadataManager.load_metadata(file_path) if metadata_obj: @@ -120,7 +120,7 @@ class MetadataManager: return model_data @staticmethod - async def save_metadata(path: str, metadata: Union[BaseModelMetadata, Dict]) -> bool: + async def save_metadata(path: str, metadata: Union[BaseModelMetadata, Dict[str, Any]]) -> bool: """ Save metadata with atomic write operations. @@ -217,7 +217,7 @@ class MetadataManager: # Create instance based on model type if model_class.__name__ == "CheckpointMetadata": - metadata = model_class( + metadata = cast(Type[CheckpointMetadata], model_class)( file_name=base_name, model_name=base_name, file_path=normalize_path(file_path), @@ -232,7 +232,7 @@ class MetadataManager: from_civitai=True ) elif model_class.__name__ == "EmbeddingMetadata": - metadata = model_class( + metadata = cast(Type[EmbeddingMetadata], model_class)( file_name=base_name, model_name=base_name, file_path=normalize_path(file_path), @@ -247,7 +247,7 @@ class MetadataManager: from_civitai=True ) else: # Default to LoraMetadata - metadata = model_class( + metadata = cast(Type[LoraMetadata], model_class)( file_name=base_name, model_name=base_name, file_path=normalize_path(file_path), diff --git a/py/utils/models.py b/py/utils/models.py index 587b243c..cfe96585 100644 --- a/py/utils/models.py +++ b/py/utils/models.py @@ -1,5 +1,5 @@ from dataclasses import dataclass, asdict, field -from typing import Dict, Optional, List, Any +from typing import Callable, Dict, Optional, List, Any from datetime import datetime import os from .constants import INVALID_AUTOV3_EMPTY_HASH @@ -23,7 +23,7 @@ def normalize_autov3(value: Any) -> Optional[str]: return None -def autov3_from_civitai_files(civitai_data: Optional[Dict], sha256: str) -> Optional[str]: +def autov3_from_civitai_files(civitai_data: Optional[Dict[str, Any]], sha256: str) -> Optional[str]: """Extract the AutoV3 hash from Civitai metadata for the matching file. Civitai versions can ship multiple files; the AutoV3 hash is only valid @@ -64,7 +64,7 @@ class BaseModelMetadata: civitai: Dict[str, Any] = field( default_factory=dict ) # Civitai API data if available - tags: List[str] = None # Model tags + tags: List[str] = field(default_factory=list) # Model tags modelDescription: str = "" # Full model description civitai_deleted: bool = False # Whether deleted from Civitai favorite: bool = False # Whether the model is a favorite @@ -96,7 +96,7 @@ class BaseModelMetadata: self.trainedWords = [] @classmethod - def from_dict(cls, data: Dict) -> "BaseModelMetadata": + def from_dict(cls, data: Dict[str, Any]) -> "BaseModelMetadata": """Create instance from dictionary""" data_copy = data.copy() @@ -136,7 +136,7 @@ class BaseModelMetadata: return instance - def to_dict(self) -> Dict: + def to_dict(self) -> Dict[str, Any]: """Convert to dictionary for JSON serialization""" result = asdict(self) @@ -158,7 +158,7 @@ class BaseModelMetadata: return result - def update_civitai_info(self, civitai_data: Dict) -> None: + def update_civitai_info(self, civitai_data: Dict[str, Any]) -> None: """Update Civitai information. Civitai's AutoV3 is the authoritative hash for recipe matching, so @@ -192,7 +192,7 @@ class BaseModelMetadata: @staticmethod def generate_unique_filename( - target_dir: str, base_name: str, extension: str, hash_provider: callable = None + target_dir: str, base_name: str, extension: str, hash_provider: Optional[Callable[[], str]] = None ) -> str: """Generate a unique filename to avoid conflicts @@ -243,7 +243,7 @@ class LoraMetadata(BaseModelMetadata): @classmethod def from_civitai_info( - cls, version_info: Dict, file_info: Dict, save_path: str + cls, version_info: Dict[str, Any], file_info: Dict[str, Any], save_path: str ) -> "LoraMetadata": """Create LoraMetadata instance from Civitai version info""" file_name = file_info.get("name", "") @@ -287,7 +287,7 @@ class CheckpointMetadata(BaseModelMetadata): @classmethod def from_civitai_info( - cls, version_info: Dict, file_info: Dict, save_path: str + cls, version_info: Dict[str, Any], file_info: Dict[str, Any], save_path: str ) -> "CheckpointMetadata": """Create CheckpointMetadata instance from Civitai version info""" file_name = file_info.get("name", "") @@ -332,7 +332,7 @@ class EmbeddingMetadata(BaseModelMetadata): @classmethod def from_civitai_info( - cls, version_info: Dict, file_info: Dict, save_path: str + cls, version_info: Dict[str, Any], file_info: Dict[str, Any], save_path: str ) -> "EmbeddingMetadata": """Create EmbeddingMetadata instance from Civitai version info""" file_name = file_info.get("name", "") diff --git a/py/utils/preview_selection.py b/py/utils/preview_selection.py index 91cfbd30..378c65d7 100644 --- a/py/utils/preview_selection.py +++ b/py/utils/preview_selection.py @@ -15,7 +15,7 @@ def _extract_nsfw_level(entry: Mapping[str, object]) -> int: value = entry.get("nsfwLevel", 0) try: - return int(value) # type: ignore[return-value] + return int(value) # pyright: ignore[reportArgumentType, reportReturnType] except (TypeError, ValueError): return 0 diff --git a/py/utils/usage_stats.py b/py/utils/usage_stats.py index 774b887e..577849db 100644 --- a/py/utils/usage_stats.py +++ b/py/utils/usage_stats.py @@ -6,7 +6,7 @@ import asyncio import logging import datetime import shutil -from typing import Dict, Set +from typing import Any, Awaitable, Dict, Set, cast from ..config import config from ..services.service_registry import ServiceRegistry @@ -68,7 +68,7 @@ class UsageStats: return # Initialize stats storage - self.stats = { + self.stats: Dict[str, Any] = { "checkpoints": {}, # sha256 -> { total: count, history: { date: count } } "loras": {}, # sha256 -> { total: count, history: { date: count } } "embeddings": {}, # sha256 -> { total: count, history: { date: count } } @@ -297,8 +297,8 @@ class UsageStats: # Process each prompt_id try: - registry = MetadataRegistry() - except NameError: + registry = MetadataRegistry() # pyright: ignore[reportPossiblyUnboundVariable] + except (ImportError, NameError): # MetadataRegistry not available (standalone mode) registry = None @@ -374,7 +374,7 @@ class UsageStats: if not callable(get_cached_data): return None - cache = await get_cached_data() + cache = await cast(Awaitable[Any], get_cached_data()) raw_data = getattr(cache, "raw_data", None) if not isinstance(raw_data, list): return None @@ -404,7 +404,7 @@ class UsageStats: if not callable(get_model_roots): return None - roots = [root for root in get_model_roots() if root] + roots = [root for root in cast(Any, get_model_roots()) if root] if not roots: return None @@ -486,7 +486,7 @@ class UsageStats: model_filename, file_path, ) - calculated_hash = await calculate_hash(file_path) + calculated_hash = await cast(Awaitable[Any], calculate_hash(file_path)) if calculated_hash: return calculated_hash @@ -557,7 +557,7 @@ class UsageStats: logger.error(f"Error processing LoRA usage: {e}", exc_info=True) @staticmethod - def _extract_embedding_names(prompt_text: str) -> set: + def _extract_embedding_names(prompt_text: str) -> set[str]: """Parse embedding:name references from prompt text. ComfyUI's SDTokenizer resolves ``embedding:`` during tokenization @@ -605,7 +605,7 @@ class UsageStats: except Exception as e: logger.error("Error processing embedding usage: %s", e, exc_info=True) - async def get_stats(self): + async def get_stats(self) -> Dict[str, Any]: """Get current usage statistics""" return self.stats @@ -633,7 +633,7 @@ class UsageStats: try: # Process metadata for this prompt_id - registry = MetadataRegistry() + registry = MetadataRegistry() # pyright: ignore[reportPossiblyUnboundVariable] metadata = registry.get_metadata(prompt_id) if metadata: await self._process_metadata(metadata) diff --git a/py/utils/utils.py b/py/utils/utils.py index 4ea8c8bf..cbf67d3d 100644 --- a/py/utils/utils.py +++ b/py/utils/utils.py @@ -261,7 +261,7 @@ def get_checkpoint_info_absolute(checkpoint_name): return asyncio.run(_get_checkpoint_info_absolute_async()) -def _format_model_name_for_comfyui(file_path: str, model_roots: list) -> str: +def _format_model_name_for_comfyui(file_path: str, model_roots: list[str]) -> str: """Format file path to ComfyUI-style model name (relative path with extension) Example: /path/to/checkpoints/Illustrious/model.safetensors -> Illustrious/model.safetensors @@ -470,7 +470,7 @@ def calculate_recipe_fingerprint(loras): def calculate_relative_path_for_model( - model_data: Dict, model_type: str = "lora" + model_data: Dict[str, Any], model_type: str = "lora" ) -> str: """Calculate relative path for existing model using template from settings diff --git a/standalone.py b/standalone.py index 6b20d996..737f5f61 100644 --- a/standalone.py +++ b/standalone.py @@ -1,6 +1,7 @@ import os import sys import json +from typing import Any, cast # Ensure the script's directory is on sys.path so that py.* imports resolve # regardless of the current working directory (e.g. when launched via # ComfyUI's python_embeded from the ComfyUI root directory). @@ -19,7 +20,7 @@ def mock_nodes_directory(): nodes_dir = os.path.join(os.path.dirname(__file__), "py", "nodes") if os.path.exists(nodes_dir): # Create a mock module for the nodes package itself - sys.modules["py.nodes"] = type("MockNodesModule", (), {}) + sys.modules["py.nodes"] = type("MockNodesModule", (), {}) # pyright: ignore[reportArgumentType] # Create mock modules for all Python files in the nodes directory for file in os.listdir(nodes_dir): @@ -27,7 +28,7 @@ def mock_nodes_directory(): module_name = file[:-3] # Remove .py extension full_module_name = f"py.nodes.{module_name}" # Create empty module object - sys.modules[full_module_name] = type( + sys.modules[full_module_name] = type( # pyright: ignore[reportArgumentType] f"Mock{module_name.capitalize()}Module", (), {} ) print(f"Created mock module for: {full_module_name}") @@ -91,6 +92,10 @@ class MockFolderPaths: # Create mock server module with PromptServer class MockPromptServer: + last_prompt_id: Any = None + last_node_id: Any = None + client_id: Any = None + def __init__(self): self.app = None @@ -108,9 +113,9 @@ class MockMetadataCollector: # Initialize basic mocks before any imports -sys.modules["folder_paths"] = MockFolderPaths() -sys.modules["server"] = type("server", (), {"PromptServer": MockPromptServer()}) -sys.modules["py.metadata_collector"] = MockMetadataCollector() +sys.modules["folder_paths"] = MockFolderPaths() # pyright: ignore[reportArgumentType] +sys.modules["server"] = type("server", (), {"PromptServer": MockPromptServer()}) # pyright: ignore[reportArgumentType] +sys.modules["py.metadata_collector"] = MockMetadataCollector() # pyright: ignore[reportArgumentType] # Now we can safely import modules that depend on folder_paths and server import argparse @@ -159,10 +164,15 @@ from py.config import config class StandaloneServer: """Server implementation for standalone mode""" + last_prompt_id: Any = None + last_node_id: Any = None + client_id: Any = None + def __init__(self): + middlewares: list[Any] = [api_json_error, cache_control] self.app = web.Application( logger=logger, - middlewares=[api_json_error, cache_control], + middlewares=middlewares, client_max_size=256 * 1024 * 1024, handler_args={ "max_field_size": HEADER_SIZE_LIMIT, @@ -303,8 +313,9 @@ class StandaloneLoraManager(LoraManager): """Extended LoraManager for standalone mode""" @classmethod - def add_routes(cls, server_instance): + def add_routes(cls, server_instance=None): """Initialize and register all routes for standalone mode""" + server_instance = cast(Any, server_instance) app = server_instance.app # Store app in a global-like location for compatibility