Compare commits

...

12 Commits

Author SHA1 Message Date
Will Miao 27027c4497 refactor(recipes): reuse shared local hash cache in create-from-example 2026-08-08 22:45:29 +08:00
Will Miao 86c85c08ec feat(recipes): pass local hash cache to remote and url recipe imports 2026-08-08 22:13:19 +08:00
Will Miao 196c8ffc3e feat(recipes): match civitai image hash sections against local hash cache 2026-08-08 22:12:47 +08:00
Will Miao cfc95ee02a feat(recipes): pass local hash cache through analysis recipe parsing 2026-08-08 22:11:50 +08:00
Will Miao 479fa36997 feat(recipes): add version-cached local hash cache builder 2026-08-08 22:04:09 +08:00
Will Miao 3e1216e9bc feat(recipes): add cache version counter to model scanners 2026-08-08 21:57:36 +08:00
Will Miao 007883b7d1 fix(recipes): backfill lora cache item by autov2/autov3 hash too 2026-08-08 21:51:34 +08:00
Will Miao dc9200a12c fix(recipes): match recipe-format lora cache item by autov2/autov3 hash 2026-08-08 21:50:33 +08:00
Will Miao d2f955266d fix(types): resolve pre-existing basedpyright errors in tests
Fix ~790 basedpyright errors across the test suite:
- Type stub subclasses of real production classes with super().__init__()
- Add missing generic type arguments and Dict[str, Any] annotations
- Add None guards before subscript/member access
- Adapt tests to production API changes (removed dead handlers,
  PersistentModelCache.get_default, _i18n_filter_added location)
2026-08-08 20:12:59 +08:00
Will Miao 8e724538bd fix(types): resolve pre-existing basedpyright errors in py/ and standalone.py
Fix ~950 basedpyright errors across the backend:
- Convert ineffective # type: ignore comments to # pyright: ignore[rule]
- Add missing generic type arguments (Dict[str, Any], list[Any], ...)
- Annotate dynamic dict literals and runtime-initialized attributes
- Widen CivitAI provider tuple signatures in recipe parsers
- Remove dead LoraRoutes handlers calling nonexistent LoraService methods
- Suppress unavoidable ServiceRegistry import cycles (basedpyright counts
  function-local imports as cycle edges)
2026-08-08 20:12:52 +08:00
Will Miao 6fcdeb799d feat(metadata): resolve AutoV3 at download time without waiting for backfill
- Read AutoV3 directly from the downloaded file's own file_info hashes
  (no SHA256 cross-matching against version_info.files, so the value is
  captured even when the API omits SHA256)
- Extract normalize_autov3() validation helper shared with the
  sha256-matching autov3_from_civitai_files path
- Fall back to the embedded safetensors header hash at download
  completion; mark '' (checked-unavailable) so the startup backfill
  query (autov3 IS NULL) never revisits the row
- Clear archive-level AutoV3 for zip-extracted models so per-file
  header resolution applies to every extracted model
2026-08-08 15:20:43 +08:00
Will Miao 97b9b1f62b feat(metadata): add CivitAI AutoV3 hash support across all storage layers
- Three-state autov3 field (not-checked / checked-unavailable / 12-hex value)
  in .metadata.json sidecars, in-memory ModelHashIndex, and SQLite
  (models.autov3 column + autov3_index table) with column-presence migration
- Background self-terminating backfill for legacy rows: per-model-type
  concurrency guard, executor-offloaded I/O, Civitai-first resolution
  (SHA256-matched version file) falling back to the embedded safetensors
  header hash
- Civitai-first propagation on metadata refresh, scan, and download paths;
  reject the empty-string SHA256 placeholder and strip OneTrainer 0x prefix
- List API hash filters and hash index lookups accept 12-char AutoV3
- Cap safetensors header reads at 64 MiB to prevent crafted-file allocation
- Prevent stale AutoV3 mappings on file replacement while preserving them on
  same-file re-registration (lazy-hash completion)
