mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-09 15:30:16 -03:00
feat(recipes): add local-only recipe rematch to scanner
This commit is contained in:
@@ -11,8 +11,10 @@ import os
|
||||
import time
|
||||
from typing import Any, Callable, Dict, Iterable, List, Optional, Set, Tuple, Union, cast
|
||||
from ..config import config
|
||||
from ..utils.constants import VALID_CHECKPOINT_SUB_TYPES, VALID_LORA_TYPES
|
||||
from ..utils.file_utils import calculate_autov3
|
||||
from .recipe_cache import RecipeCache
|
||||
from .recipes.errors import RecipeNotFoundError
|
||||
from .recipes.errors import RecipeNotFoundError, RecipePersistenceError
|
||||
from natsort import natsorted
|
||||
import sys
|
||||
import re
|
||||
@@ -26,6 +28,12 @@ if TYPE_CHECKING:
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Rematch type-gate alias map: Civitai model types are lowercased before the
|
||||
# VALID_CHECKPOINT_SUB_TYPES membership check, and raw "DiffusionModel" would
|
||||
# lowercase to "diffusionmodel", which is not a valid sub-type. Map it
|
||||
# explicitly to "diffusion_model" (mirrors Oracle R2-F1).
|
||||
_CHECKPOINT_MODEL_TYPE_ALIASES = {"diffusionmodel": "diffusion_model"}
|
||||
|
||||
|
||||
class RecipeScanner:
|
||||
"""Service for scanning and managing recipe images"""
|
||||
@@ -99,6 +107,13 @@ class RecipeScanner:
|
||||
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()
|
||||
# Computed autov3 map (absent/None-autov3 items only), rebuilt only
|
||||
# when either model scanner's cache_version changes — the
|
||||
# safetensors headers are read once per library scan, not once per
|
||||
# recipe. Mirrors the build_local_hash_cache version pattern.
|
||||
self._rematch_autov3_cache: dict[str, dict[str, Any]] | None = None
|
||||
self._rematch_autov3_versions: tuple[int, int] | None = None
|
||||
self._rematch_autov3_lock = asyncio.Lock()
|
||||
self._initialized = True
|
||||
|
||||
async def build_local_hash_cache(self) -> dict[str, dict[str, Any]]:
|
||||
@@ -145,6 +160,125 @@ class RecipeScanner:
|
||||
self._local_hash_cache_versions = versions
|
||||
return cache
|
||||
|
||||
def _is_rematch_candidate(self, entry: dict[str, Any]) -> bool:
|
||||
"""Return True when a recipe entry is eligible for local re-matching."""
|
||||
if not isinstance(entry, dict):
|
||||
return False
|
||||
unresolved = (
|
||||
entry.get("isDeleted") or not entry.get("hash") or not entry.get("file_name")
|
||||
)
|
||||
has_identifier = (
|
||||
entry.get("hash") or entry.get("modelVersionId") or entry.get("id")
|
||||
)
|
||||
return bool(unresolved and has_identifier)
|
||||
|
||||
async def _build_rematch_autov3_cache(self) -> dict[str, dict[str, Any]]:
|
||||
"""Build a version-cached map of computed AutoV3 hashes to local items.
|
||||
|
||||
Only absent/``None`` autov3 items are computed; ``''`` is the terminal
|
||||
"checked but unavailable" state and is never recomputed. The dict is
|
||||
reused while both scanners' cache_version values are unchanged, so the
|
||||
safetensors headers are read once per library scan rather than once per
|
||||
recipe. Computed values are lookup keys only — never persisted, never
|
||||
written to items.
|
||||
"""
|
||||
async with self._rematch_autov3_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._rematch_autov3_cache is not None
|
||||
and self._rematch_autov3_versions == versions
|
||||
):
|
||||
return self._rematch_autov3_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:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if "autov3" in item and item.get("autov3") is not None:
|
||||
continue
|
||||
file_path = item.get("file_path")
|
||||
if not file_path:
|
||||
continue
|
||||
computed = await asyncio.to_thread(calculate_autov3, file_path)
|
||||
key = (computed or "").lower()
|
||||
if key:
|
||||
cache[key] = item
|
||||
|
||||
self._rematch_autov3_cache = cache
|
||||
self._rematch_autov3_versions = versions
|
||||
return cache
|
||||
|
||||
async def _match_rematch_entry(
|
||||
self,
|
||||
entry: dict[str, Any],
|
||||
local_cache: dict[str, Any],
|
||||
autov3_cache: dict[str, Any],
|
||||
*,
|
||||
is_checkpoint: bool,
|
||||
) -> Optional[dict[str, Any]]:
|
||||
"""Match a recipe entry against local models across three levels.
|
||||
|
||||
L1 looks the stored hash up in the type-blind local hash cache; L2
|
||||
falls back to the version index via ``modelVersionId`` or ``id``; L3
|
||||
resolves 12-char hashes through the computed AutoV3 cache. Matched
|
||||
items are type-verified against the entry kind before being returned.
|
||||
"""
|
||||
entry_hash = (entry.get("hash") or "").lower()
|
||||
|
||||
item = local_cache.get(entry_hash)
|
||||
|
||||
if item is None:
|
||||
version_id = entry.get("modelVersionId") or entry.get("id")
|
||||
if version_id is not None:
|
||||
if is_checkpoint:
|
||||
item = self._get_checkpoint_from_version_index(str(version_id))
|
||||
else:
|
||||
item = self._get_lora_from_version_index(str(version_id))
|
||||
|
||||
if item is None and len(entry_hash) == 12:
|
||||
item = autov3_cache.get(entry_hash)
|
||||
|
||||
if item is None:
|
||||
return None
|
||||
|
||||
# Type gate: the L1 cache merges lora and checkpoint items and is
|
||||
# type-blind, so a match must be verified against the entry kind.
|
||||
sub_type = (item.get("sub_type") or "").lower()
|
||||
if sub_type:
|
||||
valid = (
|
||||
VALID_CHECKPOINT_SUB_TYPES if is_checkpoint else VALID_LORA_TYPES
|
||||
)
|
||||
if sub_type not in valid:
|
||||
return None
|
||||
else:
|
||||
civitai_type = (
|
||||
(item.get("civitai") or {}).get("model", {}) or {}
|
||||
).get("type", "")
|
||||
if civitai_type:
|
||||
normalized = civitai_type.lower()
|
||||
if is_checkpoint:
|
||||
normalized = _CHECKPOINT_MODEL_TYPE_ALIASES.get(
|
||||
normalized, normalized
|
||||
)
|
||||
valid = VALID_CHECKPOINT_SUB_TYPES
|
||||
else:
|
||||
valid = VALID_LORA_TYPES
|
||||
if normalized not in valid:
|
||||
return None
|
||||
|
||||
return item
|
||||
|
||||
def on_library_changed(self) -> None:
|
||||
"""Reset cached state when the active library changes."""
|
||||
|
||||
@@ -411,6 +545,369 @@ class RecipeScanner:
|
||||
|
||||
return False
|
||||
|
||||
async def rematch_recipe_by_id(self, recipe_id: str) -> Dict[str, Any]:
|
||||
"""Rematch a single recipe's deleted lora/checkpoint entries locally.
|
||||
|
||||
Match snapshots (local hash cache + computed autov3 cache) are built
|
||||
BEFORE acquiring the mutation lock — both are read-only snapshots and
|
||||
the version-cached hash dict would otherwise rebuild mid-run if a scan
|
||||
bumps a scanner's cache_version while we hold the lock.
|
||||
|
||||
Args:
|
||||
recipe_id: ID of the recipe to rematch
|
||||
|
||||
Returns:
|
||||
Dict summary of the rematch result (success/rematched/skipped).
|
||||
Raises RecipeNotFoundError when the recipe is missing.
|
||||
"""
|
||||
local_cache = await self.build_local_hash_cache()
|
||||
autov3_cache = await self._build_rematch_autov3_cache()
|
||||
|
||||
async with self._mutation_lock:
|
||||
# Get raw recipe from cache directly to avoid formatted fields
|
||||
cache = await self.get_cached_data()
|
||||
recipe = next(
|
||||
(r for r in cache.raw_data if str(r.get("id", "")) == recipe_id), None
|
||||
)
|
||||
|
||||
if not recipe:
|
||||
raise RecipeNotFoundError(f"Recipe {recipe_id} not found")
|
||||
|
||||
try:
|
||||
rematched, _errors = await self._rematch_single_recipe(
|
||||
recipe, local_cache, autov3_cache
|
||||
)
|
||||
except RecipePersistenceError as exc:
|
||||
return {
|
||||
"success": False,
|
||||
"errors": 1,
|
||||
"rematched": 0,
|
||||
"skipped": 0,
|
||||
"recipe": recipe,
|
||||
"error": str(exc),
|
||||
}
|
||||
|
||||
if rematched == 0:
|
||||
return {
|
||||
"success": True,
|
||||
"rematched": 0,
|
||||
"skipped": 1,
|
||||
"recipe": recipe,
|
||||
}
|
||||
|
||||
# Enriched re-fetch so the frontend receives file_url/preview fields.
|
||||
return {
|
||||
"success": True,
|
||||
"rematched": rematched,
|
||||
"skipped": 0,
|
||||
"recipe": await self.get_recipe_by_id(recipe_id),
|
||||
}
|
||||
|
||||
async def _rematch_single_recipe(
|
||||
self,
|
||||
recipe: Dict[str, Any],
|
||||
local_cache: dict[str, dict[str, Any]],
|
||||
autov3_cache: dict[str, dict[str, Any]],
|
||||
) -> Tuple[int, int]:
|
||||
"""Rematch a single recipe's lora/checkpoint entries against local models.
|
||||
|
||||
Shared per-recipe helper used by ``rematch_recipe_by_id`` and the bulk
|
||||
rematch entry points. Mutates the recipe dict in place, recomputes the
|
||||
fingerprint and persists via ``_save_recipe_persistently`` when any
|
||||
entry changed. ``_schedule_resort`` is deliberately NOT called here —
|
||||
it is hoisted to the public entry points.
|
||||
|
||||
Args:
|
||||
recipe: The recipe dictionary to rematch (modified in-place)
|
||||
local_cache: L1 hash cache snapshot (build_local_hash_cache)
|
||||
autov3_cache: L3 computed-autov3 cache snapshot
|
||||
|
||||
Returns:
|
||||
Tuple of (rematched_entries, errors). The errors element is always
|
||||
0 on a normal return — a persistence failure RAISES
|
||||
``RecipePersistenceError`` so callers can count it.
|
||||
|
||||
Raises:
|
||||
RecipePersistenceError: when the recipe changed but
|
||||
``_save_recipe_persistently`` returned False.
|
||||
"""
|
||||
rematched = 0
|
||||
|
||||
# Lora entries
|
||||
loras = recipe.get("loras", [])
|
||||
if isinstance(loras, list):
|
||||
for entry in loras:
|
||||
if not self._is_rematch_candidate(entry):
|
||||
continue
|
||||
item = await self._match_rematch_entry(
|
||||
entry, local_cache, autov3_cache, is_checkpoint=False
|
||||
)
|
||||
if item is None:
|
||||
continue
|
||||
self._write_rematch_lora_entry(entry, item)
|
||||
rematched += 1
|
||||
|
||||
# Checkpoint entry (dict only — legacy string checkpoints are skipped
|
||||
# silently since ``entry.get`` on a str would raise AttributeError).
|
||||
checkpoint = recipe.get("checkpoint")
|
||||
if isinstance(checkpoint, dict):
|
||||
if self._is_rematch_candidate(checkpoint):
|
||||
item = await self._match_rematch_entry(
|
||||
checkpoint, local_cache, autov3_cache, is_checkpoint=True
|
||||
)
|
||||
if item is not None:
|
||||
self._write_rematch_checkpoint_entry(checkpoint, item)
|
||||
rematched += 1
|
||||
|
||||
if rematched == 0:
|
||||
return (0, 0)
|
||||
|
||||
from ..utils.utils import calculate_recipe_fingerprint
|
||||
|
||||
recipe["fingerprint"] = calculate_recipe_fingerprint(recipe.get("loras", []))
|
||||
|
||||
saved = await self._save_recipe_persistently(recipe)
|
||||
if not saved:
|
||||
raise RecipePersistenceError(
|
||||
f"Failed to persist recipe {recipe.get('id')} after rematch"
|
||||
)
|
||||
|
||||
self._update_fts_index_for_recipe(recipe, "update")
|
||||
return (rematched, 0)
|
||||
|
||||
async def rematch_all_recipes(
|
||||
self, progress_callback: Optional[Callable[[Dict[str, Any]], Any]] = None
|
||||
) -> Dict[str, Any]:
|
||||
"""Rematch every recipe's deleted lora/checkpoint entries locally.
|
||||
|
||||
Match snapshots (local hash cache + computed autov3 cache) are built
|
||||
ONCE before the loop — both are read-only and the version-cached hash
|
||||
dict would otherwise rebuild mid-run if a scan bumps a scanner's
|
||||
cache_version while the mutation lock is held. ``_schedule_resort`` is
|
||||
called exactly once after the loop: it spawns an asyncio task per call,
|
||||
so per-recipe calls would race one resort task per recipe.
|
||||
|
||||
Args:
|
||||
progress_callback: Optional callback for progress updates
|
||||
(started/processing/cancelled/completed events).
|
||||
|
||||
Returns:
|
||||
Dict summary of the rematch run
|
||||
(success/status/rematched/skipped/errors/total).
|
||||
"""
|
||||
if progress_callback:
|
||||
await progress_callback({"status": "started"})
|
||||
|
||||
# Match snapshots built once and shared by every recipe in the loop.
|
||||
local_cache = await self.build_local_hash_cache()
|
||||
autov3_cache = await self._build_rematch_autov3_cache()
|
||||
|
||||
async with self._mutation_lock:
|
||||
cache = await self.get_cached_data()
|
||||
all_recipes = list(cache.raw_data)
|
||||
total = len(all_recipes)
|
||||
rematched_count = 0
|
||||
skipped_count = 0
|
||||
errors_count = 0
|
||||
|
||||
for i, recipe in enumerate(all_recipes):
|
||||
if self.is_cancelled():
|
||||
logger.info("Recipe rematch cancelled by user")
|
||||
if progress_callback:
|
||||
await progress_callback(
|
||||
{
|
||||
"status": "cancelled",
|
||||
"current": i,
|
||||
"total": total,
|
||||
"rematched": rematched_count,
|
||||
"skipped": skipped_count,
|
||||
"errors": errors_count,
|
||||
}
|
||||
)
|
||||
return {
|
||||
"success": False,
|
||||
"status": "cancelled",
|
||||
"rematched": rematched_count,
|
||||
"skipped": skipped_count,
|
||||
"errors": errors_count,
|
||||
"total": total,
|
||||
}
|
||||
|
||||
try:
|
||||
# Report progress
|
||||
if progress_callback:
|
||||
await progress_callback(
|
||||
{
|
||||
"status": "processing",
|
||||
"current": i + 1,
|
||||
"total": total,
|
||||
"recipe_name": recipe.get("name", "Unknown"),
|
||||
}
|
||||
)
|
||||
|
||||
rematched, _ = await self._rematch_single_recipe(
|
||||
recipe, local_cache, autov3_cache
|
||||
)
|
||||
if rematched > 0:
|
||||
rematched_count += 1
|
||||
else:
|
||||
skipped_count += 1
|
||||
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
f"Error rematching recipe {recipe.get('file_path')}: {exc}"
|
||||
)
|
||||
errors_count += 1
|
||||
|
||||
# Hoisted to one call — _schedule_resort spawns an asyncio task
|
||||
# per call, so per-recipe calls would race 5k resort tasks.
|
||||
self._schedule_resort()
|
||||
|
||||
# Final progress update
|
||||
if progress_callback:
|
||||
await progress_callback(
|
||||
{
|
||||
"status": "completed",
|
||||
"rematched": rematched_count,
|
||||
"skipped": skipped_count,
|
||||
"errors": errors_count,
|
||||
"total": total,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"rematched": rematched_count,
|
||||
"skipped": skipped_count,
|
||||
"errors": errors_count,
|
||||
"total": total,
|
||||
}
|
||||
|
||||
async def rematch_recipes_bulk(self, recipe_ids: List[str]) -> Dict[str, Any]:
|
||||
"""Rematch a set of recipes by their IDs.
|
||||
|
||||
Iterates ``rematch_recipe_by_id`` over each id: not-found ids are
|
||||
counted as skipped, and unexpected per-recipe exceptions are counted as
|
||||
errors with the loop continuing so partial results are never lost.
|
||||
Persist failures are already converted to the by_id return shape and
|
||||
are counted via its ``errors`` field only — never double-counted here.
|
||||
|
||||
Args:
|
||||
recipe_ids: List of recipe ids to rematch.
|
||||
|
||||
Returns:
|
||||
Dict summary of the bulk run
|
||||
(success/total/rematched/skipped/errors/recipes).
|
||||
"""
|
||||
total = len(recipe_ids)
|
||||
rematched = 0
|
||||
skipped = 0
|
||||
errors = 0
|
||||
recipes: List[Dict[str, Any]] = []
|
||||
|
||||
for recipe_id in recipe_ids:
|
||||
try:
|
||||
result = await self.rematch_recipe_by_id(recipe_id)
|
||||
if result.get("success"):
|
||||
rematched += result.get("rematched", 0)
|
||||
skipped += result.get("skipped", 0)
|
||||
if result.get("recipe"):
|
||||
recipes.append(result["recipe"])
|
||||
else:
|
||||
errors += result.get("errors", 0)
|
||||
except RecipeNotFoundError:
|
||||
skipped += 1
|
||||
except Exception as exc:
|
||||
logger.error(f"Error rematching recipe {recipe_id}: {exc}")
|
||||
errors += 1
|
||||
|
||||
self._schedule_resort()
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"total": total,
|
||||
"rematched": rematched,
|
||||
"skipped": skipped,
|
||||
"errors": errors,
|
||||
"recipes": recipes,
|
||||
}
|
||||
|
||||
def _write_rematch_lora_entry(
|
||||
self, entry: Dict[str, Any], item: Dict[str, Any]
|
||||
) -> None:
|
||||
"""Write back a matched local model to a lora recipe entry."""
|
||||
entry["isDeleted"] = False
|
||||
|
||||
# Only truthy hashes are written — pending/failed items carry an empty
|
||||
# sha256 and an unconditional write would wipe a valid stored hash.
|
||||
new_hash = (item.get("sha256") or "").lower()
|
||||
if new_hash:
|
||||
entry["hash"] = new_hash
|
||||
|
||||
if item.get("file_name"):
|
||||
entry["file_name"] = item["file_name"]
|
||||
|
||||
civitai = item.get("civitai")
|
||||
if isinstance(civitai, dict):
|
||||
if civitai.get("id") is not None:
|
||||
entry["modelVersionId"] = civitai["id"]
|
||||
# modelName comes from the item, NOT civitai.model.name — the slim
|
||||
# civitai payload drops model.name entirely.
|
||||
if item.get("model_name"):
|
||||
entry["modelName"] = item["model_name"]
|
||||
if civitai.get("name"):
|
||||
entry["modelVersionName"] = civitai["name"]
|
||||
|
||||
def _write_rematch_checkpoint_entry(
|
||||
self, entry: Dict[str, Any], item: Dict[str, Any]
|
||||
) -> None:
|
||||
"""Write back a matched local model to a checkpoint recipe entry.
|
||||
|
||||
Follows the pinned stored key set: parser-style entries carry
|
||||
name/version/id/type/baseModel/file_name/hash; widget-style entries
|
||||
additionally carry modelName/modelVersionName. Keys are only updated
|
||||
when they already exist on the entry (or written fresh for the
|
||||
identifier key when neither identifier form exists).
|
||||
"""
|
||||
entry["isDeleted"] = False
|
||||
|
||||
new_hash = (item.get("sha256") or "").lower()
|
||||
if new_hash:
|
||||
entry["hash"] = new_hash
|
||||
|
||||
if item.get("file_name"):
|
||||
entry["file_name"] = item["file_name"]
|
||||
|
||||
civitai = item.get("civitai")
|
||||
civ_name = civitai.get("name") if isinstance(civitai, dict) else None
|
||||
civ_id = civitai.get("id") if isinstance(civitai, dict) else None
|
||||
item_name = item.get("model_name")
|
||||
item_base_model = item.get("base_model")
|
||||
|
||||
# Backfill name/version/baseModel only when the entry already has them.
|
||||
if "name" in entry and item_name:
|
||||
entry["name"] = item_name
|
||||
if "version" in entry and civ_name:
|
||||
entry["version"] = civ_name
|
||||
if "baseModel" in entry and item_base_model:
|
||||
entry["baseModel"] = item_base_model
|
||||
|
||||
# Widget-style entries (modelName/modelVersionName) get stale values
|
||||
# refreshed; parser-style entries never gain them.
|
||||
if "modelName" in entry and item_name:
|
||||
entry["modelName"] = item_name
|
||||
if "modelVersionName" in entry and civ_name:
|
||||
entry["modelVersionName"] = civ_name
|
||||
|
||||
# Identifier key updated per the entry's existing convention.
|
||||
if civ_id is not None:
|
||||
if "modelVersionId" in entry:
|
||||
entry["modelVersionId"] = civ_id
|
||||
elif "id" in entry:
|
||||
entry["id"] = civ_id
|
||||
else:
|
||||
entry["modelVersionId"] = civ_id
|
||||
|
||||
async def _save_recipe_persistently(self, recipe: Dict[str, Any]) -> bool:
|
||||
"""Helper to save a recipe to both JSON and EXIF metadata."""
|
||||
recipe_id = recipe.get("id")
|
||||
|
||||
@@ -20,3 +20,12 @@ class RecipeDownloadError(RecipeServiceError):
|
||||
|
||||
class RecipeConflictError(RecipeServiceError):
|
||||
"""Raised when a conflicting recipe state is detected."""
|
||||
|
||||
|
||||
class RecipePersistenceError(RecipeServiceError):
|
||||
"""Raised when a rematched recipe cannot be persisted to disk.
|
||||
|
||||
Raised by the recipe rematch path when ``_save_recipe_persistently``
|
||||
returns False (JSON/EXIF/SQLite write failure). Callers translate it
|
||||
into a ``success: False`` summary with an ``error`` message key.
|
||||
"""
|
||||
|
||||
Reference in New Issue
Block a user