mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-25 06:51:26 -03:00
Compare commits
5 Commits
ab4154c57d
...
186ef4da78
| Author | SHA1 | Date | |
|---|---|---|---|
| 186ef4da78 | |||
| dc674098e7 | |||
| 9087b4b07c | |||
| 8e45c22d7a | |||
| 191c4e03cd |
+2227
-2223
File diff suppressed because it is too large
Load Diff
@@ -773,6 +773,8 @@
|
|||||||
"deleteAll": "Delete Selected",
|
"deleteAll": "Delete Selected",
|
||||||
"downloadMissingLoras": "Download Missing LoRAs",
|
"downloadMissingLoras": "Download Missing LoRAs",
|
||||||
"downloadExamples": "Download Example Images",
|
"downloadExamples": "Download Example Images",
|
||||||
|
"downloadMissingExamples": "Download Missing",
|
||||||
|
"reprocessExamples": "Re-process All",
|
||||||
"clear": "Clear Selection",
|
"clear": "Clear Selection",
|
||||||
"skipMetadataRefreshCount": "Skip ({count} models)",
|
"skipMetadataRefreshCount": "Skip ({count} models)",
|
||||||
"resumeMetadataRefreshCount": "Resume ({count} models)",
|
"resumeMetadataRefreshCount": "Resume ({count} models)",
|
||||||
@@ -808,6 +810,8 @@
|
|||||||
"sendToWorkflowReplace": "Send to Workflow (Replace)",
|
"sendToWorkflowReplace": "Send to Workflow (Replace)",
|
||||||
"openExamples": "Open Examples Folder",
|
"openExamples": "Open Examples Folder",
|
||||||
"downloadExamples": "Download Example Images",
|
"downloadExamples": "Download Example Images",
|
||||||
|
"downloadMissingExamples": "Download Missing",
|
||||||
|
"reprocessExamples": "Re-process All",
|
||||||
"replacePreview": "Replace Preview",
|
"replacePreview": "Replace Preview",
|
||||||
"setContentRating": "Set Content Rating",
|
"setContentRating": "Set Content Rating",
|
||||||
"moveToFolder": "Move to Folder",
|
"moveToFolder": "Move to Folder",
|
||||||
|
|||||||
+2227
-2223
File diff suppressed because it is too large
Load Diff
+2227
-2223
File diff suppressed because it is too large
Load Diff
+2227
-2223
File diff suppressed because it is too large
Load Diff
+2227
-2223
File diff suppressed because it is too large
Load Diff
+2227
-2223
File diff suppressed because it is too large
Load Diff
+2227
-2223
File diff suppressed because it is too large
Load Diff
+2227
-2223
File diff suppressed because it is too large
Load Diff
+2227
-2223
File diff suppressed because it is too large
Load Diff
@@ -2,7 +2,8 @@ import json
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
|
||||||
from .constants import CLIP_SKIP_SENTINEL, MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE, METADATA_OVERWRITE_FIELDS
|
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE
|
||||||
|
from .overwrite_utils import collect_overwrite_params
|
||||||
|
|
||||||
|
|
||||||
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
||||||
@@ -1233,14 +1234,7 @@ class MetadataOverwriteExtractor(NodeMetadataExtractor):
|
|||||||
if not inputs:
|
if not inputs:
|
||||||
return
|
return
|
||||||
|
|
||||||
overwrite_params = {}
|
overwrite_params = collect_overwrite_params(inputs)
|
||||||
for key in METADATA_OVERWRITE_FIELDS:
|
|
||||||
value = inputs.get(key)
|
|
||||||
if key == "clip_skip":
|
|
||||||
if value != CLIP_SKIP_SENTINEL:
|
|
||||||
overwrite_params[key] = value
|
|
||||||
elif value: # truthy — only overwrite when user provided a real value
|
|
||||||
overwrite_params[key] = value
|
|
||||||
|
|
||||||
if overwrite_params:
|
if overwrite_params:
|
||||||
metadata.setdefault(OVERWRITE, {})
|
metadata.setdefault(OVERWRITE, {})
|
||||||
|
|||||||
@@ -0,0 +1,42 @@
|
|||||||
|
"""Shared helpers for Metadata Overwrite node metadata collection.
|
||||||
|
|
||||||
|
Used by both the MetadataOverwriteLM node (execution time) and the
|
||||||
|
MetadataOverwriteExtractor (hook time) so the conversion/filtering logic
|
||||||
|
cannot drift between the two paths.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Any, Dict
|
||||||
|
|
||||||
|
from ..utils.utils import model_patcher_to_name
|
||||||
|
from .constants import CLIP_SKIP_SENTINEL, METADATA_OVERWRITE_FIELDS
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def collect_overwrite_params(values: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Convert node input values into non-default overwrite parameters.
|
||||||
|
|
||||||
|
For most fields, a falsy value (empty string, 0) means "not set" and is
|
||||||
|
skipped. clip_skip uses a dedicated sentinel (-25) so that a wired value
|
||||||
|
of 0 is preserved. The ``model`` field accepts either a manual string or
|
||||||
|
a wired MODEL (ModelPatcher) connection; in the latter case the source
|
||||||
|
model name is extracted from the patcher's ``cached_patcher_init`` and
|
||||||
|
stored as a ComfyUI-style relative path.
|
||||||
|
"""
|
||||||
|
result: Dict[str, Any] = {}
|
||||||
|
for key in METADATA_OVERWRITE_FIELDS:
|
||||||
|
value = values.get(key)
|
||||||
|
if key == "model" and not isinstance(value, str):
|
||||||
|
value = model_patcher_to_name(value)
|
||||||
|
if value is None:
|
||||||
|
logger.warning(
|
||||||
|
"Could not extract model name from wired MODEL input "
|
||||||
|
"(no cached_patcher_init); model metadata overwrite skipped"
|
||||||
|
)
|
||||||
|
if key == "clip_skip":
|
||||||
|
if value != CLIP_SKIP_SENTINEL:
|
||||||
|
result[key] = value
|
||||||
|
elif value:
|
||||||
|
result[key] = value
|
||||||
|
return result
|
||||||
@@ -9,10 +9,8 @@ but users may wire 0 to express "no clip skip / default".
|
|||||||
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from ..metadata_collector.constants import (
|
from ..metadata_collector.constants import CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL
|
||||||
CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL,
|
from ..metadata_collector.overwrite_utils import collect_overwrite_params
|
||||||
METADATA_OVERWRITE_FIELDS,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class MetadataOverwriteLM:
|
class MetadataOverwriteLM:
|
||||||
@@ -87,12 +85,16 @@ class MetadataOverwriteLM:
|
|||||||
},
|
},
|
||||||
),
|
),
|
||||||
"model": (
|
"model": (
|
||||||
"STRING",
|
"STRING,MODEL",
|
||||||
{
|
{
|
||||||
"default": "",
|
"default": "",
|
||||||
|
"widgetType": "STRING",
|
||||||
"tooltip": (
|
"tooltip": (
|
||||||
"The checkpoint or diffusion model (UNet) used "
|
"The checkpoint or diffusion model (UNet) used "
|
||||||
"for generation. Only overwrites when non-empty."
|
"for generation. Fill in the name manually or "
|
||||||
|
"connect a MODEL output — the model name is then "
|
||||||
|
"extracted automatically. Only overwrites when "
|
||||||
|
"non-empty."
|
||||||
),
|
),
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
@@ -158,13 +160,10 @@ class MetadataOverwriteLM:
|
|||||||
For most fields, a falsy value (empty string, 0) means "not set"
|
For most fields, a falsy value (empty string, 0) means "not set"
|
||||||
and is skipped. clip_skip uses a dedicated sentinel (-25) so that
|
and is skipped. clip_skip uses a dedicated sentinel (-25) so that
|
||||||
a wired value of 0 is preserved and reaches the metadata pipeline.
|
a wired value of 0 is preserved and reaches the metadata pipeline.
|
||||||
|
|
||||||
|
The ``model`` field accepts either a manual string or a wired MODEL
|
||||||
|
(ModelPatcher) connection; in the latter case the underlying model
|
||||||
|
name is extracted from the patcher's ``cached_patcher_init`` and
|
||||||
|
stored as a ComfyUI-style relative path.
|
||||||
"""
|
"""
|
||||||
result: dict[str, Any] = {}
|
return (collect_overwrite_params(kwargs),)
|
||||||
for key in METADATA_OVERWRITE_FIELDS:
|
|
||||||
value = kwargs.get(key)
|
|
||||||
if key == "clip_skip":
|
|
||||||
if value != _CLIP_SKIP_SENTINEL:
|
|
||||||
result[key] = value
|
|
||||||
elif value:
|
|
||||||
result[key] = value
|
|
||||||
return (result,)
|
|
||||||
|
|||||||
@@ -7,6 +7,21 @@ from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_c
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _reload_gguf_unet(
|
||||||
|
unet_path: str, weight_dtype: str, disable_dynamic: bool = False
|
||||||
|
) -> object:
|
||||||
|
"""Reload a GGUF diffusion model from disk (cached_patcher_init factory).
|
||||||
|
|
||||||
|
Mirrors the GGUF branch of UNETLoaderLM.load_unet so ModelPatcher
|
||||||
|
deepclone/dynamic machinery can rebuild GGUF models with the correct
|
||||||
|
GGMLOps. ``disable_dynamic`` is accepted for signature compatibility
|
||||||
|
with core ComfyUI loaders.
|
||||||
|
"""
|
||||||
|
loader = UNETLoaderLM()
|
||||||
|
model, = loader._load_gguf_unet(unet_path, unet_path, weight_dtype)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
class UNETLoaderLM:
|
class UNETLoaderLM:
|
||||||
"""UNET Loader with support for extra folder paths
|
"""UNET Loader with support for extra folder paths
|
||||||
|
|
||||||
@@ -196,6 +211,12 @@ class UNETLoaderLM:
|
|||||||
# Wrap with GGUFModelPatcher
|
# Wrap with GGUFModelPatcher
|
||||||
model = GGUFModelPatcher.clone(model)
|
model = GGUFModelPatcher.clone(model)
|
||||||
|
|
||||||
|
# Register a reload factory so the MODEL carries its source path
|
||||||
|
# (cached_patcher_init) like core ComfyUI loaders do — required
|
||||||
|
# for model-name extraction downstream and for ModelPatcher
|
||||||
|
# deepclone/dynamic machinery.
|
||||||
|
model.cached_patcher_init = (_reload_gguf_unet, (unet_path, weight_dtype))
|
||||||
|
|
||||||
return (model,)
|
return (model,)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -172,6 +172,7 @@ class DownloadManager:
|
|||||||
model_types = data.get("model_types", ["lora", "checkpoint"])
|
model_types = data.get("model_types", ["lora", "checkpoint"])
|
||||||
delay = float(data.get("delay", 0.2))
|
delay = float(data.get("delay", 0.2))
|
||||||
force = data.get("force", False)
|
force = data.get("force", False)
|
||||||
|
model_hashes = data.get("model_hashes", [])
|
||||||
|
|
||||||
# Step 2: Validate configuration (fast lookup)
|
# Step 2: Validate configuration (fast lookup)
|
||||||
settings_manager = get_settings_manager()
|
settings_manager = get_settings_manager()
|
||||||
@@ -241,6 +242,7 @@ class DownloadManager:
|
|||||||
delay,
|
delay,
|
||||||
active_library,
|
active_library,
|
||||||
force,
|
force,
|
||||||
|
model_hashes,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -577,8 +579,9 @@ class DownloadManager:
|
|||||||
delay,
|
delay,
|
||||||
library_name,
|
library_name,
|
||||||
force: bool = False,
|
force: bool = False,
|
||||||
|
model_hashes: list[str] | None = None,
|
||||||
):
|
):
|
||||||
"""Download example images for all models."""
|
"""Download example images for all models (or only the given hashes)."""
|
||||||
|
|
||||||
downloader = await get_downloader()
|
downloader = await get_downloader()
|
||||||
|
|
||||||
@@ -606,6 +609,18 @@ class DownloadManager:
|
|||||||
if model.get("sha256"):
|
if model.get("sha256"):
|
||||||
all_models.append((scanner_type, model, scanner))
|
all_models.append((scanner_type, model, scanner))
|
||||||
|
|
||||||
|
# Restrict to the requested hashes when provided (empty = all models).
|
||||||
|
# Explicit targets are a directed user request, so previously failed
|
||||||
|
# models are retried instead of skipped.
|
||||||
|
explicit_targets = bool(model_hashes)
|
||||||
|
if model_hashes:
|
||||||
|
hash_set = {h.lower() for h in model_hashes}
|
||||||
|
all_models = [
|
||||||
|
(scanner_type, model, scanner)
|
||||||
|
for scanner_type, model, scanner in all_models
|
||||||
|
if model.get("sha256", "").lower() in hash_set
|
||||||
|
]
|
||||||
|
|
||||||
# Update total count
|
# Update total count
|
||||||
self._progress["total"] = len(all_models)
|
self._progress["total"] = len(all_models)
|
||||||
logger.debug(f"Found {self._progress['total']} models to process")
|
logger.debug(f"Found {self._progress['total']} models to process")
|
||||||
@@ -629,6 +644,7 @@ class DownloadManager:
|
|||||||
downloader,
|
downloader,
|
||||||
library_name,
|
library_name,
|
||||||
force,
|
force,
|
||||||
|
explicit_targets,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Update progress
|
# Update progress
|
||||||
@@ -725,6 +741,7 @@ class DownloadManager:
|
|||||||
downloader,
|
downloader,
|
||||||
library_name,
|
library_name,
|
||||||
force: bool = False,
|
force: bool = False,
|
||||||
|
explicit_targets: bool = False,
|
||||||
):
|
):
|
||||||
"""Process a single model download."""
|
"""Process a single model download."""
|
||||||
|
|
||||||
@@ -747,8 +764,9 @@ class DownloadManager:
|
|||||||
self._progress["current_model"] = f"{model_name} ({model_hash[:8]})"
|
self._progress["current_model"] = f"{model_name} ({model_hash[:8]})"
|
||||||
await self._broadcast_progress(status="running")
|
await self._broadcast_progress(status="running")
|
||||||
|
|
||||||
# Skip if already in failed models (unless force mode is enabled)
|
# Skip if already in failed models (unless force mode is enabled or
|
||||||
if not force and model_hash in self._progress["failed_models"]:
|
# the model was explicitly targeted by hash)
|
||||||
|
if not force and not explicit_targets and model_hash in self._progress["failed_models"]:
|
||||||
logger.debug(f"Skipping known failed model: {model_name}")
|
logger.debug(f"Skipping known failed model: {model_name}")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -757,30 +775,34 @@ class DownloadManager:
|
|||||||
)
|
)
|
||||||
existing_files = _model_directory_has_files(model_dir)
|
existing_files = _model_directory_has_files(model_dir)
|
||||||
|
|
||||||
# Skip if already processed AND directory exists with files
|
# Model-level guard: a populated folder counts as done. Explicitly
|
||||||
if model_hash in self._progress["processed_models"]:
|
# targeted models bypass it so the per-image existence pre-check can
|
||||||
if existing_files:
|
# fill individual gaps without re-fetching existing files.
|
||||||
logger.debug(f"Skipping already processed model: {model_name}")
|
if not explicit_targets:
|
||||||
|
# Skip if already processed AND directory exists with files
|
||||||
|
if model_hash in self._progress["processed_models"]:
|
||||||
|
if existing_files:
|
||||||
|
logger.debug(f"Skipping already processed model: {model_name}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
"Model %s (%s) marked as processed but folder empty or missing, reprocessing triggered",
|
||||||
|
model_name,
|
||||||
|
model_hash,
|
||||||
|
)
|
||||||
|
# Track that we are reprocessing this model for summary logging
|
||||||
|
self._progress["reprocessed_models"].add(model_hash)
|
||||||
|
# Remove from processed models since we need to reprocess
|
||||||
|
self._progress["processed_models"].discard(model_hash)
|
||||||
|
|
||||||
|
if existing_files and model_hash not in self._progress["processed_models"]:
|
||||||
|
logger.debug(
|
||||||
|
"Model folder already populated for %s, marking as processed without download",
|
||||||
|
model_name,
|
||||||
|
)
|
||||||
|
self._progress["processed_models"].add(model_hash)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
logger.debug(
|
|
||||||
"Model %s (%s) marked as processed but folder empty or missing, reprocessing triggered",
|
|
||||||
model_name,
|
|
||||||
model_hash,
|
|
||||||
)
|
|
||||||
# Track that we are reprocessing this model for summary logging
|
|
||||||
self._progress["reprocessed_models"].add(model_hash)
|
|
||||||
# Remove from processed models since we need to reprocess
|
|
||||||
self._progress["processed_models"].discard(model_hash)
|
|
||||||
|
|
||||||
if existing_files and model_hash not in self._progress["processed_models"]:
|
|
||||||
logger.debug(
|
|
||||||
"Model folder already populated for %s, marking as processed without download",
|
|
||||||
model_name,
|
|
||||||
)
|
|
||||||
self._progress["processed_models"].add(model_hash)
|
|
||||||
return False
|
|
||||||
|
|
||||||
if not model_dir:
|
if not model_dir:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Unable to resolve example images folder for model %s (%s)",
|
"Unable to resolve example images folder for model %s (%s)",
|
||||||
@@ -884,7 +906,7 @@ class DownloadManager:
|
|||||||
model_name,
|
model_name,
|
||||||
)
|
)
|
||||||
# Clear failed_models so non-force runs can retry
|
# Clear failed_models so non-force runs can retry
|
||||||
if force and model_hash in self._progress["failed_models"]:
|
if (force or explicit_targets) and model_hash in self._progress["failed_models"]:
|
||||||
self._progress["failed_models"].discard(model_hash)
|
self._progress["failed_models"].discard(model_hash)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Removed {model_name} from failed_models after force retry with rate-limited images"
|
f"Removed {model_name} from failed_models after force retry with rate-limited images"
|
||||||
@@ -904,7 +926,7 @@ class DownloadManager:
|
|||||||
)
|
)
|
||||||
elif success:
|
elif success:
|
||||||
self._progress["processed_models"].add(model_hash)
|
self._progress["processed_models"].add(model_hash)
|
||||||
if force and model_hash in self._progress["failed_models"]:
|
if (force or explicit_targets) and model_hash in self._progress["failed_models"]:
|
||||||
self._progress["failed_models"].discard(model_hash)
|
self._progress["failed_models"].discard(model_hash)
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Removed {model_name} from failed_models after successful force retry"
|
f"Removed {model_name} from failed_models after successful force retry"
|
||||||
|
|||||||
@@ -113,6 +113,26 @@ class ExampleImagesProcessor:
|
|||||||
message = str(error).lower()
|
message = str(error).lower()
|
||||||
return '404' in message or 'file not found' in message
|
return '404' in message or 'file not found' in message
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _example_image_file_exists(model_dir: str, index: int, media_type_hint: str | None = None) -> bool:
|
||||||
|
"""Return True when the file that would be written for a media index already exists.
|
||||||
|
|
||||||
|
The final filename (``image_{index}{extension}``) depends on the downloaded
|
||||||
|
content, so the extension cannot be known ahead of time. The post-download
|
||||||
|
check skips the write when the exact target file exists; this pre-check
|
||||||
|
approximates that with the candidate extensions for the media type (videos
|
||||||
|
only when the metadata hints at a video) so the network request is avoided
|
||||||
|
for files that already exist on disk.
|
||||||
|
"""
|
||||||
|
if media_type_hint == "video":
|
||||||
|
extensions = SUPPORTED_MEDIA_EXTENSIONS['videos']
|
||||||
|
else:
|
||||||
|
extensions = SUPPORTED_MEDIA_EXTENSIONS['images']
|
||||||
|
return any(
|
||||||
|
os.path.exists(os.path.join(model_dir, f"image_{index}{ext}"))
|
||||||
|
for ext in extensions
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def download_model_images(model_hash, model_name, model_images, model_dir, optimize, downloader):
|
async def download_model_images(model_hash, model_name, model_images, model_dir, optimize, downloader):
|
||||||
"""Download images for a single model
|
"""Download images for a single model
|
||||||
@@ -140,6 +160,11 @@ class ExampleImagesProcessor:
|
|||||||
if optimize and 'civitai.com' in image_url:
|
if optimize and 'civitai.com' in image_url:
|
||||||
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
||||||
|
|
||||||
|
# Skip the download when the file already exists on disk
|
||||||
|
if ExampleImagesProcessor._example_image_file_exists(model_dir, i, image.get("type")):
|
||||||
|
logger.debug("File already exists, skipping download for %s", image_url)
|
||||||
|
continue
|
||||||
|
|
||||||
# Download the file first to determine the actual file type
|
# Download the file first to determine the actual file type
|
||||||
try:
|
try:
|
||||||
logger.debug(f"Downloading media file {i} for {model_name}")
|
logger.debug(f"Downloading media file {i} for {model_name}")
|
||||||
@@ -229,6 +254,11 @@ class ExampleImagesProcessor:
|
|||||||
if optimize and 'civitai.com' in image_url:
|
if optimize and 'civitai.com' in image_url:
|
||||||
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
||||||
|
|
||||||
|
# Skip the download when the file already exists on disk
|
||||||
|
if ExampleImagesProcessor._example_image_file_exists(model_dir, i, image.get("type")):
|
||||||
|
logger.debug("File already exists, skipping download for %s", image_url)
|
||||||
|
continue
|
||||||
|
|
||||||
async def _attempt_download() -> tuple:
|
async def _attempt_download() -> tuple:
|
||||||
logger.debug("Downloading media file %s for %s", i, model_name)
|
logger.debug("Downloading media file %s for %s", i, model_name)
|
||||||
return await downloader.download_to_memory(
|
return await downloader.download_to_memory(
|
||||||
|
|||||||
+48
-1
@@ -1,7 +1,7 @@
|
|||||||
from difflib import SequenceMatcher
|
from difflib import SequenceMatcher
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
from typing import Dict
|
from typing import Any, Dict, List, Optional
|
||||||
from ..services.service_registry import ServiceRegistry
|
from ..services.service_registry import ServiceRegistry
|
||||||
from ..config import config
|
from ..config import config
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
@@ -294,6 +294,53 @@ def _format_model_name_for_comfyui(file_path: str, model_roots: list) -> str:
|
|||||||
return os.path.basename(file_path)
|
return os.path.basename(file_path)
|
||||||
|
|
||||||
|
|
||||||
|
def model_patcher_to_name(model_patcher: Any) -> Optional[str]:
|
||||||
|
"""Extract a ComfyUI-style model name from a MODEL (ModelPatcher) object.
|
||||||
|
|
||||||
|
Core ComfyUI loaders record the absolute weight file path on the patcher's
|
||||||
|
``cached_patcher_init`` attribute:
|
||||||
|
- load_checkpoint_guess_config -> (fn, (ckpt_path, ...), index)
|
||||||
|
- load_diffusion_model -> (fn, (unet_path, model_options))
|
||||||
|
Patcher clones (LoRA loaders, model merges, ...) preserve the attribute,
|
||||||
|
so the name is recoverable anywhere downstream of a core loader — including
|
||||||
|
from LoRA Manager's own loaders (CheckpointLoaderLM / UNETLoaderLM), which
|
||||||
|
call the same core load functions.
|
||||||
|
|
||||||
|
The absolute path is converted to the ComfyUI-style relative name used by
|
||||||
|
the metadata pipeline (covering standard ComfyUI roots and LoRA Manager
|
||||||
|
extra folder paths).
|
||||||
|
|
||||||
|
Returns None when the path cannot be recovered (e.g. third-party loaders
|
||||||
|
that never set ``cached_patcher_init``).
|
||||||
|
"""
|
||||||
|
init = getattr(model_patcher, "cached_patcher_init", None)
|
||||||
|
if not isinstance(init, (tuple, list)) or len(init) < 2:
|
||||||
|
return None
|
||||||
|
args = init[1]
|
||||||
|
abs_path = args[0] if args else None
|
||||||
|
if not isinstance(abs_path, str) or not abs_path:
|
||||||
|
return None
|
||||||
|
return _abs_model_path_to_name(abs_path)
|
||||||
|
|
||||||
|
|
||||||
|
def _abs_model_path_to_name(abs_path: str) -> str:
|
||||||
|
"""Convert an absolute model path to a ComfyUI-style relative name.
|
||||||
|
|
||||||
|
Tries standard ComfyUI model roots plus LoRA Manager extra folder paths;
|
||||||
|
falls back to the bare filename.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
roots: List[str] = list(config.base_models_roots or [])
|
||||||
|
roots.extend(config.extra_checkpoints_roots or [])
|
||||||
|
roots.extend(config.extra_unet_roots or [])
|
||||||
|
formatted = _format_model_name_for_comfyui(abs_path, roots)
|
||||||
|
if formatted:
|
||||||
|
return formatted
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return os.path.basename(abs_path)
|
||||||
|
|
||||||
|
|
||||||
def fuzzy_match(text: str, pattern: str, threshold: float = 0.85) -> bool:
|
def fuzzy_match(text: str, pattern: str, threshold: float = 0.85) -> bool:
|
||||||
"""
|
"""
|
||||||
Check if text matches pattern using fuzzy matching.
|
Check if text matches pattern using fuzzy matching.
|
||||||
|
|||||||
@@ -184,7 +184,8 @@ export const DOWNLOAD_ENDPOINTS = {
|
|||||||
downloadGet: '/api/lm/download-model-get',
|
downloadGet: '/api/lm/download-model-get',
|
||||||
cancelGet: '/api/lm/cancel-download-get',
|
cancelGet: '/api/lm/cancel-download-get',
|
||||||
progress: '/api/lm/download-progress',
|
progress: '/api/lm/download-progress',
|
||||||
exampleImages: '/api/lm/force-download-example-images' // New endpoint for downloading example images
|
exampleImages: '/api/lm/force-download-example-images', // Re-process example images ignoring previous status
|
||||||
|
exampleImagesMissing: '/api/lm/download-example-images' // Download only missing example images
|
||||||
};
|
};
|
||||||
|
|
||||||
// Hugging Face API endpoints
|
// Hugging Face API endpoints
|
||||||
|
|||||||
@@ -1641,7 +1641,7 @@ export class BaseModelApiClient {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async downloadExampleImages(modelHashes, modelTypes = null) {
|
async downloadExampleImages(modelHashes, modelTypes = null, { force = true } = {}) {
|
||||||
let ws = null;
|
let ws = null;
|
||||||
|
|
||||||
await state.loadingManager.showWithProgress(async (loading) => {
|
await state.loadingManager.showWithProgress(async (loading) => {
|
||||||
@@ -1700,8 +1700,13 @@ export class BaseModelApiClient {
|
|||||||
// Determine optimize setting
|
// Determine optimize setting
|
||||||
const optimize = state.global?.settings?.optimize_example_images ?? true;
|
const optimize = state.global?.settings?.optimize_example_images ?? true;
|
||||||
|
|
||||||
|
// force=false routes to the regular endpoint, which skips already-processed models
|
||||||
|
const endpoint = force
|
||||||
|
? DOWNLOAD_ENDPOINTS.exampleImages
|
||||||
|
: DOWNLOAD_ENDPOINTS.exampleImagesMissing;
|
||||||
|
|
||||||
// Make the API request to start the download process
|
// Make the API request to start the download process
|
||||||
const response = await fetch(DOWNLOAD_ENDPOINTS.exampleImages, {
|
const response = await fetch(endpoint, {
|
||||||
method: 'POST',
|
method: 'POST',
|
||||||
headers: {
|
headers: {
|
||||||
'Content-Type': 'application/json'
|
'Content-Type': 'application/json'
|
||||||
@@ -1710,6 +1715,7 @@ export class BaseModelApiClient {
|
|||||||
model_hashes: modelHashes,
|
model_hashes: modelHashes,
|
||||||
output_dir: outputDir,
|
output_dir: outputDir,
|
||||||
optimize: optimize,
|
optimize: optimize,
|
||||||
|
force: force,
|
||||||
model_types: modelTypes || [this.apiConfig.config.singularName]
|
model_types: modelTypes || [this.apiConfig.config.singularName]
|
||||||
})
|
})
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -137,11 +137,10 @@ export class BulkContextMenu extends BaseContextMenu {
|
|||||||
downloadMissingLorasItem.style.display = currentModelType === 'recipes' ? 'flex' : 'none';
|
downloadMissingLorasItem.style.display = currentModelType === 'recipes' ? 'flex' : 'none';
|
||||||
}
|
}
|
||||||
|
|
||||||
const downloadExampleImagesItem = this.menu.querySelector('[data-action="download-example-images"]');
|
const downloadExampleImagesSubmenu = this.menu.querySelector('[data-has-submenu="download-example-images"]');
|
||||||
if (downloadExampleImagesItem) {
|
if (downloadExampleImagesSubmenu) {
|
||||||
// Show on model pages (loras, checkpoints, embeddings), hide on recipes
|
// Show on model pages (loras, checkpoints, embeddings), hide on recipes
|
||||||
const modelPages = ['loras', 'checkpoints', 'embeddings'];
|
downloadExampleImagesSubmenu.style.display = ['loras', 'checkpoints', 'embeddings'].includes(currentModelType) ? 'flex' : 'none';
|
||||||
downloadExampleImagesItem.style.display = modelPages.includes(currentModelType) ? 'flex' : 'none';
|
|
||||||
}
|
}
|
||||||
|
|
||||||
const skipMetadataRefreshItem = this.menu.querySelector('[data-action="skip-metadata-refresh"]');
|
const skipMetadataRefreshItem = this.menu.querySelector('[data-action="skip-metadata-refresh"]');
|
||||||
@@ -294,8 +293,11 @@ export class BulkContextMenu extends BaseContextMenu {
|
|||||||
case 'download-missing-loras':
|
case 'download-missing-loras':
|
||||||
this.handleDownloadMissingLoras();
|
this.handleDownloadMissingLoras();
|
||||||
break;
|
break;
|
||||||
|
case 'download-missing-example-images':
|
||||||
|
this.handleDownloadExampleImages({ force: false });
|
||||||
|
break;
|
||||||
case 'download-example-images':
|
case 'download-example-images':
|
||||||
this.handleDownloadExampleImages();
|
this.handleDownloadExampleImages({ force: true });
|
||||||
break;
|
break;
|
||||||
case 'clear':
|
case 'clear':
|
||||||
bulkManager.clearSelection();
|
bulkManager.clearSelection();
|
||||||
@@ -340,7 +342,7 @@ export class BulkContextMenu extends BaseContextMenu {
|
|||||||
await bulkMissingLoraDownloadManager.downloadMissingLoras(selectedRecipes);
|
await bulkMissingLoraDownloadManager.downloadMissingLoras(selectedRecipes);
|
||||||
}
|
}
|
||||||
|
|
||||||
async handleDownloadExampleImages() {
|
async handleDownloadExampleImages({ force = true } = {}) {
|
||||||
if (state.selectedModels.size === 0) {
|
if (state.selectedModels.size === 0) {
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
@@ -361,7 +363,7 @@ export class BulkContextMenu extends BaseContextMenu {
|
|||||||
|
|
||||||
try {
|
try {
|
||||||
const apiClient = getModelApiClient();
|
const apiClient = getModelApiClient();
|
||||||
await apiClient.downloadExampleImages([...hashes]);
|
await apiClient.downloadExampleImages([...hashes], null, { force });
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Bulk download example images failed:', error);
|
console.error('Bulk download example images failed:', error);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -347,7 +347,10 @@ export const ModelContextMenuMixin = {
|
|||||||
openExampleImagesFolder(this.currentCard.dataset.sha256);
|
openExampleImagesFolder(this.currentCard.dataset.sha256);
|
||||||
return true;
|
return true;
|
||||||
case 'download-examples':
|
case 'download-examples':
|
||||||
this.downloadExampleImages();
|
this.downloadExampleImages(false);
|
||||||
|
return true;
|
||||||
|
case 'download-examples-force':
|
||||||
|
this.downloadExampleImages(true);
|
||||||
return true;
|
return true;
|
||||||
case 'civitai':
|
case 'civitai':
|
||||||
if (this.currentCard.dataset.from_civitai === 'true') {
|
if (this.currentCard.dataset.from_civitai === 'true') {
|
||||||
@@ -378,7 +381,7 @@ export const ModelContextMenuMixin = {
|
|||||||
},
|
},
|
||||||
|
|
||||||
// Download example images method
|
// Download example images method
|
||||||
async downloadExampleImages() {
|
async downloadExampleImages(force = false) {
|
||||||
const modelHash = this.currentCard.dataset.sha256;
|
const modelHash = this.currentCard.dataset.sha256;
|
||||||
if (!modelHash) {
|
if (!modelHash) {
|
||||||
showToast('toast.contextMenu.missingHash', {}, 'error');
|
showToast('toast.contextMenu.missingHash', {}, 'error');
|
||||||
@@ -387,7 +390,7 @@ export const ModelContextMenuMixin = {
|
|||||||
|
|
||||||
try {
|
try {
|
||||||
const apiClient = getModelApiClient();
|
const apiClient = getModelApiClient();
|
||||||
await apiClient.downloadExampleImages([modelHash]);
|
await apiClient.downloadExampleImages([modelHash], null, { force });
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Error downloading example images:', error);
|
console.error('Error downloading example images:', error);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -90,7 +90,7 @@ export class BulkManager {
|
|||||||
moveAll: true,
|
moveAll: true,
|
||||||
autoOrganize: false,
|
autoOrganize: false,
|
||||||
deleteAll: true,
|
deleteAll: true,
|
||||||
setContentRating: false,
|
setContentRating: true,
|
||||||
skipMetadataRefresh: false,
|
skipMetadataRefresh: false,
|
||||||
setFavorite: true,
|
setFavorite: true,
|
||||||
unfavorite: true,
|
unfavorite: true,
|
||||||
@@ -1528,14 +1528,18 @@ export class BulkManager {
|
|||||||
let failureCount = 0;
|
let failureCount = 0;
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const apiClient = getModelApiClient();
|
const isRecipesPage = state.currentPageType === 'recipes';
|
||||||
for (const filePath of targets) {
|
for (const filePath of targets) {
|
||||||
if (cancelled) {
|
if (cancelled) {
|
||||||
showToast('toast.api.operationCancelled', {}, 'info');
|
showToast('toast.api.operationCancelled', {}, 'info');
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
try {
|
try {
|
||||||
await apiClient.saveModelMetadata(filePath, { preview_nsfw_level: level });
|
if (isRecipesPage) {
|
||||||
|
await updateRecipeMetadata(filePath, { preview_nsfw_level: level });
|
||||||
|
} else {
|
||||||
|
await getModelApiClient().saveModelMetadata(filePath, { preview_nsfw_level: level });
|
||||||
|
}
|
||||||
successCount++;
|
successCount++;
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
failureCount++;
|
failureCount++;
|
||||||
|
|||||||
@@ -32,7 +32,12 @@
|
|||||||
<div class="context-menu-separator menu-section-break"></div>
|
<div class="context-menu-separator menu-section-break"></div>
|
||||||
<!-- Media / Preview -->
|
<!-- Media / Preview -->
|
||||||
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
|
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
|
||||||
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }}</div>
|
<div class="context-menu-item has-submenu" data-has-submenu="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }} <i class="fas fa-chevron-right submenu-arrow"></i>
|
||||||
|
<div class="context-submenu">
|
||||||
|
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadMissingExamples') }}</div>
|
||||||
|
<div class="context-menu-item" data-action="download-examples-force"><i class="fas fa-redo-alt"></i> {{ t('loras.contextMenu.reprocessExamples') }}</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
|
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
|
||||||
<div class="context-menu-separator menu-section-break"></div>
|
<div class="context-menu-separator menu-section-break"></div>
|
||||||
<!-- Attributes -->
|
<!-- Attributes -->
|
||||||
|
|||||||
@@ -44,8 +44,18 @@
|
|||||||
<div class="context-menu-item" data-action="preview">
|
<div class="context-menu-item" data-action="preview">
|
||||||
<i class="fas fa-folder-open"></i> <span>{{ t('loras.contextMenu.openExamples') }}</span>
|
<i class="fas fa-folder-open"></i> <span>{{ t('loras.contextMenu.openExamples') }}</span>
|
||||||
</div>
|
</div>
|
||||||
<div class="context-menu-item" data-action="download-examples">
|
<div class="context-menu-item has-submenu" data-has-submenu="download-examples">
|
||||||
<i class="fas fa-download"></i> <span>{{ t('loras.contextMenu.downloadExamples') }}</span>
|
<i class="fas fa-download"></i>
|
||||||
|
<span>{{ t('loras.contextMenu.downloadExamples') }}</span>
|
||||||
|
<i class="fas fa-chevron-right submenu-arrow"></i>
|
||||||
|
<div class="context-submenu">
|
||||||
|
<div class="context-menu-item" data-action="download-examples">
|
||||||
|
<i class="fas fa-download"></i> <span>{{ t('loras.contextMenu.downloadMissingExamples') }}</span>
|
||||||
|
</div>
|
||||||
|
<div class="context-menu-item" data-action="download-examples-force">
|
||||||
|
<i class="fas fa-redo-alt"></i> <span>{{ t('loras.contextMenu.reprocessExamples') }}</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="context-menu-item" data-action="replace-preview">
|
<div class="context-menu-item" data-action="replace-preview">
|
||||||
<i class="fas fa-image"></i> <span>{{ t('loras.contextMenu.replacePreview') }}</span>
|
<i class="fas fa-image"></i> <span>{{ t('loras.contextMenu.replacePreview') }}</span>
|
||||||
@@ -136,8 +146,18 @@
|
|||||||
</div>
|
</div>
|
||||||
<div class="context-menu-section" data-section="download">
|
<div class="context-menu-section" data-section="download">
|
||||||
<div class="context-menu-section-header">{{ t('loras.bulkOperations.sections.download') }}</div>
|
<div class="context-menu-section-header">{{ t('loras.bulkOperations.sections.download') }}</div>
|
||||||
<div class="context-menu-item" data-action="download-example-images">
|
<div class="context-menu-item has-submenu" data-has-submenu="download-example-images">
|
||||||
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadExamples') }}</span>
|
<i class="fas fa-download"></i>
|
||||||
|
<span>{{ t('loras.bulkOperations.downloadExamples') }}</span>
|
||||||
|
<i class="fas fa-chevron-right submenu-arrow"></i>
|
||||||
|
<div class="context-submenu">
|
||||||
|
<div class="context-menu-item" data-action="download-missing-example-images">
|
||||||
|
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingExamples') }}</span>
|
||||||
|
</div>
|
||||||
|
<div class="context-menu-item" data-action="download-example-images">
|
||||||
|
<i class="fas fa-redo-alt"></i> <span>{{ t('loras.bulkOperations.reprocessExamples') }}</span>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="context-menu-item" data-action="download-missing-loras">
|
<div class="context-menu-item" data-action="download-missing-loras">
|
||||||
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingLoras') }}</span>
|
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingLoras') }}</span>
|
||||||
|
|||||||
@@ -32,7 +32,12 @@
|
|||||||
<div class="context-menu-separator menu-section-break"></div>
|
<div class="context-menu-separator menu-section-break"></div>
|
||||||
<!-- Media / Preview -->
|
<!-- Media / Preview -->
|
||||||
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
|
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
|
||||||
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }}</div>
|
<div class="context-menu-item has-submenu" data-has-submenu="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }} <i class="fas fa-chevron-right submenu-arrow"></i>
|
||||||
|
<div class="context-submenu">
|
||||||
|
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadMissingExamples') }}</div>
|
||||||
|
<div class="context-menu-item" data-action="download-examples-force"><i class="fas fa-redo-alt"></i> {{ t('loras.contextMenu.reprocessExamples') }}</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
|
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
|
||||||
<div class="context-menu-separator menu-section-break"></div>
|
<div class="context-menu-separator menu-section-break"></div>
|
||||||
<!-- Attributes -->
|
<!-- Attributes -->
|
||||||
|
|||||||
@@ -2155,4 +2155,35 @@ describe('Interaction-level regression coverage', () => {
|
|||||||
excludedItem.dispatchEvent(new Event('click', { bubbles: true }));
|
excludedItem.dispatchEvent(new Event('click', { bubbles: true }));
|
||||||
expect(window.pageControls.enterExcludedView).toHaveBeenCalledTimes(1);
|
expect(window.pageControls.enterExcludedView).toHaveBeenCalledTimes(1);
|
||||||
});
|
});
|
||||||
|
|
||||||
|
it('routes single-model example downloads to missing-only and force paths', async () => {
|
||||||
|
document.body.innerHTML = `
|
||||||
|
<div id="loraContextMenu" class="context-menu">
|
||||||
|
<div class="context-menu-item has-submenu" data-has-submenu="download-examples">
|
||||||
|
<div class="context-submenu">
|
||||||
|
<div class="context-menu-item" data-action="download-examples"></div>
|
||||||
|
<div class="context-menu-item" data-action="download-examples-force"></div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
`;
|
||||||
|
|
||||||
|
const { LoraContextMenu } = await import('../../../static/js/components/ContextMenu/LoraContextMenu.js');
|
||||||
|
const contextMenu = new LoraContextMenu();
|
||||||
|
|
||||||
|
const card = document.createElement('div');
|
||||||
|
card.className = 'model-card';
|
||||||
|
card.dataset.filepath = '/models/test.safetensors';
|
||||||
|
card.dataset.sha256 = 'abc123hash';
|
||||||
|
document.body.appendChild(card);
|
||||||
|
|
||||||
|
contextMenu.showMenu(100, 100, card);
|
||||||
|
|
||||||
|
document.querySelector('[data-action="download-examples"]').dispatchEvent(new Event('click', { bubbles: true }));
|
||||||
|
expect(downloadExampleImagesApiMock).toHaveBeenCalledWith(['abc123hash'], null, { force: false });
|
||||||
|
|
||||||
|
contextMenu.showMenu(100, 100, card);
|
||||||
|
document.querySelector('[data-action="download-examples-force"]').dispatchEvent(new Event('click', { bubbles: true }));
|
||||||
|
expect(downloadExampleImagesApiMock).toHaveBeenCalledWith(['abc123hash'], null, { force: true });
|
||||||
|
});
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -0,0 +1,133 @@
|
|||||||
|
import { describe, it, beforeEach, expect, vi } from 'vitest';
|
||||||
|
|
||||||
|
const showToastMock = vi.fn();
|
||||||
|
const translateMock = vi.fn((key, params, fallback) => (typeof fallback === 'string' ? fallback : key));
|
||||||
|
const getNSFWLevelNameMock = vi.fn((level) => {
|
||||||
|
if (level >= 16) return 'XXX';
|
||||||
|
if (level >= 8) return 'X';
|
||||||
|
if (level >= 4) return 'R';
|
||||||
|
if (level >= 2) return 'PG13';
|
||||||
|
if (level >= 1) return 'PG';
|
||||||
|
return 'Unknown';
|
||||||
|
});
|
||||||
|
|
||||||
|
const loadingManagerStub = {
|
||||||
|
showSimpleLoading: vi.fn(),
|
||||||
|
showCancelButton: vi.fn(),
|
||||||
|
hide: vi.fn(),
|
||||||
|
};
|
||||||
|
|
||||||
|
const stateStub = {
|
||||||
|
currentPageType: 'recipes',
|
||||||
|
bulkMode: false,
|
||||||
|
selectedModels: new Set(),
|
||||||
|
loadingManager: loadingManagerStub,
|
||||||
|
virtualScroller: { updateSingleItem: vi.fn() },
|
||||||
|
global: { settings: {} },
|
||||||
|
};
|
||||||
|
|
||||||
|
const saveModelMetadataMock = vi.fn();
|
||||||
|
const getModelApiClientMock = vi.fn(() => ({ saveModelMetadata: saveModelMetadataMock }));
|
||||||
|
const updateRecipeMetadataMock = vi.fn(() => Promise.resolve({ success: true }));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/state/index.js', () => ({
|
||||||
|
state: stateStub,
|
||||||
|
getCurrentPageState: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
|
||||||
|
showToast: showToastMock,
|
||||||
|
copyToClipboard: vi.fn(),
|
||||||
|
sendLoraToWorkflow: vi.fn(),
|
||||||
|
sendEmbeddingToWorkflow: vi.fn(),
|
||||||
|
buildLoraSyntax: vi.fn(),
|
||||||
|
getNSFWLevelName: getNSFWLevelNameMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/modelApiFactory.js', () => ({
|
||||||
|
getModelApiClient: getModelApiClientMock,
|
||||||
|
resetAndReload: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/recipeApi.js', () => ({
|
||||||
|
RecipeSidebarApiClient: class {},
|
||||||
|
updateRecipeMetadata: updateRecipeMetadataMock,
|
||||||
|
extractRecipeId: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/api/apiConfig.js', () => ({
|
||||||
|
MODEL_TYPES: { LORA: 'loras', CHECKPOINT: 'checkpoints', EMBEDDING: 'embeddings' },
|
||||||
|
MODEL_CONFIG: {},
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/managers/ModalManager.js', () => ({
|
||||||
|
modalManager: { showModal: vi.fn(), closeModal: vi.fn() },
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/shared/ModelCard.js', () => ({
|
||||||
|
updateCardsForBulkMode: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/i18nHelpers.js', () => ({
|
||||||
|
translate: translateMock,
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/utils/priorityTagHelpers.js', () => ({
|
||||||
|
getPriorityTagSuggestions: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
vi.mock('../../../static/js/components/shared/NsfwLevelSelector.js', () => ({
|
||||||
|
getNsfwLevelSelector: vi.fn(),
|
||||||
|
}));
|
||||||
|
|
||||||
|
describe('BulkManager bulk content rating', () => {
|
||||||
|
beforeEach(() => {
|
||||||
|
vi.clearAllMocks();
|
||||||
|
stateStub.currentPageType = 'recipes';
|
||||||
|
stateStub.bulkMode = false;
|
||||||
|
stateStub.selectedModels.clear();
|
||||||
|
saveModelMetadataMock.mockResolvedValue(undefined);
|
||||||
|
updateRecipeMetadataMock.mockResolvedValue({ success: true });
|
||||||
|
});
|
||||||
|
|
||||||
|
async function createBulkManager() {
|
||||||
|
const { BulkManager } = await import('../../../static/js/managers/BulkManager.js');
|
||||||
|
return new BulkManager();
|
||||||
|
}
|
||||||
|
|
||||||
|
it('exposes the content rating action on the recipes page action config', async () => {
|
||||||
|
const bulk = await createBulkManager();
|
||||||
|
expect(bulk.actionConfig.recipes.setContentRating).toBe(true);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('persists the rating through the recipe API when on the recipes page', async () => {
|
||||||
|
const bulk = await createBulkManager();
|
||||||
|
stateStub.currentPageType = 'recipes';
|
||||||
|
stateStub.selectedModels.add('/recipes/test.webp');
|
||||||
|
|
||||||
|
const ok = await bulk.setBulkContentRating(4, ['/recipes/test.webp']);
|
||||||
|
|
||||||
|
expect(ok).toBe(true);
|
||||||
|
expect(updateRecipeMetadataMock).toHaveBeenCalledWith('/recipes/test.webp', { preview_nsfw_level: 4 });
|
||||||
|
expect(updateRecipeMetadataMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(saveModelMetadataMock).not.toHaveBeenCalled();
|
||||||
|
expect(showToastMock).toHaveBeenCalledWith(
|
||||||
|
'toast.models.bulkContentRatingSet',
|
||||||
|
{ count: 1, level: 'R' },
|
||||||
|
'success'
|
||||||
|
);
|
||||||
|
});
|
||||||
|
|
||||||
|
it('persists the rating through the model API on model pages', async () => {
|
||||||
|
const bulk = await createBulkManager();
|
||||||
|
stateStub.currentPageType = 'loras';
|
||||||
|
stateStub.selectedModels.add('/models/test.safetensors');
|
||||||
|
|
||||||
|
const ok = await bulk.setBulkContentRating(8, ['/models/test.safetensors']);
|
||||||
|
|
||||||
|
expect(ok).toBe(true);
|
||||||
|
expect(saveModelMetadataMock).toHaveBeenCalledWith('/models/test.safetensors', { preview_nsfw_level: 8 });
|
||||||
|
expect(saveModelMetadataMock).toHaveBeenCalledTimes(1);
|
||||||
|
expect(updateRecipeMetadataMock).not.toHaveBeenCalled();
|
||||||
|
});
|
||||||
|
});
|
||||||
@@ -529,7 +529,8 @@ async def test_not_found_example_images_are_cleaned(
|
|||||||
|
|
||||||
model_dir = images_root / model_hash
|
model_dir = images_root / model_hash
|
||||||
model_dir.mkdir(parents=True, exist_ok=True)
|
model_dir.mkdir(parents=True, exist_ok=True)
|
||||||
(model_dir / "image_0.png").write_bytes(b"first")
|
# Pre-existing file collides with the valid image index (1) so the
|
||||||
|
# pre-download existence check must skip it without a network request
|
||||||
(model_dir / "image_1.png").write_bytes(b"second")
|
(model_dir / "image_1.png").write_bytes(b"second")
|
||||||
|
|
||||||
async def fake_process_local_examples(*_args, **_kwargs):
|
async def fake_process_local_examples(*_args, **_kwargs):
|
||||||
@@ -608,11 +609,188 @@ async def test_not_found_example_images_are_cleaned(
|
|||||||
]
|
]
|
||||||
|
|
||||||
files = sorted(p.name for p in model_dir.iterdir())
|
files = sorted(p.name for p in model_dir.iterdir())
|
||||||
assert files == ["image_0.png", "image_1.png"]
|
assert files == ["image_1.png"]
|
||||||
assert (model_dir / "image_0.png").read_bytes() == b"first"
|
|
||||||
assert (model_dir / "image_1.png").read_bytes() == b"second"
|
assert (model_dir / "image_1.png").read_bytes() == b"second"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_failed_models_retried_when_explicitly_targeted(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path,
|
||||||
|
settings_manager,
|
||||||
|
):
|
||||||
|
ws_manager = RecordingWebSocketManager()
|
||||||
|
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||||
|
|
||||||
|
images_root = tmp_path / "examples"
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(images_root))
|
||||||
|
|
||||||
|
model_hash = "a" * 64
|
||||||
|
model_path = tmp_path / "model.safetensors"
|
||||||
|
model_path.write_text("data", encoding="utf-8")
|
||||||
|
|
||||||
|
model_metadata = {
|
||||||
|
"sha256": model_hash,
|
||||||
|
"model_name": "Failed Example",
|
||||||
|
"file_path": str(model_path),
|
||||||
|
"file_name": "model.safetensors",
|
||||||
|
"civitai": {"images": [{"url": "https://example.com/valid.png"}]},
|
||||||
|
}
|
||||||
|
|
||||||
|
scanner = StubScanner([model_metadata.copy()])
|
||||||
|
_patch_scanner(monkeypatch, scanner)
|
||||||
|
|
||||||
|
# Persist a previous failure so the skip path is exercised
|
||||||
|
images_root.mkdir(parents=True, exist_ok=True)
|
||||||
|
(images_root / ".download_progress.json").write_text(
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"failed_models": [model_hash],
|
||||||
|
"processed_models": [],
|
||||||
|
"rate_limited_models": [],
|
||||||
|
}
|
||||||
|
),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def fake_process_local_examples(*_args, **_kwargs):
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def fake_get_updated_model(model_hash_arg, _scanner):
|
||||||
|
return model_metadata
|
||||||
|
|
||||||
|
class DownloaderStub:
|
||||||
|
def __init__(self):
|
||||||
|
self.calls: list[str] = []
|
||||||
|
|
||||||
|
async def download_to_memory(self, url, *_args, **_kwargs):
|
||||||
|
self.calls.append(url)
|
||||||
|
return True, b"\x89PNG\r\n\x1a\n", {"content-type": "image/png"}
|
||||||
|
|
||||||
|
downloader = DownloaderStub()
|
||||||
|
|
||||||
|
async def fake_get_downloader():
|
||||||
|
return downloader
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module.ExampleImagesProcessor,
|
||||||
|
"process_local_examples",
|
||||||
|
staticmethod(fake_process_local_examples),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module.MetadataUpdater,
|
||||||
|
"get_updated_model",
|
||||||
|
staticmethod(fake_get_updated_model),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(download_module, "get_downloader", fake_get_downloader)
|
||||||
|
|
||||||
|
# Without explicit hashes the previously failed model is skipped
|
||||||
|
skipped_manager = download_module.DownloadManager(ws_manager=RecordingWebSocketManager())
|
||||||
|
result = await skipped_manager.start_download({"model_types": ["lora"], "delay": 0})
|
||||||
|
assert result["success"] is True
|
||||||
|
if skipped_manager._download_task is not None:
|
||||||
|
await asyncio.wait_for(skipped_manager._download_task, timeout=1)
|
||||||
|
assert downloader.calls == []
|
||||||
|
|
||||||
|
# With explicit hashes the previously failed model is retried and cleared
|
||||||
|
result = await manager.start_download(
|
||||||
|
{"model_types": ["lora"], "delay": 0, "model_hashes": [model_hash]}
|
||||||
|
)
|
||||||
|
assert result["success"] is True
|
||||||
|
if manager._download_task is not None:
|
||||||
|
await asyncio.wait_for(manager._download_task, timeout=1)
|
||||||
|
assert downloader.calls == ["https://example.com/valid.png"]
|
||||||
|
assert manager._progress["failed_models"] == set()
|
||||||
|
assert model_hash in manager._progress["processed_models"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_explicit_targets_fill_partial_example_gaps(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
tmp_path,
|
||||||
|
settings_manager,
|
||||||
|
):
|
||||||
|
ws_manager = RecordingWebSocketManager()
|
||||||
|
|
||||||
|
images_root = tmp_path / "examples"
|
||||||
|
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(images_root))
|
||||||
|
|
||||||
|
model_hash = "b" * 64
|
||||||
|
model_path = tmp_path / "model.safetensors"
|
||||||
|
model_path.write_text("data", encoding="utf-8")
|
||||||
|
|
||||||
|
model_metadata = {
|
||||||
|
"sha256": model_hash,
|
||||||
|
"model_name": "Partial Example",
|
||||||
|
"file_path": str(model_path),
|
||||||
|
"file_name": "model.safetensors",
|
||||||
|
"civitai": {
|
||||||
|
"images": [
|
||||||
|
{"url": "https://example.com/first.png"},
|
||||||
|
{"url": "https://example.com/second.png"},
|
||||||
|
]
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
scanner = StubScanner([model_metadata.copy()])
|
||||||
|
_patch_scanner(monkeypatch, scanner)
|
||||||
|
|
||||||
|
# Simulate a partially populated folder: index 0 already downloaded
|
||||||
|
model_dir = images_root / model_hash
|
||||||
|
model_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
(model_dir / "image_0.png").write_bytes(b"existing")
|
||||||
|
|
||||||
|
async def fake_process_local_examples(*_args, **_kwargs):
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def fake_get_updated_model(model_hash_arg, _scanner):
|
||||||
|
return model_metadata
|
||||||
|
|
||||||
|
class DownloaderStub:
|
||||||
|
def __init__(self):
|
||||||
|
self.calls: list[str] = []
|
||||||
|
|
||||||
|
async def download_to_memory(self, url, *_args, **_kwargs):
|
||||||
|
self.calls.append(url)
|
||||||
|
return True, b"\x89PNG\r\n\x1a\n", {"content-type": "image/png"}
|
||||||
|
|
||||||
|
downloader = DownloaderStub()
|
||||||
|
|
||||||
|
async def fake_get_downloader():
|
||||||
|
return downloader
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module.ExampleImagesProcessor,
|
||||||
|
"process_local_examples",
|
||||||
|
staticmethod(fake_process_local_examples),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module.MetadataUpdater,
|
||||||
|
"get_updated_model",
|
||||||
|
staticmethod(fake_get_updated_model),
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(download_module, "get_downloader", fake_get_downloader)
|
||||||
|
|
||||||
|
# Untargeted run treats the populated folder as done
|
||||||
|
untargeted = download_module.DownloadManager(ws_manager=RecordingWebSocketManager())
|
||||||
|
result = await untargeted.start_download({"model_types": ["lora"], "delay": 0})
|
||||||
|
assert result["success"] is True
|
||||||
|
if untargeted._download_task is not None:
|
||||||
|
await asyncio.wait_for(untargeted._download_task, timeout=1)
|
||||||
|
assert downloader.calls == []
|
||||||
|
|
||||||
|
# Explicitly targeted run fills only the missing index, skipping the
|
||||||
|
# existing file without a network request
|
||||||
|
targeted = download_module.DownloadManager(ws_manager=ws_manager)
|
||||||
|
result = await targeted.start_download(
|
||||||
|
{"model_types": ["lora"], "delay": 0, "model_hashes": [model_hash]}
|
||||||
|
)
|
||||||
|
assert result["success"] is True
|
||||||
|
if targeted._download_task is not None:
|
||||||
|
await asyncio.wait_for(targeted._download_task, timeout=1)
|
||||||
|
assert downloader.calls == ["https://example.com/second.png"]
|
||||||
|
assert (model_dir / "image_1.png").exists()
|
||||||
|
assert (model_dir / "image_0.png").read_bytes() == b"existing"
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def settings_manager():
|
def settings_manager():
|
||||||
return get_settings_manager()
|
return get_settings_manager()
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ async def test_start_download_bootstraps_progress_and_task(
|
|||||||
release = asyncio.Event()
|
release = asyncio.Event()
|
||||||
|
|
||||||
async def fake_download(
|
async def fake_download(
|
||||||
self, output_dir, optimize, model_types, delay, library_name, force=False
|
self, output_dir, optimize, model_types, delay, library_name, force=False, model_hashes=None
|
||||||
):
|
):
|
||||||
started.set()
|
started.set()
|
||||||
await release.wait()
|
await release.wait()
|
||||||
@@ -93,6 +93,44 @@ async def test_start_download_bootstraps_progress_and_task(
|
|||||||
assert manager._progress["status"] == "completed"
|
assert manager._progress["status"] == "completed"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_start_download_forwards_model_hashes(
|
||||||
|
monkeypatch: pytest.MonkeyPatch, tmp_path
|
||||||
|
) -> None:
|
||||||
|
settings_manager = get_settings_manager()
|
||||||
|
settings_manager.settings["example_images_path"] = str(tmp_path)
|
||||||
|
settings_manager.settings["libraries"] = {"default": {}}
|
||||||
|
settings_manager.settings["active_library"] = "default"
|
||||||
|
|
||||||
|
manager = download_module.DownloadManager(ws_manager=RecordingWebSocketManager())
|
||||||
|
|
||||||
|
received: Dict[str, Any] = {}
|
||||||
|
|
||||||
|
async def fake_download(
|
||||||
|
self, output_dir, optimize, model_types, delay, library_name, force=False, model_hashes=None
|
||||||
|
):
|
||||||
|
received["model_hashes"] = model_hashes
|
||||||
|
async with self._state_lock:
|
||||||
|
self._is_downloading = False
|
||||||
|
self._download_task = None
|
||||||
|
self._progress["status"] = "completed"
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
download_module.DownloadManager,
|
||||||
|
"_download_all_example_images",
|
||||||
|
fake_download,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await manager.start_download(
|
||||||
|
{"model_types": ["lora"], "delay": 0, "model_hashes": ["abc123", "def456"]}
|
||||||
|
)
|
||||||
|
assert result["success"] is True
|
||||||
|
|
||||||
|
task = manager._download_task
|
||||||
|
assert task is not None
|
||||||
|
await asyncio.wait_for(task, timeout=1)
|
||||||
|
assert received["model_hashes"] == ["abc123", "def456"]
|
||||||
|
|
||||||
|
|
||||||
async def test_pause_and_resume_flow(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None:
|
async def test_pause_and_resume_flow(monkeypatch: pytest.MonkeyPatch, tmp_path) -> None:
|
||||||
settings_manager = get_settings_manager()
|
settings_manager = get_settings_manager()
|
||||||
settings_manager.settings["example_images_path"] = str(tmp_path)
|
settings_manager.settings["example_images_path"] = str(tmp_path)
|
||||||
|
|||||||
@@ -100,6 +100,54 @@ def test_get_file_extension_media_type_hint_low_priority() -> None:
|
|||||||
assert ext == ".mp4"
|
assert ext == ".mp4"
|
||||||
|
|
||||||
|
|
||||||
|
def test_example_image_file_exists_checks_plausible_extensions(tmp_path) -> None:
|
||||||
|
proc = processor_module.ExampleImagesProcessor
|
||||||
|
assert proc._example_image_file_exists(str(tmp_path), 0) is False
|
||||||
|
Path(tmp_path, "image_0.webp").write_bytes(b"x")
|
||||||
|
assert proc._example_image_file_exists(str(tmp_path), 0) is True
|
||||||
|
assert proc._example_image_file_exists(str(tmp_path), 1) is False
|
||||||
|
|
||||||
|
|
||||||
|
def test_example_image_file_exists_video_hint_only_checks_video_extensions(tmp_path) -> None:
|
||||||
|
proc = processor_module.ExampleImagesProcessor
|
||||||
|
Path(tmp_path, "image_2.jpg").write_bytes(b"x")
|
||||||
|
# An existing image file must not satisfy a video-hinted lookup
|
||||||
|
assert proc._example_image_file_exists(str(tmp_path), 2, "video") is False
|
||||||
|
Path(tmp_path, "image_2.mp4").write_bytes(b"x")
|
||||||
|
assert proc._example_image_file_exists(str(tmp_path), 2, "video") is True
|
||||||
|
|
||||||
|
|
||||||
|
async def test_download_model_images_with_tracking_skips_existing_files(tmp_path) -> None:
|
||||||
|
proc = processor_module.ExampleImagesProcessor
|
||||||
|
images = [
|
||||||
|
{"url": "https://image.civitai.com/a/b", "type": "image"},
|
||||||
|
{"url": "https://image.civitai.com/c/d", "type": "image"},
|
||||||
|
]
|
||||||
|
Path(tmp_path, "image_0.jpg").write_bytes(b"existing")
|
||||||
|
|
||||||
|
class RecordingDownloader:
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.calls: list[str] = []
|
||||||
|
|
||||||
|
async def download_to_memory(self, url, use_auth=False, return_headers=False):
|
||||||
|
self.calls.append(url)
|
||||||
|
return True, b"\xff\xd8\xff" + b"data", {}
|
||||||
|
|
||||||
|
downloader = RecordingDownloader()
|
||||||
|
success, is_stale, failed, rate_limited = await proc.download_model_images_with_tracking(
|
||||||
|
"hash", "model", images, str(tmp_path), False, downloader
|
||||||
|
)
|
||||||
|
|
||||||
|
assert success is True
|
||||||
|
assert is_stale is False
|
||||||
|
assert failed == []
|
||||||
|
assert rate_limited == []
|
||||||
|
# Only the missing image is requested; the existing one is skipped without a network call
|
||||||
|
assert len(downloader.calls) == 1
|
||||||
|
assert "c/d" in downloader.calls[0]
|
||||||
|
assert Path(tmp_path, "image_1.jpg").exists()
|
||||||
|
|
||||||
|
|
||||||
class StubScanner:
|
class StubScanner:
|
||||||
def __init__(self, models: list[Dict[str, Any]]) -> None:
|
def __init__(self, models: list[Dict[str, Any]]) -> None:
|
||||||
self._cache = SimpleNamespace(raw_data=models)
|
self._cache = SimpleNamespace(raw_data=models)
|
||||||
|
|||||||
Reference in New Issue
Block a user