2026-08-08 14:30:34 +08:00
206 changed files with 6565 additions and 1783 deletions
+4
View File
@@ -39,6 +39,7 @@ These fields are present in all model metadata files.
| `metadata_source` | string\|null | ❌ No | ✅ Yes | Last provider that supplied metadata (see below) |
| `last_checked_at` | float | ❌ No (default: `0`) | ✅ Yes | Unix timestamp of last metadata check |
| `hash_status` | string | ❌ No (default: `"completed"`) | ✅ Yes | Hash calculation status: `"pending"`, `"calculating"`, `"completed"`, `"failed"` |
| `autov3` | string\|null | ❌ No | ✅ Yes | CivitAI AutoV3 hash (first 12 chars, lowercase hex) sourced from the safetensors embedded metadata (`sshs_model_hash` / `modelspec.hash_sha256`). **Absent** = not yet checked (may be backfilled later); **`null`** = checked but unavailable (header has no recognized hash); **12-char hex string** = value |
---
@@ -287,6 +288,7 @@ These fields are automatically synchronized with the filesystem:
- `preview_url` — Updated if preview file is moved/removed
- `sha256` — Updated during hash calculation (when `hash_status="pending"`)
- `hash_status` — Updated during hash calculation
- `autov3` — Set when metadata is first created (from safetensors header); may be backfilled later for entries where it is absent
- `last_checked_at` — Timestamp of scan
- `metadata_source` — Set based on metadata provider
@@ -345,6 +347,7 @@ These fields can be edited by users at any time through the Lora Manager UI or b
| `metadata_source` | `null` |
| `last_checked_at` | `0` |
| `hash_status` | `"completed"` |
| `autov3` | absent (not checked) or `null` (checked, no value) |
| `usage_tips` | `"{}"` (LoRA only) |
| `model_type` | `"checkpoint"` or `"embedding"` (not present in LoRA models) |
@@ -354,6 +357,7 @@ These fields can be edited by users at any time through the Lora Manager UI or b
| Version | Date | Changes |
|---------|------|---------|
| 1.1 | 2026-08 | Added `autov3` field (CivitAI AutoV3 hash with three-state semantics) |
| 1.0 | 2026-03 | Initial schema documentation |
---
+15 -10
View File
@@ -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
+1 -1
View File
@@ -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 (
+2 -2
View File
@@ -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 {}
+1 -1
View File
@@ -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:
+11 -1
View File
@@ -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)
+2 -2
View File
@@ -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
+6 -6
View File
@@ -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:
+2 -2
View File
@@ -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")
+2 -2
View File
@@ -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 (
+1 -1
View File
@@ -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": {},
+12 -13
View File
@@ -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]] = []
+2 -2
View File
@@ -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)
+7 -7
View File
@@ -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)
+6 -6
View File
@@ -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:
+7 -3
View File
@@ -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 <lora:name:strength> syntax from text input into a list of dicts.
Each entry contains: name, model_strength, clip_strength.
+17 -5
View File
@@ -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
@@ -175,10 +179,18 @@ class RecipeMetadataParser(ABC):
lora_entry['localPath'] = local_path
lora_entry['file_name'] = os.path.splitext(os.path.basename(local_path))[0]
# Get thumbnail from local preview if available
# Get thumbnail from local preview if available.
# Match the cache item by local path first (get_path_by_hash
# cascade: 10-char autov2 / 12-char autov3), then by hash.
lora_cache = await lora_scanner.get_cached_data()
lora_item = next((item for item in lora_cache.raw_data
if item['sha256'].lower() == lora_entry['hash'].lower()), None)
h = (lora_entry.get("hash") or "").lower()
lora_item = next((item for item in lora_cache.raw_data
if (item.get("file_path") or "") == local_path), None)
if lora_item is None:
lora_item = next((item for item in lora_cache.raw_data
if (item.get("sha256") or "").lower() == h
or (item.get("autov3") or "").lower() == h
or (item.get("sha256") or "")[:10].lower() == h), None)
if lora_item and 'preview_url' in lora_item:
lora_entry['thumbnailUrl'] = config.get_preview_static_url(lora_item['preview_url'])
except Exception as e:
@@ -194,7 +206,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
+4
View File
@@ -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
+3 -1
View File
@@ -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}")
+1 -1
View File
@@ -52,7 +52,7 @@ class AutomaticMetadataParser(RecipeMetadataParser):
negative_and_params = ""
# Initialize metadata
metadata = {
metadata: Dict[str, Any] = {
"prompt": prompt,
"loras": []
}
+117 -56
View File
@@ -4,7 +4,7 @@ import json
import logging
from typing import Dict, Any, Union
from ..base import RecipeMetadataParser
from ..constants import GEN_PARAM_KEYS
from ..constants import GEN_PARAM_KEYS, VALID_LORA_TYPES
from ...services.metadata_service import get_default_metadata_provider
from ...config import config
@@ -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"),
@@ -216,7 +216,8 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
# Try to look up base model from the checkpoint hash
cp_hash = checkpoint_entry.get("hash")
if cp_hash and metadata_provider:
local_cached = local_cache.get(cp_hash) if local_cache else None
# local_cache keys are stored lowercase
local_cached = local_cache.get(cp_hash.lower()) if local_cache else None
if local_cached:
self._populate_entry_from_cache(
checkpoint_entry, local_cached
@@ -294,8 +295,15 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
# Try to get info from Civitai if hash is available
if lora_hash and metadata_provider:
local_cached = local_cache.get(lora_hash) if local_cache else None
# local_cache keys are stored lowercase
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
if local_cached:
cached_type = self._cache_item_model_type(local_cached)
if cached_type and cached_type not in VALID_LORA_TYPES:
logger.debug(
f"Skipping non-LoRA cache item for hash {lora_hash}"
)
continue
self._populate_entry_from_cache(
lora_entry, local_cached
)
@@ -304,6 +312,12 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
added_loras[str(lora_entry["id"])] = len(
result["loras"]
)
# Mirror base.py:150-151 counts for API-path loras
bm = local_cached.get("base_model") or ""
if bm:
base_model_counts[bm] = base_model_counts.get(
bm, 0
) + 1
else:
try:
civitai_info = (
@@ -649,30 +663,47 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
}
if metadata_provider:
try:
civitai_info = await metadata_provider.get_model_by_hash(
lora_hash
)
populated_entry = await self.populate_lora_from_civitai(
lora_entry,
civitai_info,
recipe_scanner,
base_model_counts,
lora_hash,
)
if populated_entry is None:
# local_cache keys are stored lowercase
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
if local_cached:
cached_type = self._cache_item_model_type(local_cached)
if cached_type and cached_type not in VALID_LORA_TYPES:
logger.debug(
f"Skipping non-LoRA cache item for hash {lora_hash}"
)
continue
lora_entry = populated_entry
self._populate_entry_from_cache(lora_entry, local_cached)
# Mirror base.py:150-151 counts for API-path loras
bm = local_cached.get("base_model") or ""
if bm:
base_model_counts[bm] = base_model_counts.get(bm, 0) + 1
if "id" in lora_entry and lora_entry["id"]:
added_loras[str(lora_entry["id"])] = len(result["loras"])
except Exception as e:
logger.error(
f"Error fetching Civitai info for LoRA hash {lora_hash}: {e}"
)
else:
try:
civitai_info = await metadata_provider.get_model_by_hash(
lora_hash
)
populated_entry = await self.populate_lora_from_civitai(
lora_entry,
civitai_info,
recipe_scanner,
base_model_counts,
lora_hash,
)
if populated_entry is None:
continue
lora_entry = populated_entry
if "id" in lora_entry and lora_entry["id"]:
added_loras[str(lora_entry["id"])] = len(result["loras"])
except Exception as e:
logger.error(
f"Error fetching Civitai info for LoRA hash {lora_hash}: {e}"
)
added_loras[lora_hash] = len(result["loras"])
result["loras"].append(lora_entry)
@@ -711,32 +742,51 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
# Try to get info from Civitai if hash is available
if lora_entry["hash"] and metadata_provider:
try:
civitai_info = await metadata_provider.get_model_by_hash(
lora_hash
)
populated_entry = await self.populate_lora_from_civitai(
lora_entry,
civitai_info,
recipe_scanner,
base_model_counts,
lora_hash,
)
if populated_entry is None:
# local_cache keys are stored lowercase
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
if local_cached:
cached_type = self._cache_item_model_type(local_cached)
if cached_type and cached_type not in VALID_LORA_TYPES:
logger.debug(
f"Skipping non-LoRA cache item for hash {lora_hash}"
)
lora_index += 1
continue # Skip invalid LoRA types
lora_entry = populated_entry
continue # Skip non-LoRA cache items
self._populate_entry_from_cache(lora_entry, local_cached)
# Mirror base.py:150-151 counts for API-path loras
bm = local_cached.get("base_model") or ""
if bm:
base_model_counts[bm] = base_model_counts.get(bm, 0) + 1
# If we have a version ID from Civitai, track it for deduplication
if "id" in lora_entry and lora_entry["id"]:
added_loras[str(lora_entry["id"])] = len(result["loras"])
except Exception as e:
logger.error(
f"Error fetching Civitai info for LoRA hash {lora_entry['hash']}: {e}"
)
else:
try:
civitai_info = await metadata_provider.get_model_by_hash(
lora_hash
)
populated_entry = await self.populate_lora_from_civitai(
lora_entry,
civitai_info,
recipe_scanner,
base_model_counts,
lora_hash,
)
if populated_entry is None:
lora_index += 1
continue # Skip invalid LoRA types
lora_entry = populated_entry
# If we have a version ID from Civitai, track it for deduplication
if "id" in lora_entry and lora_entry["id"]:
added_loras[str(lora_entry["id"])] = len(result["loras"])
except Exception as e:
logger.error(
f"Error fetching Civitai info for LoRA hash {lora_entry['hash']}: {e}"
)
# Track by hash if we have it
if lora_hash:
@@ -795,3 +845,14 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
base_model = cache_item.get("base_model", "")
if base_model:
entry["baseModel"] = base_model
@staticmethod
def _cache_item_model_type(cache_item: dict[str, Any]) -> str:
"""Lowercased civitai.model.type of a cache item, or '' when unknown."""
civ = cache_item.get("civitai")
if not isinstance(civ, dict):
return ""
model_info = civ.get("model")
if not isinstance(model_info, dict):
return ""
return (model_info.get("type") or "").lower()
+1 -1
View File
@@ -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:
+10 -2
View File
@@ -91,7 +91,15 @@ class RecipeFormatParser(RecipeMetadataParser):
exists_locally = lora_scanner.has_hash(lora['hash'])
if exists_locally:
lora_cache = await lora_scanner.get_cached_data()
lora_item = next((item for item in lora_cache.raw_data if item['sha256'].lower() == lora['hash'].lower()), None)
# Cascade match: full sha256, stored autov3, or autov2 (sha256[:10]).
h = (lora.get('hash') or '').lower()
lora_item = next(
(item for item in lora_cache.raw_data
if (item.get("sha256") or "").lower() == h
or (item.get("autov3") or "").lower() == h
or (item.get("sha256") or "")[:10].lower() == h),
None
)
if lora_item:
lora_entry['existsLocally'] = True
lora_entry['inLibrary'] = True
@@ -148,7 +156,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'),
+7 -7
View File
@@ -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)
+13 -9
View File
@@ -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,
+9 -9
View File
@@ -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:
+4 -4
View File
@@ -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)
+8 -4
View File
@@ -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:
@@ -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 {
+53 -45
View File
@@ -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,11 @@ 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):
scanner.bump_cache_version()
persist: Any = getattr(scanner, "_persist_current_cache", None)
if persist:
await persist()
history_service = await self._get_download_history_service()
@@ -2649,13 +2653,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 +2711,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 +2728,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 +2791,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 +2938,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 +3557,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 +3681,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 +3748,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
+20 -18
View File
@@ -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 = (
+109 -51
View File
@@ -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:
@@ -1118,6 +1118,12 @@ class RecipeManagementHandler:
_original_image_url,
) = await self._download_remote_media(image_url)
# Build a version-cached map of local model hashes to cache items so
# CivitaiApiMetadataParser can skip CivitAI API calls for models that
# exist on disk. Built once and shared by every parse pass below.
local_cache = await recipe_scanner.build_local_hash_cache()
from ...recipes.parsers.civitai_image import CivitaiApiMetadataParser
# Extract embedded EXIF metadata (offloaded to thread pool in this call)
embedded_gen_params = {}
parsed_embedded = None
@@ -1139,9 +1145,16 @@ class RecipeManagementHandler:
)
)
if parser:
parsed_embedded = await parser.parse_metadata(
raw_embedded, recipe_scanner=recipe_scanner
)
if isinstance(parser, CivitaiApiMetadataParser):
parsed_embedded = await parser.parse_metadata(
raw_embedded,
recipe_scanner=recipe_scanner,
local_cache=local_cache,
)
else:
parsed_embedded = await parser.parse_metadata(
raw_embedded, recipe_scanner=recipe_scanner
)
if parsed_embedded and "gen_params" in parsed_embedded:
embedded_gen_params = parsed_embedded["gen_params"]
else:
@@ -1172,9 +1185,16 @@ class RecipeManagementHandler:
civitai_inner_meta
)
if parser:
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta, recipe_scanner=recipe_scanner
)
if isinstance(parser, CivitaiApiMetadataParser):
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta,
recipe_scanner=recipe_scanner,
local_cache=local_cache,
)
else:
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta, recipe_scanner=recipe_scanner
)
if civitai_parsed and "gen_params" in civitai_parsed:
# Merge: API gen_params override EXIF at field level,
# EXIF fills in fields the API doesn't have.
@@ -1678,7 +1698,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]
@@ -1798,6 +1818,12 @@ class RecipeManagementHandler:
await self._download_remote_media(image_url)
)
# Build a version-cached map of local model hashes to cache items so
# CivitaiApiMetadataParser can skip CivitAI API calls for models that
# exist on disk. Built once and shared by every parse pass below.
local_cache = await recipe_scanner.build_local_hash_cache()
from ...recipes.parsers.civitai_image import CivitaiApiMetadataParser
# Extract embedded EXIF metadata
embedded_gen_params = {}
parsed_embedded = None
@@ -1819,9 +1845,16 @@ class RecipeManagementHandler:
)
)
if parser:
parsed_embedded = await parser.parse_metadata(
raw_embedded, recipe_scanner=recipe_scanner
)
if isinstance(parser, CivitaiApiMetadataParser):
parsed_embedded = await parser.parse_metadata(
raw_embedded,
recipe_scanner=recipe_scanner,
local_cache=local_cache,
)
else:
parsed_embedded = await parser.parse_metadata(
raw_embedded, recipe_scanner=recipe_scanner
)
if parsed_embedded and "gen_params" in parsed_embedded:
embedded_gen_params = parsed_embedded["gen_params"]
finally:
@@ -1859,9 +1892,16 @@ class RecipeManagementHandler:
)
)
if parser:
parsed_embedded = await parser.parse_metadata(
raw_orig, recipe_scanner=recipe_scanner
)
if isinstance(parser, CivitaiApiMetadataParser):
parsed_embedded = await parser.parse_metadata(
raw_orig,
recipe_scanner=recipe_scanner,
local_cache=local_cache,
)
else:
parsed_embedded = await parser.parse_metadata(
raw_orig, recipe_scanner=recipe_scanner
)
if (
parsed_embedded
and "gen_params" in parsed_embedded
@@ -1895,9 +1935,16 @@ class RecipeManagementHandler:
civitai_inner_meta
)
if parser:
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta, recipe_scanner=recipe_scanner
)
if isinstance(parser, CivitaiApiMetadataParser):
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta,
recipe_scanner=recipe_scanner,
local_cache=local_cache,
)
else:
civitai_parsed = await parser.parse_metadata(
civitai_inner_meta, recipe_scanner=recipe_scanner
)
if civitai_parsed and "gen_params" in civitai_parsed:
# Merge: API gen_params override EXIF at field level,
# EXIF fills in fields the API doesn't have.
@@ -2109,33 +2156,44 @@ class RecipeManagementHandler:
parsed_input = {**image_data, **inner_meta}
parsed_input.pop("meta", None)
# Build a local cache of {hash cache_item} so the parser can
# skip CivitAI API calls for models that exist on disk.
local_cache: Dict[str, Dict[str, Any]] = {}
lora_scanner = getattr(recipe_scanner, "_lora_scanner", None)
if lora_scanner and model_hash:
try:
parent_cache_data = await lora_scanner.get_cached_data()
for item in getattr(parent_cache_data, "raw_data", []):
if item.get("sha256", "").lower() == model_hash.lower():
local_cache[model_hash.lower()] = item
# Compute AutoV3 so the parser can also match on
# that hash type (CivitAI metadata resources use
# AutoV3).
file_path = item.get("file_path")
if file_path and os.path.exists(file_path):
try:
from ...utils.file_utils import (
calculate_autov3,
)
autov3 = calculate_autov3(file_path)
if autov3:
local_cache[autov3.lower()] = item
except Exception:
pass
break
except Exception:
pass
# Build the shared local hash cache so the parser can skip CivitAI
# API calls for models that exist on disk.
local_cache: Dict[str, Dict[str, Any]] = (
await recipe_scanner.build_local_hash_cache()
)
# Bounded supplement for un-backfilled parents. The shared builder
# never computes autov3; when the parent model exists on disk but
# its cached entry has no stored AutoV3, compute it for that single
# file and register the AutoV3 key so the parser can also match on
# that hash type (CivitAI metadata resources use AutoV3). This runs
# whenever the parent is found with an empty autov3, independent of
# whether the sha256 key is already present in the shared cache.
if model_hash:
lora_scanner = getattr(recipe_scanner, "_lora_scanner", None)
if lora_scanner:
try:
parent_cache_data = await lora_scanner.get_cached_data()
for item in getattr(parent_cache_data, "raw_data", []):
if item.get("sha256", "").lower() == model_hash.lower():
autov3 = (item.get("autov3") or "").lower()
if not autov3:
file_path = item.get("file_path")
if file_path and os.path.exists(file_path):
try:
from ...utils.file_utils import (
calculate_autov3,
)
autov3 = (
calculate_autov3(file_path) or ""
).lower()
except Exception:
pass
if autov3:
local_cache[autov3] = item
break
except Exception:
pass
parser = self._analysis_service._recipe_parser_factory.create_parser(
parsed_input
@@ -2167,10 +2225,10 @@ class RecipeManagementHandler:
parent_model_id: int | None = None
parent_version_name: str | None = None
parent_model_name: str | None = None
# Prefer sha256 key; fall back to any cached entry.
# Resolve the parent strictly by its sha256 key. There is no
# arbitrary fallback: with a full-library cache, picking any entry
# would corrupt the isDeleted reconciliation below.
parent_item = local_cache.get(model_hash.lower()) if model_hash else None
if parent_item is None and local_cache:
parent_item = next(iter(local_cache.values()))
if parent_item:
civ = parent_item.get("civitai") or {}
if isinstance(civ, dict):
@@ -2386,7 +2444,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()
+6 -71
View File
@@ -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
+2 -2
View File
@@ -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)
+1 -1
View File
@@ -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,
+4 -4
View File
@@ -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)
+2 -2
View File
@@ -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)
+11 -10
View File
@@ -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:
+16 -13
View File
@@ -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}"
+1 -1
View File
@@ -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)):
+1 -1
View File
@@ -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:
+1 -1
View File
@@ -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).
@@ -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 ``<img src=\"...\">`` 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'<img\s[^>]*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 ``<img>`` tags exclusively.
"""
def _img_to_md(match: re.Match) -> str:
def _img_to_md(match: re.Match[str]) -> str:
"""Convert an ``<img>`` 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:
+7 -3
View File
@@ -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",
+3 -3
View File
@@ -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:
+144
View File
@@ -0,0 +1,144 @@
# 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
from it have a NULL ``autov3`` column (the "not checked yet" state). This
service computes the embedded AutoV3 hash for each such model once per
process and persists it through the scanner's single write path
(:meth:`ModelScanner.update_autov3_for_model`), marking every visited row so a
subsequent run finds nothing left to do.
Three-state contract honored here:
- ``NULL`` (sqlite) / absent (dict) = not checked yet backfill computes it
- ``''`` (sqlite/dict) / JSON null = checked, no value available never recompute
- 12-char lowercase hex = value never recompute
"""
from __future__ import annotations
import asyncio
import json
import logging
import os
import threading
from typing import TYPE_CHECKING, Optional
if TYPE_CHECKING: # pragma: no cover - type-check only; runtime imports are local
from .model_scanner import ModelScanner
logger = logging.getLogger(__name__)
def _resolve_autov3(file_path: str) -> str:
"""Resolve the AutoV3 hash for a model file.
Prefers the Civitai AutoV3 reported for the file whose SHA256 matches
(the authoritative value for recipe matching); falls back to the embedded
safetensors header hash. Returns ``''`` when neither is available.
"""
try:
metadata_path = f"{os.path.splitext(file_path)[0]}.metadata.json"
if os.path.exists(metadata_path):
with open(metadata_path, "r", encoding="utf-8") as handle:
payload = json.load(handle)
if isinstance(payload, dict):
from ..utils.models import autov3_from_civitai_files # local import avoids cycles
sha256 = (payload.get("sha256") or "").lower()
civitai_autov3 = autov3_from_civitai_files(payload.get("civitai"), sha256)
if civitai_autov3:
return civitai_autov3
except Exception:
pass
from ..utils.file_utils import calculate_autov3 # local import avoids cycles
return calculate_autov3(file_path) or ""
class Autov3BackfillService:
"""Compute and persist AutoV3 hashes for models missing a checked state."""
_instance: Optional["Autov3BackfillService"] = None
_instance_lock = threading.Lock()
def __init__(self) -> None:
# Re-entrancy guard per model type: scanners for different model types
# 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[str] = set()
@classmethod
def get_instance(cls) -> "Autov3BackfillService":
"""Return the process-wide singleton instance."""
if cls._instance is None:
with cls._instance_lock:
if cls._instance is None:
cls._instance = cls()
return cls._instance
async def backfill(self, scanner: "ModelScanner") -> int:
"""Compute AutoV3 for every un-checked model of ``scanner.model_type``.
Each candidate file is read once via :func:`~py.utils.file_utils.calculate_autov3`
(cheap: safetensors header only) and the result is persisted through
``scanner.update_autov3_for_model``. Files that no longer exist on
disk are skipped they are intentionally NOT marked, because scanner
cleanup removes the stale row later.
Returns:
The number of models successfully updated. Never raises; on any
failure a warning is logged and ``0`` is returned. A duplicate
trigger for a model type that is already being backfilled returns
``0`` immediately; different model types run concurrently.
"""
model_type = scanner.model_type
if model_type in self._running_types:
return 0
self._running_types.add(model_type)
try:
# Local imports avoid import cycles at module load time.
from .persistent_model_cache import get_persistent_cache
from ..utils.file_utils import calculate_autov3
persistent = getattr(scanner, "_persistent_cache", None) or get_persistent_cache()
paths = persistent.get_models_missing_autov3(model_type)
loop = asyncio.get_running_loop()
count = 0
for path in paths:
# A file that no longer exists must not be marked; scanner
# cleanup removes the stale row later. The existence check and
# hash resolution run in the executor so the loop stays
# responsive to API requests while the backfill iterates a
# large library.
if not await loop.run_in_executor(None, os.path.exists, path):
continue
autov3 = await loop.run_in_executor(None, _resolve_autov3, path)
if await scanner.update_autov3_for_model(model_type, path, autov3):
count += 1
if paths:
logger.info(
"AutoV3 backfill: updated %d/%d models for %s",
count,
len(paths),
model_type,
)
else:
# Steady state after the first run: nothing left to backfill.
logger.debug("AutoV3 backfill: nothing to process for %s", model_type)
return count
except Exception as exc:
logger.warning(
"AutoV3 backfill failed for %s: %s",
getattr(scanner, "model_type", "?"),
exc,
)
return 0
finally:
self._running_types.discard(model_type)
+4
View File
@@ -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
+81 -70
View File
@@ -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,39 +444,50 @@ class BaseModelService(ABC):
return entry
async def _apply_hash_filters(
self, data: List[Dict], hash_filters: Dict
) -> List[Dict]:
"""Apply hash-based filtering"""
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[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``
field, both case-insensitively.
"""
if item.get("sha256", "").lower() in hash_set:
return True
autov3 = item.get("autov3", "")
return bool(autov3) and autov3.lower() in hash_set
single_hash = hash_filters.get("single_hash")
multiple_hashes = hash_filters.get("multiple_hashes")
if single_hash:
# Filter by single hash
single_hash = single_hash.lower()
# Filter by single hash (SHA256 or AutoV3)
return [
item for item in data if item.get("sha256", "").lower() == single_hash
item for item in data if matches_hash_set(item, {single_hash.lower()})
]
elif multiple_hashes:
# Filter by multiple hashes
hash_set = set(hash.lower() for hash in multiple_hashes)
return [item for item in data if item.get("sha256", "").lower() in hash_set]
# Filter by multiple hashes (SHA256 or AutoV3)
hash_set = {hash.lower() for hash in multiple_hashes}
return [item for item in data if matches_hash_set(item, hash_set)]
return data
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(
@@ -495,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:
@@ -542,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:
@@ -575,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.
@@ -591,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)
@@ -628,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()
@@ -648,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",
@@ -714,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):
@@ -727,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:
@@ -741,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
@@ -754,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
@@ -767,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
@@ -819,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
@@ -834,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
@@ -843,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)
@@ -920,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 {}
@@ -946,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()
@@ -975,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()
@@ -1004,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
@@ -1136,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
@@ -1232,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.
@@ -1265,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,
+30 -2
View File
@@ -59,6 +59,7 @@ class CacheEntryValidator:
'notes': ('', False),
'usage_tips': ('', False),
'hash_status': ('completed', False),
'autov3': (None, False),
}
@classmethod
@@ -119,8 +120,13 @@ class CacheEntryValidator:
if is_required:
errors.append(f"Required field '{field_name}' is missing or None")
if auto_repair:
working_entry[field_name] = cls._get_default_copy(default_value)
repaired = True
# A missing optional field whose default is None is already
# semantically equal to its default (e.g. autov3: absent
# means "not checked") — writing None back is a no-op, not
# a repair.
if default_value is not None:
working_entry[field_name] = cls._get_default_copy(default_value)
repaired = True
continue
# Validate field type and value
@@ -175,6 +181,15 @@ class CacheEntryValidator:
# that invalidates the entry, but we also don't mark it repaired.
pass
# Normalize autov3 to lowercase if needed (optional field, never stripped).
autov3 = working_entry.get('autov3')
if isinstance(autov3, str) and autov3:
normalized_autov3 = autov3.lower()
if normalized_autov3 != autov3:
if auto_repair:
working_entry['autov3'] = normalized_autov3
repaired = True
# Determine if entry is valid
# Entry is valid if no critical required field errors remain after repair
# Critical fields are file_path and sha256
@@ -242,6 +257,19 @@ class CacheEntryValidator:
"""
expected_type = type(default_value)
# Special case: autov3 is optional with a three-state contract.
# None = not checked, "" = checked but unavailable, otherwise a
# 12-character hex string (case-insensitive here; normalized to
# lowercase separately).
if field_name == 'autov3':
if value is None or value == "":
return None
if not isinstance(value, str):
return f"Field 'autov3' should be string or None, got {type(value).__name__}"
if len(value) != 12 or any(c not in '0123456789abcdefABCDEF' for c in value):
return "Field 'autov3' should be a 12-character hex string"
return None
# Special handling for numeric types
if expected_type == int:
if not isinstance(value, (int, float)):
+33 -6
View File
@@ -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
@@ -6,7 +10,7 @@ from datetime import datetime
from typing import Any, Dict, List, Optional
from ..utils.models import CheckpointMetadata
from ..utils.file_utils import find_preview_file, normalize_path
from ..utils.file_utils import find_preview_file, normalize_path, calculate_autov3
from ..utils.metadata_manager import MetadataManager
from ..config import config
from .model_scanner import ModelScanner
@@ -62,6 +66,11 @@ class CheckpointScanner(ModelScanner):
# Find preview image
preview_url = find_preview_file(base_name, dir_path)
# AutoV3 reads only the safetensors header, so it is cheap even for
# large checkpoints; record the checked state at creation time ("" =
# checked but unavailable).
autov3 = calculate_autov3(real_path)
# Create metadata WITHOUT calculating hash
metadata = CheckpointMetadata(
file_name=base_name,
@@ -77,6 +86,7 @@ class CheckpointScanner(ModelScanner):
sub_type="checkpoint",
from_civitai=False, # Mark as local model since no hash yet
hash_status="pending", # Mark hash as pending
autov3=autov3 or "",
)
# Save the created metadata
@@ -120,7 +130,11 @@ class CheckpointScanner(ModelScanner):
# that queries get_hash_by_filename first) will miss on every
# lookup and keep calling back into this method, creating a
# tight loop that never populates the index.
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
self._hash_index.add_entry(
metadata.sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
return metadata.sha256
async with self._hash_calculation_lock:
@@ -132,7 +146,11 @@ class CheckpointScanner(ModelScanner):
and metadata.hash_status == "completed"
and metadata.sha256
):
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
self._hash_index.add_entry(
metadata.sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
return metadata.sha256
task = self._hash_calculation_tasks.get(real_path)
@@ -185,7 +203,11 @@ class CheckpointScanner(ModelScanner):
if metadata.hash_status == "completed" and metadata.sha256:
# Populate the in-memory hash index even for pre-computed
# hashes, mirroring the fix in calculate_hash_for_model.
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
self._hash_index.add_entry(
metadata.sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
return metadata.sha256
# Update status to calculating
@@ -202,7 +224,11 @@ class CheckpointScanner(ModelScanner):
await MetadataManager.save_metadata(file_path, metadata)
# Update hash index
self._hash_index.add_entry(sha256.lower(), file_path)
self._hash_index.add_entry(
sha256.lower(),
file_path,
getattr(metadata, "autov3", None) or None,
)
# Update the in-memory cache entry so that subsequent
# _persist_current_cache / _save_persistent_cache calls
@@ -216,6 +242,7 @@ class CheckpointScanner(ModelScanner):
if entry.get("file_path") == file_path:
entry["sha256"] = sha256.lower()
entry["hash_status"] = "completed"
self.bump_cache_version()
break
logger.info(f"Hash calculated for checkpoint: {file_path}")
@@ -405,7 +432,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:
+28 -28
View File
@@ -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", "<unknown>"),
model_data.get("file_name", "<unknown>"),
)
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()
+41 -36
View File
@@ -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)
+1 -1
View File
@@ -283,7 +283,7 @@ class CivitaiBaseModelService:
return None
if isinstance(result, str):
data = json.loads(result)
data: Any = json.loads(result)
else:
data = result
+40 -35
View File
@@ -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)
+1 -1
View File
@@ -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
+113 -76
View File
@@ -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 (
@@ -18,7 +22,7 @@ from ..utils.constants import (
VALID_LORA_TYPES,
)
from ..utils.civitai_utils import normalize_civitai_download_url, rewrite_preview_url
from ..utils.file_utils import calculate_sha256
from ..utils.file_utils import calculate_sha256, calculate_autov3
from ..utils.preview_selection import resolve_mature_threshold, select_preview_media
from ..utils.utils import sanitize_folder_name
from ..utils.exif_utils import ExifUtils
@@ -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(
@@ -2160,6 +2183,10 @@ class DownloadManager:
"error": f"Zip archive does not contain any supported model files ({supported_text})",
}
actual_file_paths = extracted_paths
# The archive entry's AutoV3 (if any) describes the zip itself,
# not the extracted models; clear it so per-file header
# resolution applies to every extracted model.
metadata.autov3 = None
try:
os.remove(save_path)
except OSError as exc:
@@ -2235,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 = (
@@ -2355,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
@@ -2374,6 +2401,16 @@ class DownloadManager:
sha256 = await calculate_sha256(file_path)
if sha256:
entry.sha256 = sha256.lower()
# AutoV3: the Civitai-reported value for the downloaded file (set
# by from_civitai_info) takes precedence. Only the un-checked
# state (None) triggers a header read; '' (checked-unavailable)
# is never re-read, honoring the three-state contract so rows
# marked at download time stay untouched by later passes.
if entry.autov3 is None:
autov3 = await asyncio.get_running_loop().run_in_executor(
None, calculate_autov3, file_path
)
entry.autov3 = (autov3 or "").lower()
entries.append(entry)
return entries
@@ -2392,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 []
@@ -2451,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:
@@ -2533,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
@@ -2616,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()
@@ -2663,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()
@@ -2680,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():
@@ -2807,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:
@@ -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
+13 -8
View File
@@ -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
+1 -1
View File
@@ -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:
+28 -28
View File
@@ -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", "<unknown>"),
model_data.get("file_name", "<unknown>"),
)
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()
@@ -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,
+14 -4
View File
@@ -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")
+41 -41
View File
@@ -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", "<unknown>"),
model_data.get("file_name", "<unknown>"),
)
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.
+5 -1
View File
@@ -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()
+20 -7
View File
@@ -6,25 +6,26 @@ 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
from ..utils.model_utils import determine_base_model
from ..utils.models import autov3_from_civitai_files
from .connectivity_guard import OFFLINE_FRIENDLY_MESSAGE, is_expected_offline_error
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]]:
...
@@ -38,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
@@ -152,6 +153,18 @@ class MetadataSyncService:
civitai_metadata.get("baseModel")
)
# Civitai-first AutoV3 propagation: the freshly fetched version
# metadata may report an AutoV3 for the file whose SHA256 matches the
# local model. Persist it now so recipe matching sees it immediately —
# no full rescan or restart required (the header is never re-read to
# upgrade the checked-unavailable '' state).
sha256_value = (local_metadata.get("sha256") or "").lower()
civitai_autov3 = autov3_from_civitai_files(
local_metadata.get("civitai"), sha256_value
)
if civitai_autov3:
local_metadata["autov3"] = civitai_autov3
await self._preview_service.ensure_preview_for_metadata(
metadata_path, local_metadata, civitai_metadata.get("images", [])
)
@@ -479,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": {},
+21 -16
View File
@@ -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:
+3 -1
View File
@@ -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
+53 -4
View File
@@ -8,11 +8,12 @@ class ModelHashIndex:
self._hash_to_path: Dict[str, str] = {}
self._filename_to_hash: Dict[str, str] = {}
self._autov2_to_path: Dict[str, str] = {}
self._autov3_to_path: Dict[str, str] = {}
# New data structures for tracking duplicates
self._duplicate_hashes: Dict[str, List[str]] = {} # sha256 -> list of paths
self._duplicate_filenames: Dict[str, List[str]] = {} # filename -> list of paths
def add_entry(self, sha256: str, file_path: str) -> None:
def add_entry(self, sha256: str, file_path: str, autov3: Optional[str] = None) -> None:
"""Add or update hash index entry"""
if not sha256 or not file_path:
return
@@ -33,9 +34,14 @@ class ModelHashIndex:
self._duplicate_hashes.setdefault(sha256, []).append(file_path)
# 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)
# Same path registered again (e.g. a file replaced in place with
# new content) — used below to drop its stale autov3 mapping.
is_re_registration = existing_path == file_path
# If this is a different file with the same filename
if existing_path and existing_path != file_path:
@@ -67,12 +73,36 @@ class ModelHashIndex:
# AutoV2 = first 10 chars of SHA256
if len(sha256) >= 10:
self._autov2_to_path[sha256[:10]] = file_path
# AutoV3 is an independent hash (not derived from SHA256), stored as-is.
# Drop stale mappings for a path when it is re-registered with a NEW
# sha256 (file replaced in place) or with an explicit new autov3 value
# (correction). Re-registering the SAME file with the same sha256 and
# no autov3 (e.g. lazy-hash completion) must never clear its existing
# mapping. First-time registrations stay O(1).
if autov3:
autov3 = autov3.lower()
if is_re_registration and (existing_hash != sha256 or autov3):
stale_autov3_keys = [
key for key, mapped_path in self._autov3_to_path.items()
if mapped_path == file_path and key != autov3
]
for key in stale_autov3_keys:
del self._autov3_to_path[key]
if autov3:
self._autov3_to_path[autov3] = file_path
def add_autov3(self, autov3: str, file_path: str) -> None:
"""Add or update an AutoV3-only index entry (used when only AutoV3 is known)"""
if not autov3:
return
autov3 = autov3.lower()
self._autov3_to_path[autov3] = file_path
def _get_filename_from_path(self, file_path: str) -> str:
"""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)
@@ -167,6 +197,11 @@ class ModelHashIndex:
for k in autov2_keys_to_remove:
del self._autov2_to_path[k]
# Remove from AutoV3 index
autov3_keys_to_remove = [k for k, v in self._autov3_to_path.items() if v == file_path]
for k in autov3_keys_to_remove:
del self._autov3_to_path[k]
def remove_by_hash(self, sha256: str) -> None:
"""Remove entry by hash"""
sha256 = sha256.lower()
@@ -189,6 +224,11 @@ class ModelHashIndex:
autov2_key = sha256[:10]
if autov2_key in self._autov2_to_path:
del self._autov2_to_path[autov2_key]
# Remove AutoV3 entries pointing to any removed path
autov3_keys_to_remove = [k for k, v in self._autov3_to_path.items() if v in paths_to_remove]
for k in autov3_keys_to_remove:
del self._autov3_to_path[k]
# Update filename-to-hash and duplicate filenames for all paths
for path_to_remove in paths_to_remove:
@@ -209,22 +249,26 @@ class ModelHashIndex:
del self._duplicate_filenames[fname]
def has_hash(self, hash_value: str) -> bool:
"""Check if hash exists in index (SHA256 or AutoV2)"""
"""Check if hash exists in index (SHA256, AutoV2, or AutoV3)"""
normalized = hash_value.lower()
if normalized in self._hash_to_path:
return True
if len(normalized) == 10:
return normalized in self._autov2_to_path
if len(normalized) == 12:
return normalized in self._autov3_to_path
return False
def get_path(self, hash_value: str) -> Optional[str]:
"""Get file path for a hash (SHA256 or AutoV2)"""
"""Get file path for a hash (SHA256, AutoV2, or AutoV3)"""
normalized = hash_value.lower()
path = self._hash_to_path.get(normalized)
if path is not None:
return path
if len(normalized) == 10:
return self._autov2_to_path.get(normalized)
if len(normalized) == 12:
return self._autov3_to_path.get(normalized)
return None
def get_hash(self, file_path: str) -> Optional[str]:
@@ -243,6 +287,7 @@ class ModelHashIndex:
self._hash_to_path.clear()
self._filename_to_hash.clear()
self._autov2_to_path.clear()
self._autov3_to_path.clear()
self._duplicate_hashes.clear()
self._duplicate_filenames.clear()
@@ -253,6 +298,10 @@ class ModelHashIndex:
def get_all_filenames(self) -> Set[str]:
"""Get all filenames in the index"""
return set(self._filename_to_hash.keys())
def get_all_autov3(self) -> Dict[str, str]:
"""Get a snapshot of all AutoV3 hashes mapped to their file paths"""
return dict(self._autov3_to_path)
def get_duplicate_hashes(self) -> Dict[str, List[str]]:
"""Get dictionary of duplicate hashes and their paths"""
+13 -6
View File
@@ -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
@@ -138,6 +138,9 @@ class ModelLifecycleService:
item for item in cache.raw_data if item.get("file_path") != file_path
]
await cache.resort()
bump_cache_version = getattr(self._scanner, "bump_cache_version", None)
if callable(bump_cache_version):
bump_cache_version()
if hasattr(self._scanner, "_hash_index") and self._scanner._hash_index:
self._scanner._hash_index.remove_by_path(file_path)
@@ -146,7 +149,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}
@@ -244,6 +247,9 @@ class ModelLifecycleService:
item for item in cache.raw_data if item["file_path"] != file_path
]
await cache.resort()
bump_cache_version = getattr(self._scanner, "bump_cache_version", None)
if callable(bump_cache_version):
bump_cache_version()
excluded = getattr(self._scanner, "_excluded_models", None)
if isinstance(excluded, list):
@@ -252,7 +258,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 +363,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
+52 -52
View File
@@ -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": <str|None>}`` 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:
+2 -1
View File
@@ -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"
+249 -35
View File
@@ -5,11 +5,11 @@ 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
from ..utils.models import BaseModelMetadata, autov3_from_civitai_files
from ..config import config
from ..utils.file_utils import find_preview_file, get_preview_extension, calculate_sha256
from ..utils.file_utils import find_preview_file, get_preview_extension, calculate_sha256, calculate_autov3
from ..utils.metadata_manager import MetadataManager
from ..utils.civitai_utils import resolve_license_info
from .model_cache import ModelCache
@@ -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,8 @@ 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._cache_version: int = 0
self._hash_index = hash_index or ModelHashIndex()
self._tags_count = {} # Dictionary to store tag counts
self._is_initializing = False # Flag to track initialization state
@@ -86,6 +87,7 @@ class ModelScanner:
self._persistent_cache = get_persistent_cache()
self._name_display_mode = self._resolve_name_display_mode()
self._cancel_requested = False # Flag for cancellation
self._autov3_backfill_scheduled = False # One-time AutoV3 backfill trigger per process
try:
loop = asyncio.get_running_loop()
except RuntimeError:
@@ -97,6 +99,25 @@ class ModelScanner:
# Register this service
asyncio.create_task(self._register_service())
@property
def cache_version(self) -> int:
"""Monotonic version counter for the in-memory cache.
Every write path that mutates scanner cache state calls
:meth:`bump_cache_version`, so consumers (e.g. RecipeScanner) can
detect when a cached derivation of the raw data is stale. Reads never
bump.
"""
return self._cache_version
def bump_cache_version(self) -> None:
"""Invalidate derived caches by incrementing the cache version.
Public because external services (model lifecycle, route handlers)
rewrite scanner raw_data directly and must be able to invalidate it.
"""
self._cache_version += 1
def on_library_changed(self) -> None:
"""Reset caches when the active library changes."""
self._persistent_cache = get_persistent_cache()
@@ -106,6 +127,7 @@ class ModelScanner:
self._excluded_models = []
self._is_initializing = False
self._name_display_mode = self._resolve_name_display_mode()
self.bump_cache_version()
try:
loop = asyncio.get_running_loop()
@@ -182,7 +204,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()
@@ -225,6 +247,19 @@ class ModelScanner:
if not isinstance(notes, str):
notes = str(notes)
# AutoV3 three-state contract: absent key / None = "not checked yet",
# "" = "checked but unavailable" (never re-read the header), else the
# 12-char lowercase hex value. A metadata object already follows the
# contract and is passed through unchanged; a payload dict only carries
# an explicit checked state when the key is present.
if is_mapping:
if 'autov3' in source:
entry_autov3 = source['autov3'] or ''
else:
entry_autov3 = None
else:
entry_autov3 = get_value('autov3', None)
entry: Dict[str, Any] = {
'file_path': normalized_path,
# file_name is always stored WITHOUT extension (e.g. "OWSMianne_ANIMA_V1",
@@ -238,6 +273,7 @@ class ModelScanner:
'size': int(get_value('size', 0) or 0),
'modified': float(get_value('modified', 0.0) or 0.0),
'sha256': (get_value('sha256', '') or '').lower(),
'autov3': entry_autov3,
'base_model': get_value('base_model', '') or '',
'preview_url': preview_url,
'preview_nsfw_level': int(get_value('preview_nsfw_level', 0) or 0),
@@ -473,6 +509,13 @@ class ModelScanner:
if sha_value and path:
hash_index.add_entry(sha_value.lower(), path)
# Rebuild the AutoV3 index from the persisted autov3_index rows. These
# cover every known autov3 -> path mapping regardless of whether a
# sha256 row also exists for the same file.
for autov3_value, path in persisted.autov3_hash_rows:
if autov3_value and path:
hash_index.add_autov3(autov3_value.lower(), path)
tags_count: Dict[str, int] = {}
adjusted_raw_data: List[Dict[str, Any]] = []
for item in persisted.raw_data:
@@ -541,8 +584,30 @@ class ModelScanner:
'scanner_type': self.model_type,
'pageType': page_type
})
# Schedule the one-time AutoV3 backfill task (at most once per process)
# so entries loaded from a persisted snapshot that predates autov3 get
# their checked state computed in the background. The task never blocks
# or crashes the load path.
if not self._autov3_backfill_scheduled:
self._autov3_backfill_scheduled = True
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = None
if loop is not None:
loop.create_task(self._run_autov3_backfill())
return True
async def _run_autov3_backfill(self) -> None:
"""Backfill autov3 for entries loaded from the persisted cache that lack it."""
try:
from ..services.autov3_backfill_service import Autov3BackfillService # lazy import (module created by another unit)
await Autov3BackfillService.get_instance().backfill(self)
except Exception as exc:
logger.warning("AutoV3 backfill failed: %s", exc)
async def _save_persistent_cache(self, scan_result: CacheBuildResult) -> None:
if not scan_result or not getattr(self, '_persistent_cache', None):
return
@@ -555,6 +620,7 @@ class ModelScanner:
return
hash_snapshot = self._build_hash_index_snapshot(scan_result.hash_index)
autov3_snapshot = self._build_autov3_index_snapshot(scan_result.hash_index)
loop = asyncio.get_event_loop()
try:
await loop.run_in_executor(
@@ -563,7 +629,8 @@ class ModelScanner:
self.model_type,
list(scan_result.raw_data),
hash_snapshot,
list(scan_result.excluded_models)
list(scan_result.excluded_models),
autov3_snapshot,
)
except Exception as exc:
logger.warning("%s Scanner: Failed to persist cache: %s", self.model_type.capitalize(), exc)
@@ -589,6 +656,20 @@ class ModelScanner:
bucket.append(path)
return snapshot
def _build_autov3_index_snapshot(self, hash_index: Optional[ModelHashIndex]) -> Dict[str, List[str]]:
"""Build the autov3 -> [paths] snapshot for the persisted cache."""
snapshot: Dict[str, List[str]] = {}
if not hash_index:
return snapshot
for autov3_value, path in hash_index.get_all_autov3().items():
if not autov3_value or not path:
continue
bucket = snapshot.setdefault(autov3_value.lower(), [])
if path not in bucket:
bucket.append(path)
return snapshot
async def _persist_current_cache(self) -> None:
if self._cache is None or not getattr(self, '_persistent_cache', None):
return
@@ -712,7 +793,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"""
@@ -872,6 +953,8 @@ class ModelScanner:
)
continue
model_data = validation_result.entry
if model_data is None:
continue
self._ensure_license_flags(model_data)
# Add to cache
@@ -880,7 +963,11 @@ class ModelScanner:
# Update hash index if available
if 'sha256' in model_data and 'file_path' in model_data:
self._hash_index.add_entry(model_data['sha256'].lower(), model_data['file_path'])
self._hash_index.add_entry(
model_data['sha256'].lower(),
model_data['file_path'],
model_data.get('autov3') or None
)
# Update tags count
if 'tags' in model_data and model_data['tags']:
@@ -928,8 +1015,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:
@@ -964,6 +1051,7 @@ class ModelScanner:
logger.error(f"{self.model_type.capitalize()} Scanner: Error reconciling cache: {e}", exc_info=True)
finally:
self._is_initializing = False # Unset flag
self.bump_cache_version()
def is_initializing(self) -> bool:
"""Check if the scanner is currently initializing"""
@@ -1044,7 +1132,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
@@ -1068,7 +1156,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)")
@@ -1105,6 +1193,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)
@@ -1130,6 +1220,36 @@ class ModelScanner:
except Exception as e:
logger.error(f"Failed to compute SHA256 for {file_path}: {e}")
# AutoV3 resolution: prefer the Civitai AutoV3 reported for the file
# whose SHA256 matches (authoritative for recipe matching), falling
# back to the embedded safetensors header hash only for models never
# checked before (autov3 is None). A checked-unavailable state ('')
# is only upgraded by Civitai data — the header is never re-read.
current_autov3 = model_data.get('autov3')
if current_autov3 in (None, ''):
try:
civitai_data = None
if isinstance(metadata, BaseModelMetadata):
civitai_data = metadata.civitai
elif isinstance(metadata, dict):
civitai_data = metadata.get("civitai")
autov3 = autov3_from_civitai_files(
civitai_data, model_data.get("sha256") or ""
) or ""
if not autov3 and current_autov3 is None:
autov3 = (calculate_autov3(os.path.realpath(file_path)) or '').lower()
if autov3 != current_autov3:
model_data['autov3'] = autov3
if isinstance(metadata, BaseModelMetadata):
metadata.autov3 = autov3
await MetadataManager.save_metadata(file_path, metadata)
elif isinstance(metadata, dict):
# Dict payload: JSON null encodes the checked-unavailable state.
metadata['autov3'] = autov3 or None
await MetadataManager.save_metadata(file_path, metadata)
except Exception as e:
logger.error(f"Failed to resolve AutoV3 for {file_path}: {e}")
# Skip excluded models
if model_data.get('exclude', False):
excluded_models.append(model_data['file_path'])
@@ -1169,6 +1289,8 @@ class ModelScanner:
self._log_duplicate_filename_summary()
self.bump_cache_version()
def _log_duplicate_filename_summary(self) -> None:
"""Log a batched summary of duplicate filename conflicts once per scan."""
# Duplicate filename detection is only relevant for LoRAs, which use
@@ -1202,7 +1324,7 @@ class ModelScanner:
async def _sync_download_history(
self,
raw_data: List[Mapping[str, Any]],
raw_data: Sequence[Mapping[str, Any]],
*,
source: str,
) -> None:
@@ -1251,7 +1373,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] = []
@@ -1315,6 +1437,8 @@ class ModelScanner:
)
continue
result = validation_result.entry
if result is None:
continue
self._ensure_license_flags(result)
raw_data.append(result)
@@ -1322,7 +1446,7 @@ class ModelScanner:
sha_value = result.get('sha256')
model_path = result.get('file_path')
if sha_value and model_path:
hash_index.add_entry(sha_value.lower(), model_path)
hash_index.add_entry(sha_value.lower(), model_path, result.get('autov3') or None)
for tag in result.get('tags') or []:
tags_count[tag] = tags_count.get(tag, 0) + 1
@@ -1354,7 +1478,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:
@@ -1367,7 +1491,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
@@ -1391,14 +1516,19 @@ class ModelScanner:
await self._cache.resort()
# Update the hash index
self._hash_index.add_entry(metadata_dict['sha256'], metadata_dict['file_path'])
self._hash_index.add_entry(
metadata_dict['sha256'],
metadata_dict['file_path'],
metadata_dict.get('autov3') or None,
)
await self._persist_current_cache()
self.bump_cache_version()
return True
except Exception as e:
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:
@@ -1432,7 +1562,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, '/')
@@ -1480,7 +1610,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)
@@ -1498,7 +1628,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:
@@ -1524,7 +1654,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()
@@ -1547,6 +1677,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, '/')
@@ -1578,7 +1709,11 @@ class ModelScanner:
sha_value = cache_entry.get('sha256')
if sha_value:
self._hash_index.add_entry(sha_value.lower(), normalized_new_path)
self._hash_index.add_entry(
sha_value.lower(),
normalized_new_path,
cache_entry.get('autov3') or None,
)
all_folders = set(item['folder'] for item in cache.raw_data)
cache.folders = sorted(list(all_folders), key=lambda x: x.lower())
@@ -1592,8 +1727,11 @@ class ModelScanner:
if cache_modified:
await self._persist_current_cache()
self.bump_cache_version()
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]
@@ -1716,10 +1854,11 @@ class ModelScanner:
# ---- In-place update of the cache entry ----
existing_entry.clear()
existing_entry.update(desired_entry)
self.bump_cache_version()
# ---- 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:
@@ -1736,7 +1875,11 @@ class ModelScanner:
if old_sha:
self._hash_index.remove_by_path(file_path)
if new_sha:
self._hash_index.add_entry(new_sha, file_path)
self._hash_index.add_entry(
new_sha,
file_path,
desired_entry.get('autov3') or None,
)
# ---- Incremental version index update ----
new_civitai = desired_entry.get("civitai")
@@ -1787,6 +1930,75 @@ class ModelScanner:
return True
async def update_autov3_for_model(self, model_type: str, file_path: str, autov3: str) -> bool:
"""Persist an AutoV3 hash for a single model (single write path used by the backfill service).
Locates the in-memory cache entry by ``file_path`` and updates only its
``autov3`` field: the in-memory hash index, the SQLite snapshot via
:meth:`PersistentModelCache.update_single_model`, and the
``.metadata.json`` sidecar. sha256, tags, and every other field are
left untouched, so the persistent delta only ever differs in autov3.
Returns:
``True`` when the entry was found and updated, ``False`` otherwise.
Never raises failures are logged and swallowed.
"""
try:
if self._cache is None:
return False
entry = next(
(item for item in self._cache.raw_data if item.get('file_path') == file_path),
None,
)
if entry is None:
return False
# Normalize once so the memory entry, sidecar, and SQLite row agree.
autov3 = (autov3 or "").lower()
# Capture the pre-mutation state so update_single_model only sees
# an autov3 delta between old and new.
old_item = dict(entry)
entry['autov3'] = autov3 or ''
# Prefer add_entry when a sha256 is known so the sha256 and autov3
# maps stay in sync; fall back to an autov3-only registration.
sha_value = entry.get('sha256')
checked_autov3 = entry.get('autov3') or None
if sha_value:
self._hash_index.add_entry(sha_value.lower(), file_path, checked_autov3)
elif checked_autov3:
self._hash_index.add_autov3(checked_autov3, file_path)
persistent = getattr(self, '_persistent_cache', None)
if persistent is not None:
await asyncio.get_event_loop().run_in_executor(
None,
persistent.update_single_model,
model_type,
entry,
old_item,
)
# Sidecar write-back: JSON null encodes the checked-unavailable
# state. Skip silently when the sidecar does not exist.
metadata_path = f"{os.path.splitext(file_path)[0]}.metadata.json"
if os.path.exists(metadata_path):
with open(metadata_path, 'r', encoding='utf-8') as handle:
payload = json.load(handle)
if not isinstance(payload, dict):
payload = {}
payload['autov3'] = entry['autov3'] or None
await MetadataManager.save_metadata(metadata_path, payload)
self.bump_cache_version()
return True
except Exception as exc:
logger.warning("Failed to update AutoV3 for %s: %s", file_path, exc)
return False
@staticmethod
def _cache_entries_differ(a: Dict[str, Any], b: Dict[str, Any]) -> bool:
"""Return ``True`` when two cache-entry dicts differ in any field.
@@ -1846,7 +2058,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()
@@ -1862,7 +2074,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``
@@ -1885,7 +2097,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()
@@ -1966,7 +2178,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:
@@ -2114,6 +2326,8 @@ class ModelScanner:
await self._persist_current_cache()
self.bump_cache_version()
return True
except Exception as e:
@@ -2164,7 +2378,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:
+6 -6
View File
@@ -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:
+29 -24
View File
@@ -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):
+159 -21
View File
@@ -3,8 +3,8 @@ import logging
import os
import sqlite3
import threading
from dataclasses import dataclass
from typing import Dict, List, Mapping, Optional, Sequence, Tuple
from dataclasses import dataclass, field
from typing import Any, Dict, List, Mapping, Optional, Sequence, Tuple
from ..utils.cache_paths import CacheType, resolve_cache_path_with_migration
@@ -15,9 +15,10 @@ 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)
DEFAULT_LICENSE_FLAGS = 127 # 127 (0b1111111) encodes default CivitAI permissions with all commercial modes enabled.
@@ -36,6 +37,7 @@ class PersistentModelCache:
"size",
"modified",
"sha256",
"autov3",
"base_model",
"preview_url",
"preview_nsfw_level",
@@ -68,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
@@ -118,6 +120,10 @@ class PersistentModelCache:
"SELECT sha256, file_path FROM hash_index WHERE model_type = ?",
(model_type,),
).fetchall()
autov3_rows = conn.execute(
"SELECT autov3, file_path FROM autov3_index WHERE model_type = ?",
(model_type,),
).fetchall()
excluded = conn.execute(
"SELECT file_path FROM excluded_models WHERE model_type = ?",
(model_type,),
@@ -128,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 = []
@@ -139,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")
@@ -191,6 +197,8 @@ class PersistentModelCache:
"hash_status": row["hash_status"] or "completed",
"hf_url": row["hf_url"] or "",
}
if row["autov3"] is not None:
item["autov3"] = (row["autov3"] or "").lower()
raw_data.append(item)
hash_pairs = [(entry["sha256"].lower(), entry["file_path"]) for entry in hash_rows if entry["sha256"]]
@@ -201,10 +209,21 @@ class PersistentModelCache:
if sha_value:
hash_pairs.append((sha_value.lower(), item["file_path"]))
excluded_paths = [row["file_path"] for row in excluded]
return PersistedCacheData(raw_data=raw_data, hash_rows=hash_pairs, excluded_models=excluded_paths)
autov3_pairs = [
(entry["autov3"].lower(), entry["file_path"])
for entry in autov3_rows
if entry["autov3"]
]
def save_cache(self, model_type: str, raw_data: Sequence[Dict], hash_index: Dict[str, List[str]], excluded_models: Sequence[str]) -> None:
excluded_paths = [row["file_path"] for row in excluded]
return PersistedCacheData(
raw_data=raw_data,
hash_rows=hash_pairs,
excluded_models=excluded_paths,
autov3_hash_rows=autov3_pairs,
)
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:
@@ -219,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
}
@@ -251,13 +270,17 @@ class PersistentModelCache:
"DELETE FROM hash_index WHERE model_type = ? AND file_path = ?",
to_remove_models,
)
conn.executemany(
"DELETE FROM autov3_index WHERE model_type = ? AND file_path = ?",
to_remove_models,
)
conn.executemany(
"DELETE FROM excluded_models WHERE model_type = ? AND file_path = ?",
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)
@@ -289,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:
@@ -332,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:
@@ -373,6 +396,52 @@ class PersistentModelCache:
hash_inserts,
)
if autov3_hash_index is not None:
existing_autov3_rows = conn.execute(
"SELECT autov3, file_path FROM autov3_index WHERE model_type = ?",
(model_type,),
).fetchall()
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[str]] = {}
for autov3_value, paths in autov3_hash_index.items():
normalized_autov3 = (autov3_value or "").lower()
if not normalized_autov3:
continue
bucket = new_autov3_map.setdefault(normalized_autov3, set())
for path in paths:
if path:
bucket.add(path)
autov3_inserts: List[Tuple[str, str, str]] = []
autov3_deletes: List[Tuple[str, str, str]] = []
all_autov3 = set(existing_autov3_map.keys()) | set(new_autov3_map.keys())
for autov3_value in all_autov3:
existing_paths = existing_autov3_map.get(autov3_value, set())
new_paths = new_autov3_map.get(autov3_value, set())
for path in existing_paths - new_paths:
autov3_deletes.append((model_type, autov3_value, path))
for path in new_paths - existing_paths:
autov3_inserts.append((model_type, autov3_value, path))
if autov3_deletes:
conn.executemany(
"DELETE FROM autov3_index WHERE model_type = ? AND autov3 = ? AND file_path = ?",
autov3_deletes,
)
if autov3_inserts:
conn.executemany(
"INSERT OR IGNORE INTO autov3_index (model_type, autov3, file_path) VALUES (?, ?, ?)",
autov3_inserts,
)
existing_excluded_rows = conn.execute(
"SELECT file_path FROM excluded_models WHERE model_type = ?",
(model_type,),
@@ -435,6 +504,7 @@ class PersistentModelCache:
size INTEGER,
modified REAL,
sha256 TEXT,
autov3 TEXT,
base_model TEXT,
preview_url TEXT,
preview_nsfw_level INTEGER,
@@ -472,6 +542,13 @@ class PersistentModelCache:
PRIMARY KEY (model_type, sha256, file_path)
);
CREATE TABLE IF NOT EXISTS autov3_index (
model_type TEXT NOT NULL,
autov3 TEXT NOT NULL,
file_path TEXT NOT NULL,
PRIMARY KEY (model_type, autov3, file_path)
);
CREATE TABLE IF NOT EXISTS excluded_models (
model_type TEXT NOT NULL,
file_path TEXT NOT NULL,
@@ -504,6 +581,7 @@ class PersistentModelCache:
"license_flags": f"INTEGER DEFAULT {DEFAULT_LICENSE_FLAGS}",
"hash_status": "TEXT DEFAULT 'completed'",
"hf_url": "TEXT DEFAULT ''",
"autov3": "TEXT",
}
for column, definition in required_columns.items():
@@ -522,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):
@@ -549,6 +627,12 @@ class PersistentModelCache:
if license_flags is None:
license_flags = DEFAULT_LICENSE_FLAGS
autov3_value = item.get("autov3")
if autov3_value is None:
autov3_column = None
else:
autov3_column = (autov3_value or "").lower()
return (
model_type,
item.get("file_path"),
@@ -558,6 +642,7 @@ class PersistentModelCache:
int(item.get("size") or 0),
float(item.get("modified") or 0.0),
(item.get("sha256") or "").lower() or None,
autov3_column,
item.get("base_model") or "",
item.get("preview_url") or "",
int(item.get("preview_nsfw_level") or 0),
@@ -590,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.
@@ -630,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
@@ -663,6 +748,25 @@ class PersistentModelCache:
(model_type, new_sha, file_path),
)
# --- autov3_index ---
new_autov3: Optional[str] = new_item.get("autov3")
if new_autov3 is not None:
new_autov3 = (new_autov3 or "").lower()
old_autov3: Optional[str] = (old_item.get("autov3") if old_item else None)
if old_autov3 is not None:
old_autov3 = (old_autov3 or "").lower()
if new_autov3 != old_autov3:
if old_autov3:
conn.execute(
"DELETE FROM autov3_index WHERE model_type = ? AND autov3 = ? AND file_path = ?",
(model_type, old_autov3, file_path),
)
if new_autov3:
conn.execute(
"INSERT OR IGNORE INTO autov3_index (model_type, autov3, file_path) VALUES (?, ?, ?)",
(model_type, new_autov3, file_path),
)
conn.execute("COMMIT")
except Exception:
conn.execute("ROLLBACK")
@@ -676,6 +780,40 @@ class PersistentModelCache:
exc,
)
def get_models_missing_autov3(self, model_type: str) -> List[str]:
"""Return file paths whose models lack an AutoV3 checked state.
Only rows with a completed sha256 and a NULL autov3 column qualify
rows with '' (checked-unavailable) or a value are never returned, so
the backfill query self-terminates.
"""
if not self.is_enabled():
return []
if not self._schema_initialized:
self._initialize_schema()
if not self._schema_initialized:
return []
try:
with self._db_lock:
conn = self._connect(readonly=True)
try:
rows = conn.execute(
"SELECT file_path FROM models "
"WHERE model_type = ? AND autov3 IS NULL "
"AND sha256 IS NOT NULL AND sha256 != ''",
(model_type,),
).fetchall()
finally:
conn.close()
return [row["file_path"] for row in rows]
except Exception as exc:
logger.warning(
"Failed to query models missing autov3 for %s: %s",
model_type,
exc,
)
return []
def _load_tags(self, conn: sqlite3.Connection, model_type: str) -> Dict[str, List[str]]:
tag_rows = conn.execute(
"SELECT file_path, tag FROM model_tags WHERE model_type = ?",
+12 -8
View File
@@ -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"]:
+3 -1
View File
@@ -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)
+13 -12
View File
@@ -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:
+2 -2
View File
@@ -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', ''))
+142 -63
View File
@@ -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,24 +77,74 @@ 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
if checkpoint_scanner:
self._checkpoint_scanner = checkpoint_scanner
# Local hash cache (sha256 / autov2 / stored autov3 -> cache item),
# rebuilt only when either model scanner's cache_version changes.
self._local_hash_cache: dict[str, dict[str, Any]] | None = None
self._local_hash_cache_versions: tuple[int, int] | None = None
self._local_hash_cache_lock = asyncio.Lock()
self._initialized = True
async def build_local_hash_cache(self) -> dict[str, dict[str, Any]]:
"""Build a version-cached map of local model hashes to cache items.
Keys are the lowercase full sha256, the first 10 chars of the sha256
(autov2), and the stored lowercase autov3 value when present. An empty
autov3 is the "checked but unavailable" state and never produces a key.
Items without a sha256 are skipped. The dict is reused while both
scanners' cache_version values are unchanged; concurrent callers share
a single build via the lock.
"""
async with self._local_hash_cache_lock:
lora_scanner = self._lora_scanner
checkpoint_scanner = self._checkpoint_scanner
versions = (
lora_scanner.cache_version if lora_scanner is not None else 0,
checkpoint_scanner.cache_version
if checkpoint_scanner is not None
else 0,
)
if (
self._local_hash_cache is not None
and self._local_hash_cache_versions == versions
):
return self._local_hash_cache
cache: dict[str, dict[str, Any]] = {}
for scanner in (lora_scanner, checkpoint_scanner):
if scanner is None:
continue
data = await scanner.get_cached_data()
for item in data.raw_data:
sha256 = (item.get("sha256") or "").lower()
if not sha256:
continue
cache[sha256] = item
cache[sha256[:10]] = item
autov3 = (item.get("autov3") or "").lower()
if autov3:
cache[autov3] = item
self._local_hash_cache = cache
self._local_hash_cache_versions = versions
return cache
def on_library_changed(self) -> None:
"""Reset cached state when the active library changes."""
@@ -123,6 +173,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 +192,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 +211,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 +393,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 +546,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 +564,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 +651,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 +668,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 +685,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 +756,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 +784,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 +793,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 +816,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 +895,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 +930,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 +1014,8 @@ class RecipeScanner:
return
try:
from .recipe_fts_index import RecipeFTSIndex
self._fts_index = RecipeFTSIndex()
# Check if existing index is valid
@@ -987,7 +1053,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 +1068,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 +1099,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 +1119,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 +1137,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 +1198,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 +1207,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 +1233,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 +1277,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 +1313,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 +1364,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 +1415,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 +1483,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 +1509,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 +1590,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 +1623,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 +1650,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 +1720,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 +1810,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 +2033,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 +2297,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 +2389,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 +2437,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 +2542,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 +2775,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 +2806,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 +2832,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:
+17 -3
View File
@@ -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,14 +414,27 @@ 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
payload["extension"] = extension
return AnalysisResult(payload)
result = await parser.parse_metadata(metadata, recipe_scanner=recipe_scanner)
# Only the Civitai image parser accepts a local_cache parameter;
# passing it to other parsers would raise TypeError. Lazy import
# mirrors the repo style used in recipe_handlers.
from ...recipes.parsers.civitai_image import CivitaiApiMetadataParser
if isinstance(parser, CivitaiApiMetadataParser):
local_cache = await recipe_scanner.build_local_hash_cache()
result = await parser.parse_metadata(
metadata, recipe_scanner=recipe_scanner, local_cache=local_cache
)
else:
result = await parser.parse_metadata(
metadata, recipe_scanner=recipe_scanner
)
if include_image_base64 and image_path:
result["image_base64"] = self._encode_file(image_path)
@@ -494,7 +508,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()
+6 -2
View File
@@ -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",
+2 -2
View File
@@ -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())
+4
View File
@@ -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
+21 -12
View File
@@ -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",
+6 -5
View File
@@ -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 = (
+2 -1
View File
@@ -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] = []
@@ -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):
@@ -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,
@@ -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:
+15 -15
View File
@@ -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.
+2 -2
View File
@@ -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)
+1 -1
View File
@@ -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
+11
View File
@@ -80,6 +80,17 @@ CIVITAI_USER_MODEL_TYPES = [
# Default chunk size in megabytes used for hashing large files.
DEFAULT_HASH_CHUNK_SIZE_MB = 4
# Upper bound for a safetensors header block (bytes). Real headers are at most
# a few MB (tensor name/shape lists); the cap prevents a crafted file with an
# absurd 64-bit header length from forcing a multi-GB allocation during scan.
MAX_SAFETENSORS_HEADER_BYTES = 64 * 1024 * 1024
# First 12 chars of the SHA256 of an empty byte string. Some (re-packaging)
# training tools write this placeholder into safetensors metadata instead of a
# real hash; it must never be treated as a valid AutoV3 — several broken
# models sharing it would collide in the hash index and falsely match recipes.
INVALID_AUTOV3_EMPTY_HASH = "e3b0c44298fc"
# Auto-organize settings
AUTO_ORGANIZE_BATCH_SIZE = (
50 # Process models in batches to avoid overwhelming the system
+12 -12
View File
@@ -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:
+12 -10
View File
@@ -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),
)
+1 -1
View File
@@ -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
+2 -1
View File
@@ -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,
+31 -21
View File
@@ -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
+14 -2
View File
@@ -9,6 +9,8 @@ from typing import Any
from .constants import (
CARD_PREVIEW_WIDTH,
DEFAULT_HASH_CHUNK_SIZE_MB,
INVALID_AUTOV3_EMPTY_HASH,
MAX_SAFETENSORS_HEADER_BYTES,
PREVIEW_EXTENSIONS,
)
from .exif_utils import ExifUtils
@@ -90,6 +92,8 @@ def read_safetensors_metadata(file_path: str) -> dict[str, Any]:
if len(header_len_bytes) < 8:
return {}
header_len = struct.unpack("<Q", header_len_bytes)[0]
if header_len > MAX_SAFETENSORS_HEADER_BYTES:
return {}
header_bytes = f.read(header_len)
if len(header_bytes) < header_len:
return {}
@@ -123,8 +127,16 @@ def calculate_autov3(file_path: str) -> str | None:
return None
embedded_hash = metadata.get("sshs_model_hash") or metadata.get("modelspec.hash_sha256")
if embedded_hash and isinstance(embedded_hash, str) and len(embedded_hash) >= 12:
return embedded_hash[:12]
if embedded_hash and isinstance(embedded_hash, str):
# OneTrainer writes modelspec.hash_sha256 with a "0x" prefix.
embedded_hash = embedded_hash.strip().removeprefix("0x").removeprefix("0X")
if len(embedded_hash) >= 12:
autov3 = embedded_hash[:12].lower()
# The empty-string SHA256 placeholder written by some repackaging
# tools is not a real hash; treat it as unavailable so broken
# models never share one bogus value.
if autov3 != INVALID_AUTOV3_EMPTY_HASH:
return autov3
return None

Some files were not shown because too many files have changed in this diff Show More