Compare commits

...

20 Commits

Author SHA1 Message Date
Will Miao 24f5f7df5d feat(llm): add Gemini as a preset AI provider 2026-08-07 10:27:51 +08:00
Will Miao daf01fb1d6 feat(downloads): show batch download summary with failure details and retry 2026-08-07 10:23:17 +08:00
Will Miao 0f11b6def9 fix(recipes): allow recipes storage path on a different drive (Windows)
os.path.commonpath raises ValueError for paths on different Windows
drives. Treat that as no common root so cross-drive recipes migrations
succeed instead of failing with 'Invalid recipes path change'.
2026-08-06 22:18:24 +08:00
Will Miao 7df83f44b8 feat(SaveImageLM): add add_loras_to_prompt toggle to restore legacy lora syntax line in metadata 2026-08-06 15:58:18 +08:00
Will Miao 169fa7bed6 fix(vue-widgets): resolve pre-existing typecheck errors 2026-08-06 15:33:02 +08:00
Will Miao 027b504fe8 refactor(autocomplete): remove unused custom_words and embeddings modelTypes 2026-08-06 15:28:58 +08:00
Will Miao 186ef4da78 refactor(ui): group example image download actions into a submenu
Move the 'Download Missing' / 'Re-process All' example image actions
under a single 'Download Example Images' submenu item in the single-model
and bulk context menus, matching the existing send-to-workflow submenu
pattern. Shorten the submenu labels and update all locale translations.
2026-08-03 21:18:05 +08:00
pixelpaws dc674098e7 Merge pull request #1050 from willmiao/fix/recipes-bulk-content-rating
fix(recipes): enable bulk content rating for selected recipes
2026-08-03 20:58:24 +08:00
Will Miao 9087b4b07c feat(example-images): add missing-only download path and skip existing files
Split the single-model and bulk context menu actions into 'Download
Missing Example Images' (regular endpoint, skips already-processed
models) and 'Re-process Example Images' (force endpoint, retries
failed models).

- start_download accepts model_hashes so a selected subset can be
  processed with the progress-aware skip logic; explicitly targeted
  models bypass the failed/processed model-level guards so per-image
  gaps are filled
- pre-download existence check in the processor skips network requests
  for image files already on disk across all download paths
- force download retries previously failed models and clears their
  failed status on success
- add i18n keys for the new menu items across all locales
2026-08-03 20:52:46 +08:00
Will Miao 8e45c22d7a fix(recipes): enable bulk content rating for selected recipes 2026-08-03 19:31:58 +08:00
Will Miao 191c4e03cd feat(metadata-overwrite): support wired MODEL input on model field
The model field now accepts either a manual string or a MODEL connection.
When wired, the model name is extracted from the patcher's
cached_patcher_init (registered by core loaders load_checkpoint_guess_config
and load_diffusion_model, preserved through LoRA clones) and converted to a
ComfyUI-style relative name via config model roots.

- model input declared as "STRING,MODEL" with widgetType STRING, so the
  text widget and the dual-type connection slot coexist; non-STRING/MODEL
  links are rejected by frontend and backend type validation
- UNETLoaderLM GGUF branch now registers a custom cached_patcher_init reload
  factory so GGUF models participate in name extraction and ModelPatcher
  deepclone/dynamic machinery
- shared collect_overwrite_params() helper keeps the node and the metadata
  extractor conversion logic in sync; extraction failures are logged instead
  of silently dropping the overwrite
