mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-09-21 11:11:26 -03:00
Compare commits
12 Commits
4bf9a4b640
...
27027c4497
| Author | SHA1 | Date | |
|---|---|---|---|
| 27027c4497 | |||
| 86c85c08ec | |||
| 196c8ffc3e | |||
| cfc95ee02a | |||
| 479fa36997 | |||
| 3e1216e9bc | |||
| 007883b7d1 | |||
| dc9200a12c | |||
| d2f955266d | |||
| 8e724538bd | |||
| 6fcdeb799d | |||
| 97b9b1f62b |
@@ -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
@@ -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
@@ -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 (
|
||||
|
||||
@@ -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 {}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -52,7 +52,7 @@ class AutomaticMetadataParser(RecipeMetadataParser):
|
||||
negative_and_params = ""
|
||||
|
||||
# Initialize metadata
|
||||
metadata = {
|
||||
metadata: Dict[str, Any] = {
|
||||
"prompt": prompt,
|
||||
"loras": []
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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'),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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}"
|
||||
|
||||
@@ -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)):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 ```` 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 ````."""
|
||||
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:
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -283,7 +283,7 @@ class CivitaiBaseModelService:
|
||||
return None
|
||||
|
||||
if isinstance(result, str):
|
||||
data = json.loads(result)
|
||||
data: Any = json.loads(result)
|
||||
else:
|
||||
data = result
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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,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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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 = ?",
|
||||
|
||||
@@ -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"]:
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
Reference in New Issue
Block a user