2026-08-03 16:44:03 +08:00
Will Miao ab4154c57d feat(ui): add seeded random sort option to model pages (#1049) 2026-08-03 15:02:49 +08:00
Will Miao 28e93d12ff fix(example-images): use in-place cache sync and bulk pending-check index for large libraries 2026-08-03 12:04:56 +08:00
Will Miao 75e63c758b feat(api): add cursor-based pagination to civitai user-models endpoint 2026-08-03 11:07:06 +08:00
Will Miao 823f71f269 feat(nodes): make Lora Stack Combiner inputs dynamic 2026-08-02 22:04:40 +08:00
Will Miao 042dd4088d fix(nodes): make Lora Stack Combiner inputs optional 2026-08-01 17:14:00 +08:00
willmiao eaa791a9eb docs: auto-update supporters list in README 2026-07-31 13:25:56 +00:00
Will Miao 2228627ff4 chore(release): bump version to v1.2.0 2026-07-31 21:25:38 +08:00
Will Miao 4c647ad9c8 fix(update): throttle nightly update badge to once per day 2026-07-31 21:18:58 +08:00
Will Miao 8ca3e6c33f fix(ui): guard marquee bulk-mode entry against click jitter and stale drag state 2026-07-31 18:40:14 +08:00
79 changed files with 24741 additions and 20566 deletions
+2 -2
View File
File diff suppressed because one or more lines are too long
+313 -291
View File
File diff suppressed because it is too large Load Diff
+2243 -2221
View File
File diff suppressed because it is too large Load Diff
+23 -1
View File
@@ -678,6 +678,7 @@
"deepseek": "DeepSeek", "deepseek": "DeepSeek",
"groq": "Groq", "groq": "Groq",
"openrouter": "OpenRouter", "openrouter": "OpenRouter",
"google": "Gemini",
"opencode-go": "OpenCode Go", "opencode-go": "OpenCode Go",
"custom": "Custom (OpenAI-compatible)" "custom": "Custom (OpenAI-compatible)"
}, },
@@ -714,7 +715,9 @@
"versionsCount": "Local Versions", "versionsCount": "Local Versions",
"versionsCountDesc": "Most versions first", "versionsCountDesc": "Most versions first",
"versionsCountAsc": "Fewest versions first", "versionsCountAsc": "Fewest versions first",
"versionIdDesc": "Newest version first" "versionIdDesc": "Newest version first",
"random": "Random",
"randomAction": "Randomize (shuffle)"
}, },
"refresh": { "refresh": {
"title": "Refresh model list", "title": "Refresh model list",
@@ -771,6 +774,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)",
@@ -806,6 +811,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",
@@ -1575,6 +1582,21 @@
"downloadCsv": "Download CSV", "downloadCsv": "Download CSV",
"columnModelName": "Model Name", "columnModelName": "Model Name",
"columnError": "Error" "columnError": "Error"
},
"downloadBatchSummary": {
"title": "Batch Download Summary",
"statSuccess": "Success",
"statFailed": "Failed",
"statTotal": "Total",
"successMessage": "All {count} models downloaded successfully",
"completedWithErrors": "Completed with errors",
"failed": "Download failed",
"failedItems": "Failed Items ({count})",
"columnName": "Model Name",
"columnError": "Error",
"close": "Close",
"copyReport": "Copy Report",
"retryFailed": "Retry Failed ({count})"
} }
}, },
"modelTags": { "modelTags": {
+2243 -2221
View File
File diff suppressed because it is too large Load Diff
+2243 -2221
View File
File diff suppressed because it is too large Load Diff
+2243 -2221
View File
File diff suppressed because it is too large Load Diff
+2243 -2221
View File
File diff suppressed because it is too large Load Diff
+2243 -2221
View File
File diff suppressed because it is too large Load Diff
+2243 -2221
View File
File diff suppressed because it is too large Load Diff
+2243 -2221
View File
File diff suppressed because it is too large Load Diff
+2243 -2221
View File
File diff suppressed because it is too large Load Diff
+3 -9
View File
@@ -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, {})
+42
View File
@@ -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
+86 -10
View File
@@ -1,26 +1,102 @@
from __future__ import annotations
import inspect
import re
from typing import Any
_STACK_INPUT_PATTERN = re.compile(r"^lora_stack(?:_([ab])|(\d+))$")
def _is_stack_input(name: str) -> bool:
return bool(_STACK_INPUT_PATTERN.match(name))
def _stack_slot_number(name: str) -> int:
"""Numeric slot used to order stack inputs; legacy a/b map to 1/2."""
match = _STACK_INPUT_PATTERN.match(name)
if not match:
return -1
letter, digits = match.group(1), match.group(2)
if digits is not None:
return int(digits)
return 1 if letter == "a" else 2
class _LoraStackOptionalInputs:
"""Lookup that preserves explicit optional inputs and dynamic lora_stack slots."""
def __init__(self, explicit_inputs: dict[str, tuple[str, dict[str, Any]]]) -> None:
self._explicit_inputs = explicit_inputs
def __contains__(self, item: object) -> bool:
if not isinstance(item, str):
return False
return item in self._explicit_inputs or _is_stack_input(item)
def __getitem__(self, key: str) -> tuple[str, dict[str, Any]]:
if key in self._explicit_inputs:
return self._explicit_inputs[key]
if _is_stack_input(key):
return (
"LORA_STACK",
{
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
},
)
raise KeyError(key)
class LoraStackCombinerLM: class LoraStackCombinerLM:
NAME = "Lora Stack Combiner (LoraManager)" NAME = "Lora Stack Combiner (LoraManager)"
CATEGORY = "Lora Manager/stackers" CATEGORY = "Lora Manager/stackers"
DESCRIPTION = (
"Combines multiple LoRA stacks into a single stack. "
"Supports dynamic inputs: connect a stack to add more inputs."
)
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
optional_inputs: dict[str, tuple[str, dict[str, Any]]] = {
"lora_stack1": (
"LORA_STACK",
{
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
},
),
"lora_stack2": (
"LORA_STACK",
{
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
},
),
}
stack = inspect.stack()
if len(stack) > 2 and stack[2].function == "get_input_info":
optional_inputs = _LoraStackOptionalInputs(optional_inputs) # type: ignore[assignment]
return { return {
"required": { "required": {},
"lora_stack_a": ("LORA_STACK",), "optional": optional_inputs,
"lora_stack_b": ("LORA_STACK",),
},
} }
RETURN_TYPES = ("LORA_STACK",) RETURN_TYPES = ("LORA_STACK",)
RETURN_NAMES = ("LORA_STACK",) RETURN_NAMES = ("LORA_STACK",)
FUNCTION = "combine_stacks" FUNCTION = "combine_stacks"
def combine_stacks(self, lora_stack_a, lora_stack_b): def combine_stacks(self, lora_stack1=None, lora_stack2=None, **kwargs):
combined_stack = [] stacks = {
"lora_stack1": lora_stack1,
"lora_stack2": lora_stack2,
}
for key, value in kwargs.items():
if _is_stack_input(key) and value is not None:
stacks[key] = value
if lora_stack_a: combined_stack = []
combined_stack.extend(lora_stack_a) for key in sorted(stacks, key=_stack_slot_number):
if lora_stack_b: stack = stacks[key]
combined_stack.extend(lora_stack_b) if stack:
combined_stack.extend(stack)
return (combined_stack,) return (combined_stack,)
+14 -15
View File
@@ -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,)
+16 -3
View File
@@ -252,6 +252,13 @@ class SaveImageLM:
"tooltip": "When enabled, embeds generation parameters into the saved image metadata. Disable to skip writing generation metadata.", "tooltip": "When enabled, embeds generation parameters into the saved image metadata. Disable to skip writing generation metadata.",
}, },
), ),
"add_loras_to_prompt": (
"BOOLEAN",
{
"default": False,
"tooltip": "When enabled, appends the LoRA syntax line (e.g. <lora:name:strength>) after the positive prompt in the saved metadata.",
},
),
"add_counter_to_filename": ( "add_counter_to_filename": (
"BOOLEAN", "BOOLEAN",
{ {
@@ -348,7 +355,7 @@ class SaveImageLM:
type_lower = model_type.lower() if model_type else "other" type_lower = model_type.lower() if model_type else "other"
return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}" return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}"
def format_metadata(self, metadata_dict: dict) -> str: def format_metadata(self, metadata_dict: dict, add_loras_to_prompt: bool = False) -> str:
"""Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources.""" """Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources."""
if not metadata_dict: return "" if not metadata_dict: return ""
@@ -458,7 +465,10 @@ class SaveImageLM:
scheduler_name = scheduler_mapping.get(scheduler, scheduler) if scheduler else None scheduler_name = scheduler_mapping.get(scheduler, scheduler) if scheduler else None
# Build output lines # Build output lines
lines = [prompt] if prompt else [""] prompt_line = prompt if prompt else ""
if add_loras_to_prompt and loras_text:
prompt_line = f"{prompt_line}\n{loras_text}" if prompt_line else loras_text
lines = [prompt_line] if prompt_line else [""]
if negative_prompt: if negative_prompt:
lines.append(f"Negative prompt: {negative_prompt}") lines.append(f"Negative prompt: {negative_prompt}")
@@ -793,6 +803,7 @@ class SaveImageLM:
save_with_metadata=True, save_with_metadata=True,
add_counter_to_filename=True, add_counter_to_filename=True,
save_as_recipe=False, save_as_recipe=False,
add_loras_to_prompt=False,
): ):
"""Save images with metadata""" """Save images with metadata"""
results = [] results = []
@@ -801,7 +812,7 @@ class SaveImageLM:
raw_metadata = get_metadata() raw_metadata = get_metadata()
metadata_dict = MetadataProcessor.to_dict(raw_metadata, id) metadata_dict = MetadataProcessor.to_dict(raw_metadata, id)
metadata = self.format_metadata(metadata_dict) metadata = self.format_metadata(metadata_dict, add_loras_to_prompt)
# Process filename_prefix with pattern substitution # Process filename_prefix with pattern substitution
filename_prefix = self.format_filename(filename_prefix, metadata_dict) filename_prefix = self.format_filename(filename_prefix, metadata_dict)
@@ -943,6 +954,7 @@ class SaveImageLM:
save_with_metadata=True, save_with_metadata=True,
add_counter_to_filename=True, add_counter_to_filename=True,
save_as_recipe=False, save_as_recipe=False,
add_loras_to_prompt=False,
): ):
"""Process and save image with metadata""" """Process and save image with metadata"""
# Make sure the output directory exists # Make sure the output directory exists
@@ -974,6 +986,7 @@ class SaveImageLM:
save_with_metadata, save_with_metadata,
add_counter_to_filename, add_counter_to_filename,
save_as_recipe, save_as_recipe,
add_loras_to_prompt,
) )
return { return {
+21
View File
@@ -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:
+37 -3
View File
@@ -2590,6 +2590,8 @@ class ModelLibraryHandler:
status=400, status=400,
) )
cursor = request.query.get("cursor")
metadata_provider = await self._metadata_provider_factory() metadata_provider = await self._metadata_provider_factory()
if not metadata_provider: if not metadata_provider:
return web.json_response( return web.json_response(
@@ -2598,7 +2600,7 @@ class ModelLibraryHandler:
) )
try: try:
models = await metadata_provider.get_user_models(username) result = await metadata_provider.get_user_models(username, cursor)
except NotImplementedError: except NotImplementedError:
return web.json_response( return web.json_response(
{ {
@@ -2608,14 +2610,35 @@ class ModelLibraryHandler:
status=501, status=501,
) )
if models is None: if result is None:
return web.json_response( return web.json_response(
{"success": False, "error": "Failed to fetch user models"}, {"success": False, "error": "Failed to fetch user models"},
status=502, status=502,
) )
if isinstance(result, dict):
models = result.get("items")
next_cursor = result.get("nextCursor")
else:
# Defensive: tolerate providers that still return a raw list
models = result
next_cursor = None
if not isinstance(models, list): if not isinstance(models, list):
models = [] models = []
if next_cursor is not None and not isinstance(next_cursor, str):
next_cursor = str(next_cursor)
estimated_total = None
if cursor is None:
get_count = getattr(metadata_provider, "get_creator_model_count", None)
if get_count is not None:
try:
estimated_total = await get_count(username)
except Exception: # best-effort only
estimated_total = None
if not isinstance(estimated_total, int):
estimated_total = None
lora_scanner = await self._service_registry.get_lora_scanner() lora_scanner = await self._service_registry.get_lora_scanner()
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner() checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
@@ -2635,6 +2658,7 @@ class ModelLibraryHandler:
versions: list[dict] = [] versions: list[dict] = []
history_service = await self._get_download_history_service() history_service = await self._get_download_history_service()
model_ids: list[int] = [] model_ids: list[int] = []
model_count = 0
for model in models: for model in models:
try: try:
model_ids.append(int(model.get("id"))) model_ids.append(int(model.get("id")))
@@ -2668,6 +2692,8 @@ class ModelLibraryHandler:
if model_type not in normalized_allowed_types: if model_type not in normalized_allowed_types:
continue continue
model_count += 1
scanner = type_scanner_map.get(model_type) scanner = type_scanner_map.get(model_type)
if scanner is None: if scanner is None:
return web.json_response( return web.json_response(
@@ -2733,7 +2759,15 @@ class ModelLibraryHandler:
) )
return web.json_response( return web.json_response(
{"success": True, "username": username, "versions": versions} {
"success": True,
"username": username,
"versions": versions,
"modelCount": model_count,
"nextCursor": next_cursor,
"hasMore": next_cursor is not None,
"estimatedTotal": estimated_total,
}
) )
except Exception as exc: # pragma: no cover - defensive logging except Exception as exc: # pragma: no cover - defensive logging
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True) logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
+7
View File
@@ -1,6 +1,7 @@
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
import asyncio import asyncio
import re import re
import random
from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING
import logging import logging
import os import os
@@ -390,6 +391,12 @@ class BaseModelService(ABC):
(item.get("model_name") or item.get("file_name") or "").lower(), (item.get("model_name") or item.get("file_name") or "").lower(),
item.get("file_path", "").lower(), item.get("file_path", "").lower(),
) )
elif key_name == "random":
# Seeded random shuffle: same seed -> same order (stable pagination)
rng = random.Random(sort_params.seed or "random")
result = list(data)
rng.shuffle(result)
return result
elif key_name == "size": elif key_name == "size":
key_fn = lambda item: ( key_fn = lambda item: (
int(item.get("size", 0) or 0), int(item.get("size", 0) or 0),
+88 -5
View File
@@ -2,6 +2,7 @@ import asyncio
import copy import copy
import logging import logging
import os import os
import time
from collections import OrderedDict from collections import OrderedDict
from typing import Any, Optional, Dict, Tuple, List, Sequence from typing import Any, Optional, Dict, Tuple, List, Sequence
from .connectivity_guard import ( from .connectivity_guard import (
@@ -19,6 +20,12 @@ from ..utils.civitai_utils import resolve_license_payload
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Best-effort cache for creator model counts, keyed by lowercase username.
# Values are (monotonic timestamp, count or None); None results are cached
# too so repeated failures don't hammer the API.
_CREATOR_COUNT_CACHE_TTL_SECONDS = 600
_creator_model_count_cache: Dict[str, Tuple[float, Optional[int]]] = {}
class CivitaiClient: class CivitaiClient:
_instance = None _instance = None
@@ -743,17 +750,34 @@ class CivitaiClient:
return all_versions if all_versions else None return all_versions if all_versions else None
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(
"""Fetch all models for a specific Civitai user.""" self, username: str, cursor: Optional[str] = None
) -> Optional[Dict[str, Any]]:
"""Fetch one page (up to 100 models) for a specific Civitai user.
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
or None on failure. Pass ``cursor`` (from a previous response's
``nextCursor``) to fetch subsequent pages.
"""
if not username: if not username:
return None return None
params: Dict[str, Any] = {
"username": username,
"nsfw": "true",
"limit": 100,
"sort": "Newest",
"period": "AllTime",
}
if cursor:
params["cursor"] = cursor
try: try:
success, result = await self._make_request( success, result = await self._make_request(
"GET", "GET",
f"{self.base_url}/models", f"{self.base_url}/models",
use_auth=True, use_auth=True,
params={"username": username, "nsfw": "true"}, params=params,
) )
if not success: if not success:
@@ -765,7 +789,7 @@ class CivitaiClient:
items = result.get("items") if isinstance(result, dict) else None items = result.get("items") if isinstance(result, dict) else None
if not isinstance(items, list): if not isinstance(items, list):
return [] items = []
for model in items: for model in items:
versions = model.get("modelVersions") versions = model.get("modelVersions")
@@ -774,9 +798,68 @@ class CivitaiClient:
for version in versions: for version in versions:
self._remove_comfy_metadata(version) self._remove_comfy_metadata(version)
return items next_cursor: Optional[str] = None
metadata = result.get("metadata") if isinstance(result, dict) else None
if isinstance(metadata, dict):
raw_cursor = metadata.get("nextCursor")
if raw_cursor is not None:
next_cursor = str(raw_cursor)
return {"items": items, "nextCursor": next_cursor}
except RateLimitError: except RateLimitError:
raise raise
except Exception as exc: # pragma: no cover - defensive logging except Exception as exc: # pragma: no cover - defensive logging
logger.error("Error fetching models for %s: %s", username, exc) logger.error("Error fetching models for %s: %s", username, exc)
return None return None
async def get_creator_model_count(self, username: str) -> Optional[int]:
"""Best-effort lookup of a creator's published model count.
Uses the ``/creators`` endpoint (a contains-match query), picking the
entry whose username matches exactly (case-insensitive). Returns None
on any failure; never raises. Results (including None) are cached
for ``_CREATOR_COUNT_CACHE_TTL_SECONDS``.
"""
if not username:
return None
cache_key = username.lower()
cached = _creator_model_count_cache.get(cache_key)
if cached is not None:
cached_at, cached_count = cached
if time.monotonic() - cached_at < _CREATOR_COUNT_CACHE_TTL_SECONDS:
return cached_count
count: Optional[int] = None
try:
success, result = await self._make_request(
"GET",
f"{self.base_url}/creators",
use_auth=True,
params={"query": username, "limit": 10},
)
if success and isinstance(result, dict):
creators = result.get("items")
if isinstance(creators, list):
for creator in creators:
if not isinstance(creator, dict):
continue
creator_name = creator.get("username")
if not isinstance(creator_name, str):
continue
if creator_name.lower() != cache_key:
continue
model_count = creator.get("modelCount")
if isinstance(model_count, (int, float)) and not isinstance(
model_count, bool
):
count = int(model_count)
break
except Exception as exc: # best-effort only, never propagate
logger.debug(
"Failed to fetch creator model count for %s: %s", username, exc
)
_creator_model_count_cache[cache_key] = (time.monotonic(), count)
return count
+5
View File
@@ -201,6 +201,11 @@ PROVIDER_PRESETS: Dict[str, Dict[str, Any]] = {
"api_base": "https://openrouter.ai/api/v1", "api_base": "https://openrouter.ai/api/v1",
"requires_key": True, "requires_key": True,
}, },
"google": {
"name": "Gemini",
"api_base": "https://generativelanguage.googleapis.com/v1beta/openai",
"requires_key": True,
},
"opencode-go": { "opencode-go": {
"name": "OpenCode Go", "name": "OpenCode Go",
"api_base": "https://opencode.ai/zen/go/v1", "api_base": "https://opencode.ai/zen/go/v1",
+21 -12
View File
@@ -1,6 +1,7 @@
import asyncio import asyncio
import time import time
import logging import logging
import random
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from typing import Any, Dict, List, Optional, Tuple from typing import Any, Dict, List, Optional, Tuple
@@ -38,8 +39,8 @@ class ModelCache:
def __post_init__(self): def __post_init__(self):
self._lock = asyncio.Lock() self._lock = asyncio.Lock()
# Cache for last sort: (sort_key, order) -> sorted list # Cache for last sort: (sort_key, order, seed) -> sorted list
self._last_sort: Tuple[str, str] = (None, None) self._last_sort: Tuple[Optional[str], str, Optional[str]] = (None, "asc", None)
self._last_sorted_data: List[Dict] = [] self._last_sorted_data: List[Dict] = []
self._normalize_raw_data() self._normalize_raw_data()
self.name_display_mode = self._normalize_display_mode(self.name_display_mode) self.name_display_mode = self._normalize_display_mode(self.name_display_mode)
@@ -203,9 +204,9 @@ class ModelCache:
async def resort(self): async def resort(self):
"""Resort cached data according to last sort mode if set""" """Resort cached data according to last sort mode if set"""
async with self._lock: async with self._lock:
if self._last_sort != (None, None): if self._last_sort[0] is not None:
sort_key, order = self._last_sort sort_key, order, seed = self._last_sort
sorted_data = self._sort_data(self.raw_data, sort_key, order) sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
self._last_sorted_data = sorted_data self._last_sorted_data = sorted_data
# Update folder list # Update folder list
# else: do nothing # else: do nothing
@@ -218,7 +219,7 @@ class ModelCache:
self.folders = sorted(list(all_folders), key=lambda x: x.lower()) self.folders = sorted(list(all_folders), key=lambda x: x.lower())
self.rebuild_version_index() self.rebuild_version_index()
def _sort_data(self, data: List[Dict], sort_key: str, order: str) -> List[Dict]: def _sort_data(self, data: List[Dict], sort_key: str, order: str, seed: Optional[str] = None) -> List[Dict]:
"""Sort data by sort_key and order""" """Sort data by sort_key and order"""
start_time = time.perf_counter() start_time = time.perf_counter()
reverse = (order == 'desc') reverse = (order == 'desc')
@@ -265,6 +266,13 @@ class ModelCache:
), ),
reverse=reverse reverse=reverse
) )
elif sort_key == 'random':
# Random shuffle seeded for stable pagination: the same seed
# always yields the same order, so successive page requests
# stay consistent while browsing.
rng = random.Random(seed or 'random')
result = list(data)
rng.shuffle(result)
elif sort_key == 'versions_count': elif sort_key == 'versions_count':
# Pre-dedup sort: fall back to name sort. # Pre-dedup sort: fall back to name sort.
# Actual re-sort by version_count happens in get_paginated_data after dedup. # Actual re-sort by version_count happens in get_paginated_data after dedup.
@@ -285,15 +293,16 @@ class ModelCache:
logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration) logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration)
return result return result
async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc') -> List[Dict]: async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc', seed: Optional[str] = None) -> List[Dict]:
"""Get sorted data by sort_key and order, using cache if possible""" """Get sorted data by sort_key and order, using cache if possible"""
async with self._lock: async with self._lock:
if (sort_key, order) == self._last_sort: cache_key = (sort_key, order, seed)
if cache_key == self._last_sort:
return self._last_sorted_data return self._last_sorted_data
start_time = time.perf_counter() start_time = time.perf_counter()
sorted_data = self._sort_data(self.raw_data, sort_key, order) sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
self._last_sort = (sort_key, order) self._last_sort = cache_key
self._last_sorted_data = sorted_data self._last_sorted_data = sorted_data
duration = time.perf_counter() - start_time duration = time.perf_counter() - start_time
@@ -313,8 +322,8 @@ class ModelCache:
self.name_display_mode = normalized self.name_display_mode = normalized
if self._last_sort[0] == 'name': if self._last_sort[0] == 'name':
sort_key, order = self._last_sort sort_key, order, seed = self._last_sort
self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order) 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: async def update_preview_url(self, file_path: str, preview_url: str, preview_nsfw_level: int) -> bool:
"""Update preview_url for a specific model in all cached data """Update preview_url for a specific model in all cached data
+50 -11
View File
@@ -143,10 +143,18 @@ class ModelMetadataProvider(ABC):
pass pass
@abstractmethod @abstractmethod
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
"""Fetch models owned by the specified user""" """Fetch one page of models owned by the specified user.
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
or None when unsupported/failed. ``cursor`` continues a previous page.
"""
pass pass
async def get_creator_model_count(self, username: str) -> Optional[int]:
"""Published model count for the user; None when unsupported."""
return None
class CivitaiModelMetadataProvider(ModelMetadataProvider): class CivitaiModelMetadataProvider(ModelMetadataProvider):
"""Provider that uses Civitai API for metadata""" """Provider that uses Civitai API for metadata"""
@@ -175,8 +183,11 @@ class CivitaiModelMetadataProvider(ModelMetadataProvider):
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], Optional[str]]:
return await self.client.get_model_version_info(version_id) return await self.client.get_model_version_info(version_id)
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
return await self.client.get_user_models(username) return await self.client.get_user_models(username, cursor)
async def get_creator_model_count(self, username: str) -> Optional[int]:
return await self.client.get_creator_model_count(username)
class CivArchiveModelMetadataProvider(ModelMetadataProvider): class CivArchiveModelMetadataProvider(ModelMetadataProvider):
"""Provider that uses CivArchive API for metadata""" """Provider that uses CivArchive API for metadata"""
@@ -196,7 +207,7 @@ class CivArchiveModelMetadataProvider(ModelMetadataProvider):
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], Optional[str]]:
return await self.client.get_model_version_info(version_id) return await self.client.get_model_version_info(version_id)
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
"""Not supported by CivArchive provider""" """Not supported by CivArchive provider"""
return None return None
@@ -347,7 +358,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
version_data = await self._get_version_with_model_data(db, model_id, version_id) version_data = await self._get_version_with_model_data(db, model_id, version_id)
return version_data, None return version_data, None
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
"""Listing models by username is not supported for archive database""" """Listing models by username is not supported for archive database"""
return None return None
@@ -602,13 +613,14 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue continue
return None return None
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
for provider, label in self._iter_providers(): for provider, label in self._iter_providers():
try: try:
result = await self._call_with_rate_limit( result = await self._call_with_rate_limit(
label, label,
provider.get_user_models, provider.get_user_models,
username, username,
cursor=cursor,
) )
if result is not None: if result is not None:
return result return result
@@ -624,6 +636,19 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue continue
return None return None
async def get_creator_model_count(self, username: str) -> Optional[int]:
for provider, label in self._iter_providers():
try:
result = await provider.get_creator_model_count(username)
if result is not None:
return result
except Exception as e:
logger.debug(
"Provider %s failed for get_creator_model_count: %s", label, e
)
continue
return None
def _iter_providers(self): def _iter_providers(self):
return zip(self.providers, self._provider_labels) return zip(self.providers, self._provider_labels)
@@ -704,13 +729,17 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
version_id, version_id,
) )
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
return await self._rate_limit_helper.run( return await self._rate_limit_helper.run(
self._label, self._label,
self._provider.get_user_models, self._provider.get_user_models,
username, username,
cursor=cursor,
) )
async def get_creator_model_count(self, username: str) -> Optional[int]:
return await self._provider.get_creator_model_count(username)
class ModelMetadataProviderManager: class ModelMetadataProviderManager:
"""Manager for selecting and using model metadata providers""" """Manager for selecting and using model metadata providers"""
@@ -776,10 +805,20 @@ class ModelMetadataProviderManager:
except NotImplementedError: except NotImplementedError:
return None return None
async def get_user_models(self, username: str, provider_name: str = None) -> Optional[List[Dict]]: async def get_user_models(
"""Fetch models owned by the specified user""" self,
username: str,
provider_name: str = None,
cursor: Optional[str] = None,
) -> Optional[Dict]:
"""Fetch one page of models owned by the specified user"""
provider = self._get_provider(provider_name) provider = self._get_provider(provider_name)
return await provider.get_user_models(username) return await provider.get_user_models(username, cursor)
async def get_creator_model_count(self, username: str, provider_name: 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: str = None) -> ModelMetadataProvider:
"""Get provider by name or default provider""" """Get provider by name or default provider"""
+11 -3
View File
@@ -85,6 +85,7 @@ class SortParams:
key: str key: str
order: str order: str
seed: Optional[str] = None
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -116,7 +117,7 @@ class ModelCacheRepository:
async def fetch_sorted(self, params: SortParams) -> List[Dict[str, Any]]: async def fetch_sorted(self, params: SortParams) -> List[Dict[str, Any]]:
"""Fetch cached data pre-sorted according to ``params``.""" """Fetch cached data pre-sorted according to ``params``."""
cache = await self.get_cache() cache = await self.get_cache()
return await cache.get_sorted_data(params.key, params.order) return await cache.get_sorted_data(params.key, params.order, params.seed)
@staticmethod @staticmethod
def parse_sort(sort_by: str) -> SortParams: def parse_sort(sort_by: str) -> SortParams:
@@ -132,10 +133,17 @@ class ModelCacheRepository:
sort_key = sort_by.strip().lower() or "name" sort_key = sort_by.strip().lower() or "name"
order = "asc" order = "asc"
if order not in ("asc", "desc"): seed = None
if sort_key == "random":
# Random sort: the portion after ':' is the shuffle seed.
# A stable seed keeps paginated requests consistent; order is
# meaningless for a random shuffle.
seed = order if order and order not in ("asc", "desc") else None
order = "asc"
elif order not in ("asc", "desc"):
order = "asc" order = "asc"
return SortParams(key=sort_key, order=order) return SortParams(key=sort_key, order=order, seed=seed)
class ModelFilterSet: class ModelFilterSet:
+1 -1
View File
@@ -1752,7 +1752,7 @@ class ModelScanner:
# ---- Conditional resort (only when sort-key fields changed) ---- # ---- Conditional resort (only when sort-key fields changed) ----
need_resort = False need_resort = False
_last = cache._last_sort _last = cache._last_sort
sort_key: Optional[str] = _last[0] if _last != (None, None) else None sort_key: Optional[str] = _last[0] if _last[0] is not None else None
if sort_key == "name": if sort_key == "name":
if ( if (
old_model_name != desired_entry.get("model_name", "") old_model_name != desired_entry.get("model_name", "")
+5 -3
View File
@@ -1473,10 +1473,12 @@ class SettingsManager:
try: try:
common_root = os.path.commonpath([source, target]) common_root = os.path.commonpath([source, target])
except ValueError as exc: except ValueError:
raise ValueError("Invalid recipes path change") from exc # Windows: paths on different drives share no common root.
# A cross-drive move is valid, so treat it as no common root.
common_root = None
if common_root == source: if common_root is not None and common_root == source:
raise ValueError("Recipes path cannot be moved into a nested directory") raise ValueError("Recipes path cannot be moved into a nested directory")
planned_recipe_updates: Dict[str, Dict[str, Any]] = {} planned_recipe_updates: Dict[str, Dict[str, Any]] = {}
+132 -33
View File
@@ -14,11 +14,16 @@ from ..services.service_registry import ServiceRegistry
from ..utils.example_images_paths import ( from ..utils.example_images_paths import (
ExampleImagePathResolver, ExampleImagePathResolver,
ensure_library_root_exists, ensure_library_root_exists,
get_example_images_root,
is_hash_folder,
uses_library_scoped_folders, uses_library_scoped_folders,
) )
from ..utils.metadata_manager import MetadataManager from ..utils.metadata_manager import MetadataManager
from .example_images_processor import ExampleImagesProcessor from .example_images_processor import ExampleImagesProcessor
from .example_images_metadata import MetadataUpdater from .example_images_metadata import (
MetadataUpdater,
update_cache_from_metadata,
)
from ..services.downloader import get_downloader from ..services.downloader import get_downloader
from ..services.settings_manager import get_settings_manager from ..services.settings_manager import get_settings_manager
@@ -87,6 +92,13 @@ class _DownloadProgress(dict):
return snapshot return snapshot
# When fewer candidates than this remain in check_pending_models, probe each
# model folder directly (preserving legacy-folder migration semantics). Above
# it, build a folder index with a single directory scan so libraries with
# 100k+ models do not pay one syscall per candidate.
_BULK_LOOKUP_THRESHOLD = 1000
def _model_directory_has_files(path: str) -> bool: def _model_directory_has_files(path: str) -> bool:
"""Return True when the provided directory exists and contains entries.""" """Return True when the provided directory exists and contains entries."""
@@ -103,6 +115,36 @@ def _model_directory_has_files(path: str) -> bool:
return False return False
def _build_example_folder_index(output_dir: str) -> dict[str, bool]:
"""Build a ``{hash: has_files}`` index for a library's example-image folders.
A single directory scan over the library root replaces ``O(candidates)``
per-folder ``os.scandir`` calls, which is required for libraries with
100k+ models. Each hash folder is classified by whether it contains any
entries, matching the semantics of ``_model_directory_has_files``.
"""
index: dict[str, bool] = {}
if not output_dir or not os.path.isdir(output_dir):
return index
try:
with os.scandir(output_dir) as entries:
for entry in entries:
name = entry.name
if not entry.is_dir() or not is_hash_folder(name):
continue
try:
with os.scandir(entry.path) as subentries:
index[name.lower()] = any(subentries)
except OSError:
index[name.lower()] = False
except OSError:
pass
return index
class DownloadManager: class DownloadManager:
"""Manages downloading example images for models.""" """Manages downloading example images for models."""
@@ -130,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()
@@ -199,6 +242,7 @@ class DownloadManager:
delay, delay,
active_library, active_library,
force, force,
model_hashes,
) )
) )
@@ -410,14 +454,49 @@ class DownloadManager:
# Calculate pending count: check which models actually need processing. # Calculate pending count: check which models actually need processing.
# A model is pending if it has a hash, is not already processed or known-failed, # A model is pending if it has a hash, is not already processed or known-failed,
# and its folder doesn't exist or is empty. # and its folder doesn't exist or is empty.
pending_hashes = set() candidate_hashes = [
for model_hash, model_name in all_models_with_hash: model_hash
if model_hash not in processed_models and model_hash not in failed_models: for model_hash, _ in all_models_with_hash
if model_hash not in processed_models
and model_hash not in failed_models
]
pending_hashes: set[str] = set()
# For small candidate counts the existing per-folder check is fine
# and handles legacy folder migration.
# For large libraries, scan the library root once and do set lookups.
if len(candidate_hashes) <= _BULK_LOOKUP_THRESHOLD or not output_dir:
for model_hash in candidate_hashes:
model_dir = ExampleImagePathResolver.get_model_folder( model_dir = ExampleImagePathResolver.get_model_folder(
model_hash, active_library model_hash, active_library
) )
if not _model_directory_has_files(model_dir): if not _model_directory_has_files(model_dir):
pending_hashes.add(model_hash) pending_hashes.add(model_hash)
else:
folder_index = await asyncio.get_event_loop().run_in_executor(
None, _build_example_folder_index, output_dir
)
# In multi-library mode, folders that have not been consolidated
# into the library root yet (startup migration skipped, failed
# move, or created at the legacy path afterwards) still live at
# the legacy root/<hash> location. Only scan that root when at
# least one candidate is missing from the library-root index, so
# the fully-consolidated case does not pay an extra directory
# pass on every call.
if uses_library_scoped_folders() and any(
not folder_index.get(model_hash, False)
for model_hash in candidate_hashes
):
legacy_root = get_example_images_root()
if legacy_root and legacy_root != output_dir:
legacy_index = await asyncio.get_event_loop().run_in_executor(
None, _build_example_folder_index, legacy_root
)
for hash_key, has_files in legacy_index.items():
folder_index.setdefault(hash_key, has_files)
for model_hash in candidate_hashes:
if not folder_index.get(model_hash, False):
pending_hashes.add(model_hash)
pending_count = len(pending_hashes) pending_count = len(pending_hashes)
@@ -500,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()
@@ -529,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")
@@ -552,6 +644,7 @@ class DownloadManager:
downloader, downloader,
library_name, library_name,
force, force,
explicit_targets,
) )
# Update progress # Update progress
@@ -648,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."""
@@ -670,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
@@ -680,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)",
@@ -807,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"
@@ -827,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"
@@ -1343,8 +1442,8 @@ class DownloadManager:
await MetadataManager.save_metadata(file_path, model_copy) await MetadataManager.save_metadata(file_path, model_copy)
try: try:
await scanner.update_single_model_cache( await update_cache_from_metadata(
file_path, file_path, model_data scanner, file_path, model_copy
) )
except AttributeError: except AttributeError:
logger.debug( logger.debug(
+40 -8
View File
@@ -1,3 +1,4 @@
import inspect
import logging import logging
import os import os
import re import re
@@ -28,6 +29,31 @@ if TYPE_CHECKING: # pragma: no cover - import for type checkers only
from ..services.settings_manager import SettingsManager from ..services.settings_manager import SettingsManager
async def update_cache_from_metadata(
scanner: Any, file_path: str, metadata: Dict[str, Any]
) -> bool:
"""Update the scanner cache from a metadata dict using the in-place sync path.
``sync_cache_from_metadata`` patches the existing cache entry incrementally
(tag/hash/version indexes, targeted single-row SQL update) and only resorts
when a sort-key field changed. This avoids the ``O(n)`` full-list resort and
full cache rewrite that ``update_single_model_cache`` performs on every call,
which is critical for libraries with 100k+ models.
Falls back to the legacy full update when the scanner does not expose an
async ``sync_cache_from_metadata`` method.
Returns:
``True`` if the cache entry was updated, ``False`` otherwise.
"""
sync_method = getattr(scanner, "sync_cache_from_metadata", None)
if inspect.iscoroutinefunction(sync_method):
return await sync_method(file_path, metadata)
return await scanner.update_single_model_cache(file_path, file_path, metadata)
def _build_metadata_sync_service(settings_manager: "SettingsManager") -> MetadataSyncService: def _build_metadata_sync_service(settings_manager: "SettingsManager") -> MetadataSyncService:
"""Construct a metadata sync service bound to the provided settings.""" """Construct a metadata sync service bound to the provided settings."""
@@ -103,7 +129,7 @@ class MetadataUpdater:
progress['refreshed_models'].add(model_hash) progress['refreshed_models'].add(model_hash)
async def update_cache_func(old_path, new_path, metadata): async def update_cache_func(old_path, new_path, metadata):
return await scanner.update_single_model_cache(old_path, new_path, metadata) return await update_cache_from_metadata(scanner, new_path, metadata)
await MetadataManager.hydrate_model_data(model_data) await MetadataManager.hydrate_model_data(model_data)
success, error = await _get_metadata_sync_service().fetch_and_update_model( success, error = await _get_metadata_sync_service().fetch_and_update_model(
@@ -234,6 +260,7 @@ class MetadataUpdater:
# Save metadata to .metadata.json file # Save metadata to .metadata.json file
file_path = model.get('file_path') file_path = model.get('file_path')
model_copy: Optional[Dict[str, Any]] = None
try: try:
model_copy = model.copy() model_copy = model.copy()
model_copy.pop('folder', None) model_copy.pop('folder', None)
@@ -242,13 +269,17 @@ class MetadataUpdater:
except Exception as e: except Exception as e:
logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}") logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}")
# Save updated metadata to scanner cache # Save updated metadata to scanner cache. sync_cache_from_metadata
success = await scanner.update_single_model_cache(file_path, file_path, model) # returns False both for "already in sync" and for actual failures,
if success: # so the cache sync result is deliberately not treated as an error;
# the return value reflects whether the metadata was persisted.
if file_path and model_copy is not None:
await update_cache_from_metadata(scanner, file_path, model_copy)
logger.info(f"Successfully updated metadata for {model.get('model_name')} with {len(images)} local examples") logger.info(f"Successfully updated metadata for {model.get('model_name')} with {len(images)} local examples")
return True return True
else:
logger.warning(f"Failed to update metadata for {model.get('model_name')}") logger.warning(f"Failed to update metadata for {model.get('model_name')}")
return False
return False return False
except Exception as e: except Exception as e:
@@ -336,6 +367,7 @@ class MetadataUpdater:
# Save metadata to .metadata.json file # Save metadata to .metadata.json file
file_path = model_data.get('file_path') file_path = model_data.get('file_path')
model_copy: Optional[Dict[str, Any]] = None
if file_path: if file_path:
try: try:
model_copy = model_data.copy() model_copy = model_data.copy()
@@ -346,8 +378,8 @@ class MetadataUpdater:
logger.error(f"Failed to save metadata: {str(e)}") logger.error(f"Failed to save metadata: {str(e)}")
# Save updated metadata to scanner cache # Save updated metadata to scanner cache
if file_path: if file_path and model_copy is not None:
await scanner.update_single_model_cache(file_path, file_path, model_data) await update_cache_from_metadata(scanner, file_path, model_copy)
# Get regular images array (might be None) # Get regular images array (might be None)
regular_images = civitai_data.get('images', []) regular_images = civitai_data.get('images', [])
+2 -1
View File
@@ -15,6 +15,7 @@ from ..utils.example_images_paths import (
) )
from ..utils.metadata_manager import MetadataManager from ..utils.metadata_manager import MetadataManager
from ..utils.example_images_processor import ExampleImagesProcessor from ..utils.example_images_processor import ExampleImagesProcessor
from ..utils.example_images_metadata import update_cache_from_metadata
from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -421,7 +422,7 @@ class ExampleImagesMigration:
await MetadataManager.save_metadata(file_path, model_copy) await MetadataManager.save_metadata(file_path, model_copy)
# Update scanner cache # Update scanner cache
await scanner.update_single_model_cache(file_path, file_path, model_metadata) await update_cache_from_metadata(scanner, file_path, model_copy)
updated_models += 1 updated_models += 1
except Exception as e: except Exception as e:
+33 -3
View File
@@ -9,7 +9,7 @@ from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
from ..services.service_registry import ServiceRegistry from ..services.service_registry import ServiceRegistry
from ..services.settings_manager import get_settings_manager from ..services.settings_manager import get_settings_manager
from ..utils.example_images_paths import get_model_folder, get_model_relative_path from ..utils.example_images_paths import get_model_folder, get_model_relative_path
from .example_images_metadata import MetadataUpdater from .example_images_metadata import MetadataUpdater, update_cache_from_metadata
from ..utils.metadata_manager import MetadataManager from ..utils.metadata_manager import MetadataManager
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -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(
@@ -644,7 +674,7 @@ class ExampleImagesProcessor:
}, status=500) }, status=500)
# Update cache # Update cache
await scanner.update_single_model_cache(file_path, file_path, model_data) await update_cache_from_metadata(scanner, file_path, model_data)
# Get regular images array (might be None) # Get regular images array (might be None)
regular_images = civitai_data.get('images', []) regular_images = civitai_data.get('images', [])
@@ -759,7 +789,7 @@ class ExampleImagesProcessor:
model_copy = model_data.copy() model_copy = model_data.copy()
model_copy.pop('folder', None) model_copy.pop('folder', None)
await MetadataManager.save_metadata(file_path, model_copy) await MetadataManager.save_metadata(file_path, model_copy)
await scanner.update_single_model_cache(file_path, file_path, model_data) await update_cache_from_metadata(scanner, file_path, model_copy)
return web.json_response({ return web.json_response({
'success': True, 'success': True,
+48 -1
View File
@@ -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.
+1 -1
View File
@@ -1,7 +1,7 @@
[project] [project]
name = "comfyui-lora-manager" name = "comfyui-lora-manager"
description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!" description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!"
version = "1.1.9" version = "1.2.0"
license = {file = "LICENSE"} license = {file = "LICENSE"}
dependencies = [ dependencies = [
"aiohttp", "aiohttp",
@@ -0,0 +1,67 @@
/* Batch Download Summary Modal component styles only.
Stat cards and failure table styles are shared with the metadata refresh
result modal (metadata-refresh-result.css) and are not redefined here. */
.download-batch-summary-modal {
max-width: 700px;
}
.summary-header {
display: flex;
align-items: center;
gap: var(--space-2);
margin: var(--space-2) 0;
}
.summary-header i {
font-size: 1.4em;
}
.summary-header.success i {
color: var(--color-success);
}
.summary-header.warning i {
color: var(--color-warning);
}
.summary-header.error i {
color: var(--color-error);
}
.summary-title {
font-weight: var(--weight-semibold);
color: var(--lora-text);
}
.summary-hint {
margin-left: auto;
font-size: var(--text-xs);
color: var(--text-secondary);
}
.btn-retry {
display: inline-flex;
align-items: center;
gap: var(--space-1);
background: var(--lora-accent, #4f46e5);
color: #fff;
border: none;
border-radius: var(--border-radius-sm);
padding: var(--space-2) var(--space-3);
cursor: pointer;
font-weight: var(--weight-semibold);
}
.btn-retry:hover {
background: var(--lora-accent-hover, #4338ca);
}
.failure-link {
color: var(--lora-accent, #4f46e5);
text-decoration: none;
}
.failure-link:hover {
text-decoration: underline;
}
+1
View File
@@ -41,6 +41,7 @@
@import 'components/sidebar.css'; /* Add sidebar component */ @import 'components/sidebar.css'; /* Add sidebar component */
@import 'components/media-viewer.css'; @import 'components/media-viewer.css';
@import 'components/metadata-refresh-result.css'; @import 'components/metadata-refresh-result.css';
@import 'components/download-batch-summary.css';
.initialization-notice { .initialization-notice {
display: flex; display: flex;
+2 -1
View File
@@ -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
+8 -2
View File
@@ -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);
} }
@@ -0,0 +1,340 @@
import { translate } from '../utils/i18nHelpers.js';
import { showToast, openHuggingFace } from '../utils/uiHelpers.js';
/**
* Escape HTML entities in a string to prevent injection when interpolating into innerHTML.
* Safe for both text content and attribute values (quotes are escaped too).
* @param {string} str - The string to escape
* @returns {string} - The escaped string
*/
function _escapeHtml(str) {
if (!str) return '';
const div = document.createElement('div');
div.textContent = str;
return div.innerHTML.replace(/"/g, '&quot;').replace(/'/g, '&#39;');
}
/**
* Resolve the display name of a failed download entry.
* Prefers the resolved name carried on the entry, then known item fields,
* then derives a name from the item URL as a last resort.
* @param {Object} entry - The failed entry ({ item, error, name? })
* @returns {string} - The best available display name
*/
function _resolveItemName(entry) {
if (entry?.name) {
return entry.name;
}
const item = entry?.item ?? entry;
const direct = item?.displayName || item?.name || item?.file_name || item?.filename || item?.selectedVersion?.name;
if (direct) {
return direct;
}
if (item?.url) {
try {
const segments = new URL(item.url).pathname.split('/').filter(Boolean);
if (segments.length > 0) {
return decodeURIComponent(segments[segments.length - 1]);
}
} catch (e) {
// Unparseable URL — fall through to 'Unknown'
}
}
return 'Unknown';
}
/**
* Resolve the URL to open for a failed item always the original item URL.
* @param {Object} item - The failed item payload
* @returns {string|null} - A URL string, or null when nothing is available
*/
function _resolveItemUrl(item) {
return item?.url || null;
}
/**
* Format a raw failure error into a concise human-readable message.
* Unwraps JSON envelopes and extracts HTTP status/body details when present.
* @param {*} error - The raw error (usually a string)
* @returns {string} - The formatted error message
*/
function _formatError(error) {
if (!error) {
return 'Unknown error';
}
let base = typeof error === 'string' ? error : String(error);
// Unwrap JSON envelope: { "success": false, "error": "...", ... }
try {
const parsed = JSON.parse(base);
if (parsed && typeof parsed.error === 'string' && parsed.error) {
base = parsed.error;
}
} catch (e) {
// Not a JSON envelope — keep the raw string
}
// Extract HTTP status and JSON body details, e.g. "status=403 body={...}"
let result = base;
const statusMatch = base.match(/status=(\d{3})/);
const bodyMatch = base.match(/body=(\{.*\})/s);
if (bodyMatch) {
try {
const body = JSON.parse(bodyMatch[1]);
const detail = (typeof body?.message === 'string' && body.message)
|| (typeof body?.error === 'string' && body.error)
|| null;
if (detail) {
const status = statusMatch ? statusMatch[1] : null;
result = `${status ? `HTTP ${status}` : ''}${detail}`;
}
} catch (e) {
// Body is not valid JSON — keep the base string
}
}
// Truncate overly long messages
if (result.length > 220) {
result = result.slice(0, 220) + '…';
}
return result;
}
/**
* Build a plain-text report of the batch download results.
* @param {number} total - Total number of models attempted
* @param {number} completed - Number of models successfully downloaded
* @param {Array} failedItems - Array of failed items ({ item, error })
* @returns {string} - The report text
*/
function _buildReportText(total, completed, failedItems) {
const lines = [
'=== Batch Download Report ===',
`Date: ${new Date().toLocaleString()}`,
`Total: ${total}`,
`Successfully downloaded: ${completed}`,
`Failed: ${failedItems.length}`,
'',
];
if (failedItems.length > 0) {
lines.push('--- Failed Items ---');
failedItems.forEach((entry, i) => {
const name = _resolveItemName(entry);
const error = _formatError(entry?.error);
lines.push(`${i + 1}. ${name}${error}`);
const itemUrl = _resolveItemUrl(entry?.item ?? entry);
if (itemUrl) {
lines.push(` URL: ${itemUrl}`);
}
});
lines.push('');
}
lines.push('====================');
return lines.join('\n');
}
/**
* Handle a successful clipboard write: confirm via toast and briefly swap the
* trigger button to a "Copied!" state.
* @param {HTMLElement|null} btn - The button that triggered the copy action
*/
function _onCopyReportSuccess(btn) {
showToast('toast.api.copiedToClipboard', {}, 'success');
if (btn) {
const origHTML = btn.innerHTML;
btn.innerHTML = '<i class="fas fa-check"></i> Copied!';
setTimeout(() => { btn.innerHTML = origHTML; }, 2000);
}
}
/**
* Fallback for environments without the async Clipboard API (e.g. insecure
* contexts over LAN http where `navigator.clipboard` is undefined): copy via a
* hidden textarea and `document.execCommand('copy')`.
* @param {string} text - The report text to copy
*/
function _copyReportWithExecCommand(text) {
const textarea = document.createElement('textarea');
textarea.value = text;
document.body.appendChild(textarea);
textarea.select();
document.execCommand('copy');
document.body.removeChild(textarea);
showToast('toast.api.copiedToClipboard', {}, 'success');
}
/**
* Copy the batch download report to the clipboard.
* Uses the async Clipboard API when available, otherwise falls back to a hidden
* textarea + execCommand so the action still works in insecure contexts.
* @param {HTMLElement} btn - The button that triggered the copy action
* @param {number} total - Total number of models attempted
* @param {number} completed - Number of models successfully downloaded
* @param {Array} failedItems - Array of failed items
*/
function _copyReport(btn, total, completed, failedItems) {
const text = _buildReportText(total, completed, failedItems);
if (navigator.clipboard && typeof navigator.clipboard.writeText === 'function') {
navigator.clipboard.writeText(text)
.then(() => _onCopyReportSuccess(btn))
.catch(() => _copyReportWithExecCommand(text));
} else {
_copyReportWithExecCommand(text);
}
}
/**
* Show the batch download summary modal after a batch download completes.
* Mirrors the Metadata Fetch Summary modal lifecycle: the modal element is
* appended directly to document.body and removed on close; it is not
* registered with ModalManager.
* @param {Object} options - Summary options
* @param {number} options.total - Total number of models attempted
* @param {number} options.completed - Number of models successfully downloaded
* @param {Array} options.failedItems - Array of failed items ({ item, error })
* @param {Function} options.onRetry - Callback invoked with failedItems to retry the failed subset
*/
export function showDownloadBatchSummary({ total, completed, failedItems, onRetry }) {
const failures = failedItems || [];
const failedCount = failures.length;
// 3-state summary header semantics (mirrors BatchImportManager results header)
let headerState;
let headerIcon;
let headerText;
if (completed === 0) {
headerState = 'error';
headerIcon = 'fa-times-circle';
headerText = translate('modals.downloadBatchSummary.failed', {}, 'Download failed');
} else if (failedCount > 0) {
headerState = 'warning';
headerIcon = 'fa-exclamation-circle';
headerText = translate('modals.downloadBatchSummary.completedWithErrors', {}, 'Completed with errors');
} else {
headerState = 'success';
headerIcon = 'fa-check-circle';
headerText = translate('modals.downloadBatchSummary.successMessage', { count: completed }, 'All ' + completed + ' models downloaded successfully');
}
// Build failure table rows
const failureRows = failures.map((entry, i) => {
const item = entry?.item ?? entry;
const name = _resolveItemName(entry);
const itemUrl = _resolveItemUrl(item);
const rawError = entry?.error ? String(entry.error) : '';
const error = _formatError(entry?.error);
const nameCell = itemUrl
? `<td class="failure-name"><a href="#" class="failure-link" data-action="open-model" data-index="${i}" title="${_escapeHtml(itemUrl)}">${_escapeHtml(name)}</a></td>`
: `<td class="failure-name" title="${_escapeHtml(name)}">${_escapeHtml(name)}</td>`;
return `<tr>
<td class="failure-index">${i + 1}</td>
${nameCell}
<td class="failure-error" title="${_escapeHtml(rawError)}">${_escapeHtml(error)}</td>
</tr>`;
}).join('');
const modalHtml = `
<div id="downloadBatchSummaryModal" class="modal" style="display: block;">
<div class="modal-content download-batch-summary-modal">
<button class="close" data-action="close-modal">&times;</button>
<h2>${translate('modals.downloadBatchSummary.title', {}, 'Batch Download Summary')}</h2>
<div class="summary-header ${headerState}">
<i class="fas ${headerIcon}"></i>
<span class="summary-title">${headerText}</span>
<span class="summary-hint">${completed}/${total}</span>
</div>
<div class="refresh-summary-stats">
<div class="stat-card stat-card-success">
<div class="stat-card-body">
<span class="stat-card-label">${translate('modals.downloadBatchSummary.statSuccess', {}, 'Success')}</span>
<span class="stat-card-value">${completed}</span>
</div>
</div>
<div class="stat-card stat-card-failure">
<div class="stat-card-body">
<span class="stat-card-label">${translate('modals.downloadBatchSummary.statFailed', {}, 'Failed')}</span>
<span class="stat-card-value">${failedCount}</span>
</div>
</div>
<div class="stat-card stat-card-total">
<div class="stat-card-body">
<span class="stat-card-label">${translate('modals.downloadBatchSummary.statTotal', {}, 'Total')}</span>
<span class="stat-card-value">${total}</span>
</div>
</div>
</div>
${failedCount > 0 ? `
<div class="refresh-failures-section">
<h4><i class="fas fa-exclamation-triangle"></i> ${translate('modals.downloadBatchSummary.failedItems', { count: failedCount }, 'Failed Items (' + failedCount + ')')}</h4>
<div class="failure-table-wrapper">
<table class="failure-table">
<thead>
<tr>
<th>#</th>
<th>${translate('modals.downloadBatchSummary.columnName', {}, 'Model Name')}</th>
<th>${translate('modals.downloadBatchSummary.columnError', {}, 'Error')}</th>
</tr>
</thead>
<tbody>${failureRows}</tbody>
</table>
</div>
</div>
` : `
<div class="refresh-success-message">
<i class="fas fa-check-circle"></i> ${translate('modals.downloadBatchSummary.successMessage', { count: completed }, 'All ' + completed + ' models downloaded successfully')}
</div>
`}
<div class="modal-actions">
${failedCount > 0 ? `
<button class="btn-retry" data-action="retry-failed"><i class="fas fa-redo"></i> ${translate('modals.downloadBatchSummary.retryFailed', { count: failedCount }, 'Retry Failed (' + failedCount + ')')}</button>
<button class="secondary-btn" data-action="copy-report"><i class="fas fa-copy"></i> ${translate('modals.downloadBatchSummary.copyReport', {}, 'Copy Report')}</button>
` : ''}
<button class="cancel-btn" data-action="close-modal">${translate('modals.downloadBatchSummary.close', {}, 'Close')}</button>
</div>
</div>
</div>
`;
const existing = document.getElementById('downloadBatchSummaryModal');
if (existing) existing.remove();
const container = document.createElement('div');
container.innerHTML = modalHtml;
const modal = container.firstElementChild;
document.body.appendChild(modal);
modal.addEventListener('click', (e) => {
const actionEl = e.target.closest('[data-action]');
const action = actionEl?.dataset.action;
if (!action) return;
e.preventDefault();
switch (action) {
case 'close-modal':
modal.remove();
break;
case 'retry-failed':
modal.remove();
if (typeof onRetry === 'function') {
onRetry(failures);
}
break;
case 'copy-report':
_copyReport(actionEl, total, completed, failures);
break;
case 'open-model': {
// Keep the modal open; just open the item's original URL in a new tab
const entry = failures[Number(actionEl.dataset.index)];
const item = entry?.item;
if (!item?.url) break;
openHuggingFace(item.url);
break;
}
}
});
}
+56 -17
View File
@@ -108,10 +108,20 @@ export class PageControls {
const sortSelect = document.getElementById('sortSelect'); const sortSelect = document.getElementById('sortSelect');
if (sortSelect) { if (sortSelect) {
initSortDropdown(sortSelect); initSortDropdown(sortSelect);
sortSelect.value = this.pageState.sortBy; this.applySortToSelect(this.pageState.sortBy);
sortSelect.addEventListener('change', async (e) => { sortSelect.addEventListener('change', async (e) => {
this.pageState.sortBy = e.target.value; let value = e.target.value;
this.saveSortPreference(e.target.value); if (value.startsWith('random')) {
// Every pick of Random reshuffles the list: generate a
// fresh seed so the backend keeps a stable order across
// paginated requests.
value = this._randomizeSortValue();
}
this.pageState.sortBy = value;
this.saveSortPreference(value);
// Reset the seeded Random option when switching away from
// Random, or re-apply the fresh seed when picking it again.
this.applySortToSelect(value);
await this.resetAndReload(); await this.resetAndReload();
}); });
} }
@@ -312,6 +322,44 @@ export class PageControls {
} }
} }
/**
* Apply a sort value to the native sort <select>, keeping the Random
* option's value in sync when the persisted value carries a seed
* (e.g. "random:abc123"). Must be used instead of assigning
* sortSelect.value directly whenever the value may be a seeded random
* sort, otherwise the native select has no matching option.
* @param {string} sortValue - Sort value like "name:asc" or "random:<seed>"
*/
applySortToSelect(sortValue) {
const sortSelect = document.getElementById('sortSelect');
if (!sortSelect) return;
const randomOpt = sortSelect.querySelector('option[value="random"], option[value^="random:"]');
if (randomOpt) {
randomOpt.value = String(sortValue).startsWith('random') ? sortValue : 'random';
}
sortSelect.value = sortValue;
}
/**
* Generate a fresh seeded random sort value ("random:<seed>") and keep
* the native <select> in sync so its value matches the persisted sort
* string and the dropdown shows the selected label.
* @returns {string} The new sort value, e.g. "random:abc123xyz"
*/
_randomizeSortValue() {
const seed = Math.random().toString(36).slice(2, 12);
const value = `random:${seed}`;
const sortSelect = document.getElementById('sortSelect');
if (sortSelect) {
const randomOpt = sortSelect.querySelector('option[value="random"], option[value^="random:"]');
if (randomOpt) {
randomOpt.value = value;
}
sortSelect.value = value;
}
return value;
}
/** /**
* Load sort preference from storage * Load sort preference from storage
*/ */
@@ -326,10 +374,7 @@ export class PageControls {
// Handle legacy format conversion // Handle legacy format conversion
const convertedSort = this.convertLegacySortFormat(savedSort); const convertedSort = this.convertLegacySortFormat(savedSort);
this.pageState.sortBy = convertedSort; this.pageState.sortBy = convertedSort;
const sortSelect = document.getElementById('sortSelect'); this.applySortToSelect(convertedSort);
if (sortSelect) {
sortSelect.value = convertedSort;
}
} }
} }
@@ -523,9 +568,9 @@ export class PageControls {
this.pageState.sortBy = restoredSort; this.pageState.sortBy = restoredSort;
this.saveSortPreference(restoredSort); this.saveSortPreference(restoredSort);
this._removeVlmSortOption(); this._removeVlmSortOption();
this.applySortToSelect(restoredSort);
const sortSelect = document.getElementById('sortSelect'); const sortSelect = document.getElementById('sortSelect');
if (sortSelect) { if (sortSelect) {
sortSelect.value = restoredSort;
sortSelect.disabled = false; sortSelect.disabled = false;
} }
} }
@@ -575,10 +620,7 @@ export class PageControls {
const savedGroupedSort = getStorageItem(groupedKey); const savedGroupedSort = getStorageItem(groupedKey);
if (savedGroupedSort) { if (savedGroupedSort) {
this.pageState.sortBy = savedGroupedSort; this.pageState.sortBy = savedGroupedSort;
const sortSelect = document.getElementById('sortSelect'); this.applySortToSelect(savedGroupedSort);
if (sortSelect) {
sortSelect.value = savedGroupedSort;
}
} }
} else { } else {
// Leaving group mode: persist current sort for next time, restore non-group sort // Leaving group mode: persist current sort for next time, restore non-group sort
@@ -586,10 +628,7 @@ export class PageControls {
const savedNormalSort = getStorageItem(`${this.pageType}_sort`); const savedNormalSort = getStorageItem(`${this.pageType}_sort`);
if (savedNormalSort) { if (savedNormalSort) {
this.pageState.sortBy = savedNormalSort; this.pageState.sortBy = savedNormalSort;
const sortSelect = document.getElementById('sortSelect'); this.applySortToSelect(savedNormalSort);
if (sortSelect) {
sortSelect.value = savedNormalSort;
}
} }
} }
} }
@@ -874,7 +913,7 @@ export class PageControls {
} }
if (sortSelect) { if (sortSelect) {
sortSelect.value = this.pageState.sortBy; this.applySortToSelect(this.pageState.sortBy);
} }
if (searchInput) { if (searchInput) {
searchInput.value = this.pageState.filters?.search || ''; searchInput.value = this.pageState.filters?.search || '';
+13 -3
View File
@@ -96,7 +96,16 @@ export function initSortDropdown(select) {
}; };
const choose = (value) => { const choose = (value) => {
if (select.value === value) return; if (select.value === value) {
// Re-picking the already-selected option is normally a no-op,
// matching native <select> behavior. The seeded Random sort is
// the exception: clicking it again should reshuffle, so let the
// change handler (PageControls) generate a fresh seed.
if (String(value).startsWith('random')) {
select.dispatchEvent(new Event('change', { bubbles: true }));
}
return;
}
select.value = value; select.value = value;
select.dispatchEvent(new Event('change', { bubbles: true })); select.dispatchEvent(new Event('change', { bubbles: true }));
}; };
@@ -277,9 +286,10 @@ export function initSortDropdown(select) {
} }
// Rebuild the menu when <option>s change (VLM adds/removes a temporary // Rebuild the menu when <option>s change (VLM adds/removes a temporary
// option at runtime). // option at runtime, and the seeded Random sort option gets a new value
// attribute each time it is picked).
const observer = new MutationObserver(() => buildMenu()); const observer = new MutationObserver(() => buildMenu());
observer.observe(select, { childList: true }); observer.observe(select, { childList: true, subtree: true, attributes: true, attributeFilter: ['value'] });
buildMenu(); buildMenu();
group.dataset.sortReady = '1'; group.dataset.sortReady = '1';
+48 -4
View File
@@ -27,6 +27,8 @@ export class BulkManager {
// Drag detection properties // Drag detection properties
this.dragThreshold = 5; // Pixels to move before considering it a drag this.dragThreshold = 5; // Pixels to move before considering it a drag
this.dragDelayMs = 100; // Minimum hold time before a drag is treated as a marquee
this.minMarqueeSize = 10; // Minimum drag box (px) before a marquee counts as a selection
this.mouseDownTime = 0; this.mouseDownTime = 0;
this.mouseDownPosition = { x: 0, y: 0 }; this.mouseDownPosition = { x: 0, y: 0 };
@@ -88,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,
@@ -173,6 +175,19 @@ export class BulkManager {
}); });
eventManager.addHandler('mousemove', 'bulkManager-marquee-move', (e) => { eventManager.addHandler('mousemove', 'bulkManager-marquee-move', (e) => {
// Only track marquee/drag while the left button is physically held.
// mouseup can be missed (release outside the window, focus loss, driver quirks),
// so mousemove must verify the button state itself instead of relying on it.
if (!(e.buttons & 1)) {
if (this.isMarqueeActive) {
this.endMarqueeSelection(e);
} else {
this.mouseDownTime = 0;
this.isDragging = false;
}
return false;
}
if (this.isMarqueeActive) { if (this.isMarqueeActive) {
this.lastClientX = e.clientX; this.lastClientX = e.clientX;
this.lastClientY = e.clientY; this.lastClientY = e.clientY;
@@ -184,7 +199,10 @@ export class BulkManager {
const dy = e.clientY - this.mouseDownPosition.y; const dy = e.clientY - this.mouseDownPosition.y;
const distance = Math.sqrt(dx * dx + dy * dy); const distance = Math.sqrt(dx * dx + dy * dy);
if (distance >= this.dragThreshold) { // Require both enough movement AND enough hold time so quick
// click jitter from micro-movement input devices is not a marquee.
const heldTime = Date.now() - this.mouseDownTime;
if (heldTime >= this.dragDelayMs && distance >= this.dragThreshold) {
this.isDragging = true; this.isDragging = true;
this.startMarqueeSelection(e, true); this.startMarqueeSelection(e, true);
} }
@@ -1510,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++;
@@ -1958,9 +1980,31 @@ export class BulkManager {
// Remove visual feedback class // Remove visual feedback class
document.body.classList.remove('marquee-selecting'); document.body.classList.remove('marquee-selecting');
// Compute the actual drag box size in document coordinates, matching how
// updateMarqueeSelectionFromPosition tracks the rectangle. Client-space
// size would wrongly flag auto-scroll marquees (tiny pointer movement,
// large document-space box) as accidental clicks.
const container = document.querySelector('.page-content');
const scrollX = container?.scrollLeft || 0;
const scrollY = container?.scrollTop || 0;
const dragWidth = Math.abs((e.clientX + scrollX) - this.marqueeStartDoc.x);
const dragHeight = Math.abs((e.clientY + scrollY) - this.marqueeStartDoc.y);
const isTinyMarquee = dragWidth < this.minMarqueeSize && dragHeight < this.minMarqueeSize;
// Get selection count // Get selection count
const selectionCount = state.selectedModels.size; const selectionCount = state.selectedModels.size;
// A tiny box (e.g. click jitter that happened to graze a card) is treated
// as an accidental click: undo any selection and leave bulk mode.
if (isTinyMarquee) {
this.clearSelection();
if (state.bulkMode) {
this.toggleBulkMode();
}
this.initialSelectedModels.clear();
return;
}
// If no models were selected, exit bulk mode // If no models were selected, exit bulk mode
if (selectionCount === 0) { if (selectionCount === 0) {
if (state.bulkMode) { if (state.bulkMode) {
+16 -3
View File
@@ -8,6 +8,7 @@ import { FolderTreeManager } from '../components/FolderTreeManager.js';
import { translate } from '../utils/i18nHelpers.js'; import { translate } from '../utils/i18nHelpers.js';
import { extractCivitaiModelUrlParts } from '../utils/civitaiUtils.js'; import { extractCivitaiModelUrlParts } from '../utils/civitaiUtils.js';
import { formatFileSize } from '../utils/formatters.js'; import { formatFileSize } from '../utils/formatters.js';
import { showDownloadBatchSummary } from '../components/DownloadBatchSummaryModal.js';
export class DownloadManager { export class DownloadManager {
constructor() { constructor() {
@@ -1548,6 +1549,10 @@ export class DownloadManager {
modalManager.closeModal('downloadModal'); modalManager.closeModal('downloadModal');
return this.executeBatchDownload(downloadItems, { modelRoot, targetFolder, useDefaultPaths });
}
async executeBatchDownload(downloadItems, { modelRoot, targetFolder, useDefaultPaths }) {
const batchDownloadId = Date.now().toString(); const batchDownloadId = Date.now().toString();
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://'; const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
const ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${batchDownloadId}`); const ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${batchDownloadId}`);
@@ -1558,6 +1563,7 @@ export class DownloadManager {
let completedDownloads = 0; let completedDownloads = 0;
let failedDownloads = 0; let failedDownloads = 0;
let cancelled = false; let cancelled = false;
const failedItems = [];
loadingManager.showCancelButton(async () => { loadingManager.showCancelButton(async () => {
if (cancelled) return; if (cancelled) return;
@@ -1658,6 +1664,7 @@ export class DownloadManager {
if (!response.success) { if (!response.success) {
failedDownloads++; failedDownloads++;
failedItems.push({ item, error: response.error || 'Unknown error', name });
} else { } else {
completedDownloads++; completedDownloads++;
updateProgress(100, completedDownloads, ''); updateProgress(100, completedDownloads, '');
@@ -1666,6 +1673,7 @@ export class DownloadManager {
if (!cancelled) { if (!cancelled) {
console.error(`Failed to download ${name}:`, err); console.error(`Failed to download ${name}:`, err);
failedDownloads++; failedDownloads++;
failedItems.push({ item, error: err?.message || 'Unknown error', name });
} }
} }
} }
@@ -1679,10 +1687,15 @@ export class DownloadManager {
} else if (failedDownloads === 0) { } else if (failedDownloads === 0) {
showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success'); showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success');
} else { } else {
showToast('toast.loras.downloadPartialSuccess', { showDownloadBatchSummary({
completed: completedDownloads,
total: downloadItems.length, total: downloadItems.length,
}, 'warning'); completed: completedDownloads,
failedItems,
onRetry: (failed) => this.executeBatchDownload(
failed.map((f) => f.item),
{ modelRoot, targetFolder, useDefaultPaths }
),
});
} }
await resetAndReload(true); await resetAndReload(true);
+33 -3
View File
@@ -27,6 +27,8 @@ export class UpdateService {
this.isUpdating = false; this.isUpdating = false;
this.channelMode = null; this.channelMode = null;
this.hasGit = false; this.hasGit = false;
this.nightlyNotifyDate = getStorageItem('nightly_notify_date', '');
this.nightlyBadgeShown = false;
this.progressKeepVisible = false; this.progressKeepVisible = false;
this.currentVersionInfo = null; this.currentVersionInfo = null;
this.versionMismatch = false; this.versionMismatch = false;
@@ -552,6 +554,11 @@ export class UpdateService {
this.updateAvailable = data.update_available; this.updateAvailable = data.update_available;
// Nightly channel: surface the update badge at most once per calendar day.
if (this.updateAvailable && this.channelMode === 'nightly' && this.nightlyNotifyDate !== this._getTodayKey()) {
this._markNightlyNotified();
}
this.lastCheckTime = now; this.lastCheckTime = now;
setStorageItem('last_update_check', now.toString()); setStorageItem('last_update_check', now.toString());
@@ -602,6 +609,28 @@ export class UpdateService {
return false; return false;
} }
_getTodayKey() {
const now = new Date();
const month = String(now.getMonth() + 1).padStart(2, '0');
const day = String(now.getDate()).padStart(2, '0');
return `${now.getFullYear()}-${month}-${day}`;
}
_isNightlyBadgeAllowed() {
if (this.channelMode !== 'nightly') {
return true;
}
// Keep the badge visible for the rest of the session once shown, but do
// not show it again on later sessions within the same calendar day.
return this.nightlyNotifyDate !== this._getTodayKey() || this.nightlyBadgeShown;
}
_markNightlyNotified() {
this.nightlyNotifyDate = this._getTodayKey();
this.nightlyBadgeShown = true;
setStorageItem('nightly_notify_date', this.nightlyNotifyDate);
}
updateBadgeVisibility() { updateBadgeVisibility() {
const updateToggle = document.querySelector('.update-toggle'); const updateToggle = document.querySelector('.update-toggle');
const updateBadge = document.querySelector('.update-toggle .update-badge'); const updateBadge = document.querySelector('.update-toggle .update-badge');
@@ -609,9 +638,12 @@ export class UpdateService {
? bannerService.getUnreadBannerCount() ? bannerService.getUnreadBannerCount()
: 0; : 0;
// Force updating badges visibility based on current state
const shouldShowUpdate = this.updateNotificationsEnabled && this.updateAvailable && this._isNightlyBadgeAllowed();
if (updateToggle) { if (updateToggle) {
let tooltipKey = 'header.actions.notifications'; let tooltipKey = 'header.actions.notifications';
if (this.updateNotificationsEnabled && this.updateAvailable) { if (shouldShowUpdate) {
tooltipKey = 'update.updateAvailable'; tooltipKey = 'update.updateAvailable';
} else if (unreadBanners > 0) { } else if (unreadBanners > 0) {
tooltipKey = 'update.tabs.messages'; tooltipKey = 'update.tabs.messages';
@@ -619,8 +651,6 @@ export class UpdateService {
updateToggle.title = translate(tooltipKey); updateToggle.title = translate(tooltipKey);
} }
// Force updating badges visibility based on current state
const shouldShowUpdate = this.updateNotificationsEnabled && this.updateAvailable;
const shouldShow = shouldShowUpdate || unreadBanners > 0; const shouldShow = shouldShowUpdate || unreadBanners > 0;
if (updateBadge) { if (updateBadge) {
+6 -1
View File
@@ -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 -->
+24 -4
View File
@@ -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>
+5
View File
@@ -48,6 +48,11 @@
<option value="versions_count:asc">{{ t('loras.controls.sort.versionsCountAsc', default='Fewest versions first') }}</option> <option value="versions_count:asc">{{ t('loras.controls.sort.versionsCountAsc', default='Fewest versions first') }}</option>
</optgroup> </optgroup>
{% endif %} {% endif %}
{% if page_id != 'recipes' %}
<optgroup label="{{ t('loras.controls.sort.random', default='Random') }}">
<option value="random">{{ t('loras.controls.sort.randomAction', default='Randomize (shuffle)') }}</option>
</optgroup>
{% endif %}
{% if page_id == 'recipes' %} {% if page_id == 'recipes' %}
<optgroup label="{{ t('recipes.controls.sort.lorasCount') }}"> <optgroup label="{{ t('recipes.controls.sort.lorasCount') }}">
<option value="loras_count:desc">{{ t('recipes.controls.sort.lorasCountDesc') }}</option> <option value="loras_count:desc">{{ t('recipes.controls.sort.lorasCountDesc') }}</option>
@@ -202,6 +202,7 @@
<option value="deepseek">{{ t('settings.aiProvider.providerOptions.deepseek') }}</option> <option value="deepseek">{{ t('settings.aiProvider.providerOptions.deepseek') }}</option>
<option value="groq">{{ t('settings.aiProvider.providerOptions.groq') }}</option> <option value="groq">{{ t('settings.aiProvider.providerOptions.groq') }}</option>
<option value="openrouter">{{ t('settings.aiProvider.providerOptions.openrouter') }}</option> <option value="openrouter">{{ t('settings.aiProvider.providerOptions.openrouter') }}</option>
<option value="google">{{ t('settings.aiProvider.providerOptions.google') }}</option>
<option value="opencode-go">{{ t('settings.aiProvider.providerOptions.opencode-go') }}</option> <option value="opencode-go">{{ t('settings.aiProvider.providerOptions.opencode-go') }}</option>
<option value="custom">{{ t('settings.aiProvider.providerOptions.custom') }}</option> <option value="custom">{{ t('settings.aiProvider.providerOptions.custom') }}</option>
</select> </select>
+6 -1
View File
@@ -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,438 @@
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest';
const {
SUMMARY_MODULE,
I18N_HELPERS_MODULE,
UI_HELPERS_MODULE,
} = vi.hoisted(() => ({
SUMMARY_MODULE: new URL('../../../static/js/components/DownloadBatchSummaryModal.js', import.meta.url).pathname,
I18N_HELPERS_MODULE: new URL('../../../static/js/utils/i18nHelpers.js', import.meta.url).pathname,
UI_HELPERS_MODULE: new URL('../../../static/js/utils/uiHelpers.js', import.meta.url).pathname,
}));
const showToastMock = vi.hoisted(() => vi.fn());
const openHuggingFaceMock = vi.hoisted(() => vi.fn());
vi.mock(I18N_HELPERS_MODULE, () => ({
translate: vi.fn((_key, _params, fallback) => fallback ?? ''),
}));
vi.mock(UI_HELPERS_MODULE, () => ({
showToast: showToastMock,
openHuggingFace: openHuggingFaceMock,
}));
// A realistic failure payload from the backend: a JSON envelope whose `error`
// field embeds an HTTP status and a nested JSON body (Civitai Early Access).
const REAL_ERROR = '{"success": false, "error": "Failed to resolve authenticated Civitai redirect: status=403 body={\\"error\\":\\"Early Access\\",\\"deadline\\":\\"2026-08-12T08:18:36.063Z\\",\\"message\\":\\"This asset is in Early Access. You can use Buzz access it now!\\"}", "download_id": "1786065633067"}';
// The human-readable error the component should derive from REAL_ERROR.
const FORMATTED_REAL_ERROR = 'HTTP 403 — This asset is in Early Access. You can use Buzz access it now!';
describe('DownloadBatchSummaryModal', () => {
let showDownloadBatchSummary;
beforeEach(async () => {
document.body.innerHTML = '';
showToastMock.mockClear();
openHuggingFaceMock.mockClear();
({ showDownloadBatchSummary } = await import(SUMMARY_MODULE));
});
afterEach(() => {
document.body.innerHTML = '';
delete navigator.clipboard;
delete document.execCommand;
vi.restoreAllMocks();
vi.useRealTimers();
});
it('renders a warning summary with stat cards and a failure table on partial success', () => {
showDownloadBatchSummary({
total: 3,
completed: 2,
failedItems: [
{ item: { displayName: 'LoraA' }, error: 'timeout' },
{ item: { name: 'LoraB' }, error: '404' },
],
onRetry: vi.fn(),
});
const modal = document.getElementById('downloadBatchSummaryModal');
expect(modal).not.toBeNull();
expect(modal.querySelector('.summary-header').classList.contains('warning')).toBe(true);
// Success / Failed / Total stat cards.
const statValues = Array.from(modal.querySelectorAll('.stat-card-value')).map(el => el.textContent);
expect(statValues).toEqual(['2', '2', '3']);
const rows = modal.querySelectorAll('.failure-table tbody tr');
expect(rows).toHaveLength(2);
expect(rows[0].querySelector('.failure-name').textContent).toBe('LoraA');
expect(rows[0].querySelector('.failure-error').textContent).toBe('timeout');
expect(rows[1].querySelector('.failure-name').textContent).toBe('LoraB');
expect(rows[1].querySelector('.failure-error').textContent).toBe('404');
expect(modal.querySelector('[data-action="retry-failed"]').textContent).toContain('Retry Failed (2)');
expect(modal.querySelector('[data-action="copy-report"]')).not.toBeNull();
});
it('renders an error header when every download failed', () => {
showDownloadBatchSummary({
total: 2,
completed: 0,
failedItems: [
{ item: { displayName: 'LoraA' }, error: 'timeout' },
{ item: { displayName: 'LoraB' }, error: '404' },
],
onRetry: vi.fn(),
});
const modal = document.getElementById('downloadBatchSummaryModal');
expect(modal.querySelector('.summary-header').classList.contains('error')).toBe(true);
expect(modal.querySelector('.summary-title').textContent).toBe('Download failed');
});
it('renders a success summary without a failure table or retry button', () => {
showDownloadBatchSummary({ total: 2, completed: 2, failedItems: [], onRetry: vi.fn() });
const modal = document.getElementById('downloadBatchSummaryModal');
expect(modal.querySelector('.summary-header').classList.contains('success')).toBe(true);
expect(modal.querySelector('.failure-table')).toBeNull();
expect(modal.querySelector('[data-action="retry-failed"]')).toBeNull();
expect(modal.querySelector('.refresh-success-message')).not.toBeNull();
});
it('escapes HTML in failed item names and errors', () => {
showDownloadBatchSummary({
total: 1,
completed: 0,
failedItems: [
{ item: { name: '<img src=x onerror=alert(1)>', url: 'https://example.com/xss-model' }, error: '<script>bad()</script>' },
],
onRetry: vi.fn(),
});
const nameCell = document.querySelector('.failure-name');
const errorCell = document.querySelector('.failure-error');
// The URL resolves, so the name renders inside the failure link; the
// escaped entities must render back to the literal payload as text...
expect(nameCell.querySelector('a.failure-link')).not.toBeNull();
expect(nameCell.textContent).toContain('<img src=x onerror=alert(1)>');
expect(errorCell.textContent).toContain('<script>bad()</script>');
// ...and never as live DOM nodes.
expect(document.querySelector('.failure-table img')).toBeNull();
expect(document.querySelector('.failure-table script')).toBeNull();
expect(nameCell.innerHTML).toContain('&lt;img');
});
it('removes the modal and invokes onRetry with the original failed items', () => {
const onRetry = vi.fn();
const failedItems = [{ item: { displayName: 'LoraA' }, error: 'timeout' }];
showDownloadBatchSummary({ total: 3, completed: 2, failedItems, onRetry });
document.querySelector('[data-action="retry-failed"]').click();
expect(document.getElementById('downloadBatchSummaryModal')).toBeNull();
expect(onRetry).toHaveBeenCalledTimes(1);
expect(onRetry).toHaveBeenCalledWith(failedItems);
// Same object references, not copies.
expect(onRetry.mock.calls[0][0][0]).toBe(failedItems[0]);
});
it('closes the modal via the close action without retrying', () => {
const onRetry = vi.fn();
showDownloadBatchSummary({
total: 2,
completed: 1,
failedItems: [{ item: { name: 'LoraA' }, error: 'timeout' }],
onRetry,
});
document.querySelector('.cancel-btn[data-action="close-modal"]').click();
expect(document.getElementById('downloadBatchSummaryModal')).toBeNull();
expect(onRetry).not.toHaveBeenCalled();
});
it('copies a plain-text batch report to the clipboard', async () => {
const writeText = vi.fn().mockResolvedValue(undefined);
Object.defineProperty(navigator, 'clipboard', { value: { writeText }, configurable: true });
showDownloadBatchSummary({
total: 3,
completed: 2,
failedItems: [
{ item: { displayName: 'LoraA', url: 'https://civitai.red/models/111/lora-a?modelVersionId=222' }, error: 'timeout' },
{ item: { name: 'LoraB', url: 'https://example.com/lora-b' }, error: '404' },
],
onRetry: vi.fn(),
});
document.querySelector('[data-action="copy-report"]').click();
// writeText is invoked synchronously by the click handler.
expect(writeText).toHaveBeenCalledTimes(1);
const text = writeText.mock.calls[0][0];
expect(text).toContain('Batch Download Report');
expect(text).toContain('Total: 3');
expect(text).toContain('LoraA — timeout');
expect(text).toContain('LoraB — 404');
// Each failed item with a URL gets an indented URL line right after it.
expect(text).toContain(' URL: https://civitai.red/models/111/lora-a?modelVersionId=222');
expect(text).toContain(' URL: https://example.com/lora-b');
// Exactly the two URLs from the failed items — nothing more, no undefined.
expect(text.match(/^\s+URL:/gm)).toHaveLength(2);
expect(text).not.toContain('URL: undefined');
// The toast fires after the mocked clipboard promise settles.
await vi.waitFor(() => expect(showToastMock).toHaveBeenCalledTimes(1));
expect(showToastMock).toHaveBeenCalledWith('toast.api.copiedToClipboard', {}, 'success');
});
it('omits the URL line for failed items without a resolvable url', async () => {
const writeText = vi.fn().mockResolvedValue(undefined);
Object.defineProperty(navigator, 'clipboard', { value: { writeText }, configurable: true });
showDownloadBatchSummary({
total: 2,
completed: 0,
failedItems: [
{ item: { name: 'WithUrl', url: 'https://example.com/with-url' }, error: 'boom' },
{ item: { name: 'NoUrl' }, error: 'boom' },
],
onRetry: vi.fn(),
});
document.querySelector('[data-action="copy-report"]').click();
const text = writeText.mock.calls[0][0];
expect(text).toContain(' URL: https://example.com/with-url');
// Only the one URL line exists — the URL-less item contributes none.
expect(text.match(/^\s+URL:/gm)).toHaveLength(1);
expect(text).not.toContain(' URL: undefined');
expect(text).not.toContain(' URL: null');
await vi.waitFor(() => expect(showToastMock).toHaveBeenCalledTimes(1));
});
it('falls back to execCommand when navigator.clipboard is unavailable', async () => {
// afterEach deletes navigator.clipboard, but be explicit so this test is
// robust even if a previous test failed before its cleanup ran.
delete navigator.clipboard;
// jsdom does not implement document.execCommand, so install a mock for the
// fallback path (removed by the afterEach cleanup above).
const execCommandMock = vi.fn(() => true);
document.execCommand = execCommandMock;
showDownloadBatchSummary({
total: 3,
completed: 2,
failedItems: [
{ item: { displayName: 'LoraA' }, error: 'timeout' },
{ item: { name: 'LoraB' }, error: '404' },
],
onRetry: vi.fn(),
});
document.querySelector('[data-action="copy-report"]').click();
// Without the async Clipboard API the fallback must run synchronously.
expect(execCommandMock).toHaveBeenCalledWith('copy');
await Promise.resolve();
await Promise.resolve();
expect(showToastMock).toHaveBeenCalledWith('toast.api.copiedToClipboard', {}, 'success');
});
it('keeps only a single modal instance across repeated calls', () => {
showDownloadBatchSummary({
total: 2,
completed: 1,
failedItems: [{ item: { name: 'A' }, error: 'e' }],
onRetry: vi.fn(),
});
showDownloadBatchSummary({ total: 3, completed: 3, failedItems: [], onRetry: vi.fn() });
expect(document.querySelectorAll('#downloadBatchSummaryModal')).toHaveLength(1);
const modal = document.getElementById('downloadBatchSummaryModal');
expect(modal.querySelector('.summary-header').classList.contains('success')).toBe(true);
});
it('resolves failure names from entry.name, item fields, URL paths, or Unknown', () => {
showDownloadBatchSummary({
total: 4,
completed: 0,
failedItems: [
{ name: 'entryName', item: { displayName: 'ItemName' }, error: 'e1' },
{ item: { selectedVersion: { name: 'v1.0' } }, error: 'e2' },
{ item: { url: 'https://civitai.red/models/837884/midjourney-artful-nsfw?modelVersionId=3153960' }, error: 'e3' },
{ item: {}, error: 'e4' },
],
onRetry: vi.fn(),
});
const names = Array.from(document.querySelectorAll('.failure-name')).map(el => el.textContent);
expect(names).toEqual(['entryName', 'v1.0', 'midjourney-artful-nsfw', 'Unknown']);
});
it('formats the real JSON failure payload into a concise HTTP error and truncates long ones', () => {
showDownloadBatchSummary({
total: 2,
completed: 0,
failedItems: [
{ item: { name: 'EarlyAccess' }, error: REAL_ERROR },
{ item: { name: 'LongError' }, error: 'x'.repeat(300) },
],
onRetry: vi.fn(),
});
const errorCells = document.querySelectorAll('.failure-error');
expect(errorCells[0].textContent).toBe(FORMATTED_REAL_ERROR);
expect(errorCells[1].textContent).toBe('x'.repeat(220) + '…');
});
it('keeps the raw error string in the error cell title for debugging', () => {
showDownloadBatchSummary({
total: 1,
completed: 0,
failedItems: [{ item: { name: 'EarlyAccess' }, error: REAL_ERROR }],
onRetry: vi.fn(),
});
const errorCell = document.querySelector('.failure-error');
expect(errorCell.getAttribute('title')).toBe(REAL_ERROR);
expect(errorCell.getAttribute('title')).not.toBe(FORMATTED_REAL_ERROR);
});
it('opens the original item url in a new tab when a failure link is clicked', () => {
showDownloadBatchSummary({
total: 1,
completed: 0,
failedItems: [{
item: {
url: 'https://civitai.red/models/837884/midjourney-artful-nsfw?modelVersionId=3153960',
modelId: '837884',
selectedVersion: { id: '3153960' },
},
error: 'rate limited',
}],
onRetry: vi.fn(),
});
document.querySelector('.failure-link').click();
expect(openHuggingFaceMock).toHaveBeenCalledTimes(1);
expect(openHuggingFaceMock).toHaveBeenCalledWith('https://civitai.red/models/837884/midjourney-artful-nsfw?modelVersionId=3153960');
// The modal stays open so the user can keep inspecting the failures.
expect(document.getElementById('downloadBatchSummaryModal')).not.toBeNull();
});
it('opens the item url directly when selectedVersion is absent', () => {
showDownloadBatchSummary({
total: 1,
completed: 0,
failedItems: [{
item: {
modelId: '837884',
modelVersionId: '3153960',
url: 'https://civitai.red/models/837884/midjourney-artful-nsfw',
},
error: 'rate limited',
}],
onRetry: vi.fn(),
});
document.querySelector('.failure-link').click();
expect(openHuggingFaceMock).toHaveBeenCalledTimes(1);
expect(openHuggingFaceMock).toHaveBeenCalledWith('https://civitai.red/models/837884/midjourney-artful-nsfw');
expect(document.getElementById('downloadBatchSummaryModal')).not.toBeNull();
});
it('opens the original huggingface url directly when a huggingface failure link is clicked', () => {
showDownloadBatchSummary({
total: 1,
completed: 0,
failedItems: [{
item: {
url: 'https://huggingface.co/user/repo',
source: 'huggingface',
repo: 'user/repo',
filename: 'model.safetensors',
revision: 'main',
},
error: 'download failed',
}],
onRetry: vi.fn(),
});
document.querySelector('.failure-link').click();
expect(openHuggingFaceMock).toHaveBeenCalledTimes(1);
expect(openHuggingFaceMock).toHaveBeenCalledWith('https://huggingface.co/user/repo');
expect(document.getElementById('downloadBatchSummaryModal')).not.toBeNull();
});
it('opens an arbitrary URL via openHuggingFace for fallback items', () => {
showDownloadBatchSummary({
total: 1,
completed: 0,
failedItems: [{ item: { url: 'https://example.com/model' }, error: 'boom' }],
onRetry: vi.fn(),
});
document.querySelector('.failure-link').click();
expect(openHuggingFaceMock).toHaveBeenCalledTimes(1);
expect(openHuggingFaceMock).toHaveBeenCalledWith('https://example.com/model');
});
it('renders the failure name as plain text when no URL can be resolved', () => {
showDownloadBatchSummary({
total: 1,
completed: 0,
failedItems: [{ item: { modelId: null }, error: 'boom' }],
onRetry: vi.fn(),
});
expect(document.querySelector('a.failure-link')).toBeNull();
expect(document.querySelector('.failure-name').textContent).toBe('Unknown');
// Without a link there is nothing to open: clicking the cell is inert.
document.querySelector('.failure-name').click();
expect(openHuggingFaceMock).not.toHaveBeenCalled();
});
it('copies formatted errors (not raw JSON) into the report text', async () => {
const writeText = vi.fn().mockResolvedValue(undefined);
Object.defineProperty(navigator, 'clipboard', { value: { writeText }, configurable: true });
showDownloadBatchSummary({
total: 1,
completed: 0,
failedItems: [{
item: {
name: 'EarlyAccess',
url: 'https://civitai.red/models/123/early-access?modelVersionId=456',
},
error: REAL_ERROR,
}],
onRetry: vi.fn(),
});
document.querySelector('[data-action="copy-report"]').click();
expect(writeText).toHaveBeenCalledTimes(1);
const text = writeText.mock.calls[0][0];
expect(text).toContain(FORMATTED_REAL_ERROR);
expect(text).toContain(' URL: https://civitai.red/models/123/early-access?modelVersionId=456');
expect(text).not.toContain('download_id');
expect(text).not.toContain('Failed to resolve authenticated Civitai redirect');
await vi.waitFor(() => expect(showToastMock).toHaveBeenCalledTimes(1));
});
});
@@ -0,0 +1,221 @@
import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest';
const resetAndReloadMock = vi.fn();
const getModelApiClientMock = vi.fn();
vi.mock('../../../static/js/api/modelApiFactory.js', () => ({
getModelApiClient: getModelApiClientMock,
resetAndReload: resetAndReloadMock,
}));
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
showToast: vi.fn(),
openCivitaiByMetadata: vi.fn(),
updatePanelPositions: vi.fn(),
}));
vi.mock('../../../static/js/managers/DownloadManager.js', () => ({
downloadManager: { showDownloadModal: vi.fn() },
}));
vi.mock('../../../static/js/components/SidebarManager.js', () => ({
sidebarManager: {
setHostPageControls: vi.fn(),
initialize: vi.fn(async () => {}),
refresh: vi.fn(async () => {}),
cleanup: vi.fn(),
isInitialized: false,
},
}));
vi.mock('../../../static/js/components/alphabet/index.js', () => ({
createAlphabetBar: vi.fn(() => ({ destroy: vi.fn() })),
}));
vi.mock('../../../static/js/utils/updateCheckHelpers.js', () => ({
performModelUpdateCheck: vi.fn(async () => ({ status: 'success', displayName: 'LoRA', records: [] })),
}));
beforeEach(() => {
vi.resetModules();
vi.clearAllMocks();
localStorage.clear();
sessionStorage.clear();
resetAndReloadMock.mockResolvedValue(undefined);
getModelApiClientMock.mockReturnValue({});
global.fetch = vi.fn().mockResolvedValue({
ok: true,
json: async () => ({ success: true, base_models: [] }),
});
});
afterEach(() => {
delete window.bulkManager;
delete window.modelDuplicatesManager;
delete global.fetch;
});
function renderControlsDom(pageKey) {
document.body.dataset.page = pageKey;
document.body.innerHTML = `
<div class="controls">
<div id="excludedViewBanner" class="excluded-view-banner hidden">
<button id="excludedViewBackBtn">Back</button>
</div>
<div class="actions">
<div class="action-buttons">
<div class="control-group">
<select id="sortSelect">
<option value="name:asc">Name Asc</option>
<option value="name:desc">Name Desc</option>
<option value="random">Randomize (shuffle)</option>
</select>
</div>
<div class="control-group dropdown-group">
<button data-action="refresh" class="dropdown-main"></button>
<button class="dropdown-toggle"></button>
<div class="dropdown-menu">
<div class="dropdown-item" data-action="full-rebuild"></div>
</div>
</div>
<div class="control-group">
<button data-action="fetch"></button>
</div>
<div class="control-group">
<button data-action="download"></button>
</div>
<div class="control-group">
<button data-action="bulk"></button>
</div>
<div class="control-group">
<button data-action="find-duplicates"></button>
</div>
<div class="control-group">
<button id="favoriteFilterBtn" class="favorite-filter"></button>
</div>
<div class="control-group dropdown-group update-filter-group">
<button id="updateFilterBtn" class="dropdown-main update-filter" aria-busy="false">
<span>Updates</span>
</button>
<button id="updateFilterMenuToggle" class="dropdown-toggle"></button>
<div class="dropdown-menu">
<div id="checkUpdatesMenuItem" class="dropdown-item" data-action="check-updates">
<span>Check updates</span>
</div>
</div>
</div>
</div>
</div>
</div>
<div id="customFilterIndicator" class="control-group hidden">
<div class="filter-active">
<span class="customFilterText" title=""></span>
<i class="fas fa-times-circle clear-filter"></i>
</div>
</div>
<div id="breadcrumbContainer"></div>
<div id="duplicatesBanner" style="display: none;"></div>
<div class="alphabet-bar-container"></div>
`;
}
async function createControls() {
const stateModule = await import('../../../static/js/state/index.js');
stateModule.initPageState('loras');
const { LorasControls } = await import('../../../static/js/components/controls/LorasControls.js');
return { stateModule, controls: new LorasControls() };
}
describe('Random sort option', () => {
it('generates a seeded sort value when Random is picked', async () => {
renderControlsDom('loras');
const { controls } = await createControls();
const sortSelect = document.getElementById('sortSelect');
const randomOpt = sortSelect.querySelector('option[value="random"]');
sortSelect.value = 'random';
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
await Promise.resolve();
expect(controls.pageState.sortBy).toMatch(/^random:[a-z0-9]+$/);
expect(localStorage.getItem('lora_manager_loras_sort')).toBe(controls.pageState.sortBy);
expect(randomOpt.value).toBe(controls.pageState.sortBy);
expect(sortSelect.value).toBe(controls.pageState.sortBy);
expect(resetAndReloadMock).toHaveBeenCalled();
});
it('reshuffles with a fresh seed every time Random is picked again', async () => {
renderControlsDom('loras');
const { controls } = await createControls();
const sortSelect = document.getElementById('sortSelect');
const randomOpt = sortSelect.querySelector('option[value="random"]');
// First pick
sortSelect.value = 'random';
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
await Promise.resolve();
const firstSeed = controls.pageState.sortBy;
// Second pick: the option now carries the seeded value, like a menu click
sortSelect.value = randomOpt.value;
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
await Promise.resolve();
expect(controls.pageState.sortBy).toMatch(/^random:[a-z0-9]+$/);
expect(controls.pageState.sortBy).not.toBe(firstSeed);
});
it('restores a persisted seeded random sort on load', async () => {
renderControlsDom('loras');
const savedSort = 'random:persistedseed';
localStorage.setItem('lora_manager_loras_sort', savedSort);
const { controls } = await createControls();
const sortSelect = document.getElementById('sortSelect');
expect(controls.pageState.sortBy).toBe(savedSort);
expect(sortSelect.value).toBe(savedSort);
expect(sortSelect.querySelector('option[value="random:persistedseed"]')).not.toBeNull();
});
it('applies a non-random sort back to the plain random option', async () => {
renderControlsDom('loras');
const { controls } = await createControls();
const sortSelect = document.getElementById('sortSelect');
const randomOpt = sortSelect.querySelector('option[value="random"]');
// Seed a random sort, then switch to a normal sort
sortSelect.value = 'random';
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
await Promise.resolve();
controls.applySortToSelect('name:desc');
expect(sortSelect.value).toBe('name:desc');
expect(randomOpt.value).toBe('random');
});
it('resets the seeded option when switching away from Random via the dropdown change handler', async () => {
renderControlsDom('loras');
const { controls } = await createControls();
const sortSelect = document.getElementById('sortSelect');
const randomOpt = sortSelect.querySelector('option[value="random"]');
// Pick Random: the option is now seeded
sortSelect.value = 'random';
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
await Promise.resolve();
expect(randomOpt.value).toMatch(/^random:[a-z0-9]+$/);
// Switch to a non-random sort through the change handler (as a menu
// click does); the option must go back to the plain "random" value
sortSelect.value = 'name:desc';
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
await Promise.resolve();
expect(controls.pageState.sortBy).toBe('name:desc');
expect(sortSelect.value).toBe('name:desc');
expect(randomOpt.value).toBe('random');
});
});
@@ -0,0 +1,68 @@
import { describe, it, beforeEach, expect } from 'vitest';
import { initSortDropdown } from '../../../static/js/components/controls/SortDropdown.js';
function renderSortDropdownDom() {
document.body.innerHTML = `
<div class="sort-dropdown-group">
<select id="sortSelect">
<option value="name:asc">Name Asc</option>
<option value="name:desc">Name Desc</option>
<option value="random" selected>Randomize (shuffle)</option>
</select>
<button class="sort-trigger" type="button">
<span class="sort-trigger__label"></span>
</button>
<div class="sort-dropdown-menu"></div>
</div>
`;
return {
select: document.getElementById('sortSelect'),
menu: document.querySelector('.sort-dropdown-menu'),
label: document.querySelector('.sort-trigger__label'),
};
}
describe('SortDropdown menu sync', () => {
let select;
let menu;
let label;
beforeEach(() => {
({ select, menu, label } = renderSortDropdownDom());
initSortDropdown(select);
});
it('rebuilds the menu and highlights the selected item when an option value attribute changes', async () => {
// The seeded Random option gets a new value each time it is picked.
// The select's value getter follows the selected option's new value.
const randomOpt = select.querySelector('option[value="random"]');
randomOpt.value = 'random:abc123';
await Promise.resolve();
const items = [...menu.querySelectorAll('.sort-option')];
expect(items.map((el) => el.dataset.value)).toContain('random:abc123');
const seededItem = items.find((el) => el.dataset.value === 'random:abc123');
expect(seededItem.classList.contains('is-selected')).toBe(true);
expect(label.textContent).toBe('Randomize (shuffle)');
});
it('drops the stale seeded item and re-selects the plain random item when the option is reset', async () => {
const randomOpt = select.querySelector('option[value="random"]');
randomOpt.value = 'random:abc123';
await Promise.resolve();
// The rebuild must have happened: the seeded item is in the menu
const seededItems = [...menu.querySelectorAll('.sort-option')]
.filter((el) => el.dataset.value === 'random:abc123');
expect(seededItems).toHaveLength(1);
// PageControls resets the option to "random" when switching away
randomOpt.value = 'random';
await Promise.resolve();
const items = [...menu.querySelectorAll('.sort-option')];
expect(items.map((el) => el.dataset.value)).not.toContain('random:abc123');
const randomItem = items.find((el) => el.dataset.value === 'random');
expect(randomItem.classList.contains('is-selected')).toBe(true);
});
});
@@ -0,0 +1,195 @@
import { beforeEach, describe, expect, it, vi } from "vitest";
const { APP_MODULE, EXTENSION_MODULE, appMock, registeredExtensions } =
vi.hoisted(() => {
const registeredExtensions = [];
const appMock = {
configuringGraph: false,
registerExtension: (ext) => registeredExtensions.push(ext),
};
return {
APP_MODULE: new URL("../../../scripts/app.js", import.meta.url).pathname,
EXTENSION_MODULE: new URL(
"../../../web/comfyui/lora_stack_dynamic_inputs.js",
import.meta.url
).pathname,
appMock,
registeredExtensions,
};
});
vi.mock(APP_MODULE, () => ({
app: appMock,
}));
describe("Lora Stack Combiner dynamic inputs", () => {
let extension;
beforeEach(async () => {
vi.resetModules();
registeredExtensions.length = 0;
appMock.configuringGraph = false;
await import(EXTENSION_MODULE);
extension = registeredExtensions.find(
(ext) => ext.name === "Comfy.LoraManager.LoraStackCombiner"
);
expect(extension).toBeDefined();
});
function createNodeType() {
const nodeType = { prototype: {} };
extension.beforeRegisterNodeDef(
nodeType,
{ name: "Lora Stack Combiner (LoraManager)" },
appMock
);
return nodeType;
}
function createNode(inputs = []) {
const node = {
comfyClass: "Lora Stack Combiner (LoraManager)",
inputs: inputs.map((name) => ({ name, type: "LORA_STACK" })),
addInput: vi.fn(function (name, type, opts) {
this.inputs.push({ name, type, ...opts });
}),
removeInput: vi.fn(function (index) {
this.inputs.splice(index, 1);
}),
};
return node;
}
function makeLinkInfo() {
return { id: 999, origin_id: 1, target_id: 2 };
}
it("adds a third input when the last slot gets connected", () => {
const nodeType = createNodeType();
const node = createNode(["lora_stack1", "lora_stack2"]);
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
node.onConnectionsChange(1, 1, true, makeLinkInfo());
expect(node.inputs.map((input) => input.name)).toEqual([
"lora_stack1",
"lora_stack2",
"lora_stack3",
]);
});
it("does not add an input when a non-last slot gets connected", () => {
const nodeType = createNodeType();
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
node.onConnectionsChange(1, 0, true, makeLinkInfo());
expect(node.inputs.map((input) => input.name)).toEqual([
"lora_stack1",
"lora_stack2",
"lora_stack3",
]);
});
it("removes a disconnected middle slot and renumbers", () => {
// Simulates a real LiteGraph disconnect event: it fires only for slots that
// had a link, and input.link has already been cleared before the event fires.
const nodeType = createNodeType();
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
node.inputs[0].link = 11;
node.inputs[1].link = null; // slot 2 was just disconnected
node.inputs[2].link = 13;
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
node.onConnectionsChange(1, 1, false, makeLinkInfo());
expect(node.inputs.map((input) => input.name)).toEqual([
"lora_stack1",
"lora_stack2",
]);
});
it("keeps the last slot when it is disconnected", () => {
const nodeType = createNodeType();
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
node.inputs[0].link = 11;
node.inputs[1].link = 12;
node.inputs[2].link = null; // last slot was just disconnected
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
node.onConnectionsChange(1, 2, false, makeLinkInfo());
expect(node.inputs.map((input) => input.name)).toEqual([
"lora_stack1",
"lora_stack2",
"lora_stack3",
]);
expect(node.removeInput).not.toHaveBeenCalled();
});
it("keeps at least two inputs when disconnecting", () => {
const nodeType = createNodeType();
const node = createNode(["lora_stack1", "lora_stack2"]);
node.inputs[0].link = 11;
node.inputs[1].link = null; // slot 2 was just disconnected
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
node.onConnectionsChange(1, 1, false, makeLinkInfo());
expect(node.inputs.map((input) => input.name)).toEqual([
"lora_stack1",
"lora_stack2",
]);
expect(node.removeInput).not.toHaveBeenCalled();
});
it("does nothing while the graph is being configured", () => {
appMock.configuringGraph = true;
const nodeType = createNodeType();
const node = createNode(["lora_stack1", "lora_stack2"]);
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
node.onConnectionsChange(1, 1, true, makeLinkInfo());
expect(node.inputs.map((input) => input.name)).toEqual([
"lora_stack1",
"lora_stack2",
]);
expect(node.addInput).not.toHaveBeenCalled();
});
it("leaves legacy lora_stack_a/b inputs untouched", () => {
const nodeType = createNodeType();
const node = createNode(["lora_stack_a", "lora_stack_b"]);
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
node.onConnectionsChange(1, 0, true, makeLinkInfo());
expect(node.inputs.map((input) => input.name)).toEqual([
"lora_stack_a",
"lora_stack_b",
]);
expect(node.addInput).not.toHaveBeenCalled();
});
it("ensures two numbered inputs exist on creation", () => {
const node = createNode([]);
extension.nodeCreated(node, {});
expect(node.inputs.map((input) => input.name)).toEqual([
"lora_stack1",
"lora_stack2",
]);
});
it("does not add numbered inputs to legacy workflows", () => {
const node = createNode(["lora_stack_a", "lora_stack_b"]);
extension.nodeCreated(node, {});
expect(node.inputs.map((input) => input.name)).toEqual([
"lora_stack_a",
"lora_stack_b",
]);
});
});
@@ -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();
});
});
@@ -0,0 +1,186 @@
import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest';
import { state } from '../../../static/js/state/index.js';
import { MODEL_TYPES } from '../../../static/js/api/apiConfig.js';
import { eventManager } from '../../../static/js/utils/EventManager.js';
import { BulkManager } from '../../../static/js/managers/BulkManager.js';
function fire(type, init = {}) {
return new MouseEvent(type, { bubbles: true, cancelable: true, ...init });
}
describe('BulkManager marquee guards', () => {
beforeEach(() => {
vi.useFakeTimers();
// jsdom may not provide requestAnimationFrame; stub it so the auto-scroll loop is a no-op.
window.requestAnimationFrame = vi.fn();
window.cancelAnimationFrame = vi.fn();
eventManager.cleanup();
state.currentPageType = MODEL_TYPES.LORA;
state.bulkMode = false;
state.selectedModels.clear();
document.body.innerHTML = '<div class="page-content"></div>';
const pageContent = document.querySelector('.page-content');
pageContent.getBoundingClientRect = () => ({
top: 0,
left: 0,
right: 1000,
bottom: 1000,
width: 1000,
height: 1000,
x: 0,
y: 0,
toJSON: () => ({}),
});
pageContent.scrollBy = vi.fn();
});
afterEach(() => {
eventManager.cleanup();
vi.useRealTimers();
document.body.innerHTML = '';
});
function createBulkManager() {
const bulk = new BulkManager();
bulk.initialize();
return bulk;
}
it('never starts a marquee when the left button is not held', () => {
const bulk = createBulkManager();
const pageContent = document.querySelector('.page-content');
pageContent.dispatchEvent(fire('mousedown', { button: 0, clientX: 10, clientY: 10 }));
document.dispatchEvent(fire('mousemove', { buttons: 0, clientX: 50, clientY: 50 }));
expect(bulk.mouseDownTime).toBe(0);
expect(bulk.isMarqueeActive).toBe(false);
expect(state.bulkMode).toBe(false);
expect(document.querySelector('.marquee-selection')).toBeNull();
});
it('requires holding the left button for the drag delay before starting a marquee', () => {
const bulk = createBulkManager();
const pageContent = document.querySelector('.page-content');
pageContent.dispatchEvent(fire('mousedown', { button: 0, clientX: 10, clientY: 10 }));
// Fast movement: far enough, but too soon after mousedown.
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 30, clientY: 10 }));
expect(state.bulkMode).toBe(false);
expect(bulk.isMarqueeActive).toBe(false);
// Once the hold time has elapsed, the same drag qualifies.
vi.advanceTimersByTime(100);
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 35, clientY: 12 }));
expect(state.bulkMode).toBe(true);
expect(bulk.isMarqueeActive).toBe(true);
expect(document.querySelector('.marquee-selection')).not.toBeNull();
});
it('ends an active marquee if the left button is released without a mouseup event', () => {
const bulk = createBulkManager();
bulk.mouseDownPosition = { x: 10, y: 10 };
bulk.startMarqueeSelection({}, true);
expect(state.bulkMode).toBe(true);
expect(document.querySelector('.marquee-selection')).not.toBeNull();
// No mouseup was dispatched; a plain move with the button released finalizes it.
document.dispatchEvent(fire('mousemove', { buttons: 0, clientX: 50, clientY: 50 }));
expect(bulk.isMarqueeActive).toBe(false);
expect(document.querySelector('.marquee-selection')).toBeNull();
expect(state.bulkMode).toBe(false); // zero selected -> auto-exit
});
it('treats a tiny marquee as an accidental click: clears selection and exits bulk mode', () => {
const bulk = createBulkManager();
const card = document.createElement('div');
card.className = 'model-card selected';
card.dataset.filepath = '/models/test.safetensors';
document.body.appendChild(card);
state.selectedModels.add('/models/test.safetensors');
bulk.mouseDownPosition = { x: 100, y: 100 };
bulk.startMarqueeSelection({}, true);
expect(state.bulkMode).toBe(true);
bulk.endMarqueeSelection({ clientX: 103, clientY: 104 });
expect(state.bulkMode).toBe(false);
expect(state.selectedModels.size).toBe(0);
expect(card.classList.contains('selected')).toBe(false);
});
it('keeps selection and bulk mode when the marquee is large enough', () => {
const bulk = createBulkManager();
const card = document.createElement('div');
card.className = 'model-card selected';
card.dataset.filepath = '/models/test.safetensors';
document.body.appendChild(card);
state.selectedModels.add('/models/test.safetensors');
bulk.mouseDownPosition = { x: 100, y: 100 };
bulk.startMarqueeSelection({}, true);
bulk.endMarqueeSelection({ clientX: 130, clientY: 140 });
expect(state.bulkMode).toBe(true);
expect(state.selectedModels.has('/models/test.safetensors')).toBe(true);
expect(card.classList.contains('selected')).toBe(true);
});
it('keeps auto-scroll marquee selections when the pointer only moved a few pixels', () => {
const bulk = createBulkManager();
const pageContent = document.querySelector('.page-content');
// Card just below the press point in document coordinates.
const card = document.createElement('div');
card.className = 'model-card';
card.dataset.filepath = '/models/off-screen.safetensors';
card.getBoundingClientRect = () => ({
top: 950,
left: 400,
right: 600,
bottom: 1050,
width: 200,
height: 100,
x: 400,
y: 950,
toJSON: () => ({}),
});
document.body.appendChild(card);
pageContent.dispatchEvent(fire('mousedown', { button: 0, clientX: 500, clientY: 900 }));
vi.advanceTimersByTime(100);
// Small pointer move: enough to start the marquee, but under minMarqueeSize.
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 506, clientY: 906 }));
expect(bulk.isMarqueeActive).toBe(true);
// Auto-scroll grows the document-space box while the pointer stays nearly still.
pageContent.scrollTop = 200;
card.getBoundingClientRect = () => ({
top: 750,
left: 400,
right: 600,
bottom: 850,
width: 200,
height: 100,
x: 400,
y: 750,
toJSON: () => ({}),
});
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 506, clientY: 906 }));
expect(state.selectedModels.has('/models/off-screen.safetensors')).toBe(true);
// Release: the client-space box is tiny, but the document-space box is not.
document.dispatchEvent(fire('mouseup', { button: 0, clientX: 506, clientY: 906 }));
expect(state.selectedModels.has('/models/off-screen.safetensors')).toBe(true);
expect(state.bulkMode).toBe(true);
});
});
@@ -0,0 +1,317 @@
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest';
const {
DOWNLOAD_MANAGER_MODULE,
MODAL_MANAGER_MODULE,
UI_HELPERS_MODULE,
STATE_MODULE,
LOADING_MANAGER_MODULE,
API_FACTORY_MODULE,
STORAGE_HELPERS_MODULE,
FOLDER_TREE_MANAGER_MODULE,
I18N_HELPERS_MODULE,
SUMMARY_MODULE,
mockApiClient,
mockLoadingManager,
showToastMock,
showDownloadBatchSummaryMock,
resetAndReloadMock,
} = vi.hoisted(() => {
// Shared API client returned by the mocked getModelApiClient factory.
const mockApiClient = {
apiConfig: {
config: {
displayName: 'LoRA',
singularName: 'lora',
},
},
downloadModel: vi.fn(),
downloadHfModel: vi.fn(),
cancelDownload: vi.fn(),
};
// Shared loading manager served both via state.loadingManager and the
// LoadingManager constructor mock.
const mockLoadingManager = {
showSimpleLoading: vi.fn(),
hide: vi.fn(),
restoreProgressBar: vi.fn(),
showDownloadProgress: vi.fn(() => vi.fn()),
setStatus: vi.fn(),
showCancelButton: vi.fn(),
};
return {
DOWNLOAD_MANAGER_MODULE: new URL('../../../static/js/managers/DownloadManager.js', import.meta.url).pathname,
MODAL_MANAGER_MODULE: new URL('../../../static/js/managers/ModalManager.js', import.meta.url).pathname,
UI_HELPERS_MODULE: new URL('../../../static/js/utils/uiHelpers.js', import.meta.url).pathname,
STATE_MODULE: new URL('../../../static/js/state/index.js', import.meta.url).pathname,
LOADING_MANAGER_MODULE: new URL('../../../static/js/managers/LoadingManager.js', import.meta.url).pathname,
API_FACTORY_MODULE: new URL('../../../static/js/api/modelApiFactory.js', import.meta.url).pathname,
STORAGE_HELPERS_MODULE: new URL('../../../static/js/utils/storageHelpers.js', import.meta.url).pathname,
FOLDER_TREE_MANAGER_MODULE: new URL('../../../static/js/components/FolderTreeManager.js', import.meta.url).pathname,
I18N_HELPERS_MODULE: new URL('../../../static/js/utils/i18nHelpers.js', import.meta.url).pathname,
SUMMARY_MODULE: new URL('../../../static/js/components/DownloadBatchSummaryModal.js', import.meta.url).pathname,
mockApiClient,
mockLoadingManager,
showToastMock: vi.fn(),
showDownloadBatchSummaryMock: vi.fn(),
resetAndReloadMock: vi.fn(),
};
});
vi.mock(MODAL_MANAGER_MODULE, () => ({
modalManager: {
showModal: vi.fn(),
closeModal: vi.fn(),
},
}));
vi.mock(UI_HELPERS_MODULE, () => ({
showToast: showToastMock,
}));
vi.mock(STATE_MODULE, () => ({
state: {
global: {
settings: {},
},
loadingManager: mockLoadingManager,
},
}));
vi.mock(LOADING_MANAGER_MODULE, () => ({
LoadingManager: vi.fn(() => mockLoadingManager),
}));
vi.mock(API_FACTORY_MODULE, () => ({
getModelApiClient: vi.fn(() => mockApiClient),
resetAndReload: resetAndReloadMock,
}));
vi.mock(STORAGE_HELPERS_MODULE, () => ({
getStorageItem: vi.fn((_key, defaultValue) => defaultValue),
setStorageItem: vi.fn(),
}));
vi.mock(FOLDER_TREE_MANAGER_MODULE, () => ({
FolderTreeManager: vi.fn(() => ({
clearSelection: vi.fn(),
init: vi.fn(),
})),
}));
vi.mock(I18N_HELPERS_MODULE, () => ({
translate: vi.fn((_, __, fallback) => fallback ?? ''),
}));
vi.mock(SUMMARY_MODULE, () => ({
showDownloadBatchSummary: showDownloadBatchSummaryMock,
}));
/**
* Fake WebSocket used by executeBatchDownload. Resolves `onopen` on the
* microtask queue right after construction (which happens after the real
* code has assigned `onopen`), so the open promise resolves deterministically
* without real timers.
*/
class FakeWebSocket {
static instances = [];
constructor(url) {
this.url = url;
this.onopen = null;
this.onmessage = null;
this.onerror = null;
this.close = vi.fn();
FakeWebSocket.instances.push(this);
queueMicrotask(() => {
if (this.onopen) this.onopen();
});
}
static get lastInstance() {
return FakeWebSocket.instances[FakeWebSocket.instances.length - 1];
}
}
describe('DownloadManager batch download summary flow', () => {
let DownloadManager;
let manager;
const options = { modelRoot: '/models/loras', targetFolder: '', useDefaultPaths: true };
const makeItem = (modelId, versionId, name) => ({
modelId,
displayName: name,
selectedVersion: { id: versionId, name, existsLocally: false },
});
const item0 = makeItem('111', 'v1', 'Model A');
const item1 = makeItem('222', 'v2', 'Model B');
beforeEach(async () => {
document.body.innerHTML = '';
FakeWebSocket.instances = [];
// Reset the shared mocks so mockResolvedValueOnce queues and call
// history never leak between tests.
mockApiClient.downloadModel.mockReset();
mockApiClient.downloadHfModel.mockReset();
mockApiClient.cancelDownload.mockReset();
showToastMock.mockClear();
showDownloadBatchSummaryMock.mockClear();
resetAndReloadMock.mockClear();
mockLoadingManager.hide.mockClear();
mockLoadingManager.setStatus.mockClear();
mockLoadingManager.showCancelButton.mockClear();
mockLoadingManager.showDownloadProgress.mockClear();
vi.stubGlobal('WebSocket', FakeWebSocket);
vi.resetModules();
({ DownloadManager } = await import(DOWNLOAD_MANAGER_MODULE));
manager = new DownloadManager();
// The constructor leaves apiClient null; executeBatchDownload reads it
// directly, so point it at the shared mocked client.
manager.apiClient = mockApiClient;
});
afterEach(() => {
document.body.innerHTML = '';
vi.unstubAllGlobals();
});
it('shows the success toast when every item downloads successfully', async () => {
mockApiClient.downloadModel.mockResolvedValue({ success: true });
await manager.executeBatchDownload([item0, item1], options);
expect(mockApiClient.downloadModel).toHaveBeenCalledTimes(2);
// Each item is downloaded with its own modelId + versionId.
expect(mockApiClient.downloadModel.mock.calls[0][0]).toBe('111');
expect(mockApiClient.downloadModel.mock.calls[0][1]).toBe('v1');
expect(mockApiClient.downloadModel.mock.calls[1][0]).toBe('222');
expect(mockApiClient.downloadModel.mock.calls[1][1]).toBe('v2');
expect(showDownloadBatchSummaryMock).not.toHaveBeenCalled();
expect(showToastMock).toHaveBeenCalledTimes(1);
expect(showToastMock).toHaveBeenCalledWith('toast.loras.allDownloadSuccessful', { count: 2 }, 'success');
expect(resetAndReloadMock).toHaveBeenCalledWith(true);
});
it('shows a partial-failure summary when some items fail', async () => {
// The failing item has no displayName/filename, so the resolved entry
// name falls back to the selected version name.
const unnamedItem = { modelId: '333', selectedVersion: { id: 'v3', name: 'V3', existsLocally: false } };
mockApiClient.downloadModel
.mockResolvedValueOnce({ success: false, error: 'rate limited' })
.mockResolvedValueOnce({ success: true });
await manager.executeBatchDownload([unnamedItem, item1], options);
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
expect(summary.total).toBe(2);
expect(summary.completed).toBe(1);
expect(summary.failedItems).toHaveLength(1);
expect(summary.failedItems[0].item).toBe(unnamedItem);
expect(summary.failedItems[0].error).toBe('rate limited');
// The resolved display name is carried on the failed entry.
expect(summary.failedItems[0].name).toBe('V3');
expect(summary.onRetry).toEqual(expect.any(Function));
// No success toast and no downloadPartialSuccess toast for this path.
expect(showToastMock).not.toHaveBeenCalledWith('toast.loras.allDownloadSuccessful', expect.anything(), 'success');
expect(showToastMock).not.toHaveBeenCalledWith('toast.loras.downloadPartialSuccess', expect.anything(), expect.anything());
});
it('shows an all-failed summary when every item fails', async () => {
mockApiClient.downloadModel.mockResolvedValue({ success: false, error: 'x' });
await manager.executeBatchDownload([item0, item1], options);
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
expect(summary.total).toBe(2);
expect(summary.completed).toBe(0);
expect(summary.failedItems).toHaveLength(2);
expect(summary.failedItems[0].item).toBe(item0);
expect(summary.failedItems[1].item).toBe(item1);
expect(showToastMock).not.toHaveBeenCalledWith('toast.loras.allDownloadSuccessful', expect.anything(), expect.anything());
});
it('records the error message when downloadModel rejects', async () => {
// The item carries a filename but no displayName, so the resolved entry
// name comes from the filename.
const filenameItem = { modelId: '444', filename: 'model.safetensors', selectedVersion: { id: 'v4' } };
mockApiClient.downloadModel.mockRejectedValue(new Error('network down'));
await manager.executeBatchDownload([filenameItem], options);
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
expect(summary.total).toBe(1);
expect(summary.completed).toBe(0);
expect(summary.failedItems).toHaveLength(1);
expect(summary.failedItems[0].item).toBe(filenameItem);
expect(summary.failedItems[0].error).toBe('network down');
expect(summary.failedItems[0].name).toBe('model.safetensors');
});
it('retries the failed subset through onRetry with unwrapped items', async () => {
mockApiClient.downloadModel
.mockResolvedValueOnce({ success: false, error: 'rate limited' })
.mockResolvedValueOnce({ success: true });
await manager.executeBatchDownload([item0, item1], options);
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(1);
const summary = showDownloadBatchSummaryMock.mock.calls[0][0];
expect(summary.failedItems).toHaveLength(1);
// Retry the exact failed subset returned by the summary. The onRetry
// callback unwraps the { item, error } entries back into raw model items
// before re-running executeBatchDownload. Make the retried item fail
// again so a second summary is produced.
mockApiClient.downloadModel.mockResolvedValueOnce({ success: false, error: 'still rate limited' });
await summary.onRetry(summary.failedItems);
// downloadModel is called a third time — only for the failed item (item0),
// NOT for the item that already succeeded (item1).
expect(mockApiClient.downloadModel).toHaveBeenCalledTimes(3);
const retryCall = mockApiClient.downloadModel.mock.calls[2];
expect(retryCall[0]).toBe(item0.modelId);
expect(retryCall[1]).toBe(item0.selectedVersion.id);
// A fresh summary is produced for the retry run (call count 1 -> 2).
expect(showDownloadBatchSummaryMock).toHaveBeenCalledTimes(2);
const retrySummary = showDownloadBatchSummaryMock.mock.calls[1][0];
expect(retrySummary.total).toBe(1);
expect(retrySummary.completed).toBe(0);
expect(retrySummary.failedItems).toHaveLength(1);
expect(retrySummary.failedItems[0].item).toBe(item0);
expect(retrySummary.failedItems[0].error).toBe('still rate limited');
});
it('stops the batch without showing a summary when cancelled before downloads start', async () => {
const downloadPromise = manager.executeBatchDownload([item0, item1], options);
// showCancelButton captured the cancel callback synchronously. Invoking it
// sets `cancelled = true` before the download loop runs (the loop only
// starts after the WebSocket open promise resolves on the microtask queue).
const cancelCallback = mockLoadingManager.showCancelButton.mock.calls[0][0];
const cancelPromise = cancelCallback();
await Promise.all([downloadPromise, cancelPromise]);
expect(mockApiClient.downloadModel).not.toHaveBeenCalled();
expect(showDownloadBatchSummaryMock).not.toHaveBeenCalled();
expect(showToastMock).toHaveBeenCalledWith(
'toast.downloads.downloadStopped',
expect.anything(),
'info',
expect.stringContaining('Download cancelled')
);
expect(resetAndReloadMock).toHaveBeenCalledWith(true);
});
});
@@ -61,3 +61,106 @@ describe('UpdateService passive checks', () => {
expect(fetchMock).toHaveBeenCalledWith('/api/lm/check-updates?nightly=false'); expect(fetchMock).toHaveBeenCalledWith('/api/lm/check-updates?nightly=false');
}); });
}); });
describe('UpdateService nightly notification throttling', () => {
let fetchMock;
let updateToggle;
let updateBadge;
function stubUpdateBadgeDom() {
updateToggle = document.createElement('div');
updateToggle.className = 'update-toggle';
updateBadge = document.createElement('span');
updateBadge.className = 'update-badge';
updateToggle.appendChild(updateBadge);
document.body.appendChild(updateToggle);
vi.spyOn(document, 'querySelector').mockImplementation((selector) => {
if (selector === '.update-toggle') return updateToggle;
if (selector === '.update-toggle .update-badge') return updateBadge;
return null;
});
}
function makeUpdateResponse(channel) {
return {
success: true,
current_version: 'v1.0.0',
latest_version: channel === 'nightly' ? 'main-abc1234' : 'v1.1.0',
update_available: true,
git_info: { short_hash: 'abc123' },
has_git: true,
nightly: channel === 'nightly',
changelog: ['test: change'],
releases: [],
behind_by: 3,
commit_date: '2026-07-31',
};
}
beforeEach(() => {
fetchMock = vi.fn().mockResolvedValue(createFetchResponse(makeUpdateResponse('release')));
global.fetch = fetchMock;
stubUpdateBadgeDom();
});
afterEach(() => {
vi.restoreAllMocks();
delete global.fetch;
});
it('shows the nightly badge once and keeps it visible for the session', async () => {
stubSettingsUpdateChannel('nightly');
fetchMock.mockResolvedValue(createFetchResponse(makeUpdateResponse('nightly')));
const service = new UpdateService();
service.updateNotificationsEnabled = true;
await service.checkForUpdates({ force: true });
expect(service.updateAvailable).toBe(true);
expect(service.nightlyBadgeShown).toBe(true);
expect(service.nightlyNotifyDate).toBe(service._getTodayKey());
expect(updateBadge.classList.contains('visible')).toBe(true);
// A repeated check within the same session keeps the badge visible.
await service.checkForUpdates({ force: true });
expect(updateBadge.classList.contains('visible')).toBe(true);
});
it('suppresses the nightly badge on a later session in the same day', async () => {
stubSettingsUpdateChannel('nightly');
fetchMock.mockResolvedValue(createFetchResponse(makeUpdateResponse('nightly')));
const firstService = new UpdateService();
firstService.updateNotificationsEnabled = true;
await firstService.checkForUpdates({ force: true });
expect(updateBadge.classList.contains('visible')).toBe(true);
// Simulate a fresh page session on the same calendar day.
const secondService = new UpdateService();
secondService.updateNotificationsEnabled = true;
await secondService.checkForUpdates({ force: true });
expect(secondService.updateAvailable).toBe(true);
expect(secondService.nightlyBadgeShown).toBe(false);
expect(updateBadge.classList.contains('visible')).toBe(false);
});
it('is not affected by the daily limit on the release channel', async () => {
stubSettingsUpdateChannel('release');
fetchMock.mockResolvedValue(createFetchResponse(makeUpdateResponse('release')));
const firstService = new UpdateService();
firstService.updateNotificationsEnabled = true;
await firstService.checkForUpdates({ force: true });
expect(updateBadge.classList.contains('visible')).toBe(true);
const secondService = new UpdateService();
secondService.updateNotificationsEnabled = true;
await secondService.checkForUpdates({ force: true });
expect(secondService.updateAvailable).toBe(true);
expect(updateBadge.classList.contains('visible')).toBe(true);
});
});
+114 -1
View File
@@ -1,4 +1,11 @@
from py.nodes.lora_stack_combiner import LoraStackCombinerLM import types
import pytest
from py.nodes.lora_stack_combiner import (
LoraStackCombinerLM,
_LoraStackOptionalInputs,
)
def test_combine_stacks_preserves_order(): def test_combine_stacks_preserves_order():
@@ -49,3 +56,109 @@ def test_combine_stacks_allows_duplicate_entries():
(combined_stack,) = node.combine_stacks([duplicate_entry], [duplicate_entry]) (combined_stack,) = node.combine_stacks([duplicate_entry], [duplicate_entry])
assert combined_stack == [duplicate_entry, duplicate_entry] assert combined_stack == [duplicate_entry, duplicate_entry]
def test_combine_stacks_returns_empty_when_both_unconnected():
node = LoraStackCombinerLM()
(combined_stack,) = node.combine_stacks()
assert combined_stack == []
def test_combine_stacks_returns_other_when_one_unconnected():
node = LoraStackCombinerLM()
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
(combined_stack_a,) = node.combine_stacks(lora_stack1=stack_a)
(combined_stack_b,) = node.combine_stacks(lora_stack2=stack_a)
assert combined_stack_a == stack_a
assert combined_stack_b == stack_a
def test_combine_stacks_with_dynamic_third_slot():
node = LoraStackCombinerLM()
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
(combined_stack,) = node.combine_stacks(
lora_stack1=stack_a, lora_stack2=stack_b, lora_stack3=stack_c
)
assert combined_stack == stack_a + stack_b + stack_c
def test_combine_stacks_orders_by_slot_number_not_call_order():
node = LoraStackCombinerLM()
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
(combined_stack,) = node.combine_stacks(
lora_stack3=stack_c, lora_stack2=stack_b, lora_stack1=stack_a
)
assert combined_stack == stack_a + stack_b + stack_c
def test_combine_stacks_accepts_only_dynamic_slot():
node = LoraStackCombinerLM()
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
(combined_stack,) = node.combine_stacks(lora_stack3=stack_c)
assert combined_stack == stack_c
def test_combine_stacks_handles_legacy_input_names():
node = LoraStackCombinerLM()
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
(combined_stack,) = node.combine_stacks(lora_stack_a=stack_a, lora_stack_b=stack_b)
assert combined_stack == stack_a + stack_b
def test_input_types_exposes_two_default_slots():
input_types = LoraStackCombinerLM.INPUT_TYPES()
assert set(input_types["optional"]) == {"lora_stack1", "lora_stack2"}
assert input_types["optional"]["lora_stack1"][0] == "LORA_STACK"
assert input_types["optional"]["lora_stack2"][0] == "LORA_STACK"
def test_input_types_recognizes_dynamic_slots_from_get_input_info(monkeypatch):
frames = [None, None, types.SimpleNamespace(function="get_input_info")]
monkeypatch.setattr(
"py.nodes.lora_stack_combiner.inspect.stack", lambda: frames
)
input_types = LoraStackCombinerLM.INPUT_TYPES()
optional = input_types["optional"]
assert "lora_stack3" in optional
assert optional["lora_stack3"][0] == "LORA_STACK"
assert "lora_stack25" in optional
assert optional["lora_stack25"][0] == "LORA_STACK"
def test_lora_stack_optional_inputs_proxy():
proxy = _LoraStackOptionalInputs({"lora_stack1": ("LORA_STACK", {})})
assert "lora_stack1" in proxy
assert "lora_stack2" in proxy
assert "lora_stack10" in proxy
assert "lora_stack_a" in proxy
assert "lora_stack" not in proxy
assert "lora_stacka" not in proxy
assert "lora_stack_1" not in proxy
assert "text" not in proxy
assert proxy["lora_stack1"][0] == "LORA_STACK"
assert proxy["lora_stack5"][0] == "LORA_STACK"
with pytest.raises(KeyError):
proxy["not_a_stack"]
+43
View File
@@ -86,6 +86,41 @@ def test_save_image_skips_png_parameters_when_metadata_disabled_and_keeps_workfl
assert img.info["workflow"] == json.dumps(workflow) assert img.info["workflow"] == json.dumps(workflow)
def test_save_image_does_not_append_loras_to_prompt_by_default(monkeypatch, tmp_path):
_configure_save_paths(monkeypatch, tmp_path)
_configure_metadata(
monkeypatch,
{"prompt": "prompt text", "seed": 123, "loras": "<lora:foo:0.7>"},
)
node = SaveImageLM()
node.save_images([_make_image()], "ComfyUI", "png", id="node-1")
image_path = tmp_path / "sample_00001_.png"
with Image.open(image_path) as img:
assert "<lora:" not in img.info["parameters"]
assert img.info["parameters"] == "prompt text\nSeed: 123, Version: ComfyUI"
def test_save_image_appends_loras_to_prompt_when_enabled(monkeypatch, tmp_path):
_configure_save_paths(monkeypatch, tmp_path)
_configure_metadata(
monkeypatch,
{"prompt": "prompt text", "seed": 123, "loras": "<lora:foo:0.7>"},
)
node = SaveImageLM()
node.save_images(
[_make_image()], "ComfyUI", "png", id="node-1", add_loras_to_prompt=True
)
image_path = tmp_path / "sample_00001_.png"
with Image.open(image_path) as img:
assert img.info["parameters"] == (
"prompt text\n<lora:foo:0.7>\nSeed: 123, Version: ComfyUI"
)
def test_save_image_skips_jpeg_metadata_when_disabled(monkeypatch, tmp_path): def test_save_image_skips_jpeg_metadata_when_disabled(monkeypatch, tmp_path):
_configure_save_paths(monkeypatch, tmp_path) _configure_save_paths(monkeypatch, tmp_path)
_configure_metadata(monkeypatch, {"prompt": "prompt text", "seed": 123}) _configure_metadata(monkeypatch, {"prompt": "prompt text", "seed": 123})
@@ -451,6 +486,14 @@ class TestParameterDefaultConsistency:
assert SaveImageLM.save_images.__defaults__[5] == 0 assert SaveImageLM.save_images.__defaults__[5] == 0
assert SaveImageLM.process_image.__defaults__[7] == 0 assert SaveImageLM.process_image.__defaults__[7] == 0
def test_add_loras_to_prompt_defaults_are_consistent(self):
input_types = SaveImageLM.INPUT_TYPES()
optional = input_types["optional"]
assert optional["add_loras_to_prompt"][1]["default"] is False
assert SaveImageLM.save_images.__defaults__[-1] is False
assert SaveImageLM.process_image.__defaults__[-1] is False
def test_png_does_not_pass_webp_method_or_jpeg_subsampling(monkeypatch, tmp_path): def test_png_does_not_pass_webp_method_or_jpeg_subsampling(monkeypatch, tmp_path):
_configure_save_paths(monkeypatch, tmp_path) _configure_save_paths(monkeypatch, tmp_path)
+97 -5
View File
@@ -900,18 +900,28 @@ class FakeMetadataProvider:
async def get_model_versions(self, _model_id): async def get_model_versions(self, _model_id):
return {"modelVersions": [], "name": "", "type": "lora"} return {"modelVersions": [], "name": "", "type": "lora"}
async def get_user_models(self, _username): async def get_user_models(self, _username, cursor=None):
return [] return {"items": [], "nextCursor": None}
async def get_creator_model_count(self, _username):
return None
class FakeUserModelsProvider(FakeMetadataProvider): class FakeUserModelsProvider(FakeMetadataProvider):
def __init__(self, models): def __init__(self, models, next_cursor=None, estimated_total=None):
self.models = models self.models = models
self.next_cursor = next_cursor
self.estimated_total = estimated_total
self.received_usernames: list[str] = [] self.received_usernames: list[str] = []
self.received_cursors: list = []
async def get_user_models(self, username): async def get_user_models(self, username, cursor=None):
self.received_usernames.append(username) self.received_usernames.append(username)
return self.models self.received_cursors.append(cursor)
return {"items": self.models, "nextCursor": self.next_cursor}
async def get_creator_model_count(self, _username):
return self.estimated_total
async def fake_metadata_provider_factory(): async def fake_metadata_provider_factory():
@@ -1286,6 +1296,88 @@ async def test_get_civitai_user_models_requires_username():
assert "username" in payload["error"].lower() assert "username" in payload["error"].lower()
@pytest.mark.asyncio
async def test_get_civitai_user_models_returns_pagination_fields():
models = [
{
"id": 1,
"name": "Model A",
"type": "LORA",
"tags": [],
"modelVersions": [
{"id": 100, "name": "v1", "images": [{"url": "http://example.com/a.jpg"}]},
],
},
{
"id": 2,
"name": "Unsupported",
"type": "Other",
"modelVersions": [{"id": 200, "name": "v1"}],
},
]
provider = FakeUserModelsProvider(models, next_cursor="cursor-token", estimated_total=2140)
async def provider_factory():
return provider
handler = ModelLibraryHandler(
ServiceRegistryAdapter(
get_lora_scanner=fake_scanner_factory,
get_checkpoint_scanner=fake_scanner_factory,
get_embedding_scanner=fake_scanner_factory,
get_downloaded_version_history_service=fake_download_history_service_factory,
),
metadata_provider_factory=provider_factory,
)
response = await handler.get_civitai_user_models(
FakeRequest(query={"username": "pixel"})
)
payload = json.loads(response.text)
assert response.status == 200
assert payload["success"] is True
# modelCount only counts models surviving the type filter
assert payload["modelCount"] == 1
assert payload["nextCursor"] == "cursor-token"
assert payload["hasMore"] is True
# first page includes the estimated total
assert payload["estimatedTotal"] == 2140
assert provider.received_cursors == [None]
@pytest.mark.asyncio
async def test_get_civitai_user_models_passes_cursor_and_omits_estimate():
provider = FakeUserModelsProvider([], next_cursor=None, estimated_total=999)
async def provider_factory():
return provider
handler = ModelLibraryHandler(
ServiceRegistryAdapter(
get_lora_scanner=fake_scanner_factory,
get_checkpoint_scanner=fake_scanner_factory,
get_embedding_scanner=fake_scanner_factory,
get_downloaded_version_history_service=fake_download_history_service_factory,
),
metadata_provider_factory=provider_factory,
)
response = await handler.get_civitai_user_models(
FakeRequest(query={"username": "pixel", "cursor": "opaque-token"})
)
payload = json.loads(response.text)
assert response.status == 200
assert payload["success"] is True
assert payload["nextCursor"] is None
assert payload["hasMore"] is False
# cursor requests must not include the estimated total
assert payload["estimatedTotal"] is None
assert provider.received_cursors == ["opaque-token"]
def test_ensure_handler_mapping_caches_result(): def test_ensure_handler_mapping_caches_result():
call_records = [] call_records = []
+1 -1
View File
@@ -183,7 +183,7 @@ class FakeCache:
def __init__(self, items): def __init__(self, items):
self.items = list(items) self.items = list(items)
async def get_sorted_data(self, sort_key, order): async def get_sorted_data(self, sort_key, order, seed=None):
if sort_key == "name": if sort_key == "name":
data = sorted(self.items, key=lambda x: x["model_name"].lower()) data = sorted(self.items, key=lambda x: x["model_name"].lower())
if order == "desc": if order == "desc":
+142
View File
@@ -363,6 +363,148 @@ async def test_check_pending_models_handles_corrupted_progress_file(
assert result["pending_count"] == 1 assert result["pending_count"] == 1
@pytest.mark.asyncio
@pytest.mark.usefixtures("tmp_path")
async def test_check_pending_models_uses_bulk_folder_index_for_large_libraries(
monkeypatch: pytest.MonkeyPatch,
tmp_path,
settings_manager,
):
"""For >1000 candidates the pre-check scans the library root once instead of
probing every folder individually."""
ws_manager = RecordingWebSocketManager()
manager = download_module.DownloadManager(ws_manager=ws_manager)
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
# 1500 unprocessed models triggers the bulk lookup path
models = [
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
for i in range(1500)
]
# Create folders with files for the first 500 models
for i in range(500):
model_dir = tmp_path / f"{i:064x}"
model_dir.mkdir()
(model_dir / "image_0.png").write_text("data")
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
per_model_checks = 0
def counting_model_directory_has_files(path: str) -> bool:
nonlocal per_model_checks
per_model_checks += 1
return False
monkeypatch.setattr(
download_module,
"_model_directory_has_files",
counting_model_directory_has_files,
)
result = await manager.check_pending_models(["lora"])
assert result["success"] is True
assert result["total_models"] == 1500
assert result["pending_count"] == 1000
assert result["needs_download"] is True
# The per-folder check should not be used once we cross the threshold.
assert per_model_checks == 0
@pytest.mark.asyncio
@pytest.mark.usefixtures("tmp_path")
async def test_check_pending_models_uses_per_folder_check_for_small_candidate_sets(
monkeypatch: pytest.MonkeyPatch,
tmp_path,
settings_manager,
):
"""For <=1000 candidates the pre-check keeps the accurate per-folder path."""
ws_manager = RecordingWebSocketManager()
manager = download_module.DownloadManager(ws_manager=ws_manager)
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
models = [
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
for i in range(500)
]
# Create folders with files for the first 200 models
for i in range(200):
model_dir = tmp_path / f"{i:064x}"
model_dir.mkdir()
(model_dir / "image_0.png").write_text("data")
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
per_model_checks = 0
original_has_files = download_module._model_directory_has_files
def counting_model_directory_has_files(path: str) -> bool:
nonlocal per_model_checks
per_model_checks += 1
return original_has_files(path)
monkeypatch.setattr(
download_module,
"_model_directory_has_files",
counting_model_directory_has_files,
)
result = await manager.check_pending_models(["lora"])
assert result["success"] is True
assert result["total_models"] == 500
assert result["pending_count"] == 300
assert result["needs_download"] is True
# Per-folder path should run once per candidate.
assert per_model_checks == 500
@pytest.mark.asyncio
@pytest.mark.usefixtures("tmp_path")
async def test_check_pending_models_bulk_index_includes_legacy_folders(
monkeypatch: pytest.MonkeyPatch,
tmp_path,
settings_manager,
):
"""In multi-library mode the bulk index also scans the legacy root so models
whose folders have not been consolidated yet are not reported pending."""
ws_manager = RecordingWebSocketManager()
manager = download_module.DownloadManager(ws_manager=ws_manager)
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
monkeypatch.setitem(settings_manager.settings, "libraries", {"default": {}, "extra": {}})
monkeypatch.setitem(settings_manager.settings, "active_library", "extra")
# 1500 unprocessed models triggers the bulk lookup path
models = [
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
for i in range(1500)
]
# Folders live at the LEGACY root/<hash> path (not yet consolidated)
for i in range(500):
model_dir = tmp_path / f"{i:064x}"
model_dir.mkdir()
(model_dir / "image_0.png").write_text("data")
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
result = await manager.check_pending_models(["lora"])
assert result["success"] is True
assert result["total_models"] == 1500
assert result["pending_count"] == 1000
assert result["needs_download"] is True
@pytest.fixture @pytest.fixture
def settings_manager(): def settings_manager():
return get_settings_manager() return get_settings_manager()
+161
View File
@@ -35,9 +35,11 @@ class DummyDownloader:
def reset_singletons(): def reset_singletons():
CivitaiClient._instance = None CivitaiClient._instance = None
ModelMetadataProviderManager._instance = None ModelMetadataProviderManager._instance = None
civitai_client_module._creator_model_count_cache.clear()
yield yield
CivitaiClient._instance = None CivitaiClient._instance = None
ModelMetadataProviderManager._instance = None ModelMetadataProviderManager._instance = None
civitai_client_module._creator_model_count_cache.clear()
@pytest.fixture @pytest.fixture
@@ -622,3 +624,162 @@ async def test_get_image_info_handles_invalid_id(monkeypatch, downloader, caplog
assert result is None assert result is None
assert "Invalid image ID format" in caplog.text assert "Invalid image ID format" in caplog.text
async def test_get_user_models_requests_first_page_with_stable_params(downloader):
request_calls = []
async def fake_make_request(method, url, use_auth=True, **kwargs):
request_calls.append({"method": method, "url": url, "kwargs": kwargs})
return True, {
"items": [
{
"id": 1,
"modelVersions": [
{"id": 100, "images": [{"meta": {"comfy": {"x": 1}}}]}
],
}
],
"metadata": {"nextCursor": "next-token"},
}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
result = await client.get_user_models("pixel")
assert result is not None
assert result["nextCursor"] == "next-token"
assert len(result["items"]) == 1
# comfy metadata is still stripped
assert "comfy" not in result["items"][0]["modelVersions"][0]["images"][0]["meta"]
call = request_calls[0]
assert call["method"] == "GET"
assert call["url"] == "https://civitai.red/api/v1/models"
params = call["kwargs"]["params"]
assert params["username"] == "pixel"
assert params["nsfw"] == "true"
assert params["limit"] == 100
assert params["sort"] == "Newest"
assert params["period"] == "AllTime"
assert "cursor" not in params
async def test_get_user_models_passes_cursor_and_stringifies_next_cursor(downloader):
request_calls = []
async def fake_make_request(method, url, use_auth=True, **kwargs):
request_calls.append(kwargs)
return True, {"items": [], "metadata": {"nextCursor": 12345}}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
result = await client.get_user_models("pixel", cursor="opaque-token")
assert request_calls[0]["params"]["cursor"] == "opaque-token"
assert result == {"items": [], "nextCursor": "12345"}
async def test_get_user_models_without_next_cursor_returns_none_cursor(downloader):
async def fake_make_request(method, url, use_auth=True, **kwargs):
return True, {"items": [{"id": 1, "modelVersions": []}], "metadata": {}}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
result = await client.get_user_models("pixel")
assert result == {"items": [{"id": 1, "modelVersions": []}], "nextCursor": None}
async def test_get_user_models_failure_returns_none(downloader):
async def fake_make_request(method, url, use_auth=True, **kwargs):
return False, "500 server error"
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
result = await client.get_user_models("pixel")
assert result is None
async def test_get_creator_model_count_matches_exact_username(downloader):
request_calls = []
async def fake_make_request(method, url, use_auth=True, **kwargs):
request_calls.append({"url": url, "kwargs": kwargs})
return True, {
"items": [
{"username": "pixelart", "modelCount": 5},
{"username": "Pixel", "modelCount": 2140},
]
}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
count = await client.get_creator_model_count("pixel")
assert count == 2140
assert request_calls[0]["url"] == "https://civitai.red/api/v1/creators"
assert request_calls[0]["kwargs"]["params"] == {"query": "pixel", "limit": 10}
async def test_get_creator_model_count_without_exact_match_returns_none(downloader):
async def fake_make_request(method, url, use_auth=True, **kwargs):
return True, {"items": [{"username": "pixelart", "modelCount": 5}]}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
count = await client.get_creator_model_count("pixel")
assert count is None
async def test_get_creator_model_count_caches_results(downloader):
request_count = 0
async def fake_make_request(method, url, use_auth=True, **kwargs):
nonlocal request_count
request_count += 1
return True, {"items": [{"username": "pixel", "modelCount": 42}]}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
assert await client.get_creator_model_count("pixel") == 42
# case-insensitive cache key, second call served from cache
assert await client.get_creator_model_count("Pixel") == 42
assert request_count == 1
async def test_get_creator_model_count_caches_failures(downloader):
request_count = 0
async def fake_make_request(method, url, use_auth=True, **kwargs):
nonlocal request_count
request_count += 1
return False, "500 server error"
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
assert await client.get_creator_model_count("pixel") is None
assert await client.get_creator_model_count("pixel") is None
assert request_count == 1
async def test_get_creator_model_count_never_raises(downloader):
async def fake_make_request(method, url, use_auth=True, **kwargs):
return True, "unexpected non-dict payload"
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
assert await client.get_creator_model_count("pixel") is None
@@ -26,6 +26,7 @@ class StubScanner:
def __init__(self, models: list[dict]) -> None: def __init__(self, models: list[dict]) -> None:
self._cache = SimpleNamespace(raw_data=models) self._cache = SimpleNamespace(raw_data=models)
self.sync_calls: list[tuple[str, dict]] = []
async def get_cached_data(self): async def get_cached_data(self):
return self._cache return self._cache
@@ -38,6 +39,14 @@ class StubScanner:
break break
return True return True
async def sync_cache_from_metadata(self, file_path: str, metadata: dict) -> bool:
self.sync_calls.append((file_path, metadata))
for index, model in enumerate(self._cache.raw_data):
if model.get("file_path") == metadata.get("file_path"):
self._cache.raw_data[index] = metadata
break
return True
def _patch_scanner(monkeypatch: pytest.MonkeyPatch, scanner: StubScanner) -> None: def _patch_scanner(monkeypatch: pytest.MonkeyPatch, scanner: StubScanner) -> None:
async def _get_lora_scanner(cls): async def _get_lora_scanner(cls):
@@ -520,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):
@@ -588,6 +598,9 @@ async def test_not_found_example_images_are_cleaned(
assert missing_url in downloader.calls assert missing_url in downloader.calls
assert manager._progress["failed_models"] == {model_hash} assert manager._progress["failed_models"] == {model_hash}
assert model_hash in manager._progress["processed_models"] assert model_hash in manager._progress["processed_models"]
assert scanner.sync_calls
assert len(scanner.sync_calls) == 1
assert scanner.sync_calls[0][0] == str(model_path)
remaining_images = model_metadata["civitai"]["images"] remaining_images = model_metadata["civitai"]["images"]
assert remaining_images == [ assert remaining_images == [
@@ -596,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()
+2 -2
View File
@@ -884,7 +884,7 @@ async def test_sync_cache_conditional_resort_skipped(tmp_path: Path, monkeypatch
raw_data=[dict(entry)], folders=[], name_display_mode="model_name" raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
) )
await scanner._cache.resort() await scanner._cache.resort()
scanner._cache._last_sort = ("name", "asc") # name sort is active scanner._cache._last_sort = ("name", "asc", None) # name sort is active
scanner._tags_count = {"alpha": 1} scanner._tags_count = {"alpha": 1}
scanner._hash_index.add_entry("abc123", "/m/a.safetensors") scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
@@ -935,7 +935,7 @@ async def test_sync_cache_conditional_resort_triggered(tmp_path: Path, monkeypat
raw_data=[dict(entry)], folders=[], name_display_mode="model_name" raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
) )
await scanner._cache.resort() await scanner._cache.resort()
scanner._cache._last_sort = ("name", "asc") scanner._cache._last_sort = ("name", "asc", None)
scanner._tags_count = {"alpha": 1} scanner._tags_count = {"alpha": 1}
scanner._hash_index.add_entry("abc123", "/m/a.safetensors") scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
+97
View File
@@ -0,0 +1,97 @@
"""Tests for sort parsing and the seeded random sort mode."""
import asyncio
import pytest
from py.services.model_cache import ModelCache
from py.services.model_query import ModelCacheRepository, SortParams
def _make_cache(items):
return ModelCache(
raw_data=[
{
"file_path": f"/models/{name}.safetensors",
"file_name": f"{name}.safetensors",
"model_name": name,
"folder": "",
"size": 100,
"modified": 0.0,
}
for name in items
],
folders=[],
)
class TestParseSort:
def test_random_with_seed(self):
params = ModelCacheRepository.parse_sort("random:abc123")
assert params == SortParams(key="random", order="asc", seed="abc123")
def test_random_without_seed(self):
params = ModelCacheRepository.parse_sort("random")
assert params == SortParams(key="random", order="asc", seed=None)
def test_random_empty_seed_falls_back_to_none(self):
params = ModelCacheRepository.parse_sort("random:")
assert params.seed is None
def test_regular_sorts_unaffected(self):
params = ModelCacheRepository.parse_sort("name:desc")
assert params == SortParams(key="name", order="desc", seed=None)
class TestRandomShuffle:
@pytest.mark.asyncio
async def test_same_seed_yields_same_order(self):
cache = _make_cache(["a", "b", "c", "d", "e"])
await asyncio.sleep(0) # allow background resort task to run
first = await cache.get_sorted_data("random", "asc", "seed1")
second = await cache.get_sorted_data("random", "asc", "seed1")
assert [item["model_name"] for item in first] == [
item["model_name"] for item in second
]
@pytest.mark.asyncio
async def test_different_seeds_yield_different_orders(self):
cache = _make_cache([f"m{i}" for i in range(20)])
await asyncio.sleep(0)
first = await cache.get_sorted_data("random", "asc", "seed-a")
second = await cache.get_sorted_data("random", "asc", "seed-b")
assert [item["model_name"] for item in first] != [
item["model_name"] for item in second
]
@pytest.mark.asyncio
async def test_shuffle_is_a_permutation(self):
cache = _make_cache(["a", "b", "c", "d", "e"])
await asyncio.sleep(0)
shuffled = await cache.get_sorted_data("random", "asc", "seed")
assert sorted(item["model_name"] for item in shuffled) == [
"a",
"b",
"c",
"d",
"e",
]
assert len({item["file_path"] for item in shuffled}) == 5
@pytest.mark.asyncio
async def test_missing_seed_is_stable(self):
cache = _make_cache(["a", "b", "c", "d", "e"])
await asyncio.sleep(0)
first = await cache.get_sorted_data("random", "asc")
second = await cache.get_sorted_data("random", "asc")
assert [item["model_name"] for item in first] == [
item["model_name"] for item in second
]
+53
View File
@@ -860,6 +860,59 @@ def test_set_recipes_path_rewrites_symlinked_recipe_metadata(manager, tmp_path):
assert not old_json_path.exists() assert not old_json_path.exists()
def test_set_recipes_path_allows_cross_drive_migration(manager, tmp_path, monkeypatch):
# Windows regression: os.path.commonpath raises ValueError for paths on
# different drives (ntpath semantics). Cross-drive moves must succeed.
lora_root = tmp_path / "loras"
old_recipes_dir = lora_root / "recipes" / "nested"
old_recipes_dir.mkdir(parents=True)
manager.set("folder_paths", {"loras": [str(lora_root)]})
recipe_id = "recipe-cross-drive"
old_image_path = old_recipes_dir / f"{recipe_id}.webp"
old_json_path = old_recipes_dir / f"{recipe_id}.recipe.json"
old_image_path.write_bytes(b"image-bytes")
old_json_path.write_text(
json.dumps(
{
"id": recipe_id,
"file_path": str(old_image_path),
"title": "Recipe Cross Drive",
}
),
encoding="utf-8",
)
new_recipes_dir = tmp_path / "N_drive" / "AI" / "Library" / "Recipes"
# The effective current recipes dir (source of the migration) is
# lora_root/recipes — the nested subdirectory holds the recipe files.
source = str(lora_root / "recipes")
target = str(new_recipes_dir)
real_commonpath = os.path.commonpath
def fake_commonpath(paths):
# Simulate ntpath on Windows: a source/target pair on different
# drives shares no common root and raises ValueError.
if {source, target} <= set(paths):
raise ValueError("Paths don't have the same drive")
return real_commonpath(paths)
monkeypatch.setattr(os.path, "commonpath", fake_commonpath)
manager.set("recipes_path", str(new_recipes_dir))
migrated_image_path = new_recipes_dir / "nested" / f"{recipe_id}.webp"
migrated_json_path = new_recipes_dir / "nested" / f"{recipe_id}.recipe.json"
assert manager.get("recipes_path") == str(new_recipes_dir.resolve())
assert migrated_image_path.read_bytes() == b"image-bytes"
migrated_payload = json.loads(migrated_json_path.read_text(encoding="utf-8"))
assert migrated_payload["file_path"] == str(migrated_image_path)
assert not old_image_path.exists()
assert not old_json_path.exists()
def test_set_recipes_path_rejects_file_target(manager, tmp_path): def test_set_recipes_path_rejects_file_target(manager, tmp_path):
lora_root = tmp_path / "loras" lora_root = tmp_path / "loras"
lora_root.mkdir() lora_root.mkdir()
@@ -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)
+8 -3
View File
@@ -15,6 +15,7 @@ class StubScanner:
def __init__(self, cache_items: List[Dict[str, Any]]) -> None: def __init__(self, cache_items: List[Dict[str, Any]]) -> None:
self.cache = SimpleNamespace(raw_data=cache_items) self.cache = SimpleNamespace(raw_data=cache_items)
self.updates: List[Tuple[str, str, Dict[str, Any]]] = [] self.updates: List[Tuple[str, str, Dict[str, Any]]] = []
self.sync_updates: List[Tuple[str, Dict[str, Any]]] = []
async def get_cached_data(self): async def get_cached_data(self):
return self.cache return self.cache
@@ -23,6 +24,10 @@ class StubScanner:
self.updates.append((old_path, new_path, metadata)) self.updates.append((old_path, new_path, metadata))
return True return True
async def sync_cache_from_metadata(self, file_path: str, metadata: Dict[str, Any]) -> bool:
self.sync_updates.append((file_path, metadata))
return True
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
def patch_metadata_manager(monkeypatch: pytest.MonkeyPatch): def patch_metadata_manager(monkeypatch: pytest.MonkeyPatch):
@@ -83,7 +88,7 @@ async def test_update_metadata_after_import_enriches_entries(monkeypatch: pytest
assert custom[0]["type"] == "image" assert custom[0]["type"] == "image"
assert Path(patch_metadata_manager[0][0]) == model_file assert Path(patch_metadata_manager[0][0]) == model_file
assert scanner.updates assert scanner.sync_updates
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -151,8 +156,8 @@ async def test_update_metadata_after_import_preserves_existing_metadata(
assert saved_payload["civitai"]["trainedWords"] == ["foo"] assert saved_payload["civitai"]["trainedWords"] == ["foo"]
assert {entry["id"] for entry in saved_payload["civitai"]["customImages"]} == {"existing-id", "new-id"} assert {entry["id"] for entry in saved_payload["civitai"]["customImages"]} == {"existing-id", "new-id"}
assert scanner.updates assert scanner.sync_updates
updated_metadata = scanner.updates[-1][2] updated_metadata = scanner.sync_updates[-1][1]
assert updated_metadata["civitai"]["images"] == existing_payload["civitai"]["images"] assert updated_metadata["civitai"]["images"] == existing_payload["civitai"]["images"]
assert {entry["id"] for entry in updated_metadata["civitai"]["customImages"]} == {"existing-id", "new-id"} assert {entry["id"] for entry in updated_metadata["civitai"]["customImages"]} == {"existing-id", "new-id"}
@@ -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)
@@ -45,7 +45,7 @@ export interface AutocompleteTextWidgetInterface {
const props = defineProps<{ const props = defineProps<{
widget: AutocompleteTextWidgetInterface widget: AutocompleteTextWidgetInterface
node: { id: number } node: { id: number }
modelType?: 'loras' | 'embeddings' | 'custom_words' | 'prompt' modelType?: 'loras' | 'prompt'
placeholder?: string placeholder?: string
showPreview?: boolean showPreview?: boolean
spellcheck?: boolean spellcheck?: boolean
@@ -98,7 +98,7 @@ interface LoraInfoWidget {
onSetValue?: (v: unknown) => void onSetValue?: (v: unknown) => void
callback?: unknown callback?: unknown
options?: { options?: {
getValue?: () => LoraInfoWidgetValue getValue?: () => unknown
setValue?: (v: unknown) => void setValue?: (v: unknown) => void
} }
node?: { widgets?: Array<{ id?: string }>; widgets_values?: Array<unknown> } node?: { widgets?: Array<{ id?: string }>; widgets_values?: Array<unknown> }
@@ -299,8 +299,12 @@ onMounted(() => {
// ComponentWidgetImpl.value getter/setter delegates to options.getValue/options.setValue. // ComponentWidgetImpl.value getter/setter delegates to options.getValue/options.setValue.
// These must be set for workflow JSON persistence (LGraphNode.serialize/configure) to work. // These must be set for workflow JSON persistence (LGraphNode.serialize/configure) to work.
props.widget.options.getValue = buildValue if (props.widget.options) {
props.widget.options.setValue = applyValue props.widget.options.getValue = buildValue
props.widget.options.setValue = applyValue
} else {
console.warn('[LoraInfoWidget] widget.options missing, value persistence disabled')
}
// Also set serializeValue for prompt/API serialization path (executionUtil.ts) // Also set serializeValue for prompt/API serialization path (executionUtil.ts)
props.widget.serializeValue = async () => buildValue() props.widget.serializeValue = async () => buildValue()
@@ -3,7 +3,7 @@ import { ref, onMounted, onUnmounted, type Ref } from 'vue'
// Dynamic import type for AutoComplete class // Dynamic import type for AutoComplete class
type AutoCompleteClass = new ( type AutoCompleteClass = new (
inputElement: HTMLTextAreaElement, inputElement: HTMLTextAreaElement,
modelType: 'loras' | 'embeddings' | 'custom_words' | 'prompt', modelType: 'loras' | 'prompt',
options?: AutocompleteOptions options?: AutocompleteOptions
) => AutoCompleteInstance ) => AutoCompleteInstance
@@ -29,7 +29,7 @@ export interface UseAutocompleteOptions {
export function useAutocomplete( export function useAutocomplete(
textareaRef: Ref<HTMLTextAreaElement | null>, textareaRef: Ref<HTMLTextAreaElement | null>,
modelType: 'loras' | 'embeddings' | 'custom_words' | 'prompt' = 'loras', modelType: 'loras' | 'prompt' = 'loras',
options: UseAutocompleteOptions = {} options: UseAutocompleteOptions = {}
) { ) {
const autocompleteInstance = ref<AutoCompleteInstance | null>(null) const autocompleteInstance = ref<AutoCompleteInstance | null>(null)
+6 -9
View File
@@ -36,6 +36,9 @@ const AUTOCOMPLETE_TEXT_MIN_HEIGHT_DEFAULT = 300
const AUTOCOMPLETE_METADATA_VERSION = 1 const AUTOCOMPLETE_METADATA_VERSION = 1
const LORA_MANAGER_WIDGET_IDS_PROPERTY = '__lm_widget_ids' const LORA_MANAGER_WIDGET_IDS_PROPERTY = '__lm_widget_ids'
// Access LiteGraph global for Vue DOM mode detection (matches AutocompleteTextWidget.vue)
declare const LiteGraph: { vueNodesMode?: boolean } | undefined
// @ts-ignore - ComfyUI external module // @ts-ignore - ComfyUI external module
import { app } from '../../../scripts/app.js' import { app } from '../../../scripts/app.js'
// @ts-ignore - ComfyUI external module // @ts-ignore - ComfyUI external module
@@ -718,7 +721,7 @@ function createLoraInfoWidget(node: any) {
function createAutocompleteTextWidgetFactory( function createAutocompleteTextWidgetFactory(
node: any, node: any,
widgetName: string, widgetName: string,
modelType: 'loras' | 'embeddings' | 'prompt', modelType: 'loras' | 'prompt',
inputOptions: { placeholder?: string } = {} inputOptions: { placeholder?: string } = {}
) { ) {
const metadataWidgetName = `__lm_autocomplete_meta_${widgetName}` const metadataWidgetName = `__lm_autocomplete_meta_${widgetName}`
@@ -835,7 +838,7 @@ function createAutocompleteTextWidgetFactory(
applyAutocompleteTextLayoutFix( applyAutocompleteTextLayoutFix(
widget, widget,
container, container,
typeof LiteGraph !== 'undefined' && LiteGraph.vueNodesMode typeof LiteGraph !== 'undefined' && LiteGraph.vueNodesMode === true
) )
} }
@@ -964,13 +967,7 @@ app.registerExtension({
const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {} const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {}
return createAutocompleteTextWidgetFactory(node, 'text', 'loras', options) return createAutocompleteTextWidgetFactory(node, 'text', 'loras', options)
}, },
// Autocomplete text widget for embeddings (used by Prompt node) // Autocomplete text widget for prompt (used by Prompt and Text nodes)
// @ts-ignore
AUTOCOMPLETE_TEXT_EMBEDDINGS(node) {
const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {}
return createAutocompleteTextWidgetFactory(node, 'text', 'embeddings', options)
},
// Autocomplete text widget for prompt (supports both embeddings and custom words)
// @ts-ignore // @ts-ignore
AUTOCOMPLETE_TEXT_PROMPT(node) { AUTOCOMPLETE_TEXT_PROMPT(node) {
const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {} const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {}
+127
View File
@@ -0,0 +1,127 @@
import { app } from "../../scripts/app.js";
/**
* Extension for LoraStackCombinerLM node to support dynamic lora_stack inputs.
* Defaults to two inputs; connecting the last slot adds a new empty one, and
* disconnecting a non-last slot removes it (at least two are always kept).
* Based on the dynamic input pattern from Impact Pack's Switch (Any) node.
*/
const STACK_INPUT_PATTERN = /^lora_stack\d+$/;
app.registerExtension({
name: "Comfy.LoraManager.LoraStackCombiner",
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.name !== "Lora Stack Combiner (LoraManager)") {
return;
}
const onConnectionsChange = nodeType.prototype.onConnectionsChange;
nodeType.prototype.onConnectionsChange = function(type, index, connected, link_info) {
// Skip while the graph is being (re)configured (load, paste, subgraph ops)
if (app.configuringGraph) {
return onConnectionsChange?.apply?.(this, arguments);
}
const stackTrace = new Error().stack;
// Skip during graph loading/pasting to avoid interference
if (stackTrace.includes('loadGraphData') || stackTrace.includes('pasteFromClipboard')) {
return onConnectionsChange?.apply?.(this, arguments);
}
// Skip subgraph operations
if (stackTrace.includes('convertToSubgraph') || stackTrace.includes('Subgraph.configure')) {
return onConnectionsChange?.apply?.(this, arguments);
}
if (!link_info) {
return onConnectionsChange?.apply?.(this, arguments);
}
// Handle input connections (type === 1)
if (type === 1) {
const input = this.inputs[index];
// Only process numbered lora_stack inputs (legacy a/b slots are left untouched)
if (!input || !STACK_INPUT_PATTERN.test(input.name)) {
return onConnectionsChange?.apply?.(this, arguments);
}
// Count existing numbered lora_stack inputs
let stackInputCount = 0;
for (const inp of this.inputs) {
if (STACK_INPUT_PATTERN.test(inp.name)) {
stackInputCount++;
}
}
// Renumber all numbered lora_stack inputs sequentially
let slotIndex = 1;
for (const inp of this.inputs) {
if (STACK_INPUT_PATTERN.test(inp.name)) {
inp.name = `lora_stack${slotIndex}`;
slotIndex++;
}
}
// Add new input slot if connected and this was the last one
if (connected) {
const lastStackIndex = stackInputCount;
if (index === lastStackIndex || index === this.inputs.findIndex(i => i.name === `lora_stack${lastStackIndex}`)) {
this.addInput(`lora_stack${slotIndex}`, "LORA_STACK", {
tooltip: "A LoRA stack to combine. Connect to add more inputs."
});
}
}
// Remove disconnected input slots (but keep at least two).
// LiteGraph fires this event only for slots that had a link, and
// it has already cleared input.link by the time the event fires,
// so the disconnected slot is always empty at this point.
if (!connected && stackInputCount > 2) {
const disconnectedInput = this.inputs[index];
if (disconnectedInput && STACK_INPUT_PATTERN.test(disconnectedInput.name)) {
// Keep the last slot so there is always an empty slot to reconnect into
const isLastStackSlot = index === this.inputs.findLastIndex(i => STACK_INPUT_PATTERN.test(i.name));
if (!isLastStackSlot) {
this.removeInput(index);
// Renumber again after removal
let newSlotIndex = 1;
for (const inp of this.inputs) {
if (STACK_INPUT_PATTERN.test(inp.name)) {
inp.name = `lora_stack${newSlotIndex}`;
newSlotIndex++;
}
}
}
}
}
}
return onConnectionsChange?.apply?.(this, arguments);
};
},
nodeCreated(node, app) {
if (node.comfyClass !== "Lora Stack Combiner (LoraManager)") {
return;
}
// Leave legacy (a/b) workflows untouched
const hasLegacyInputs = node.inputs.some(inp => inp.name === "lora_stack_a" || inp.name === "lora_stack_b");
if (hasLegacyInputs) {
return;
}
// Ensure at least two numbered lora_stack inputs exist on creation
const stackInputCount = node.inputs.filter(inp => STACK_INPUT_PATTERN.test(inp.name)).length;
for (let i = stackInputCount + 1; i <= 2; i++) {
node.addInput(`lora_stack${i}`, "LORA_STACK", {
tooltip: "A LoRA stack to combine. Connect to add more inputs."
});
}
}
});
+64 -66
View File
@@ -2118,14 +2118,14 @@ to { transform: rotate(360deg);
padding: 20px 0; padding: 20px 0;
} }
.autocomplete-text-widget[data-v-3f3d7a1a] { .autocomplete-text-widget[data-v-55e3316e] {
background: transparent; background: transparent;
height: 100%; height: 100%;
display: flex; display: flex;
flex-direction: column; flex-direction: column;
box-sizing: border-box; box-sizing: border-box;
} }
.input-wrapper[data-v-3f3d7a1a] { .input-wrapper[data-v-55e3316e] {
position: relative; position: relative;
flex: 1; flex: 1;
display: flex; display: flex;
@@ -2133,7 +2133,7 @@ to { transform: rotate(360deg);
} }
/* Canvas mode styles (default) - matches built-in comfy-multiline-input */ /* Canvas mode styles (default) - matches built-in comfy-multiline-input */
.text-input[data-v-3f3d7a1a] { .text-input[data-v-55e3316e] {
flex: 1; flex: 1;
width: 100%; width: 100%;
background-color: var(--comfy-input-bg, #222); background-color: var(--comfy-input-bg, #222);
@@ -2152,7 +2152,7 @@ to { transform: rotate(360deg);
} }
/* Vue DOM mode styles - matches built-in p-textarea in Vue DOM mode */ /* Vue DOM mode styles - matches built-in p-textarea in Vue DOM mode */
.text-input.vue-dom-mode[data-v-3f3d7a1a] { .text-input.vue-dom-mode[data-v-55e3316e] {
background-color: var(--color-charcoal-400, #313235); background-color: var(--color-charcoal-400, #313235);
color: #fff; color: #fff;
padding: 8px 12px 30px 12px; /* Reserve bottom space for clear button */ padding: 8px 12px 30px 12px; /* Reserve bottom space for clear button */
@@ -2161,12 +2161,12 @@ to { transform: rotate(360deg);
font-size: 12px; font-size: 12px;
font-family: inherit; font-family: inherit;
} }
.text-input[data-v-3f3d7a1a]:focus { .text-input[data-v-55e3316e]:focus {
outline: none; outline: none;
} }
/* Clear button styles */ /* Clear button styles */
.clear-button[data-v-3f3d7a1a] { .clear-button[data-v-55e3316e] {
position: absolute; position: absolute;
right: 6px; right: 6px;
bottom: 6px; /* Changed from top to bottom */ bottom: 6px; /* Changed from top to bottom */
@@ -2189,31 +2189,31 @@ to { transform: rotate(360deg);
} }
/* Show clear button when hovering over input wrapper */ /* Show clear button when hovering over input wrapper */
.input-wrapper:hover .clear-button[data-v-3f3d7a1a] { .input-wrapper:hover .clear-button[data-v-55e3316e] {
opacity: 0.7; opacity: 0.7;
pointer-events: auto; pointer-events: auto;
} }
.clear-button[data-v-3f3d7a1a]:hover { .clear-button[data-v-55e3316e]:hover {
opacity: 1; opacity: 1;
background: rgba(255, 100, 100, 0.8); background: rgba(255, 100, 100, 0.8);
} }
.clear-button svg[data-v-3f3d7a1a] { .clear-button svg[data-v-55e3316e] {
width: 12px; width: 12px;
height: 12px; height: 12px;
} }
/* Vue DOM mode adjustments for clear button */ /* Vue DOM mode adjustments for clear button */
.text-input.vue-dom-mode ~ .clear-button[data-v-3f3d7a1a] { .text-input.vue-dom-mode ~ .clear-button[data-v-55e3316e] {
right: 8px; right: 8px;
bottom: 10px; /* Changed from top to bottom, adjusted for Vue DOM padding */ bottom: 10px; /* Changed from top to bottom, adjusted for Vue DOM padding */
width: 20px; width: 20px;
height: 20px; height: 20px;
background: rgba(107, 114, 128, 0.6); background: rgba(107, 114, 128, 0.6);
} }
.text-input.vue-dom-mode ~ .clear-button[data-v-3f3d7a1a]:hover { .text-input.vue-dom-mode ~ .clear-button[data-v-55e3316e]:hover {
background: oklch(62% 0.18 25); background: oklch(62% 0.18 25);
} }
.text-input.vue-dom-mode ~ .clear-button svg[data-v-3f3d7a1a] { .text-input.vue-dom-mode ~ .clear-button svg[data-v-55e3316e] {
width: 14px; width: 14px;
height: 14px; height: 14px;
} }
@@ -2224,7 +2224,7 @@ to { transform: rotate(360deg);
resize: vertical !important; resize: vertical !important;
} }
.lora-info-widget[data-v-a99cc1ab] { .lora-info-widget[data-v-d7692b6f] {
padding: 12px; padding: 12px;
background: rgba(40, 44, 52, 0.6); background: rgba(40, 44, 52, 0.6);
border-radius: 4px; border-radius: 4px;
@@ -2240,45 +2240,45 @@ to { transform: rotate(360deg);
determined solely by CSS not by descendant content. This breaks the determined solely by CSS not by descendant content. This breaks the
feedback loop where content grows ResizeObserver resizes content feedback loop where content grows ResizeObserver resizes content
reflows repeat. Same technique used by tags_widget.js + lm_styles.css. */ reflows repeat. Same technique used by tags_widget.js + lm_styles.css. */
.lora-info-widget.lm-vue-node[data-v-a99cc1ab] { .lora-info-widget.lm-vue-node[data-v-d7692b6f] {
contain: layout size; contain: layout size;
} }
/* ── Tab bar ── */ /* ── Tab bar ── */
.lora-info-tabs[data-v-a99cc1ab] { .lora-info-tabs[data-v-d7692b6f] {
display: flex; display: flex;
gap: 0; gap: 0;
margin-bottom: 10px; margin-bottom: 10px;
border-bottom: 1px solid var(--border-color, #444); border-bottom: 1px solid var(--border-color, #444);
flex-shrink: 0; flex-shrink: 0;
} }
.lora-info-tab[data-v-a99cc1ab] { .lora-info-tab[data-v-d7692b6f] {
flex: 1; flex: 1;
text-align: center; text-align: center;
cursor: pointer; cursor: pointer;
padding: 6px 0; padding: 6px 0;
position: relative; position: relative;
} }
.lora-info-tab-input[data-v-a99cc1ab] { .lora-info-tab-input[data-v-d7692b6f] {
position: absolute; position: absolute;
opacity: 0; opacity: 0;
width: 0; width: 0;
height: 0; height: 0;
} }
.lora-info-tab-label[data-v-a99cc1ab] { .lora-info-tab-label[data-v-d7692b6f] {
font-size: 12px; font-size: 12px;
font-weight: 500; font-weight: 500;
color: var(--fg-color, #fff); color: var(--fg-color, #fff);
opacity: 0.5; opacity: 0.5;
transition: opacity 0.15s; transition: opacity 0.15s;
} }
.lora-info-tab:hover .lora-info-tab-label[data-v-a99cc1ab] { .lora-info-tab:hover .lora-info-tab-label[data-v-d7692b6f] {
opacity: 0.75; opacity: 0.75;
} }
.lora-info-tab.active .lora-info-tab-label[data-v-a99cc1ab] { .lora-info-tab.active .lora-info-tab-label[data-v-d7692b6f] {
opacity: 1; opacity: 1;
} }
.lora-info-tab.active[data-v-a99cc1ab]::after { .lora-info-tab.active[data-v-d7692b6f]::after {
content: ''; content: '';
position: absolute; position: absolute;
bottom: -1px; bottom: -1px;
@@ -2290,16 +2290,16 @@ to { transform: rotate(360deg);
} }
/* ── Tab content ── */ /* ── Tab content ── */
.tab-content[data-v-a99cc1ab] { .tab-content[data-v-d7692b6f] {
flex: 1; flex: 1;
min-height: 0; min-height: 0;
overflow: hidden; overflow: hidden;
} }
.notes-tab[data-v-a99cc1ab] { .notes-tab[data-v-d7692b6f] {
display: flex; display: flex;
flex-direction: column; flex-direction: column;
} }
.description-tab[data-v-a99cc1ab] { .description-tab[data-v-d7692b6f] {
display: flex; display: flex;
flex-direction: column; flex-direction: column;
overflow-y: auto; overflow-y: auto;
@@ -2307,12 +2307,12 @@ to { transform: rotate(360deg);
} }
/* ── Info fields (shared) ── */ /* ── Info fields (shared) ── */
.info-field[data-v-a99cc1ab] { .info-field[data-v-d7692b6f] {
display: flex; display: flex;
flex-direction: column; flex-direction: column;
gap: 4px; gap: 4px;
} }
.info-label[data-v-a99cc1ab] { .info-label[data-v-d7692b6f] {
font-size: 10px; font-size: 10px;
font-weight: 600; font-weight: 600;
text-transform: uppercase; text-transform: uppercase;
@@ -2320,7 +2320,7 @@ to { transform: rotate(360deg);
color: var(--fg-color, #fff); color: var(--fg-color, #fff);
opacity: 0.6; opacity: 0.6;
} }
.lora-filename[data-v-a99cc1ab] { .lora-filename[data-v-d7692b6f] {
font-size: 13px; font-size: 13px;
font-weight: 500; font-weight: 500;
color: var(--fg-color, #fff); color: var(--fg-color, #fff);
@@ -2331,11 +2331,11 @@ to { transform: rotate(360deg);
user-select: text; user-select: text;
-webkit-user-select: text; -webkit-user-select: text;
} }
.notes-field[data-v-a99cc1ab] { .notes-field[data-v-d7692b6f] {
flex: 1; flex: 1;
min-height: 0; min-height: 0;
} }
.lora-notes[data-v-a99cc1ab] { .lora-notes[data-v-d7692b6f] {
width: 100%; width: 100%;
flex: 1; flex: 1;
min-height: 60px; min-height: 60px;
@@ -2350,14 +2350,14 @@ to { transform: rotate(360deg);
font-family: inherit; font-family: inherit;
outline: none; outline: none;
} }
.lora-notes[data-v-a99cc1ab]:focus { .lora-notes[data-v-d7692b6f]:focus {
border-color: var(--comfy-input-border, #444); border-color: var(--comfy-input-border, #444);
} }
.lora-notes[data-v-a99cc1ab]:disabled { .lora-notes[data-v-d7692b6f]:disabled {
opacity: 0.6; opacity: 0.6;
cursor: not-allowed; cursor: not-allowed;
} }
.save-btn[data-v-a99cc1ab] { .save-btn[data-v-d7692b6f] {
width: 100%; width: 100%;
margin-top: 8px; margin-top: 8px;
padding: 6px 12px; padding: 6px 12px;
@@ -2371,11 +2371,11 @@ to { transform: rotate(360deg);
box-sizing: border-box; box-sizing: border-box;
flex-shrink: 0; flex-shrink: 0;
} }
.save-btn[data-v-a99cc1ab]:hover:not(:disabled) { .save-btn[data-v-d7692b6f]:hover:not(:disabled) {
background: rgba(66, 153, 225, 0.25); background: rgba(66, 153, 225, 0.25);
border-color: rgba(66, 153, 225, 0.6); border-color: rgba(66, 153, 225, 0.6);
} }
.save-btn[data-v-a99cc1ab]:disabled { .save-btn[data-v-d7692b6f]:disabled {
opacity: 0.4; opacity: 0.4;
cursor: not-allowed; cursor: not-allowed;
background: rgba(66, 153, 225, 0.05); background: rgba(66, 153, 225, 0.05);
@@ -2383,7 +2383,7 @@ to { transform: rotate(360deg);
} }
/* ── Description states ── */ /* ── Description states ── */
.description-state[data-v-a99cc1ab] { .description-state[data-v-d7692b6f] {
display: flex; display: flex;
align-items: center; align-items: center;
justify-content: center; justify-content: center;
@@ -2395,22 +2395,22 @@ to { transform: rotate(360deg);
min-height: 0; min-height: 0;
flex-shrink: 0; flex-shrink: 0;
} }
.description-state.error[data-v-a99cc1ab] { .description-state.error[data-v-d7692b6f] {
opacity: 0.7; opacity: 0.7;
color: #f87171; color: #f87171;
} }
/* ── Description content ── */ /* ── Description content ── */
.description-content[data-v-a99cc1ab] { .description-content[data-v-d7692b6f] {
min-height: 0; min-height: 0;
} }
.description-section[data-v-a99cc1ab] { .description-section[data-v-d7692b6f] {
margin-bottom: 14px; margin-bottom: 14px;
} }
.description-section[data-v-a99cc1ab]:last-child { .description-section[data-v-d7692b6f]:last-child {
margin-bottom: 0; margin-bottom: 0;
} }
.description-text[data-v-a99cc1ab] { .description-text[data-v-d7692b6f] {
padding: 8px 0; padding: 8px 0;
font-size: 12px; font-size: 12px;
line-height: 1.5; line-height: 1.5;
@@ -2422,41 +2422,41 @@ to { transform: rotate(360deg);
user-select: text; user-select: text;
-webkit-user-select: text; -webkit-user-select: text;
} }
.description-text[data-v-a99cc1ab] p { .description-text[data-v-d7692b6f] p {
margin: 0 0 8px 0; margin: 0 0 8px 0;
} }
.description-text[data-v-a99cc1ab] p:last-child { .description-text[data-v-d7692b6f] p:last-child {
margin-bottom: 0; margin-bottom: 0;
} }
.description-text[data-v-a99cc1ab] a { .description-text[data-v-d7692b6f] a {
color: rgba(66, 153, 225, 0.9); color: rgba(66, 153, 225, 0.9);
} }
.description-text[data-v-a99cc1ab] ul, .description-text[data-v-d7692b6f] ul,
.description-text[data-v-a99cc1ab] ol { .description-text[data-v-d7692b6f] ol {
padding-left: 20px; padding-left: 20px;
margin: 4px 0; margin: 4px 0;
} }
.description-text[data-v-a99cc1ab] h1, .description-text[data-v-d7692b6f] h1,
.description-text[data-v-a99cc1ab] h2, .description-text[data-v-d7692b6f] h2,
.description-text[data-v-a99cc1ab] h3 { .description-text[data-v-d7692b6f] h3 {
font-size: 13px; font-size: 13px;
margin: 10px 0 4px 0; margin: 10px 0 4px 0;
font-weight: 600; font-weight: 600;
opacity: 0.95; opacity: 0.95;
} }
.description-text[data-v-a99cc1ab] code { .description-text[data-v-d7692b6f] code {
background: rgba(255, 255, 255, 0.08); background: rgba(255, 255, 255, 0.08);
padding: 1px 4px; padding: 1px 4px;
border-radius: 3px; border-radius: 3px;
font-size: 11px; font-size: 11px;
} }
.description-text[data-v-a99cc1ab] img { .description-text[data-v-d7692b6f] img {
max-width: 100%; max-width: 100%;
border-radius: 4px; border-radius: 4px;
} }
/* ── Placeholder (shared) ── */ /* ── Placeholder (shared) ── */
.placeholder[data-v-a99cc1ab] { .placeholder[data-v-d7692b6f] {
font-style: italic; font-style: italic;
color: rgba(226, 232, 240, 0.5); color: rgba(226, 232, 240, 0.5);
text-align: center; text-align: center;
@@ -2465,10 +2465,10 @@ to { transform: rotate(360deg);
} }
/* ── Spinner (Font Awesome) ── */ /* ── Spinner (Font Awesome) ── */
.fa-spinner[data-v-a99cc1ab] { .fa-spinner[data-v-d7692b6f] {
animation: fa-spin-a99cc1ab 1s linear infinite; animation: fa-spin-d7692b6f 1s linear infinite;
} }
@keyframes fa-spin-a99cc1ab { @keyframes fa-spin-d7692b6f {
0% { transform: rotate(0deg); 0% { transform: rotate(0deg);
} }
100% { transform: rotate(360deg); 100% { transform: rotate(360deg);
@@ -15316,7 +15316,7 @@ const _sfc_main$1 = /* @__PURE__ */ defineComponent({
}; };
} }
}); });
const AutocompleteTextWidget = /* @__PURE__ */ _export_sfc(_sfc_main$1, [["__scopeId", "data-v-3f3d7a1a"]]); const AutocompleteTextWidget = /* @__PURE__ */ _export_sfc(_sfc_main$1, [["__scopeId", "data-v-55e3316e"]]);
const _hoisted_1 = { class: "lora-info-tabs" }; const _hoisted_1 = { class: "lora-info-tabs" };
const _hoisted_2 = { class: "tab-content notes-tab" }; const _hoisted_2 = { class: "tab-content notes-tab" };
const _hoisted_3 = { class: "info-field" }; const _hoisted_3 = { class: "info-field" };
@@ -15511,8 +15511,12 @@ const _sfc_main = /* @__PURE__ */ defineComponent({
if (data.filePath !== void 0) filePath.value = data.filePath; if (data.filePath !== void 0) filePath.value = data.filePath;
} }
}; };
props.widget.options.getValue = buildValue; if (props.widget.options) {
props.widget.options.setValue = applyValue; props.widget.options.getValue = buildValue;
props.widget.options.setValue = applyValue;
} else {
console.warn("[LoraInfoWidget] widget.options missing, value persistence disabled");
}
props.widget.serializeValue = async () => buildValue(); props.widget.serializeValue = async () => buildValue();
props.widget.onSetValue = applyValue; props.widget.onSetValue = applyValue;
const widgetIndex = (_b = (_a2 = props.widget.node) == null ? void 0 : _a2.widgets) == null ? void 0 : _b.findIndex( const widgetIndex = (_b = (_a2 = props.widget.node) == null ? void 0 : _a2.widgets) == null ? void 0 : _b.findIndex(
@@ -15641,7 +15645,7 @@ const _sfc_main = /* @__PURE__ */ defineComponent({
}; };
} }
}); });
const LoraInfoWidget = /* @__PURE__ */ _export_sfc(_sfc_main, [["__scopeId", "data-v-a99cc1ab"]]); const LoraInfoWidget = /* @__PURE__ */ _export_sfc(_sfc_main, [["__scopeId", "data-v-d7692b6f"]]);
function createVueWidgetCleanup(vueApp, onCleanup) { function createVueWidgetCleanup(vueApp, onCleanup) {
let didUnmount = false; let didUnmount = false;
return () => { return () => {
@@ -16637,7 +16641,7 @@ function createAutocompleteTextWidgetFactory(node, widgetName, modelType, inputO
applyAutocompleteTextLayoutFix( applyAutocompleteTextLayoutFix(
widget, widget,
container, container,
typeof LiteGraph !== "undefined" && LiteGraph.vueNodesMode typeof LiteGraph !== "undefined" && LiteGraph.vueNodesMode === true
); );
} }
const vueCleanup = createVueWidgetCleanup(vueApp, () => { const vueCleanup = createVueWidgetCleanup(vueApp, () => {
@@ -16747,13 +16751,7 @@ app$1.registerExtension({
const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {}; const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {};
return createAutocompleteTextWidgetFactory(node, "text", "loras", options); return createAutocompleteTextWidgetFactory(node, "text", "loras", options);
}, },
// Autocomplete text widget for embeddings (used by Prompt node) // Autocomplete text widget for prompt (used by Prompt and Text nodes)
// @ts-ignore
AUTOCOMPLETE_TEXT_EMBEDDINGS(node) {
const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {};
return createAutocompleteTextWidgetFactory(node, "text", "embeddings", options);
},
// Autocomplete text widget for prompt (supports both embeddings and custom words)
// @ts-ignore // @ts-ignore
AUTOCOMPLETE_TEXT_PROMPT(node) { AUTOCOMPLETE_TEXT_PROMPT(node) {
const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {}; const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {};
File diff suppressed because one or more lines are too long