Compare commits

...

43 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
Will Miao dd6bdbf297 fix(update): persist update_channel via settings.json instead of hasGit
After b464fdc3 (preserve .git on release switch), the hasGit-based
channel detection is unreliable — .git now exists for both release
and nightly installs, so page refresh always reset the channel.

- Add _resolveChannelFromSettings() with migration heuristic:
  !hasGit → release (ZIP), detached HEAD → release (on tag),
  on branch → nightly. Uses gitInfo.branch from check-updates.
- Persist resolved channel to settings.json on first load
  (one-time migration) and on explicit switchChannel.
- Add update_channel validation (release|nightly) in backend
  update_settings handler.
- Remove hasGit-based guessing from initialize(); defer to
  checkForUpdates where full gitInfo is available.
- Channel resolution runs before checkForUpdates early-returns
  to avoid null channelMode on reload-within-interval.

Tests: 361 passed.
2026-07-31 13:23:54 +08:00
Will Miao b47dde87e4 fix(settings): suppress error toasts when optional model roots are empty 2026-07-31 10:07:52 +08:00
Will Miao 99e65cccd8 fix(update): downgrade settings backup/restore logs from INFO to DEBUG 2026-07-30 20:32:00 +08:00
Will Miao 3bdacb8f46 fix(test): update release channel git test to mock _perform_git_update instead of _download_and_replace_zip 2026-07-30 18:35:43 +08:00
Will Miao b4f9c224d3 fix(example-images): move multi→single-library consolidation to startup, eliminate per-request os.listdir()
Move reverse-migration logic from get_model_folder() (hot path, called on
every metadata/example-images request) to ExampleImagesMigration, where it
runs once at startup.  On network storage this was causing 22-38s delays
per LoRA card click.

Additionally optimize prune_stale_example_images() to read the directory
listing once instead of per image entry (O(N*M) → O(M)).  Also reorder
consolidation checks so regex filters run before filesystem stat calls.
2026-07-30 18:11:40 +08:00
Will Miao 5ec0399c81 fix(i18n): remove redundant 'preserved' sentence from release channel message, sync all 10 locales 2026-07-30 16:38:04 +08:00
Will Miao b464fdc333 fix(update): preserve .git on release channel switch, use git checkout tag
Previously, switching to the release channel would delete .git/ and
fall back to a ZIP download. This broke update.bat, manual git
commands, and CM git-based update detection.

Now the release path uses git checkout <latest-tag> when .git exists,
and only falls back to ZIP when .git is absent (CM CNR installs).
.git is never deleted - the ZIP→nightly path remains a one-way
upgrade via _init_git_repo.

Also updates locale strings (en, zh-CN, zh-TW, ja) to remove the
now-inaccurate "remove the Git repository" wording.
2026-07-29 21:23:39 +08:00
Will Miao 53825500db fix(update): add staging protection to switch_channel
switch_channel has three destructive code paths (git reset + clean,
git init + checkout --force, and rmtree + ZIP replace) that were
missing the _stage_preserved_items / _restore_preserved_items safety
net already applied to perform_update.

Wrap the channel-specific logic in a try/finally so preserved user
data (settings.json, civitai/, cache/, etc.) is physically moved
outside plugin_root before any git operation and always restored.
2026-07-29 20:41:36 +08:00
Will Miao f2ac790752 fix(update): stage preserved items outside repo before git/ZIP update
Move settings.json, civitai/, wildcards/, backups/, stats/, logs/,
cache/, and model_cache/ to a temp directory before git reset/clean
or ZIP replacement, then restore them in a try/finally block.

This prevents data loss on Windows where git clean -e exclusion
patterns can fail due to path-separator mismatches or where file
locks (open SQLite/log handles) cause the restore step to be skipped
on failure.

Also unifies three hardcoded skip lists (_clean_plugin_folder,
skip_items, skip_tracked) to derive from the single _PRESERVE_DIRS
constant, fixing drift where logs/ was missing from the ZIP path.
2026-07-29 19:49:50 +08:00
Will Miao 0d8805cdee fix(recipes): update cards in-place after LoRA download, preventing scroll reset 2026-07-29 11:35:28 +08:00
pixelpaws 656e24ac9b Merge pull request #1044 from d1udiu/fix-filter
fix(filters): prevent search query from being persisted in localStorage
2026-07-29 11:30:40 +08:00
d1udiu 6718b37403 fix(filters): prevent search query from being persisted in localStorage 2026-07-29 10:12:42 +08:00
Will Miao c9e5e784fc fix(metadata-overwrite): use sentinel default for clip_skip to accept wired 0 2026-07-28 23:13:00 +08:00
Will Miao f92f958682 fix(SaveImageLM): correct scheduler mapping and deduplicate sampler map
- Fix incorrect mapping: "normal" -> "Normal" (was "Simple")
- Replace inline sampler_mapping with CIVITAI_SAMPLER_MAP reference
  to eliminate duplicate definition
2026-07-28 21:39:09 +08:00
Will Miao f63fab0676 fix(cache): deduplicate model entries on add and reconcile to prevent duplicate cards (#1041) 2026-07-28 20:44:57 +08:00
Will Miao cfc4903c0c fix(update): read ahead_by from GitHub compare API when status is ahead/diverged
The compare API URL format compare/{local_hash}...main returns
status='ahead' when main is ahead of the local commit. The count is
in the ahead_by field, not behind_by. The old code only read behind_by
which is always 0 in this case, causing the UI to show 'Up to date'
when actually several commits behind.

Also handle status='diverged' (both sides have unique commits) by
reading ahead_by for the remote-ahead count.

Frontend adds a hash comparison fallback: if behind_by is 0 but local
and remote commit hashes differ, show 'Behind main' instead of the
incorrect 'Up to date'.

Tests: _AheadCompareDownloader and _DivergedCompareDownloader mocks
for the two status paths.
2026-07-28 17:47:38 +08:00
Will Miao a527a847fe fix(download): route UNet/diffusion model downloads to unet roots in location step
When downloading a diffusion model (UNet) from the checkpoints page, the
download modal's location step always showed checkpoint roots and paths.
Now the modal detects the file subtype and switches to unet_roots endpoint,
default_unet_root key, and 'unet' path template.
2026-07-28 17:21:12 +08:00
Will Miao 91b0bf8933 fix(download_queue): deduplicate download_history rows before creating unique index (#1041) 2026-07-27 21:36:58 +08:00
Will Miao 66d1c96783 feat(update): add Release/Nightly channel switching
- Add POST /api/lm/switch-channel endpoint with git init / ZIP fallback
- Add _backup_git/_restore_git helpers with safe rollback
- Version-info endpoint now returns has_git flag for auto-detection
- Check-updates always returns releases (changelog) regardless of channel
- Nightly mode shows 'N commits behind main' with commit hash and date
- View on GitHub link points to /commits/main in nightly mode
- Channel toggle UI with pill-style buttons in update modal
- Confirmation dialog with Esc / backdrop-dismiss support
- Channel derived from has_git on every page load, no localStorage
- i18n: 11 new keys translated across 9 non-English locales
- CSS: unified card-style sections in _base.css
- Tests: 8 new tests covering switch-channel, nightly response, init_git_repo
2026-07-27 20:27:05 +08:00
Will Miao 986128076e fix(widget): guard setValue against non-array input to prevent workflow load crash (#1039) 2026-07-26 21:54:37 +08:00
Will Miao 1de0a53241 feat(grouping): version-group library cards by HuggingFace repo for non-Civitai sources (#1040) 2026-07-26 21:46:55 +08:00
Will Miao 0ec7eaf606 fix(wildcards): resolve weighted N::value syntax inside wildcard YAML lists (#1039) 2026-07-26 18:33:31 +08:00
Will Miao d9fcb0e92b fix(filter): preserve search term through filter apply/clear operations 2026-07-26 16:49:57 +08:00
103 changed files with 26566 additions and 20646 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 -2205
View File
File diff suppressed because it is too large Load Diff
+39 -1
View File
@@ -678,6 +678,7 @@
"deepseek": "DeepSeek",
"groq": "Groq",
"openrouter": "OpenRouter",
"google": "Gemini",
"opencode-go": "OpenCode Go",
"custom": "Custom (OpenAI-compatible)"
},
@@ -714,7 +715,9 @@
"versionsCount": "Local Versions",
"versionsCountDesc": "Most versions first",
"versionsCountAsc": "Fewest versions first",
"versionIdDesc": "Newest version first"
"versionIdDesc": "Newest version first",
"random": "Random",
"randomAction": "Randomize (shuffle)"
},
"refresh": {
"title": "Refresh model list",
@@ -771,6 +774,8 @@
"deleteAll": "Delete Selected",
"downloadMissingLoras": "Download Missing LoRAs",
"downloadExamples": "Download Example Images",
"downloadMissingExamples": "Download Missing",
"reprocessExamples": "Re-process All",
"clear": "Clear Selection",
"skipMetadataRefreshCount": "Skip ({count} models)",
"resumeMetadataRefreshCount": "Resume ({count} models)",
@@ -806,6 +811,8 @@
"sendToWorkflowReplace": "Send to Workflow (Replace)",
"openExamples": "Open Examples Folder",
"downloadExamples": "Download Example Images",
"downloadMissingExamples": "Download Missing",
"reprocessExamples": "Re-process All",
"replacePreview": "Replace Preview",
"setContentRating": "Set Content Rating",
"moveToFolder": "Move to Folder",
@@ -1548,6 +1555,7 @@
"empty": "No version history available for this model yet.",
"error": "Failed to load versions.",
"missingModelId": "This model is missing a Civitai model id.",
"hfGroupInfo": "This is a HuggingFace model group. Open the library to see all versions in the grid.",
"confirm": {
"delete": "Delete this version from your library?"
},
@@ -1574,6 +1582,21 @@
"downloadCsv": "Download CSV",
"columnModelName": "Model Name",
"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": {
@@ -1751,6 +1774,12 @@
"checkingMessage": "Please wait while we check for the latest version.",
"showNotifications": "Show update notifications",
"latestBadge": "Latest",
"latestMain": "Latest main",
"channel": "Update Channel",
"channels": {
"release": "Release",
"nightly": "Nightly"
},
"updateProgress": {
"preparing": "Preparing update...",
"installing": "Installing update...",
@@ -1771,6 +1800,15 @@
"warning": "Warning: Nightly builds may contain experimental features and could be unstable.",
"enable": "Enable Nightly Updates"
},
"channelSwitch": {
"nightlyTitle": "Switch to Nightly Channel",
"nightlyMessage": "Switching to Nightly will initialize a Git repository and track the latest main branch commits. Updates will be more frequent but may be unstable. You can switch back to Release at any time.",
"releaseTitle": "Switch to Release Channel",
"releaseMessage": "Switching to Release will checkout the latest stable release tag. You can switch back to Nightly at any time.",
"switching": "Switching to {channel} channel...",
"completed": "Successfully switched to {channel} channel",
"failed": "Failed to switch channel"
},
"banners": {
"recent": "Recent messages",
"empty": "No recent banners yet.",
+2243 -2205
View File
File diff suppressed because it is too large Load Diff
+2243 -2205
View File
File diff suppressed because it is too large Load Diff
+2243 -2205
View File
File diff suppressed because it is too large Load Diff
+2243 -2205
View File
File diff suppressed because it is too large Load Diff
+2243 -2205
View File
File diff suppressed because it is too large Load Diff
+2243 -2205
View File
File diff suppressed because it is too large Load Diff
+2243 -2205
View File
File diff suppressed because it is too large Load Diff
+2243 -2205
View File
File diff suppressed because it is too large Load Diff
+6
View File
@@ -1,5 +1,11 @@
"""Constants used by the metadata collector"""
# Sentinel value for clip_skip to distinguish "unconnected / widget default"
# from "user wired value 0". Both ComfyUI CLIPSetLastLayer (-24..-1) and
# A1111 conventions treat 0 as meaningless for clip skipping, but users may
# explicitly wire 0 to the overwrite node to express "no clip skip / default".
CLIP_SKIP_SENTINEL = -25
# Metadata categories
MODELS = "models"
PROMPTS = "prompts"
+6 -1
View File
@@ -678,7 +678,12 @@ class MetadataProcessor:
for overwrite_info in metadata.get(OVERWRITE, {}).values():
overwrite_params = overwrite_info.get("parameters", {})
for key, value in overwrite_params.items():
if value: # truthy check — only overwrite when user provided a real value
if key == "clip_skip":
# Accept any value from overwrite node (sentinel -25 already
# filtered upstream). Needed because falsy check treats 0
# as "not set" even though 0 is a valid wired input here.
params[key] = value
elif value: # truthy check — only overwrite when user provided a real value
params[key] = value
# Bridge: the overwrite node exposes the field as "model" (more accurate),
+3 -6
View File
@@ -2,7 +2,8 @@ import json
import os
import re
from .constants import 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):
@@ -1233,11 +1234,7 @@ class MetadataOverwriteExtractor(NodeMetadataExtractor):
if not inputs:
return
overwrite_params = {}
for key in METADATA_OVERWRITE_FIELDS:
value = inputs.get(key)
if value: # truthy — only overwrite when user provided a real value
overwrite_params[key] = value
overwrite_params = collect_overwrite_params(inputs)
if overwrite_params:
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:
NAME = "Lora Stack Combiner (LoraManager)"
CATEGORY = "Lora Manager/stackers"
DESCRIPTION = (
"Combines multiple LoRA stacks into a single stack. "
"Supports dynamic inputs: connect a stack to add more inputs."
)
@classmethod
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 {
"required": {
"lora_stack_a": ("LORA_STACK",),
"lora_stack_b": ("LORA_STACK",),
},
"required": {},
"optional": optional_inputs,
}
RETURN_TYPES = ("LORA_STACK",)
RETURN_NAMES = ("LORA_STACK",)
FUNCTION = "combine_stacks"
def combine_stacks(self, lora_stack_a, lora_stack_b):
combined_stack = []
def combine_stacks(self, lora_stack1=None, lora_stack2=None, **kwargs):
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.extend(lora_stack_a)
if lora_stack_b:
combined_stack.extend(lora_stack_b)
combined_stack = []
for key in sorted(stacks, key=_stack_slot_number):
stack = stacks[key]
if stack:
combined_stack.extend(stack)
return (combined_stack,)
+29 -17
View File
@@ -1,13 +1,16 @@
"""Metadata Overwrite node — allows users to manually specify generation parameters
that override the automatically collected/inferred metadata.
All inputs have falsy defaults: only truthy (non-empty / non-zero) values
will overwrite the corresponding field in the final metadata.
Most inputs have falsy defaults (empty string / 0) which are skipped.
clip_skip uses a sentinel default (-25) so that a wired value of 0 is
preserved both ComfyUI and A1111 conventions have no meaningful 0 value,
but users may wire 0 to express "no clip skip / default".
"""
from typing import Any
from ..metadata_collector.constants import METADATA_OVERWRITE_FIELDS
from ..metadata_collector.constants import CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL
from ..metadata_collector.overwrite_utils import collect_overwrite_params
class MetadataOverwriteLM:
@@ -82,12 +85,16 @@ class MetadataOverwriteLM:
},
),
"model": (
"STRING",
"STRING,MODEL",
{
"default": "",
"widgetType": "STRING",
"tooltip": (
"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."
),
},
),
@@ -116,10 +123,14 @@ class MetadataOverwriteLM:
"clip_skip": (
"INT",
{
"default": 0,
"min": -24,
"default": _CLIP_SKIP_SENTINEL,
"min": -25,
"max": 24,
"tooltip": "Clip skip. Only overwrites when non-zero.",
"tooltip": (
"Clip skip (ComfyUI: -24..-1, A1111: 1+). "
"Default -25 means not set — any other value "
"overwrites."
),
},
),
"additional_data": (
@@ -144,14 +155,15 @@ class MetadataOverwriteLM:
OUTPUT_NODE = True
def collect_metadata(self, **kwargs: Any) -> tuple[dict[str, Any]]:
"""Collect non-falsy input values into a metadata dict.
"""Collect non-default input values into a metadata dict.
Only values that are truthy (non-empty string, non-zero number)
are included matching the overwrite logic in the metadata pipeline.
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 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] = {}
for key in METADATA_OVERWRITE_FIELDS:
value = kwargs.get(key)
if value:
result[key] = value
return (result,)
return (collect_overwrite_params(kwargs),)
+33 -10
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.",
},
),
"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": (
"BOOLEAN",
{
@@ -348,7 +355,7 @@ class SaveImageLM:
type_lower = model_type.lower() if model_type else "other"
return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}"
def format_metadata(self, metadata_dict: dict) -> 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."""
if not metadata_dict: return ""
@@ -446,29 +453,42 @@ class SaveImageLM:
lora_resource["versionName"] = lora_civitai["name"]
civitai_resources.append(lora_resource)
sampler_display = self._get_civitai_sampler_name(sampler, scheduler)
sampler_name = CIVITAI_SAMPLER_MAP.get(sampler, sampler) if sampler else None
scheduler_mapping = {
"normal": "Normal",
"karras": "Karras",
"exponential": "Exponential",
"sgm_uniform": "SGM Uniform",
"sgm_quadratic": "SGM Quadratic",
}
scheduler_name = scheduler_mapping.get(scheduler, scheduler) if scheduler else None
# 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:
lines.append(f"Negative prompt: {negative_prompt}")
params: list[str] = []
if steps is not None:
params.append(f"Steps: {steps}")
if sampler_display:
params.append(f"Sampler: {sampler_display}")
if sampler_name:
if scheduler_name:
params.append(f"Sampler: {sampler_name} {scheduler_name}")
else:
params.append(f"Sampler: {sampler_name}")
if cfg is not None:
params.append(f"CFG scale: {cfg}")
if seed is not None:
params.append(f"Seed: {seed}")
if size:
params.append(f"Size: {size}")
if clip_skip:
if clip_skip is not None:
try:
cs = int(clip_skip)
if cs != 0:
params.append(f"Clip skip: {abs(cs)}")
params.append(f"Clip skip: {abs(int(clip_skip))}")
except (ValueError, TypeError):
pass
additional_data = metadata_dict.get("additional_data", "")
@@ -783,6 +803,7 @@ class SaveImageLM:
save_with_metadata=True,
add_counter_to_filename=True,
save_as_recipe=False,
add_loras_to_prompt=False,
):
"""Save images with metadata"""
results = []
@@ -791,7 +812,7 @@ class SaveImageLM:
raw_metadata = get_metadata()
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
filename_prefix = self.format_filename(filename_prefix, metadata_dict)
@@ -933,6 +954,7 @@ class SaveImageLM:
save_with_metadata=True,
add_counter_to_filename=True,
save_as_recipe=False,
add_loras_to_prompt=False,
):
"""Process and save image with metadata"""
# Make sure the output directory exists
@@ -964,6 +986,7 @@ class SaveImageLM:
save_with_metadata,
add_counter_to_filename,
save_as_recipe,
add_loras_to_prompt,
)
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__)
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:
"""UNET Loader with support for extra folder paths
@@ -196,6 +211,12 @@ class UNETLoaderLM:
# Wrap with GGUFModelPatcher
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,)
except Exception as e:
+42 -3
View File
@@ -1562,6 +1562,11 @@ class SettingsHandler:
{"success": False, "error": validation_error}
)
if key == "update_channel" and value not in ("release", "nightly"):
return web.json_response(
{"success": False, "error": "update_channel must be 'release' or 'nightly'"}
)
if value == "__DELETE__" and key in (
"proxy_username",
"proxy_password",
@@ -2585,6 +2590,8 @@ class ModelLibraryHandler:
status=400,
)
cursor = request.query.get("cursor")
metadata_provider = await self._metadata_provider_factory()
if not metadata_provider:
return web.json_response(
@@ -2593,7 +2600,7 @@ class ModelLibraryHandler:
)
try:
models = await metadata_provider.get_user_models(username)
result = await metadata_provider.get_user_models(username, cursor)
except NotImplementedError:
return web.json_response(
{
@@ -2603,14 +2610,35 @@ class ModelLibraryHandler:
status=501,
)
if models is None:
if result is None:
return web.json_response(
{"success": False, "error": "Failed to fetch user models"},
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):
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()
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
@@ -2630,6 +2658,7 @@ class ModelLibraryHandler:
versions: list[dict] = []
history_service = await self._get_download_history_service()
model_ids: list[int] = []
model_count = 0
for model in models:
try:
model_ids.append(int(model.get("id")))
@@ -2663,6 +2692,8 @@ class ModelLibraryHandler:
if model_type not in normalized_allowed_types:
continue
model_count += 1
scanner = type_scanner_map.get(model_type)
if scanner is None:
return web.json_response(
@@ -2728,7 +2759,15 @@ class ModelLibraryHandler:
)
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
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
+3 -1
View File
@@ -394,12 +394,14 @@ class ModelListingHandler:
)
# View-local-versions filter: show all local versions of a specific model
# Accepts either a CivitAI modelId (int) or a HF group key like "hf:user/repo"
civitai_model_id = request.query.get("civitai_model_id")
if civitai_model_id is not None:
try:
civitai_model_id = int(civitai_model_id)
except (TypeError, ValueError):
civitai_model_id = None
# Keep as string — could be an HF group key (e.g. "hf:user/repo")
pass
return {
"page": page,
+313 -45
View File
@@ -38,6 +38,84 @@ def _clean_excludes() -> List[str]:
return excludes
def _stage_preserved_items(plugin_root: str) -> tuple[str, list[str]]:
"""Move preserved user-data items to a temp directory outside *plugin_root*.
This ensures that ``git reset --hard``, ``git clean -fd``, and ZIP-based
replacement cannot touch these files even when ``-e`` exclusion patterns
are mishandled (e.g. on Windows where forward-slash patterns may not
match backslash-prefixed paths in some Git builds, or where file locks
prevent deletion/recreation).
Returns:
``(backup_root, staged_names)``: the temp directory path and the
list of item names that were successfully moved.
"""
backup_root = tempfile.mkdtemp(prefix='lora_manager_update_')
staged: list[str] = []
for name in _PRESERVE_DIRS:
src = os.path.join(plugin_root, name)
if not os.path.lexists(src):
continue
dst = os.path.join(backup_root, name)
try:
shutil.move(src, dst)
staged.append(name)
logger.debug("Staged '%s' for update safety", name)
except OSError:
# ``shutil.move`` may fail on Windows if a file handle inside
# the directory is still open (e.g. a SQLite WAL file). Fall
# back to copy-then-remove.
logger.debug("Move failed for '%s', falling back to copy", name)
try:
if os.path.isdir(src) and not os.path.islink(src):
shutil.copytree(src, dst, symlinks=True)
shutil.rmtree(src, ignore_errors=True)
else:
shutil.copy2(src, dst)
os.remove(src)
staged.append(name)
logger.info("Copied (then removed) '%s' for update safety", name)
except Exception as exc:
logger.warning(
"Could not stage '%s': %s (will rely on git -e / skip lists)", name, exc
)
return backup_root, staged
def _restore_preserved_items(plugin_root: str, backup_root: str, staged: list[str]) -> None:
"""Move staged items back from *backup_root* into *plugin_root*.
Any leftover placeholder at the destination (created by git checkout or
ZIP extraction) is removed before the move.
"""
for name in staged:
src = os.path.join(backup_root, name)
dst = os.path.join(plugin_root, name)
try:
if os.path.lexists(dst):
if os.path.isdir(dst) and not os.path.islink(dst):
shutil.rmtree(dst, ignore_errors=True)
else:
os.remove(dst)
shutil.move(src, dst)
logger.debug("Restored '%s' after update", name)
except OSError:
logger.debug("Move failed restoring '%s', falling back to copy", name)
try:
if os.path.isdir(src) and not os.path.islink(src):
shutil.copytree(src, dst, symlinks=True, dirs_exist_ok=True)
shutil.rmtree(src, ignore_errors=True)
else:
shutil.copy2(src, dst)
os.remove(src)
logger.info("Copied '%s' back after update", name)
except Exception as exc:
logger.error("Failed to restore '%s': %s", name, exc)
shutil.rmtree(backup_root, ignore_errors=True)
class UpdateRoutes:
"""Routes for handling plugin update checks"""
@@ -47,6 +125,7 @@ class UpdateRoutes:
app.router.add_get('/api/lm/check-updates', UpdateRoutes.check_updates)
app.router.add_get('/api/lm/version-info', UpdateRoutes.get_version_info)
app.router.add_post('/api/lm/perform-update', UpdateRoutes.perform_update)
app.router.add_post('/api/lm/switch-channel', UpdateRoutes.switch_channel)
@staticmethod
async def check_updates(request):
@@ -65,10 +144,17 @@ class UpdateRoutes:
# Fetch remote version from GitHub
if nightly:
remote_version, changelog = await UpdateRoutes._get_nightly_version()
releases = None
local_hash = git_info.get('short_hash', '')
nightly_version, releases_result = await asyncio.gather(
UpdateRoutes._get_nightly_version(local_hash),
UpdateRoutes._get_remote_version()
)
remote_version, _, behind_by, commit_date = nightly_version
_, changelog, releases = releases_result
else:
remote_version, changelog, releases = await UpdateRoutes._get_remote_version()
behind_by = 0
commit_date = ''
# Compare versions
if nightly:
@@ -81,6 +167,10 @@ class UpdateRoutes:
remote_version.replace('v', '')
)
current_dir = os.path.dirname(os.path.abspath(__file__))
plugin_root = os.path.dirname(os.path.dirname(current_dir))
has_git = os.path.exists(os.path.join(plugin_root, '.git'))
response_data = {
'success': True,
'current_version': local_version,
@@ -88,13 +178,13 @@ class UpdateRoutes:
'update_available': update_available,
'changelog': changelog,
'git_info': git_info,
'nightly': nightly
'nightly': nightly,
'has_git': has_git,
'releases': releases,
'behind_by': behind_by,
'commit_date': commit_date
}
# Include releases list for stable mode
if releases is not None:
response_data['releases'] = releases
return web.json_response(response_data)
except NETWORK_EXCEPTIONS as e:
@@ -126,9 +216,14 @@ class UpdateRoutes:
# Format: version-short_hash
version_string = f"{local_version}-{short_hash}"
current_dir = os.path.dirname(os.path.abspath(__file__))
plugin_root = os.path.dirname(os.path.dirname(current_dir))
has_git = os.path.exists(os.path.join(plugin_root, '.git'))
return web.json_response({
'success': True,
'version': version_string
'version': version_string,
'has_git': has_git
})
except Exception as e:
@@ -156,20 +251,22 @@ class UpdateRoutes:
if os.path.exists(settings_path):
with open(settings_path, 'r', encoding='utf-8') as f:
settings_backup = f.read()
logger.info("Backed up settings.json")
logger.debug("Backed up settings.json (%d bytes)", len(settings_backup))
git_folder = os.path.join(plugin_root, '.git')
if os.path.exists(git_folder):
# Git update
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
else:
# Fallback: Download ZIP and replace files
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
try:
git_folder = os.path.join(plugin_root, '.git')
if os.path.exists(git_folder):
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
else:
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
finally:
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
if settings_backup and success:
with open(settings_path, 'w', encoding='utf-8') as f:
f.write(settings_backup)
logger.info("Restored settings.json")
logger.debug("Restored settings.json content (%d bytes)", len(settings_backup))
if success:
return web.json_response({
@@ -190,6 +287,164 @@ class UpdateRoutes:
'error': str(e)
})
@staticmethod
async def switch_channel(request):
"""
Switch between release and nightly update channels.
ZIP/CNR install Nightly: git init + checkout main (one-way upgrade)
Git install Release: git checkout latest tag (.git preserved)
ZIP/CNR install Release: ZIP download (no .git, stays in ZIP mode)
Git install Nightly: git checkout main + pull
"""
try:
body = await request.json() if request.has_body else {}
channel = body.get('channel', '')
if channel not in ('release', 'nightly'):
return web.json_response({
'success': False,
'error': f'Invalid channel: {channel}. Must be "release" or "nightly".'
})
current_dir = os.path.dirname(os.path.abspath(__file__))
plugin_root = os.path.dirname(os.path.dirname(current_dir))
settings_path = ensure_settings_file(logger)
settings_backup = None
if os.path.exists(settings_path):
with open(settings_path, 'r', encoding='utf-8') as f:
settings_backup = f.read()
logger.debug("Backed up settings.json before channel switch (%d bytes)", len(settings_backup))
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
try:
git_folder = os.path.join(plugin_root, '.git')
if channel == 'nightly':
git_backup = None
if os.path.exists(git_folder):
git_backup = UpdateRoutes._backup_git(git_folder, 'nightly')
success = False
new_version = ''
try:
if os.path.exists(git_folder):
success, new_version = await UpdateRoutes._perform_git_update(
plugin_root, nightly=True
)
else:
success, new_version = UpdateRoutes._init_git_repo(plugin_root)
finally:
UpdateRoutes._restore_git(git_backup, git_folder, success, 'nightly')
else:
success = False
new_version = ''
if os.path.exists(git_folder):
success, new_version = await UpdateRoutes._perform_git_update(
plugin_root, nightly=False
)
else:
tracking_file = os.path.join(plugin_root, '.tracking')
if os.path.exists(tracking_file):
os.remove(tracking_file)
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
finally:
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
if settings_backup and success:
with open(settings_path, 'w', encoding='utf-8') as f:
f.write(settings_backup)
logger.debug("Restored settings.json content after channel switch (%d bytes)", len(settings_backup))
if success:
return web.json_response({
'success': True,
'channel': channel,
'new_version': new_version,
'message': f'Switched to {channel} channel'
})
else:
return web.json_response({
'success': False,
'error': f'Failed to switch to {channel} channel'
})
except Exception as e:
logger.error("Failed to switch channel: %s", e, exc_info=True)
return web.json_response({
'success': False,
'error': str(e)
})
@staticmethod
def _init_git_repo(plugin_root: str) -> tuple[bool, str]:
"""
Initialize a Git repository in a ZIP-installed plugin folder.
Clones the remote history and checks out main branch.
"""
try:
import git
except ImportError:
logger.error(
"GitPython is not available: cannot initialize git repo. "
"Install git or set $GIT_PYTHON_GIT_EXECUTABLE to the git binary path."
)
return False, ""
clean_excludes = _clean_excludes()
try:
repo = git.Repo.init(plugin_root)
origin = repo.create_remote(
'origin',
'https://github.com/willmiao/ComfyUI-Lora-Manager.git'
)
origin.fetch()
repo.create_head('main', origin.refs.main)
repo.git.checkout('main', '--force')
repo.git.reset('--hard')
repo.git.clean('-fd', *clean_excludes)
tracking_file = os.path.join(plugin_root, '.tracking')
if os.path.exists(tracking_file):
os.remove(tracking_file)
logger.info("Removed .tracking file (now in git mode)")
new_version = f"main-{repo.head.commit.hexsha[:7]}"
logger.info("Initialized git repo on main branch: %s", new_version)
return True, new_version
except Exception as e:
logger.error("Failed to initialize git repo: %s", e, exc_info=True)
return False, ""
@staticmethod
def _backup_git(git_folder, label):
try:
backup_dir = tempfile.mkdtemp()
backup = os.path.join(backup_dir, '.git')
shutil.copytree(git_folder, backup)
logger.info("Backed up .git before switching to %s", label)
return backup
except Exception as e:
logger.error("Failed to backup .git before %s switch: %s", label, e)
return None
@staticmethod
def _restore_git(git_backup, git_folder, success, label):
if git_backup and not success:
try:
if os.path.exists(git_folder):
shutil.rmtree(git_folder)
shutil.copytree(git_backup, git_folder)
logger.info("Restored .git after failed %s switch", label)
except Exception as e:
logger.error("Failed to restore .git after %s switch: %s", label, e)
if git_backup:
shutil.rmtree(os.path.dirname(git_backup), ignore_errors=True)
@staticmethod
async def _download_and_replace_zip(plugin_root: str) -> tuple[bool, str]:
"""
@@ -244,8 +499,7 @@ class UpdateRoutes:
except Exception:
logger.debug("Could not close downloaded-version history database", exc_info=True)
# Skip settings.json, civitai, model cache and runtime cache folders
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=['settings.json', 'civitai', 'model_cache', 'cache', 'wildcards', 'backups', 'stats'])
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=list(_PRESERVE_DIRS))
# Extract ZIP to temp dir
with tempfile.TemporaryDirectory() as tmp_dir:
@@ -255,7 +509,7 @@ class UpdateRoutes:
extracted_root = next(os.scandir(tmp_dir)).path
# Copy files, skipping user data that should be preserved
skip_items = {'settings.json', 'civitai', 'wildcards', 'backups', 'stats'}
skip_items = set(_PRESERVE_DIRS)
for item in os.listdir(extracted_root):
if item in skip_items:
continue
@@ -272,7 +526,7 @@ class UpdateRoutes:
# for ComfyUI Manager to work properly
tracking_info_file = os.path.join(plugin_root, '.tracking')
tracking_files = []
skip_tracked = {'civitai', 'wildcards', 'backups', 'stats'}
skip_tracked = set(_PRESERVE_DIRS) - {'settings.json'}
for root, dirs, files in os.walk(extracted_root):
# Skip user data directories and their contents
rel_root = os.path.relpath(root, extracted_root)
@@ -295,7 +549,8 @@ class UpdateRoutes:
except Exception as e:
logger.error(f"ZIP update failed: {e}", exc_info=True)
return False, ""
@staticmethod
def _clean_plugin_folder(plugin_root, skip_files=None):
skip_files = skip_files or []
for item in os.listdir(plugin_root):
@@ -308,41 +563,54 @@ class UpdateRoutes:
os.remove(path)
@staticmethod
async def _get_nightly_version() -> tuple[str, List[str]]:
"""
Fetch latest commit from main branch
"""
async def _get_nightly_version(local_hash: str = "") -> tuple[str, List[str], int, str]:
repo_owner = "willmiao"
repo_name = "ComfyUI-Lora-Manager"
# Use GitHub API to fetch the latest commit from main branch
github_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/commits/main"
try:
downloader = await get_downloader()
success, data = await downloader.make_request('GET', github_url, custom_headers={'Accept': 'application/vnd.github+json'})
success, data = await downloader.make_request(
'GET', github_url,
custom_headers={'Accept': 'application/vnd.github+json'}
)
if not success:
logger.warning(f"Failed to fetch GitHub commit: {data}")
return "main", []
commit_sha = data.get('sha', '')[:7] # Short hash
logger.warning("Failed to fetch GitHub commit: %s", data)
return "main", [], 0, ""
commit_sha = data.get('sha', '')[:7]
commit_message = data.get('commit', {}).get('message', '')
# Format as "main-{short_hash}"
commit_date = data.get('commit', {}).get('committer', {}).get('date', '')[:10]
version = f"main-{commit_sha}"
# Use commit message as changelog
changelog = [commit_message] if commit_message else []
return version, changelog
behind_by = 0
if local_hash and local_hash not in ('unknown', 'stable'):
compare_url = (
f"https://api.github.com/repos/{repo_owner}/{repo_name}"
f"/compare/{local_hash}...main"
)
c_ok, c_data = await downloader.make_request(
'GET', compare_url,
custom_headers={'Accept': 'application/vnd.github+json'}
)
if c_ok:
if c_data.get('status') in ('ahead', 'diverged'):
behind_by = c_data.get('ahead_by', 0)
else:
behind_by = c_data.get('behind_by', 0)
return version, changelog, behind_by, commit_date
except NETWORK_EXCEPTIONS as e:
logger.warning("Unable to reach GitHub for nightly version: %s", e)
return "main", []
return "main", [], 0, ""
except Exception as e:
logger.error(f"Error fetching nightly version: {e}", exc_info=True)
return "main", []
logger.error("Error fetching nightly version: %s", e, exc_info=True)
return "main", [], 0, ""
@staticmethod
def _compare_nightly_versions(local_git_info: Dict[str, str], remote_version: str) -> bool:
+52 -9
View File
@@ -1,7 +1,8 @@
from abc import ABC, abstractmethod
import asyncio
import re
from typing import Any, Dict, List, Optional, Type, TYPE_CHECKING
import random
from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING
import logging
import os
import time
@@ -109,12 +110,15 @@ class BaseModelService(ABC):
if civitai_model_id is not None:
sorted_data = [
item for item in sorted_data
if self._extract_model_id(item) == civitai_model_id
if self._extract_group_key(item) == civitai_model_id
]
# VLM mode: always sort by version ID descending (newest version first),
# regardless of the current sort_by preference.
# Fall back to modified timestamp for non-CivitAI sources.
sorted_data.sort(
key=lambda x: self._extract_version_id(x) or 0,
key=lambda x: self._extract_version_id(x)
or x.get("modified", 0)
or 0,
reverse=True,
)
@@ -129,18 +133,21 @@ class BaseModelService(ABC):
ufs = self.settings.get("version_grouping", "same_base")
group_by_base = ufs == "same_base"
dedup_map = {} # (modelId [,base_model]) -> (item, version_id)
dedup_map = {} # (modelId [,base_model]) -> (item, version_or_modified)
version_counter = {} # same-key -> count
standalone = []
for item in sorted_data:
mid = self._extract_model_id(item)
mid = self._extract_group_key(item)
if mid is None:
standalone.append(item)
continue
key = (mid, item.get("base_model") or "") if group_by_base else mid
# Count all versions per key
version_counter[key] = version_counter.get(key, 0) + 1
vid = self._extract_version_id(item) or 0
# Prefer CivitAI version_id; fall back to modified timestamp
vid = self._extract_version_id(item)
if vid is None:
vid = item.get("modified", 0) or 0
if key not in dedup_map or vid > dedup_map[key][1]:
dedup_map[key] = (item, vid)
# Attach version_count to each surviving grouped item (shallow copy
@@ -174,16 +181,19 @@ class BaseModelService(ABC):
model_groups: Dict[Any, List[Dict]] = {}
ungrouped_standalone: List[Dict] = []
for item in sorted_data:
mid = self._extract_model_id(item)
mid = self._extract_group_key(item)
if mid is None:
ungrouped_standalone.append(item)
continue
key = (mid, item.get("base_model") or "") if group_by_base else mid
model_groups.setdefault(key, []).append(item)
# Sort versions within each group by version id descending
# Sort versions within each group by version id (descending);
# fall back to modified timestamp for non-CivitAI sources.
for items in model_groups.values():
items.sort(
key=lambda x: self._extract_version_id(x) or 0,
key=lambda x: self._extract_version_id(x)
or x.get("modified", 0)
or 0,
reverse=True,
)
# Sort groups by version count
@@ -381,6 +391,12 @@ class BaseModelService(ABC):
(item.get("model_name") or item.get("file_name") or "").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":
key_fn = lambda item: (
int(item.get("size", 0) or 0),
@@ -697,6 +713,33 @@ class BaseModelService(ABC):
return annotated
@staticmethod
def _extract_hf_group_key(item: Dict) -> Optional[str]:
"""Extract `hf:{owner}/{repo}` from item's ``hf_url``, or None."""
hf_url = item.get("hf_url") if isinstance(item, dict) else None
if not hf_url or not isinstance(hf_url, str):
return None
m = re.match(
r"https?://huggingface\.co/([^/]+/[^/]+)", hf_url.strip()
)
if not m:
return None
return f"hf:{m.group(1)}"
@staticmethod
def _extract_group_key(item: Dict) -> Union[int, str, None]:
"""Return the group identity key: CivitAI modelId (int) or HF repo (str).
Preference order:
1. CivitAI ``modelId`` (int)
2. HF repo identity ``hf:{owner}/{repo}`` (str)
3. ``None`` (no known grouping source)
"""
mid = BaseModelService._extract_model_id(item)
if mid is not None:
return mid
return BaseModelService._extract_hf_group_key(item)
@staticmethod
def _extract_model_id(item: Dict) -> Optional[int]:
civitai = item.get("civitai") if isinstance(item, dict) else None
+88 -5
View File
@@ -2,6 +2,7 @@ import asyncio
import copy
import logging
import os
import time
from collections import OrderedDict
from typing import Any, Optional, Dict, Tuple, List, Sequence
from .connectivity_guard import (
@@ -19,6 +20,12 @@ from ..utils.civitai_utils import resolve_license_payload
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:
_instance = None
@@ -743,17 +750,34 @@ class CivitaiClient:
return all_versions if all_versions else None
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
"""Fetch all models for a specific Civitai user."""
async def get_user_models(
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:
return None
params: Dict[str, Any] = {
"username": username,
"nsfw": "true",
"limit": 100,
"sort": "Newest",
"period": "AllTime",
}
if cursor:
params["cursor"] = cursor
try:
success, result = await self._make_request(
"GET",
f"{self.base_url}/models",
use_auth=True,
params={"username": username, "nsfw": "true"},
params=params,
)
if not success:
@@ -765,7 +789,7 @@ class CivitaiClient:
items = result.get("items") if isinstance(result, dict) else None
if not isinstance(items, list):
return []
items = []
for model in items:
versions = model.get("modelVersions")
@@ -774,9 +798,68 @@ class CivitaiClient:
for version in versions:
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:
raise
except Exception as exc: # pragma: no cover - defensive logging
logger.error("Error fetching models for %s: %s", username, exc)
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
+34 -2
View File
@@ -31,7 +31,7 @@ class DownloadQueueService:
_instance: Optional[DownloadQueueService] = None
_class_lock: asyncio.Lock = asyncio.Lock()
_SCHEMA = """
_SCHEMA_TABLES = """
CREATE TABLE IF NOT EXISTS download_queue (
download_id TEXT PRIMARY KEY,
model_id INTEGER,
@@ -74,6 +74,9 @@ class DownloadQueueService:
);
CREATE INDEX IF NOT EXISTS idx_dh_completed ON download_history(completed_at DESC);
CREATE INDEX IF NOT EXISTS idx_dh_status ON download_history(status);
"""
_CREATE_UNIQUE_INDEX = """
CREATE UNIQUE INDEX IF NOT EXISTS idx_dh_download_id
ON download_history(download_id) WHERE download_id IS NOT NULL;
"""
@@ -115,10 +118,39 @@ class DownloadQueueService:
if self._schema_initialized:
return
with self._connect() as conn:
conn.executescript(self._SCHEMA)
conn.executescript(self._SCHEMA_TABLES)
# Creating the unique index on download_history.download_id can
# fail if pre-existing rows have duplicate values (e.g. from a
# previous version that lacked the index). Deduplicate first so
# that the migration does not crash on startup.
if not self._index_exists(conn, "idx_dh_download_id"):
self._remove_duplicate_download_ids(conn)
conn.executescript(self._CREATE_UNIQUE_INDEX)
conn.commit()
self._schema_initialized = True
@staticmethod
def _index_exists(conn: sqlite3.Connection, name: str) -> bool:
return conn.execute(
"SELECT 1 FROM sqlite_master WHERE type='index' AND name=?",
(name,),
).fetchone() is not None
@staticmethod
def _remove_duplicate_download_ids(conn: sqlite3.Connection) -> None:
conn.execute("""
DELETE FROM download_history
WHERE id NOT IN (
SELECT MIN(id)
FROM download_history
WHERE download_id IS NOT NULL
GROUP BY download_id
)
AND download_id IS NOT NULL
""")
def get_database_path(self) -> str:
"""Return the resolved database file path."""
return self._db_path
+5
View File
@@ -201,6 +201,11 @@ PROVIDER_PRESETS: Dict[str, Dict[str, Any]] = {
"api_base": "https://openrouter.ai/api/v1",
"requires_key": True,
},
"google": {
"name": "Gemini",
"api_base": "https://generativelanguage.googleapis.com/v1beta/openai",
"requires_key": True,
},
"opencode-go": {
"name": "OpenCode Go",
"api_base": "https://opencode.ai/zen/go/v1",
+21 -12
View File
@@ -1,6 +1,7 @@
import asyncio
import time
import logging
import random
logger = logging.getLogger(__name__)
from typing import Any, Dict, List, Optional, Tuple
@@ -38,8 +39,8 @@ class ModelCache:
def __post_init__(self):
self._lock = asyncio.Lock()
# Cache for last sort: (sort_key, order) -> sorted list
self._last_sort: Tuple[str, str] = (None, None)
# Cache for last sort: (sort_key, order, seed) -> sorted list
self._last_sort: Tuple[Optional[str], str, Optional[str]] = (None, "asc", None)
self._last_sorted_data: List[Dict] = []
self._normalize_raw_data()
self.name_display_mode = self._normalize_display_mode(self.name_display_mode)
@@ -203,9 +204,9 @@ class ModelCache:
async def resort(self):
"""Resort cached data according to last sort mode if set"""
async with self._lock:
if self._last_sort != (None, None):
sort_key, order = self._last_sort
sorted_data = self._sort_data(self.raw_data, sort_key, order)
if self._last_sort[0] is not None:
sort_key, order, seed = self._last_sort
sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
self._last_sorted_data = sorted_data
# Update folder list
# else: do nothing
@@ -218,7 +219,7 @@ class ModelCache:
self.folders = sorted(list(all_folders), key=lambda x: x.lower())
self.rebuild_version_index()
def _sort_data(self, data: List[Dict], sort_key: str, order: str) -> 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"""
start_time = time.perf_counter()
reverse = (order == 'desc')
@@ -265,6 +266,13 @@ class ModelCache:
),
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':
# Pre-dedup sort: fall back to name sort.
# 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)
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"""
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
start_time = time.perf_counter()
sorted_data = self._sort_data(self.raw_data, sort_key, order)
self._last_sort = (sort_key, order)
sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
self._last_sort = cache_key
self._last_sorted_data = sorted_data
duration = time.perf_counter() - start_time
@@ -313,8 +322,8 @@ class ModelCache:
self.name_display_mode = normalized
if self._last_sort[0] == 'name':
sort_key, order = self._last_sort
self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order)
sort_key, order, seed = self._last_sort
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:
"""Update preview_url for a specific model in all cached data
+50 -11
View File
@@ -143,10 +143,18 @@ class ModelMetadataProvider(ABC):
pass
@abstractmethod
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
"""Fetch models owned by the specified user"""
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
"""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
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):
"""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]]:
return await self.client.get_model_version_info(version_id)
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
return await self.client.get_user_models(username)
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
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):
"""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]]:
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"""
return None
@@ -347,7 +358,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
version_data = await self._get_version_with_model_data(db, model_id, version_id)
return version_data, None
async def get_user_models(self, username: str) -> 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"""
return None
@@ -602,13 +613,14 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue
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():
try:
result = await self._call_with_rate_limit(
label,
provider.get_user_models,
username,
cursor=cursor,
)
if result is not None:
return result
@@ -624,6 +636,19 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue
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):
return zip(self.providers, self._provider_labels)
@@ -704,13 +729,17 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
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(
self._label,
self._provider.get_user_models,
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:
"""Manager for selecting and using model metadata providers"""
@@ -776,10 +805,20 @@ class ModelMetadataProviderManager:
except NotImplementedError:
return None
async def get_user_models(self, username: str, provider_name: str = None) -> Optional[List[Dict]]:
"""Fetch models owned by the specified user"""
async def get_user_models(
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)
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:
"""Get provider by name or default provider"""
+11 -3
View File
@@ -85,6 +85,7 @@ class SortParams:
key: str
order: str
seed: Optional[str] = None
@dataclass(frozen=True)
@@ -116,7 +117,7 @@ class ModelCacheRepository:
async def fetch_sorted(self, params: SortParams) -> List[Dict[str, Any]]:
"""Fetch cached data pre-sorted according to ``params``."""
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
def parse_sort(sort_by: str) -> SortParams:
@@ -132,10 +133,17 @@ class ModelCacheRepository:
sort_key = sort_by.strip().lower() or "name"
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"
return SortParams(key=sort_key, order=order)
return SortParams(key=sort_key, order=order, seed=seed)
class ModelFilterSet:
+36 -10
View File
@@ -927,6 +927,25 @@ class ModelScanner:
# Update cache data
self._cache.raw_data = [item for item in self._cache.raw_data if item['file_path'] not in missing_files]
dedup_removed = 0
seen_paths: set = set()
deduped: list = []
for item in reversed(self._cache.raw_data):
path = item.get('file_path', '')
if path not in seen_paths:
seen_paths.add(path)
deduped.append(item)
else:
for tag in item.get('tags', []):
if tag in self._tags_count:
self._tags_count[tag] = max(0, self._tags_count[tag] - 1)
if self._tags_count[tag] == 0:
del self._tags_count[tag]
dedup_removed += 1
if dedup_removed > 0:
self._cache.raw_data = list(reversed(deduped))
total_removed += dedup_removed
# Resort cache if changes were made
if total_added > 0 or total_removed > 0:
# Update folders list
@@ -1352,18 +1371,25 @@ class ModelScanner:
# Update folder in metadata
metadata_dict['folder'] = folder
# Add to cache
self._cache.raw_data.append(metadata_dict)
self._cache.add_to_version_index(metadata_dict)
file_path = metadata_dict.get('file_path', '')
if file_path:
old_entries = [item for item in self._cache.raw_data if item.get('file_path') == file_path]
for old_entry in old_entries:
for tag in old_entry.get('tags', []):
if tag in self._tags_count:
self._tags_count[tag] = max(0, self._tags_count[tag] - 1)
if self._tags_count[tag] == 0:
del self._tags_count[tag]
self._hash_index.remove_by_path(file_path)
self._cache.raw_data = [item for item in self._cache.raw_data if item.get('file_path') != file_path]
for tag in metadata_dict.get('tags', []):
self._tags_count[tag] = self._tags_count.get(tag, 0) + 1
self._cache.raw_data.append(metadata_dict)
# Resort cache data
await self._cache.resort()
# Update folders list
all_folders = set(self._cache.folders)
all_folders.add(folder)
self._cache.folders = sorted(list(all_folders), key=lambda x: x.lower())
# Update the hash index
self._hash_index.add_entry(metadata_dict['sha256'], metadata_dict['file_path'])
await self._persist_current_cache()
@@ -1726,7 +1752,7 @@ class ModelScanner:
# ---- Conditional resort (only when sort-key fields changed) ----
need_resort = False
_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 (
old_model_name != desired_entry.get("model_name", "")
+5 -3
View File
@@ -1473,10 +1473,12 @@ class SettingsManager:
try:
common_root = os.path.commonpath([source, target])
except ValueError as exc:
raise ValueError("Invalid recipes path change") from exc
except ValueError:
# 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")
planned_recipe_updates: Dict[str, Dict[str, Any]] = {}
+36 -3
View File
@@ -19,7 +19,7 @@ logger = logging.getLogger(__name__)
_WILDCARD_PATTERN = re.compile(r"__([\w\s.\-+/*\\]+?)__")
_OPTION_PATTERN = re.compile(r"{([^{}]*?)}")
_TRIGGER_WORD_PATTERN = re.compile(r"^trigger_words\d+$")
_WEIGHTED_OPTION_PATTERN = re.compile(r"^\s*([0-9.]+)::")
_WEIGHTED_OPTION_PATTERN = re.compile(r"^\s*-?\d+(\.\d+)?::")
_NUMERIC_PATTERN = re.compile(r"^-?\d+(\.\d+)?$")
@@ -390,7 +390,7 @@ class WildcardService:
) -> str | None:
keyword = _normalize_wildcard_key(raw_key)
if keyword in wildcard_dict:
return rng.choice(wildcard_dict[keyword])
return self._pick_weighted_or_plain(wildcard_dict[keyword], rng)
if "*" in keyword:
regex_pattern = keyword.replace("*", ".*").replace("+", r"\+")
@@ -400,7 +400,7 @@ class WildcardService:
if compiled.match(key):
aggregated.extend(values)
if aggregated:
return rng.choice(aggregated)
return self._pick_weighted_or_plain(aggregated, rng)
if "/" not in keyword:
fallback_keyword = _normalize_wildcard_key(f"*/{keyword}")
@@ -409,6 +409,39 @@ class WildcardService:
return None
def _pick_weighted_or_plain(
self, values: list[str], rng: random.Random
) -> str:
"""Pick a value from the list, respecting N::weight prefix if present.
When any value in the list uses the ``N::value`` weighted syntax with a
weight different from 1, the pick uses weighted random selection. When
no such weighting is present, a plain ``rng.choice`` is used (preserving
backward compatibility for unweighted wildcard files).
In either case the ``N::`` prefix is always stripped from the returned
value, matching the behaviour of ``{...}`` option groups.
"""
# Fast path: skip weighting logic entirely when no :: syntax exists
if not any("::" in v for v in values):
return rng.choice(values)
weighted_options: list[tuple[float, str]] = []
for value in values:
weight = 1.0
parts = value.split("::", 1)
if len(parts) == 2 and _is_numeric_string(parts[0].strip()):
weight = float(parts[0].strip())
weighted_options.append((weight, value))
any_weighted = any(w != 1.0 for w, _ in weighted_options)
if any_weighted:
picked = self._weighted_choice(weighted_options, rng)
else:
picked = rng.choice(values)
return self._strip_weight_prefix(picked)
def is_trigger_words_input(name: str) -> bool:
return bool(_TRIGGER_WORD_PATTERN.match(name))
+132 -33
View File
@@ -14,11 +14,16 @@ from ..services.service_registry import ServiceRegistry
from ..utils.example_images_paths import (
ExampleImagePathResolver,
ensure_library_root_exists,
get_example_images_root,
is_hash_folder,
uses_library_scoped_folders,
)
from ..utils.metadata_manager import MetadataManager
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.settings_manager import get_settings_manager
@@ -87,6 +92,13 @@ class _DownloadProgress(dict):
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:
"""Return True when the provided directory exists and contains entries."""
@@ -103,6 +115,36 @@ def _model_directory_has_files(path: str) -> bool:
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:
"""Manages downloading example images for models."""
@@ -130,6 +172,7 @@ class DownloadManager:
model_types = data.get("model_types", ["lora", "checkpoint"])
delay = float(data.get("delay", 0.2))
force = data.get("force", False)
model_hashes = data.get("model_hashes", [])
# Step 2: Validate configuration (fast lookup)
settings_manager = get_settings_manager()
@@ -199,6 +242,7 @@ class DownloadManager:
delay,
active_library,
force,
model_hashes,
)
)
@@ -410,14 +454,49 @@ class DownloadManager:
# 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,
# and its folder doesn't exist or is empty.
pending_hashes = set()
for model_hash, model_name in all_models_with_hash:
if model_hash not in processed_models and model_hash not in failed_models:
candidate_hashes = [
model_hash
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_hash, active_library
)
if not _model_directory_has_files(model_dir):
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)
@@ -500,8 +579,9 @@ class DownloadManager:
delay,
library_name,
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()
@@ -529,6 +609,18 @@ class DownloadManager:
if model.get("sha256"):
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
self._progress["total"] = len(all_models)
logger.debug(f"Found {self._progress['total']} models to process")
@@ -552,6 +644,7 @@ class DownloadManager:
downloader,
library_name,
force,
explicit_targets,
)
# Update progress
@@ -648,6 +741,7 @@ class DownloadManager:
downloader,
library_name,
force: bool = False,
explicit_targets: bool = False,
):
"""Process a single model download."""
@@ -670,8 +764,9 @@ class DownloadManager:
self._progress["current_model"] = f"{model_name} ({model_hash[:8]})"
await self._broadcast_progress(status="running")
# Skip if already in failed models (unless force mode is enabled)
if not force and model_hash in self._progress["failed_models"]:
# Skip if already in failed models (unless force mode is enabled or
# 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}")
return False
@@ -680,30 +775,34 @@ class DownloadManager:
)
existing_files = _model_directory_has_files(model_dir)
# 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}")
# Model-level guard: a populated folder counts as done. Explicitly
# targeted models bypass it so the per-image existence pre-check can
# fill individual gaps without re-fetching existing files.
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
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:
logger.warning(
"Unable to resolve example images folder for model %s (%s)",
@@ -807,7 +906,7 @@ class DownloadManager:
model_name,
)
# 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)
logger.info(
f"Removed {model_name} from failed_models after force retry with rate-limited images"
@@ -827,7 +926,7 @@ class DownloadManager:
)
elif success:
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)
logger.info(
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)
try:
await scanner.update_single_model_cache(
file_path, file_path, model_data
await update_cache_from_metadata(
scanner, file_path, model_copy
)
except AttributeError:
logger.debug(
+61 -45
View File
@@ -1,3 +1,4 @@
import inspect
import logging
import os
import re
@@ -28,6 +29,31 @@ if TYPE_CHECKING: # pragma: no cover - import for type checkers only
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:
"""Construct a metadata sync service bound to the provided settings."""
@@ -103,8 +129,8 @@ class MetadataUpdater:
progress['refreshed_models'].add(model_hash)
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)
success, error = await _get_metadata_sync_service().fetch_and_update_model(
sha256=model_hash,
@@ -234,6 +260,7 @@ class MetadataUpdater:
# Save metadata to .metadata.json file
file_path = model.get('file_path')
model_copy: Optional[Dict[str, Any]] = None
try:
model_copy = model.copy()
model_copy.pop('folder', None)
@@ -241,14 +268,18 @@ class MetadataUpdater:
logger.info(f"Saved metadata for {model.get('model_name')}")
except Exception as e:
logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}")
# Save updated metadata to scanner cache
success = await scanner.update_single_model_cache(file_path, file_path, model)
if success:
# Save updated metadata to scanner cache. sync_cache_from_metadata
# returns False both for "already in sync" and for actual failures,
# 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")
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
except Exception as e:
@@ -336,6 +367,7 @@ class MetadataUpdater:
# Save metadata to .metadata.json file
file_path = model_data.get('file_path')
model_copy: Optional[Dict[str, Any]] = None
if file_path:
try:
model_copy = model_data.copy()
@@ -344,11 +376,11 @@ class MetadataUpdater:
logger.info(f"Saved metadata for {model_data.get('model_name')}")
except Exception as e:
logger.error(f"Failed to save metadata: {str(e)}")
# Save updated metadata to scanner cache
if file_path:
await scanner.update_single_model_cache(file_path, file_path, model_data)
if file_path and model_copy is not None:
await update_cache_from_metadata(scanner, file_path, model_copy)
# Get regular images array (might be None)
regular_images = civitai_data.get('images', [])
@@ -475,13 +507,19 @@ class MetadataUpdater:
return False
model_folder = get_model_folder(model_hash)
if not model_folder:
if not model_folder or not os.path.isdir(model_folder):
return False
civitai = getattr(metadata, "civitai", None)
if not isinstance(civitai, dict):
return False
# Read the directory listing once so every image entry reuses it.
try:
dir_entries = os.listdir(model_folder)
except OSError:
dir_entries = []
has_changes = False
custom_images = civitai.get("customImages")
@@ -493,24 +531,15 @@ class MetadataUpdater:
if not img_id:
continue
if not os.path.isdir(model_folder):
prefix = f"custom_{img_id}"
found = any(
f.startswith(prefix) and os.path.isfile(
os.path.join(model_folder, f)
)
for f in dir_entries
)
if not found:
stale.append(idx)
else:
found = False
try:
prefix = f"custom_{img_id}"
for fname in os.listdir(model_folder):
if fname.startswith(prefix) and os.path.isfile(
os.path.join(model_folder, fname)
):
found = True
break
except OSError:
stale.append(idx)
continue
if not found:
stale.append(idx)
if stale:
for idx in reversed(stale):
@@ -532,22 +561,9 @@ class MetadataUpdater:
# is gone.
continue
if not os.path.isdir(model_folder):
prefix = f"image_{idx}."
if not any(f.startswith(prefix) for f in dir_entries):
stale.append(idx)
else:
found = False
try:
prefix = f"image_{idx}."
for fname in os.listdir(model_folder):
if fname.startswith(prefix):
found = True
break
except OSError:
stale.append(idx)
continue
if not found:
stale.append(idx)
if stale:
for idx in reversed(stale):
+98 -2
View File
@@ -3,11 +3,19 @@ import logging
import os
import re
import json
import shutil
from ..services.settings_manager import get_settings_manager
from ..services.service_registry import ServiceRegistry
from ..utils.example_images_paths import iter_library_roots
from ..utils.example_images_paths import (
get_example_images_root,
is_hash_folder,
iter_library_roots,
uses_library_scoped_folders,
_library_folder_has_only_hash_dirs,
)
from ..utils.metadata_manager import MetadataManager
from ..utils.example_images_processor import ExampleImagesProcessor
from ..utils.example_images_metadata import update_cache_from_metadata
from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
logger = logging.getLogger(__name__)
@@ -36,6 +44,90 @@ settings = _SettingsProxy()
class ExampleImagesMigration:
"""Handles migrations for example images naming conventions"""
@staticmethod
def _consolidate_library_folders():
"""Move hash folders from library-named subdirectories back to root.
When a user switches from multi-library mode back to single-library
mode, example images previously stored under e.g.
``<root>/default/<hash>/`` need to be moved back to
``<root>/<hash>/``. Running this once at startup removes the need
for ``get_model_folder()`` to perform directory scans on every
request.
"""
if uses_library_scoped_folders():
return
root = get_example_images_root()
if not root or not os.path.isdir(root):
return
moved: list[str] = []
cleaned: list[str] = []
try:
for entry in os.listdir(root):
# Fast regex checks first — no filesystem I/O.
if is_hash_folder(entry) or entry == "_deleted":
continue
entry_path = os.path.join(root, entry)
if not os.path.isdir(entry_path):
continue
if not _library_folder_has_only_hash_dirs(entry_path):
continue
try:
for hash_entry in os.listdir(entry_path):
hash_path = os.path.join(entry_path, hash_entry)
if not os.path.isdir(hash_path) or not is_hash_folder(hash_entry):
continue
target = os.path.join(root, hash_entry)
if not os.path.exists(target):
try:
shutil.move(hash_path, target)
moved.append(hash_entry)
except (OSError, shutil.Error) as exc:
logger.error(
"Failed to move '%s''%s': %s",
hash_path, target, exc,
)
except OSError as exc:
logger.error(
"Failed to list library subdirectory '%s': %s",
entry_path, exc,
)
try:
remaining = os.listdir(entry_path)
except OSError:
remaining = []
if not remaining:
try:
os.rmdir(entry_path)
cleaned.append(entry)
except OSError as exc:
logger.debug(
"Could not remove empty library dir '%s': %s",
entry_path, exc,
)
except OSError as exc:
logger.error(
"Failed to list example images root during consolidation: %s",
exc,
)
if moved:
logger.info(
"Consolidated %d example image folder(s) to root",
len(moved),
)
if cleaned:
logger.info(
"Removed %d empty library directories",
len(cleaned),
)
@staticmethod
async def check_and_run_migrations():
"""Check if migrations are needed and run them in background"""
@@ -44,6 +136,10 @@ class ExampleImagesMigration:
logger.debug("No example images path configured or path doesn't exist, skipping migrations")
return
# Run library-to-root consolidation once at startup so the hot
# path (get_model_folder) stays a pure-path computation.
ExampleImagesMigration._consolidate_library_folders()
for library_name, library_path in iter_library_roots():
if not library_path or not os.path.exists(library_path):
continue
@@ -326,7 +422,7 @@ class ExampleImagesMigration:
await MetadataManager.save_metadata(file_path, model_copy)
# 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
except Exception as e:
+6 -30
View File
@@ -83,7 +83,12 @@ def ensure_library_root_exists(library_name: Optional[str] = None) -> str:
def get_model_folder(model_hash: str, library_name: Optional[str] = None) -> str:
"""Return the folder path for a model's example images."""
"""Return the folder path for a model's example images.
Multi-library single-library consolidation is handled once at startup by
``ExampleImagesMigration._consolidate_library_folders`` this function is a
pure path computation on the hot path (no directory scans).
"""
if not model_hash:
return ""
@@ -113,35 +118,6 @@ def get_model_folder(model_hash: str, library_name: Optional[str] = None) -> str
exc,
)
return legacy_folder
elif not os.path.exists(resolved_folder):
# Reverse migration: when consolidating from multi-library to
# single-library mode (e.g. after "default" was cleaned up), look
# for existing example images inside library-named subdirectories
# and bring them back to the root level.
root = get_example_images_root()
if root:
try:
for entry in os.listdir(root):
entry_path = os.path.join(root, entry)
if not os.path.isdir(entry_path):
continue
if is_hash_folder(entry) or entry == "_deleted":
continue
if not _library_folder_has_only_hash_dirs(entry_path):
continue
legacy = os.path.join(entry_path, normalized_hash)
if os.path.exists(legacy):
shutil.move(legacy, resolved_folder)
logger.info(
"Consolidated example images from '%s' to '%s'",
legacy, resolved_folder,
)
break
except OSError as exc:
logger.error(
"Failed to consolidate example images during "
"library merge: %s", exc,
)
return resolved_folder
+34 -4
View File
@@ -9,7 +9,7 @@ from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
from ..services.service_registry import ServiceRegistry
from ..services.settings_manager import get_settings_manager
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
logger = logging.getLogger(__name__)
@@ -113,6 +113,26 @@ class ExampleImagesProcessor:
message = str(error).lower()
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
async def download_model_images(model_hash, model_name, model_images, model_dir, optimize, downloader):
"""Download images for a single model
@@ -139,7 +159,12 @@ class ExampleImagesProcessor:
original_url = image_url
if optimize and 'civitai.com' in 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
try:
logger.debug(f"Downloading media file {i} for {model_name}")
@@ -229,6 +254,11 @@ class ExampleImagesProcessor:
if optimize and 'civitai.com' in 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:
logger.debug("Downloading media file %s for %s", i, model_name)
return await downloader.download_to_memory(
@@ -644,7 +674,7 @@ class ExampleImagesProcessor:
}, status=500)
# 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)
regular_images = civitai_data.get('images', [])
@@ -759,7 +789,7 @@ class ExampleImagesProcessor:
model_copy = model_data.copy()
model_copy.pop('folder', None)
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({
'success': True,
+48 -1
View File
@@ -1,7 +1,7 @@
from difflib import SequenceMatcher
import os
import re
from typing import Dict
from typing import Any, Dict, List, Optional
from ..services.service_registry import ServiceRegistry
from ..config import config
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)
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:
"""
Check if text matches pattern using fuzzy matching.
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-lora-manager"
description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!"
version = "1.1.9"
version = "1.2.0"
license = {file = "LICENSE"}
dependencies = [
"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
@@ -151,6 +151,7 @@ body.modal-open {
.support-section,
.changelog-section,
.update-info,
.update-channels,
.info-item,
.path-preview {
background: var(--surface-subtle);
+126 -9
View File
@@ -93,15 +93,13 @@
.update-content {
display: flex;
flex-direction: column;
gap: var(--space-3);
gap: var(--space-2);
}
.update-info {
display: flex;
justify-content: space-between;
align-items: center;
border-radius: var(--border-radius-sm);
padding: var(--space-3);
}
.update-info .version-info {
@@ -175,7 +173,6 @@
border: 1px solid var(--lora-border);
border-radius: var(--border-radius-sm);
padding: var(--space-2);
margin: var(--space-2) 0;
}
[data-theme="dark"] .update-progress {
@@ -233,11 +230,6 @@
}
/* Changelog section */
.changelog-section {
border-radius: var(--border-radius-sm);
padding: var(--space-3);
}
.changelog-section h3 {
margin-top: 0;
margin-bottom: var(--space-2);
@@ -349,6 +341,131 @@
text-decoration: underline;
}
/* Channel Toggle */
.update-channels {
}
.channels-label {
font-size: 0.9em;
color: var(--text-color);
opacity: 0.8;
margin-bottom: 8px;
}
.channel-toggle {
display: flex;
gap: 0;
background: var(--lora-surface);
border-radius: 8px;
padding: 3px;
width: fit-content;
}
.channel-btn {
display: flex;
align-items: center;
gap: 6px;
padding: 8px 20px;
border: none;
border-radius: 6px;
background: transparent;
color: var(--text-secondary, #999);
cursor: pointer;
font-size: 0.9em;
font-weight: 500;
transition: all 0.2s ease;
white-space: nowrap;
}
.channel-btn:hover {
color: var(--text-primary, #ddd);
background: rgba(255, 255, 255, 0.04);
}
.channel-btn.active {
background: var(--lora-accent, #4285F4);
color: #fff;
box-shadow: 0 1px 3px rgba(0, 0, 0, 0.2);
}
.channel-btn.active i {
color: #fff;
}
.channel-btn i {
font-size: 0.85em;
}
/* Channel Switch Confirmation Overlay */
.channel-switch-overlay {
position: fixed;
inset: 0;
background: rgba(0, 0, 0, 0.6);
display: flex;
align-items: center;
justify-content: center;
z-index: 10000;
backdrop-filter: blur(2px);
}
.channel-switch-dialog {
background: var(--lora-surface);
border: 1px solid var(--border-color, rgba(255, 255, 255, 0.1));
border-radius: 12px;
padding: 28px 32px;
max-width: 420px;
width: 90%;
box-shadow: 0 8px 32px rgba(0, 0, 0, 0.4);
}
.channel-switch-dialog h3 {
margin: 0 0 12px;
font-size: 1.1em;
color: var(--text-primary, #eee);
}
.channel-switch-dialog p {
margin: 0 0 24px;
font-size: 0.9em;
color: var(--text-secondary, #aaa);
line-height: 1.6;
}
.channel-switch-actions {
display: flex;
justify-content: flex-end;
gap: 10px;
}
.channel-switch-cancel {
padding: 8px 18px;
border: 1px solid var(--border-color, rgba(255, 255, 255, 0.1));
border-radius: 6px;
background: transparent;
color: var(--text-secondary, #aaa);
cursor: pointer;
font-size: 0.9em;
}
.channel-switch-cancel:hover {
background: rgba(255, 255, 255, 0.04);
}
.channel-switch-confirm {
padding: 8px 18px;
border: none;
border-radius: 6px;
background: var(--lora-accent, #4285F4);
color: #fff;
cursor: pointer;
font-size: 0.9em;
font-weight: 500;
}
.channel-switch-confirm:hover {
opacity: 0.9;
}
/* Update preferences section */
.update-preferences {
border-top: 1px solid var(--lora-border);
+1
View File
@@ -41,6 +41,7 @@
@import 'components/sidebar.css'; /* Add sidebar component */
@import 'components/media-viewer.css';
@import 'components/metadata-refresh-result.css';
@import 'components/download-batch-summary.css';
.initialization-notice {
display: flex;
+2 -1
View File
@@ -184,7 +184,8 @@ export const DOWNLOAD_ENDPOINTS = {
downloadGet: '/api/lm/download-model-get',
cancelGet: '/api/lm/cancel-download-get',
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
+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;
await state.loadingManager.showWithProgress(async (loading) => {
@@ -1700,8 +1700,13 @@ export class BaseModelApiClient {
// Determine optimize setting
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
const response = await fetch(DOWNLOAD_ENDPOINTS.exampleImages, {
const response = await fetch(endpoint, {
method: 'POST',
headers: {
'Content-Type': 'application/json'
@@ -1710,6 +1715,7 @@ export class BaseModelApiClient {
model_hashes: modelHashes,
output_dir: outputDir,
optimize: optimize,
force: force,
model_types: modelTypes || [this.apiConfig.config.singularName]
})
});
@@ -137,11 +137,10 @@ export class BulkContextMenu extends BaseContextMenu {
downloadMissingLorasItem.style.display = currentModelType === 'recipes' ? 'flex' : 'none';
}
const downloadExampleImagesItem = this.menu.querySelector('[data-action="download-example-images"]');
if (downloadExampleImagesItem) {
const downloadExampleImagesSubmenu = this.menu.querySelector('[data-has-submenu="download-example-images"]');
if (downloadExampleImagesSubmenu) {
// Show on model pages (loras, checkpoints, embeddings), hide on recipes
const modelPages = ['loras', 'checkpoints', 'embeddings'];
downloadExampleImagesItem.style.display = modelPages.includes(currentModelType) ? 'flex' : 'none';
downloadExampleImagesSubmenu.style.display = ['loras', 'checkpoints', 'embeddings'].includes(currentModelType) ? 'flex' : 'none';
}
const skipMetadataRefreshItem = this.menu.querySelector('[data-action="skip-metadata-refresh"]');
@@ -294,8 +293,11 @@ export class BulkContextMenu extends BaseContextMenu {
case 'download-missing-loras':
this.handleDownloadMissingLoras();
break;
case 'download-missing-example-images':
this.handleDownloadExampleImages({ force: false });
break;
case 'download-example-images':
this.handleDownloadExampleImages();
this.handleDownloadExampleImages({ force: true });
break;
case 'clear':
bulkManager.clearSelection();
@@ -340,7 +342,7 @@ export class BulkContextMenu extends BaseContextMenu {
await bulkMissingLoraDownloadManager.downloadMissingLoras(selectedRecipes);
}
async handleDownloadExampleImages() {
async handleDownloadExampleImages({ force = true } = {}) {
if (state.selectedModels.size === 0) {
return;
}
@@ -361,7 +363,7 @@ export class BulkContextMenu extends BaseContextMenu {
try {
const apiClient = getModelApiClient();
await apiClient.downloadExampleImages([...hashes]);
await apiClient.downloadExampleImages([...hashes], null, { force });
} catch (error) {
console.error('Bulk download example images failed:', error);
}
@@ -347,7 +347,10 @@ export const ModelContextMenuMixin = {
openExampleImagesFolder(this.currentCard.dataset.sha256);
return true;
case 'download-examples':
this.downloadExampleImages();
this.downloadExampleImages(false);
return true;
case 'download-examples-force':
this.downloadExampleImages(true);
return true;
case 'civitai':
if (this.currentCard.dataset.from_civitai === 'true') {
@@ -378,7 +381,7 @@ export const ModelContextMenuMixin = {
},
// Download example images method
async downloadExampleImages() {
async downloadExampleImages(force = false) {
const modelHash = this.currentCard.dataset.sha256;
if (!modelHash) {
showToast('toast.contextMenu.missingHash', {}, 'error');
@@ -387,7 +390,7 @@ export const ModelContextMenuMixin = {
try {
const apiClient = getModelApiClient();
await apiClient.downloadExampleImages([modelHash]);
await apiClient.downloadExampleImages([modelHash], null, { force });
} catch (error) {
console.error('Error downloading example images:', error);
}
@@ -260,8 +260,9 @@ export class RecipeContextMenu extends BaseContextMenu {
strength: lora.strength || 1.0,
// Model identifiers
modelId: lora.modelId || lora.model_id || civitaiInfo.modelId,
hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash,
modelVersionId: civitaiInfo.id || lora.modelVersionId,
id: civitaiInfo.id || lora.modelVersionId,
// Metadata
thumbnailUrl: civitaiInfo.images?.[0]?.url || '',
@@ -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;
}
}
});
}
+1
View File
@@ -1421,6 +1421,7 @@ class RecipeModal {
strength: lora.strength || 1.0,
// Model identifiers
modelId: lora.modelId || lora.model_id || civitaiInfo.modelId,
hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash,
id: civitaiInfo.id || lora.modelVersionId,
+56 -17
View File
@@ -108,10 +108,20 @@ export class PageControls {
const sortSelect = document.getElementById('sortSelect');
if (sortSelect) {
initSortDropdown(sortSelect);
sortSelect.value = this.pageState.sortBy;
this.applySortToSelect(this.pageState.sortBy);
sortSelect.addEventListener('change', async (e) => {
this.pageState.sortBy = e.target.value;
this.saveSortPreference(e.target.value);
let value = 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();
});
}
@@ -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
*/
@@ -326,10 +374,7 @@ export class PageControls {
// Handle legacy format conversion
const convertedSort = this.convertLegacySortFormat(savedSort);
this.pageState.sortBy = convertedSort;
const sortSelect = document.getElementById('sortSelect');
if (sortSelect) {
sortSelect.value = convertedSort;
}
this.applySortToSelect(convertedSort);
}
}
@@ -523,9 +568,9 @@ export class PageControls {
this.pageState.sortBy = restoredSort;
this.saveSortPreference(restoredSort);
this._removeVlmSortOption();
this.applySortToSelect(restoredSort);
const sortSelect = document.getElementById('sortSelect');
if (sortSelect) {
sortSelect.value = restoredSort;
sortSelect.disabled = false;
}
}
@@ -575,10 +620,7 @@ export class PageControls {
const savedGroupedSort = getStorageItem(groupedKey);
if (savedGroupedSort) {
this.pageState.sortBy = savedGroupedSort;
const sortSelect = document.getElementById('sortSelect');
if (sortSelect) {
sortSelect.value = savedGroupedSort;
}
this.applySortToSelect(savedGroupedSort);
}
} else {
// 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`);
if (savedNormalSort) {
this.pageState.sortBy = savedNormalSort;
const sortSelect = document.getElementById('sortSelect');
if (sortSelect) {
sortSelect.value = savedNormalSort;
}
this.applySortToSelect(savedNormalSort);
}
}
}
@@ -874,7 +913,7 @@ export class PageControls {
}
if (sortSelect) {
sortSelect.value = this.pageState.sortBy;
this.applySortToSelect(this.pageState.sortBy);
}
if (searchInput) {
searchInput.value = this.pageState.filters?.search || '';
+13 -3
View File
@@ -96,7 +96,16 @@ export function initSortDropdown(select) {
};
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.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
// 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());
observer.observe(select, { childList: true });
observer.observe(select, { childList: true, subtree: true, attributes: true, attributeFilter: ['value'] });
buildMenu();
group.dataset.sortReady = '1';
+6
View File
@@ -489,6 +489,12 @@ export function createModelCard(model, modelType) {
const modelId = civitaiData?.modelId ?? civitaiData?.model_id;
if (modelId !== undefined && modelId !== null && modelId !== '') {
card.dataset.modelId = modelId;
} else if (model.hf_url) {
// For HF-only models, derive a group key from hf_url for version grouping
const match = model.hf_url.match(/https?:\/\/huggingface\.co\/([^/]+\/[^/]+)/);
if (match) {
card.dataset.modelId = 'hf:' + match[1];
}
}
// LoRA specific data
+8 -1
View File
@@ -473,7 +473,14 @@ export async function showModelModal(model, modelType) {
const loadingExamplesText = translate('modals.model.loading.examples', {}, 'Loading examples...');
const loadingVersionsText = translate('modals.model.loading.versions', {}, 'Loading versions...');
const civitaiModelId = modelWithFullData.civitai?.modelId || '';
// Use CivitAI modelId, or derive HF group key for HF-only models
let civitaiModelId = modelWithFullData.civitai?.modelId || '';
if (!civitaiModelId && modelWithFullData.hf_url) {
const match = modelWithFullData.hf_url.match(/https?:\/\/huggingface\.co\/([^/]+\/[^/]+)/);
if (match) {
civitaiModelId = 'hf:' + match[1];
}
}
const civitaiVersionId = modelWithFullData.civitai?.id || '';
const navAriaLabel = translate('modals.model.navigation.label', {}, 'Model navigation');
const previousTitle = translate('modals.model.navigation.previousWithShortcut', {}, 'Previous model (←)');
@@ -950,6 +950,26 @@ export function initVersionsTab({
renderErrorState(container, translate('modals.model.versions.missingModelId', {}, 'This model is missing a Civitai model id.'));
return;
}
// HF group keys (e.g. "hf:user/repo") are not real CivitAI model IDs —
// skip the remote API call and show a helpful message instead.
const isHfGroupKey = typeof modelId === 'string' && modelId.startsWith('hf:');
if (isHfGroupKey) {
controller.isLoading = false;
controller.hasLoaded = true;
controller.record = null;
const hfMsg = translate(
'modals.model.versions.hfGroupInfo',
{},
'This is a HuggingFace model group. Open the library to see all versions in the grid.'
);
container.innerHTML = `
<div class="versions-empty-state">
<i class="fas fa-info-circle"></i>
<p>${escapeHtml(hfMsg)}</p>
</div>
`;
return;
}
if (controller.hasLoaded && !forceRefresh) {
return;
}
+48 -4
View File
@@ -27,6 +27,8 @@ export class BulkManager {
// Drag detection properties
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.mouseDownPosition = { x: 0, y: 0 };
@@ -88,7 +90,7 @@ export class BulkManager {
moveAll: true,
autoOrganize: false,
deleteAll: true,
setContentRating: false,
setContentRating: true,
skipMetadataRefresh: false,
setFavorite: true,
unfavorite: true,
@@ -173,6 +175,19 @@ export class BulkManager {
});
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) {
this.lastClientX = e.clientX;
this.lastClientY = e.clientY;
@@ -184,7 +199,10 @@ export class BulkManager {
const dy = e.clientY - this.mouseDownPosition.y;
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.startMarqueeSelection(e, true);
}
@@ -1510,14 +1528,18 @@ export class BulkManager {
let failureCount = 0;
try {
const apiClient = getModelApiClient();
const isRecipesPage = state.currentPageType === 'recipes';
for (const filePath of targets) {
if (cancelled) {
showToast('toast.api.operationCancelled', {}, 'info');
break;
}
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++;
} catch (error) {
failureCount++;
@@ -1958,9 +1980,31 @@ export class BulkManager {
// Remove visual feedback class
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
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 (selectionCount === 0) {
if (state.bulkMode) {
+43 -11
View File
@@ -8,6 +8,7 @@ import { FolderTreeManager } from '../components/FolderTreeManager.js';
import { translate } from '../utils/i18nHelpers.js';
import { extractCivitaiModelUrlParts } from '../utils/civitaiUtils.js';
import { formatFileSize } from '../utils/formatters.js';
import { showDownloadBatchSummary } from '../components/DownloadBatchSummaryModal.js';
export class DownloadManager {
constructor() {
@@ -158,6 +159,7 @@ export class DownloadManager {
this.modelVersionId = null;
this.source = null;
this.selectedFile = null;
this._isDiffusionModel = false;
this.selectedFolder = '';
this.batchModels = [];
@@ -787,24 +789,40 @@ export class DownloadManager {
async proceedToLocationContent() {
try {
// Fetch model roots
const rootsData = await this.apiClient.fetchModelRoots();
const _isDiffusionModel = this.selectedFile
? (this.selectedFile.type === 'UNet' || this.selectedFile.type === 'Diffusion Model')
: (this.currentVersion?.files || []).some(
f => f.type === 'UNet' || f.type === 'Diffusion Model'
);
this._isDiffusionModel = _isDiffusionModel;
let rootsData;
if (this._isDiffusionModel && this.apiClient.modelType === 'checkpoints') {
rootsData = await this.apiClient.fetchModelRoots('diffusion_model');
} else {
rootsData = await this.apiClient.fetchModelRoots();
}
const modelRoot = document.getElementById('modelRoot');
modelRoot.innerHTML = rootsData.roots.map(root =>
`<option value="${root}">${root}</option>`
).join('');
// Set default root if available
const singularType = this.apiClient.modelType.replace(/s$/, '');
const singularType = this._isDiffusionModel
? 'unet'
: this.apiClient.modelType.replace(/s$/, '');
const defaultRootKey = `default_${singularType}_root`;
const defaultRoot = state.global.settings[defaultRootKey];
console.log(`Default root for ${this.apiClient.modelType}:`, defaultRoot);
console.log(`Default root for ${singularType}:`, defaultRoot);
console.log('Available roots:', rootsData.roots);
if (defaultRoot && rootsData.roots.includes(defaultRoot)) {
console.log(`Setting default root: ${defaultRoot}`);
modelRoot.value = defaultRoot;
}
const subtypeDisplay = this._isDiffusionModel ? 'Diffusion Model' : this.apiClient.apiConfig.config.displayName;
document.getElementById('modelRootLabel').textContent =
translate('modals.download.selectTypeRoot', { type: subtypeDisplay });
// Set autocomplete="off" on folderPath input
const folderPathInput = document.getElementById('folderPath');
if (folderPathInput) {
@@ -1531,6 +1549,10 @@ export class DownloadManager {
modalManager.closeModal('downloadModal');
return this.executeBatchDownload(downloadItems, { modelRoot, targetFolder, useDefaultPaths });
}
async executeBatchDownload(downloadItems, { modelRoot, targetFolder, useDefaultPaths }) {
const batchDownloadId = Date.now().toString();
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
const ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${batchDownloadId}`);
@@ -1541,6 +1563,7 @@ export class DownloadManager {
let completedDownloads = 0;
let failedDownloads = 0;
let cancelled = false;
const failedItems = [];
loadingManager.showCancelButton(async () => {
if (cancelled) return;
@@ -1641,6 +1664,7 @@ export class DownloadManager {
if (!response.success) {
failedDownloads++;
failedItems.push({ item, error: response.error || 'Unknown error', name });
} else {
completedDownloads++;
updateProgress(100, completedDownloads, '');
@@ -1649,6 +1673,7 @@ export class DownloadManager {
if (!cancelled) {
console.error(`Failed to download ${name}:`, err);
failedDownloads++;
failedItems.push({ item, error: err?.message || 'Unknown error', name });
}
}
}
@@ -1662,10 +1687,15 @@ export class DownloadManager {
} else if (failedDownloads === 0) {
showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success');
} else {
showToast('toast.loras.downloadPartialSuccess', {
completed: completedDownloads,
showDownloadBatchSummary({
total: downloadItems.length,
}, 'warning');
completed: completedDownloads,
failedItems,
onRetry: (failed) => this.executeBatchDownload(
failed.map((f) => f.item),
{ modelRoot, targetFolder, useDefaultPaths }
),
});
}
await resetAndReload(true);
@@ -1776,13 +1806,15 @@ export class DownloadManager {
const modelRoot = document.getElementById('modelRoot').value;
const config = this.apiClient.apiConfig.config;
let fullPath = modelRoot || translate('modals.download.selectTypeRoot', { type: config.displayName });
const subtypeDisplay = this._isDiffusionModel ? 'Diffusion Model' : config.displayName;
let fullPath = modelRoot || translate('modals.download.selectTypeRoot', { type: subtypeDisplay });
if (modelRoot) {
if (this.useDefaultPath) {
// Show actual template path
try {
const singularType = this.apiClient.modelType.replace(/s$/, '');
const singularType = this._isDiffusionModel
? 'unet'
: this.apiClient.modelType.replace(/s$/, '');
const templates = state.global.settings.download_path_templates;
const template = templates[singularType];
fullPath += `/${template}`;
+6 -2
View File
@@ -729,10 +729,12 @@ export class FilterManager {
const pageState = getCurrentPageState();
const storageKey = `${this.currentPage}_filters`;
// Save filters to localStorage (exclude EMPTY_WILDCARD_MARKER)
// Save filters to localStorage (exclude EMPTY_WILDCARD_MARKER and transient search)
const filtersSnapshot = this.cloneFilters();
// Don't persist EMPTY_WILDCARD_MARKER - it's a runtime-only marker
filtersSnapshot.baseModel = filtersSnapshot.baseModel.filter(m => m !== EMPTY_WILDCARD_MARKER);
// Don't persist search - it's transient and managed by SearchManager
delete filtersSnapshot.search;
setStorageItem(storageKey, filtersSnapshot);
// Update state with current filters
@@ -984,6 +986,7 @@ export class FilterManager {
}
cloneFilters() {
const pageState = getCurrentPageState();
return {
...this.filters,
baseModel: [...(this.filters.baseModel || [])],
@@ -991,7 +994,8 @@ export class FilterManager {
autoTags: { ...(this.filters.autoTags || {}) },
license: { ...(this.filters.license || {}) },
modelTypes: [...(this.filters.modelTypes || [])],
tagLogic: this.filters.tagLogic || 'any'
tagLogic: this.filters.tagLogic || 'any',
search: pageState?.filters?.search ?? ''
};
}
+38 -17
View File
@@ -1517,11 +1517,20 @@ export class SettingsManager {
return data;
}
async loadLoraRoots() {
try {
const defaultLoraRootSelect = document.getElementById('defaultLoraRoot');
if (!defaultLoraRootSelect) return;
showNoRootsPlaceholder(select) {
select.innerHTML = '';
const option = document.createElement('option');
option.value = '';
option.textContent = translate('settings.folderSettings.noDefault', {}, 'No Default');
select.appendChild(option);
select.disabled = true;
}
async loadLoraRoots() {
const defaultLoraRootSelect = document.getElementById('defaultLoraRoot');
if (!defaultLoraRootSelect) return;
try {
// Fetch lora roots
const response = await fetch('/api/lm/loras/roots');
if (!response.ok) {
@@ -1530,10 +1539,12 @@ export class SettingsManager {
const data = await response.json();
if (!data.roots || data.roots.length === 0) {
throw new Error('No LoRA roots found');
this.showNoRootsPlaceholder(defaultLoraRootSelect);
return;
}
defaultLoraRootSelect.innerHTML = '';
defaultLoraRootSelect.disabled = false;
// Add options for each root
data.roots.forEach(root => {
@@ -1548,15 +1559,16 @@ export class SettingsManager {
} catch (error) {
console.error('Error loading LoRA roots:', error);
this.showNoRootsPlaceholder(defaultLoraRootSelect);
showToast('toast.settings.loraRootsFailed', { message: error.message }, 'error');
}
}
async loadCheckpointRoots() {
try {
const defaultCheckpointRootSelect = document.getElementById('defaultCheckpointRoot');
if (!defaultCheckpointRootSelect) return;
const defaultCheckpointRootSelect = document.getElementById('defaultCheckpointRoot');
if (!defaultCheckpointRootSelect) return;
try {
// Fetch checkpoint roots (checkpoint paths only, not unet)
const response = await fetch('/api/lm/checkpoints/checkpoints_roots');
if (!response.ok) {
@@ -1565,10 +1577,12 @@ export class SettingsManager {
const data = await response.json();
if (!data.roots || data.roots.length === 0) {
throw new Error('No checkpoint roots found');
this.showNoRootsPlaceholder(defaultCheckpointRootSelect);
return;
}
defaultCheckpointRootSelect.innerHTML = '';
defaultCheckpointRootSelect.disabled = false;
// Add options for each root
data.roots.forEach(root => {
@@ -1583,15 +1597,16 @@ export class SettingsManager {
} catch (error) {
console.error('Error loading checkpoint roots:', error);
this.showNoRootsPlaceholder(defaultCheckpointRootSelect);
showToast('toast.settings.checkpointRootsFailed', { message: error.message }, 'error');
}
}
async loadUnetRoots() {
try {
const defaultUnetRootSelect = document.getElementById('defaultUnetRoot');
if (!defaultUnetRootSelect) return;
const defaultUnetRootSelect = document.getElementById('defaultUnetRoot');
if (!defaultUnetRootSelect) return;
try {
// Fetch unet roots (diffusion model paths only)
const response = await fetch('/api/lm/checkpoints/unet_roots');
if (!response.ok) {
@@ -1600,10 +1615,12 @@ export class SettingsManager {
const data = await response.json();
if (!data.roots || data.roots.length === 0) {
throw new Error('No diffusion model roots found');
this.showNoRootsPlaceholder(defaultUnetRootSelect);
return;
}
defaultUnetRootSelect.innerHTML = '';
defaultUnetRootSelect.disabled = false;
// Add options for each root
data.roots.forEach(root => {
@@ -1618,15 +1635,16 @@ export class SettingsManager {
} catch (error) {
console.error('Error loading diffusion model roots:', error);
this.showNoRootsPlaceholder(defaultUnetRootSelect);
showToast('toast.settings.unetRootsFailed', { message: error.message }, 'error');
}
}
async loadEmbeddingRoots() {
try {
const defaultEmbeddingRootSelect = document.getElementById('defaultEmbeddingRoot');
if (!defaultEmbeddingRootSelect) return;
const defaultEmbeddingRootSelect = document.getElementById('defaultEmbeddingRoot');
if (!defaultEmbeddingRootSelect) return;
try {
// Fetch embedding roots
const response = await fetch('/api/lm/embeddings/roots');
if (!response.ok) {
@@ -1635,10 +1653,12 @@ export class SettingsManager {
const data = await response.json();
if (!data.roots || data.roots.length === 0) {
throw new Error('No embedding roots found');
this.showNoRootsPlaceholder(defaultEmbeddingRootSelect);
return;
}
defaultEmbeddingRootSelect.innerHTML = '';
defaultEmbeddingRootSelect.disabled = false;
// Add options for each root
data.roots.forEach(root => {
@@ -1653,6 +1673,7 @@ export class SettingsManager {
} catch (error) {
console.error('Error loading embedding roots:', error);
this.showNoRootsPlaceholder(defaultEmbeddingRootSelect);
showToast('toast.settings.embeddingRootsFailed', { message: error.message }, 'error');
}
}
+291 -48
View File
@@ -1,11 +1,12 @@
import { modalManager } from './ModalManager.js';
import {
getStorageItem,
setStorageItem,
getStoredVersionInfo,
import {
getStorageItem,
setStorageItem,
getStoredVersionInfo,
setStoredVersionInfo,
isVersionMatch
} from '../utils/storageHelpers.js';
import { state } from '../state/index.js';
import { bannerService } from './BannerService.js';
import { translate } from '../utils/i18nHelpers.js';
@@ -24,7 +25,11 @@ export class UpdateService {
this.updateNotificationsEnabled = getStorageItem('show_update_notifications', true);
this.lastCheckTime = parseInt(getStorageItem('last_update_check') || '0');
this.isUpdating = false;
this.nightlyMode = getStorageItem('nightly_updates', false);
this.channelMode = null;
this.hasGit = false;
this.nightlyNotifyDate = getStorageItem('nightly_notify_date', '');
this.nightlyBadgeShown = false;
this.progressKeepVisible = false;
this.currentVersionInfo = null;
this.versionMismatch = false;
this.activeNotificationTab = 'updates';
@@ -49,43 +54,180 @@ export class UpdateService {
updateBtn.addEventListener('click', () => this.performUpdate());
}
// Register event listener for nightly update toggle
const nightlyCheckbox = document.getElementById('nightlyUpdateToggle');
if (nightlyCheckbox) {
nightlyCheckbox.checked = this.nightlyMode;
nightlyCheckbox.addEventListener('change', (e) => {
this.nightlyMode = e.target.checked;
setStorageItem('nightly_updates', e.target.checked);
this.updateNightlyWarning();
this.updateModalContent();
// Re-check for updates when switching channels
this.manualCheckForUpdates();
});
this.updateNightlyWarning();
}
this.wireChannelButtons();
this.setupNotificationCenter();
window.addEventListener('lm:banner-history-updated', this.handleBannerHistoryUpdated);
this.updateTabBadges();
// Perform update check if needed
this.checkForUpdates().then(() => {
// Ensure badges are updated after checking
this.updateBadgeVisibility();
this.checkVersionInfo().then(() => {
this.checkForUpdates().then(() => {
this.updateBadgeVisibility();
});
});
// Immediately update modal content with current values (even if from default)
this.updateModalContent();
// Check version info for mismatch after loading basic info
this.checkVersionInfo();
}
updateNightlyWarning() {
const warning = document.getElementById('nightlyWarning');
if (warning) {
warning.style.display = this.nightlyMode ? 'flex' : 'none';
wireChannelButtons() {
const releaseBtn = document.getElementById('channelRelease');
const nightlyBtn = document.getElementById('channelNightly');
if (releaseBtn) {
releaseBtn.addEventListener('click', () => this.switchChannel('release'));
}
if (nightlyBtn) {
nightlyBtn.addEventListener('click', () => this.switchChannel('nightly'));
}
}
async switchChannel(channel) {
if (channel === this.channelMode) {
return;
}
if (this.isUpdating) {
return;
}
if (!this.hasGit && channel === 'nightly') {
const confirmed = await this._confirmChannelSwitch(
'update.channelSwitch.nightlyTitle',
'update.channelSwitch.nightlyMessage'
);
if (!confirmed) return;
}
if (this.hasGit && channel === 'release') {
const confirmed = await this._confirmChannelSwitch(
'update.channelSwitch.releaseTitle',
'update.channelSwitch.releaseMessage'
);
if (!confirmed) return;
}
try {
this.isUpdating = true;
this.showUpdateProgress(true);
this.updateProgress(10, translate('update.channelSwitch.switching', { channel }));
const response = await fetch('/api/lm/switch-channel', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ channel })
});
const data = await response.json();
if (data.success) {
this.channelMode = channel;
// Persist channel preference to settings.json
fetch('/api/lm/settings', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ update_channel: channel })
}).then(r => {
if (!r.ok) console.warn('Failed to persist update channel:', r.status);
}).catch(e => console.warn('Failed to persist update channel:', e));
await this.checkForUpdates({ force: true });
this.updateModalContent();
this.updateChannelUI();
this._showSwitchCompleteMessage(data.new_version);
this.progressKeepVisible = true;
} else {
throw new Error(data.error || translate('update.channelSwitch.failed'));
}
} catch (error) {
console.error('Channel switch failed:', error);
this.updateProgress(0, translate('update.channelSwitch.failed'));
} finally {
if (this.progressKeepVisible) {
this.isUpdating = false;
this.progressKeepVisible = false;
} else {
setTimeout(() => {
this.showUpdateProgress(false);
this.isUpdating = false;
}, 2000);
}
}
}
updateChannelUI() {
const releaseBtn = document.getElementById('channelRelease');
const nightlyBtn = document.getElementById('channelNightly');
if (releaseBtn) {
releaseBtn.classList.toggle('active', this.channelMode === 'release');
}
if (nightlyBtn) {
nightlyBtn.classList.toggle('active', this.channelMode === 'nightly');
}
}
_resolveChannelFromSettings() {
const stored = state?.global?.settings?.update_channel;
if (stored === 'nightly' || stored === 'release') {
return stored;
}
if (!this.hasGit) {
return 'release';
}
if (this.gitInfo?.branch === 'detached') {
return 'release';
}
return 'nightly';
}
async _confirmChannelSwitch(titleKey, messageKey) {
return new Promise((resolve) => {
const title = translate(titleKey);
const message = translate(messageKey);
const cancelText = translate('common.cancel');
const confirmText = translate('common.confirm');
const overlay = document.createElement('div');
overlay.className = 'channel-switch-overlay';
overlay.innerHTML = `
<div class="channel-switch-dialog">
<h3>${title}</h3>
<p>${message}</p>
<div class="channel-switch-actions">
<button class="secondary-btn channel-switch-cancel">${cancelText}</button>
<button class="primary-btn channel-switch-confirm">${confirmText}</button>
</div>
</div>
`;
const dismiss = (result) => {
document.removeEventListener('keydown', onKeydown);
overlay.remove();
resolve(result);
};
const onKeydown = (e) => {
if (e.key === 'Escape') {
e.stopPropagation();
e.preventDefault();
dismiss(false);
}
};
document.addEventListener('keydown', onKeydown, { capture: true });
overlay.addEventListener('click', (e) => {
if (e.target === overlay) {
dismiss(false);
}
});
overlay.querySelector('.channel-switch-cancel').addEventListener('click', () => {
dismiss(false);
});
overlay.querySelector('.channel-switch-confirm').addEventListener('click', () => {
dismiss(true);
});
document.body.appendChild(overlay);
});
}
setupNotificationCenter() {
@@ -355,6 +497,18 @@ export class UpdateService {
}
async checkForUpdates({ force = false } = {}) {
let needsMigration = false;
if (this.channelMode === null) {
const stored = state?.global?.settings?.update_channel;
if (stored === 'nightly' || stored === 'release') {
this.channelMode = stored;
} else if (!this.hasGit) {
this.channelMode = 'release';
needsMigration = true;
}
// hasGit=true with no stored value: wait for gitInfo.branch
}
if (!force && !this.updateNotificationsEnabled) {
return;
}
@@ -373,7 +527,8 @@ export class UpdateService {
try {
// Call backend API to check for updates with nightly flag
const response = await fetch(`/api/lm/check-updates?nightly=${this.nightlyMode}`);
const nightly = (this.channelMode ?? (this.hasGit ? 'nightly' : 'release')) === 'nightly';
const response = await fetch(`/api/lm/check-updates?nightly=${nightly}`);
const data = await response.json();
if (data.success) {
@@ -381,17 +536,35 @@ export class UpdateService {
this.latestVersion = data.latest_version || "v0.0.0";
this.updateInfo = data;
this.gitInfo = data.git_info || this.gitInfo;
// Explicitly set update availability based on version comparison
this.updateAvailable = this.isNewerVersion(this.latestVersion, this.currentVersion);
// Update last check time
this.hasGit = data.has_git || false;
if (needsMigration || this.channelMode === null) {
this.channelMode = this._resolveChannelFromSettings();
if (state?.global?.settings) {
state.global.settings.update_channel = this.channelMode;
}
fetch('/api/lm/settings', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ update_channel: this.channelMode })
}).then(r => {
if (!r.ok) console.warn('Failed to persist update channel:', r.status);
}).catch(e => console.warn('Failed to persist update channel:', e));
}
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;
setStorageItem('last_update_check', now.toString());
// Update UI
this.updateBadgeVisibility();
this.updateModalContent();
this.updateChannelUI();
console.log("Update check complete:", {
currentVersion: this.currentVersion,
@@ -435,6 +608,28 @@ export class UpdateService {
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() {
const updateToggle = document.querySelector('.update-toggle');
@@ -443,9 +638,12 @@ export class UpdateService {
? bannerService.getUnreadBannerCount()
: 0;
// Force updating badges visibility based on current state
const shouldShowUpdate = this.updateNotificationsEnabled && this.updateAvailable && this._isNightlyBadgeAllowed();
if (updateToggle) {
let tooltipKey = 'header.actions.notifications';
if (this.updateNotificationsEnabled && this.updateAvailable) {
if (shouldShowUpdate) {
tooltipKey = 'update.updateAvailable';
} else if (unreadBanners > 0) {
tooltipKey = 'update.tabs.messages';
@@ -453,8 +651,6 @@ export class UpdateService {
updateToggle.title = translate(tooltipKey);
}
// Force updating badges visibility based on current state
const shouldShowUpdate = this.updateNotificationsEnabled && this.updateAvailable;
const shouldShow = shouldShowUpdate || unreadBanners > 0;
if (updateBadge) {
@@ -482,8 +678,31 @@ export class UpdateService {
if (currentVersionEl) currentVersionEl.textContent = this.currentVersion;
const newVersionLabel = modal.querySelector('.new-version .label');
if (newVersionLabel) {
newVersionLabel.textContent = (this.updateInfo?.nightly)
? `${translate('update.latestMain')}:`
: `${translate('update.newVersion')}:`;
}
if (newVersionEl) {
newVersionEl.textContent = this.latestVersion;
if (this.updateInfo?.nightly) {
const behind = this.updateInfo.behind_by || 0;
const remoteHash = this.latestVersion.replace('main-', '');
const localHash = this.gitInfo.short_hash || '';
const date = this.updateInfo.commit_date || '';
const datePart = date ? ` · ${date}` : '';
if (behind > 0) {
newVersionEl.textContent = `${behind} commit${behind !== 1 ? 's' : ''} behind main (${remoteHash}${datePart})`;
} else if (localHash !== remoteHash) {
newVersionEl.textContent = `Behind main (${remoteHash}${datePart})`;
} else {
newVersionEl.textContent = `Up to date (${remoteHash}${datePart})`;
}
} else {
newVersionEl.textContent = this.latestVersion;
}
}
// Update update button state
@@ -599,8 +818,12 @@ export class UpdateService {
// Update GitHub link to point to the specific release if available
const githubLink = modal.querySelector('.update-link');
if (githubLink && this.latestVersion) {
const versionTag = this.latestVersion.replace(/^v/, '');
githubLink.href = `https://github.com/willmiao/ComfyUI-Lora-Manager/releases/tag/v${versionTag}`;
if (this.updateInfo?.nightly) {
githubLink.href = 'https://github.com/willmiao/ComfyUI-Lora-Manager/commits/main';
} else {
const versionTag = this.latestVersion.replace(/^v/, '');
githubLink.href = `https://github.com/willmiao/ComfyUI-Lora-Manager/releases/tag/v${versionTag}`;
}
}
}
@@ -623,7 +846,7 @@ export class UpdateService {
'Content-Type': 'application/json'
},
body: JSON.stringify({
nightly: this.nightlyMode
nightly: this.channelMode === 'nightly'
})
});
@@ -698,7 +921,26 @@ export class UpdateService {
progressText.textContent = text;
}
}
_showSwitchCompleteMessage(version) {
this.showUpdateProgress(true);
this.updateProgress(100, '');
const progressText = document.getElementById('updateProgressText');
if (progressText) {
progressText.innerHTML = `
<div style="text-align: center; color: var(--lora-success);">
<i class="fas fa-check-circle" style="margin-right: 8px;"></i>
${translate('update.completion.successMessage', { version })}
<br><br>
<div style="opacity: 0.95; color: var(--lora-error); font-size: 1em;">
${translate('update.completion.restartMessage')}<br>
${translate('update.completion.reloadMessage')}
</div>
</div>
`;
}
}
showUpdateCompleteMessage(newVersion) {
const modal = document.getElementById('updateModal');
if (!modal) return;
@@ -771,6 +1013,7 @@ export class UpdateService {
// Update the modal content immediately with current data
this.updateModalContent();
this.updateChannelUI();
this.renderRecentBanners();
// Show the modal with current data
@@ -801,8 +1044,8 @@ export class UpdateService {
if (data.success) {
this.currentVersionInfo = data.version;
// Check if version matches stored version
this.hasGit = data.has_git || false;
this.versionMismatch = !isVersionMatch(this.currentVersionInfo);
if (this.versionMismatch) {
+18 -3
View File
@@ -3,6 +3,7 @@ import { translate } from '../../utils/i18nHelpers.js';
import { getModelApiClient } from '../../api/modelApiFactory.js';
import { MODEL_TYPES } from '../../api/apiConfig.js';
import { getStorageItem } from '../../utils/storageHelpers.js';
import { state } from '../../state/index.js';
export class DownloadManager {
constructor(importManager) {
@@ -125,11 +126,25 @@ export class DownloadManager {
showToast('toast.recipes.nameSaved', { name: this.importManager.recipeName }, 'success');
}
// Close modal
modalManager.closeModal('importModal');
// Refresh the recipe
window.recipeManager.loadRecipes(true);
if (isDownloadOnly && state.virtualScroller) {
const recipeId = this.importManager.recipeId;
try {
const detailRes = await fetch(`/api/lm/recipe/${encodeURIComponent(recipeId)}`);
if (detailRes.ok) {
const updated = await detailRes.json();
state.virtualScroller.updateSingleItem(updated.file_path, updated);
} else {
throw new Error(`API returned ${detailRes.status}`);
}
} catch (e) {
console.warn('Failed to update recipe card in-place, falling back to reload:', e);
await window.recipeManager.loadRecipes({ resetPage: true, preserveScroll: true });
}
} else {
window.recipeManager.loadRecipes({ resetPage: true, preserveScroll: true });
}
} catch (error) {
console.error('Error:', error);
+1
View File
@@ -333,6 +333,7 @@ export const PATH_TEMPLATE_PLACEHOLDERS = [
export const DEFAULT_PATH_TEMPLATES = {
lora: '{base_model}/{first_tag}',
checkpoint: '{base_model}',
unet: '{base_model}',
embedding: '{first_tag}'
};
+6 -1
View File
@@ -32,7 +32,12 @@
<div class="context-menu-separator menu-section-break"></div>
<!-- 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="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-separator menu-section-break"></div>
<!-- Attributes -->
+24 -4
View File
@@ -44,8 +44,18 @@
<div class="context-menu-item" data-action="preview">
<i class="fas fa-folder-open"></i> <span>{{ t('loras.contextMenu.openExamples') }}</span>
</div>
<div class="context-menu-item" data-action="download-examples">
<i class="fas fa-download"></i> <span>{{ t('loras.contextMenu.downloadExamples') }}</span>
<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-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 class="context-menu-item" data-action="replace-preview">
<i class="fas fa-image"></i> <span>{{ t('loras.contextMenu.replacePreview') }}</span>
@@ -136,8 +146,18 @@
</div>
<div class="context-menu-section" data-section="download">
<div class="context-menu-section-header">{{ t('loras.bulkOperations.sections.download') }}</div>
<div class="context-menu-item" data-action="download-example-images">
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadExamples') }}</span>
<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-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 class="context-menu-item" data-action="download-missing-loras">
<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>
</optgroup>
{% 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' %}
<optgroup label="{{ t('recipes.controls.sort.lorasCount') }}">
<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="groq">{{ t('settings.aiProvider.providerOptions.groq') }}</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="custom">{{ t('settings.aiProvider.providerOptions.custom') }}</option>
</select>
@@ -19,6 +19,20 @@
<div class="notification-panels">
<div class="notification-panel active" id="updatesPanel" role="tabpanel" aria-labelledby="updatesTab" aria-hidden="false" tabindex="0" data-notification-panel="updates">
<div class="update-content">
<!-- Channel Selector -->
<div class="update-channels" id="updateChannels">
<div class="channels-label">{{ t('update.channel') }}</div>
<div class="channel-toggle">
<button type="button" class="channel-btn" data-channel="release" id="channelRelease">
<i class="fas fa-tag"></i> {{ t('update.channels.release') }}
</button>
<button type="button" class="channel-btn" data-channel="nightly" id="channelNightly">
<i class="fas fa-moon"></i> {{ t('update.channels.nightly') }}
</button>
</div>
</div>
<div class="update-info">
<div class="version-info">
<div class="current-version">
+6 -1
View File
@@ -32,7 +32,12 @@
<div class="context-menu-separator menu-section-break"></div>
<!-- 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="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-separator menu-section-break"></div>
<!-- Attributes -->
@@ -2155,4 +2155,35 @@ describe('Interaction-level regression coverage', () => {
excludedItem.dispatchEvent(new Event('click', { bubbles: true }));
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);
});
});
@@ -106,6 +106,118 @@ afterEach(() => {
});
});
describe('SettingsManager root selects', () => {
const rootCases = [
{
method: 'loadLoraRoots',
selectId: 'defaultLoraRoot',
endpoint: '/api/lm/loras/roots',
errorKey: 'toast.settings.loraRootsFailed',
},
{
method: 'loadCheckpointRoots',
selectId: 'defaultCheckpointRoot',
endpoint: '/api/lm/checkpoints/checkpoints_roots',
errorKey: 'toast.settings.checkpointRootsFailed',
},
{
method: 'loadUnetRoots',
selectId: 'defaultUnetRoot',
endpoint: '/api/lm/checkpoints/unet_roots',
errorKey: 'toast.settings.unetRootsFailed',
},
{
method: 'loadEmbeddingRoots',
selectId: 'defaultEmbeddingRoot',
endpoint: '/api/lm/embeddings/roots',
errorKey: 'toast.settings.embeddingRootsFailed',
},
];
const appendRootSelect = (id) => {
const select = document.createElement('select');
select.id = id;
document.body.appendChild(select);
return select;
};
it.each(rootCases)(
'populates the $method select with roots and keeps it enabled',
async ({ method, selectId, endpoint }) => {
const manager = createManager();
const select = appendRootSelect(selectId);
select.disabled = true;
global.fetch = vi.fn().mockResolvedValue({
ok: true,
json: async () => ({
success: true,
roots: ['/models/root-a', '/models/root-b'],
}),
});
await manager[method]();
expect(global.fetch).toHaveBeenCalledWith(endpoint);
expect(Array.from(select.options).map(option => option.value)).toEqual([
'/models/root-a',
'/models/root-b',
]);
expect(select.disabled).toBe(false);
expect(showToast).not.toHaveBeenCalled();
}
);
it.each(rootCases)(
'shows a placeholder and no error toast when $method has empty roots',
async ({ method, selectId, endpoint }) => {
const manager = createManager();
const select = appendRootSelect(selectId);
global.fetch = vi.fn().mockResolvedValue({
ok: true,
json: async () => ({
success: true,
roots: [],
}),
});
await manager[method]();
expect(global.fetch).toHaveBeenCalledWith(endpoint);
expect(select.options).toHaveLength(1);
expect(select.options[0].value).toBe('');
expect(select.options[0].textContent).toBe('No Default');
expect(select.disabled).toBe(true);
expect(showToast).not.toHaveBeenCalled();
}
);
it.each(rootCases)(
'shows an error toast when the $method roots request fails',
async ({ method, selectId, errorKey }) => {
const manager = createManager();
const select = appendRootSelect(selectId);
global.fetch = vi.fn().mockResolvedValue({
ok: false,
status: 500,
});
await manager[method]();
expect(select.options).toHaveLength(1);
expect(select.options[0].value).toBe('');
expect(select.disabled).toBe(true);
expect(showToast).toHaveBeenCalledWith(
errorKey,
expect.objectContaining({ message: expect.any(String) }),
'error',
);
}
);
});
describe('SettingsManager library controls', () => {
it('loads libraries and populates the select', async () => {
const manager = createManager();
+123 -2
View File
@@ -1,12 +1,26 @@
import { describe, beforeEach, afterEach, expect, it, vi } from 'vitest';
import { UpdateService } from '../../../static/js/managers/UpdateService.js';
import { state } from '../../../static/js/state/index.js';
function createFetchResponse(payload) {
return {
json: vi.fn().mockResolvedValue(payload)
json: vi.fn().mockResolvedValue(payload),
ok: true,
};
}
function stubSettingsUpdateChannel(channel) {
state.global = state.global || {};
state.global.settings = state.global.settings || {};
state.global.settings.update_channel = channel;
}
function clearSettingsUpdateChannel() {
if (state.global?.settings) {
delete state.global.settings.update_channel;
}
}
describe('UpdateService passive checks', () => {
let service;
let fetchMock;
@@ -16,10 +30,13 @@ describe('UpdateService passive checks', () => {
success: true,
current_version: 'v1.0.0',
latest_version: 'v1.0.0',
git_info: { short_hash: 'abc123' }
git_info: { short_hash: 'abc123' },
has_git: true,
}));
global.fetch = fetchMock;
stubSettingsUpdateChannel('release');
service = new UpdateService();
service.updateNotificationsEnabled = false;
service.lastCheckTime = 0;
@@ -28,6 +45,7 @@ describe('UpdateService passive checks', () => {
afterEach(() => {
delete global.fetch;
clearSettingsUpdateChannel();
});
it('skips passive update checks when notifications are disabled', async () => {
@@ -43,3 +61,106 @@ describe('UpdateService passive checks', () => {
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);
});
});
@@ -870,7 +870,8 @@ def test_metadata_overwrite_extractor_stores_truthy_values(metadata_registry):
assert "steps" not in params
assert "sampler" not in params
assert "scheduler" not in params
assert "clip_skip" not in params
# clip_skip=0 is now stored (0 != sentinel -25) — wired 0 is valid
assert params["clip_skip"] == 0
metadata_registry.clear_metadata()
@@ -880,8 +881,10 @@ def test_metadata_overwrite_extractor_empty_inputs(metadata_registry):
metadata_registry.start_collection("prompt-ow2")
metadata = metadata_registry.prompt_metadata["prompt-ow2"]
from py.metadata_collector.constants import CLIP_SKIP_SENTINEL
inputs = {key: "" for key in METADATA_OVERWRITE_FIELDS}
inputs.update({"seed": 0, "steps": 0, "cfg_scale": 0.0, "clip_skip": 0})
inputs.update({"seed": 0, "steps": 0, "cfg_scale": 0.0, "clip_skip": CLIP_SKIP_SENTINEL})
MetadataOverwriteExtractor.extract("ow-2", inputs, None, metadata)
@@ -950,7 +953,8 @@ def test_extract_generation_params_overwrite_falsy_skipped(metadata_registry, po
registry_obj.set_current_prompt(populated_registry["prompt"])
metadata2 = registry_obj.prompt_metadata["promptA"]
# Inject overwrite with falsy values
# Inject overwrite with falsy values (except clip_skip=0 which is now
# treated as a valid wired input thanks to the -25 sentinel)
metadata2[OVERWRITE] = {
"ow-1": {
"parameters": {
@@ -974,6 +978,9 @@ def test_extract_generation_params_overwrite_falsy_skipped(metadata_registry, po
assert params["prompt"] == "A castle on a hill"
assert params["cfg_scale"] == 7.5
# clip_skip=0 is a valid wired value (not the -25 sentinel) — should be applied
assert params["clip_skip"] == 0
registry_obj.clear_metadata()
+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():
@@ -49,3 +56,109 @@ def test_combine_stacks_allows_duplicate_entries():
(combined_stack,) = node.combine_stacks([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)
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):
_configure_save_paths(monkeypatch, tmp_path)
_configure_metadata(monkeypatch, {"prompt": "prompt text", "seed": 123})
@@ -451,6 +486,14 @@ class TestParameterDefaultConsistency:
assert SaveImageLM.save_images.__defaults__[5] == 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):
_configure_save_paths(monkeypatch, tmp_path)
+97 -5
View File
@@ -900,18 +900,28 @@ class FakeMetadataProvider:
async def get_model_versions(self, _model_id):
return {"modelVersions": [], "name": "", "type": "lora"}
async def get_user_models(self, _username):
return []
async def get_user_models(self, _username, cursor=None):
return {"items": [], "nextCursor": None}
async def get_creator_model_count(self, _username):
return None
class FakeUserModelsProvider(FakeMetadataProvider):
def __init__(self, models):
def __init__(self, models, next_cursor=None, estimated_total=None):
self.models = models
self.next_cursor = next_cursor
self.estimated_total = estimated_total
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)
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():
@@ -1286,6 +1296,88 @@ async def test_get_civitai_user_models_requires_username():
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():
call_records = []
+328 -1
View File
@@ -1,10 +1,33 @@
import logging
import os
import shutil
from aiohttp import ClientError
from aiohttp import web
import pytest
from py.routes import update_routes
def _fake_request(body=None, query_params=None):
from multidict import MultiDict
q = MultiDict(query_params or {})
req = type("Req", (), {
"has_body": body is not None,
"match_info": {},
"rel_url": type("U", (), {"query": q})(),
"query": q,
"app": {},
})()
async def _json():
return body or {}
req.json = _json
return req
class OfflineDownloader:
async def make_request(self, *_, **__):
return False, "Cannot connect to host"
@@ -53,10 +76,12 @@ async def test_get_nightly_version_network_error_logs_warning(monkeypatch, caplo
caplog.set_level(logging.WARNING)
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(RaisingDownloader()))
version, changelog = await update_routes.UpdateRoutes._get_nightly_version()
version, changelog, behind_by, commit_date = await update_routes.UpdateRoutes._get_nightly_version()
assert version == "main"
assert changelog == []
assert behind_by == 0
assert commit_date == ""
assert "Unable to reach GitHub for nightly version" in caplog.text
assert "Traceback" not in caplog.text
@@ -236,3 +261,305 @@ async def test_perform_git_update_stable_preserves_user_dirs(monkeypatch, tmp_pa
clean_args = clean_calls[0][1]
for name in update_routes._PRESERVE_DIRS:
assert name in clean_args, f"{name} missing from git clean excludes (stable)"
def test_init_git_repo_creates_valid_repo(tmp_path, monkeypatch):
if not shutil.which("git"):
pytest.skip("git executable not found")
plugin_root = tmp_path / "plugin"
plugin_root.mkdir()
(plugin_root / ".tracking").write_text("pyproject.toml")
(plugin_root / "settings.json").write_text('{"some": "value"}')
try:
success, version = update_routes.UpdateRoutes._init_git_repo(str(plugin_root))
except Exception as e:
pytest.skip(f"Network unavailable for git fetch: {e}")
assert success is True
assert version.startswith("main-")
assert len(version) > len("main-")
assert (plugin_root / ".git").is_dir()
assert not (plugin_root / ".tracking").exists()
assert (plugin_root / "settings.json").exists()
assert (plugin_root / "pyproject.toml").exists()
@pytest.mark.asyncio
async def test_switch_channel_invalid_channel_returns_error():
req = _fake_request({"channel": "bad_channel"})
resp = await update_routes.UpdateRoutes.switch_channel(req)
data = _raw_body(resp)
assert not data["success"]
assert "Invalid channel" in data["error"]
@pytest.mark.asyncio
async def test_switch_channel_to_nightly_without_git_inits_repo(monkeypatch, tmp_path):
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
routes_file.parent.mkdir(parents=True)
routes_file.write_text("")
monkeypatch.setattr(update_routes, "__file__", str(routes_file))
monkeypatch.setattr(update_routes, "ensure_settings_file", lambda logger: str(tmp_path / "settings.json"))
monkeypatch.setattr(
update_routes.UpdateRoutes,
"_init_git_repo",
staticmethod(lambda plugin_root: (True, "main-fedcba9")),
)
req = _fake_request({"channel": "nightly"})
resp = await update_routes.UpdateRoutes.switch_channel(req)
data = _raw_body(resp)
assert data["success"] is True
assert data["channel"] == "nightly"
assert data["new_version"] == "main-fedcba9"
@pytest.mark.asyncio
async def test_switch_channel_to_nightly_with_git_calls_git_update(monkeypatch, tmp_path):
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
routes_file.parent.mkdir(parents=True)
routes_file.write_text("")
monkeypatch.setattr(update_routes, "__file__", str(routes_file))
monkeypatch.setattr(update_routes, "ensure_settings_file", lambda logger: str(tmp_path / "settings.json"))
(tmp_path / ".git").mkdir()
async def _fake_git_update(*args, **kwargs):
return True, "main-1111111"
monkeypatch.setattr(
update_routes.UpdateRoutes, "_perform_git_update", _fake_git_update
)
req = _fake_request({"channel": "nightly"})
resp = await update_routes.UpdateRoutes.switch_channel(req)
data = _raw_body(resp)
assert data["success"] is True
assert data["channel"] == "nightly"
assert data["new_version"] == "main-1111111"
@pytest.mark.asyncio
async def test_switch_channel_to_release_with_git_calls_git_update(monkeypatch, tmp_path):
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
routes_file.parent.mkdir(parents=True)
routes_file.write_text("")
monkeypatch.setattr(update_routes, "__file__", str(routes_file))
monkeypatch.setattr(update_routes, "ensure_settings_file", lambda logger: str(tmp_path / "settings.json"))
(tmp_path / ".git").mkdir()
async def _fake_git_update(*args, **kwargs):
return True, "v9.9.9"
monkeypatch.setattr(
update_routes.UpdateRoutes, "_perform_git_update", _fake_git_update
)
req = _fake_request({"channel": "release"})
resp = await update_routes.UpdateRoutes.switch_channel(req)
data = _raw_body(resp)
assert data["success"] is True
assert data["channel"] == "release"
assert data["new_version"] == "v9.9.9"
@pytest.mark.asyncio
async def test_switch_channel_to_release_without_git_still_downloads_zip(monkeypatch, tmp_path):
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
routes_file.parent.mkdir(parents=True)
routes_file.write_text("")
monkeypatch.setattr(update_routes, "__file__", str(routes_file))
monkeypatch.setattr(update_routes, "ensure_settings_file", lambda logger: str(tmp_path / "settings.json"))
async def _fake_zip(*args, **kwargs):
return True, "v2.0.0"
monkeypatch.setattr(
update_routes.UpdateRoutes, "_download_and_replace_zip", _fake_zip
)
req = _fake_request({"channel": "release"})
resp = await update_routes.UpdateRoutes.switch_channel(req)
data = _raw_body(resp)
assert data["success"] is True
assert data["channel"] == "release"
assert data["new_version"] == "v2.0.0"
class _NightlyDownloader:
"""Returns a fake main-branch commit AND a compare response."""
commit_sha = "7777777"
commit_msg = "test: add nightly feature"
commit_date = "2026-07-27T12:00:00Z"
behind_by = 5
async def make_request(self, method, url, **kwargs):
if "/compare/" in url:
return True, {"behind_by": self.behind_by}
return True, {
"sha": self.commit_sha,
"commit": {
"message": self.commit_msg,
"committer": {"date": self.commit_date},
},
}
@pytest.mark.asyncio
async def test_get_nightly_version_parses_behind_by(monkeypatch):
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(_NightlyDownloader()))
version, changelog, behind_by, commit_date = await update_routes.UpdateRoutes._get_nightly_version(
local_hash="abc1234"
)
assert version == "main-7777777"
assert behind_by == 5
assert commit_date == "2026-07-27"
assert len(changelog) == 1
assert changelog[0] == "test: add nightly feature"
class _AheadCompareDownloader:
"""Fake compare API response with status='ahead' (main is ahead of local)."""
commit_sha = "9999999"
commit_msg = "latest commit"
commit_date = "2026-07-28T00:00:00Z"
ahead_by = 3
async def make_request(self, method, url, **kwargs):
if "/compare/" in url:
return True, {"status": "ahead", "ahead_by": self.ahead_by, "behind_by": 0}
return True, {
"sha": self.commit_sha,
"commit": {
"message": self.commit_msg,
"committer": {"date": self.commit_date},
},
}
@pytest.mark.asyncio
async def test_get_nightly_version_reads_ahead_by_when_ahead(monkeypatch):
"""compare/{local}...main returns status='ahead' → read ahead_by."""
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(_AheadCompareDownloader()))
version, changelog, behind_by, commit_date = await update_routes.UpdateRoutes._get_nightly_version(
local_hash="oldhash"
)
assert version == "main-9999999"
assert behind_by == 3
assert commit_date == "2026-07-28"
class _DivergedCompareDownloader:
"""Fake compare API response with status='diverged' (both have unique commits)."""
commit_sha = "aaaaaaa"
commit_msg = "diverged test"
commit_date = "2026-07-29T00:00:00Z"
async def make_request(self, method, url, **kwargs):
if "/compare/" in url:
return True, {"status": "diverged", "ahead_by": 5, "behind_by": 2}
return True, {
"sha": self.commit_sha,
"commit": {
"message": self.commit_msg,
"committer": {"date": self.commit_date},
},
}
@pytest.mark.asyncio
async def test_get_nightly_version_reads_ahead_by_when_diverged(monkeypatch):
"""compare/{local}...main returns status='diverged' → read ahead_by (remote ahead)."""
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(_DivergedCompareDownloader()))
version, changelog, behind_by, commit_date = await update_routes.UpdateRoutes._get_nightly_version(
local_hash="divhash"
)
assert behind_by == 5
class _CheckUpdatesDownloader:
"""Fake downloader returning both a release list and a nightly commit + compare."""
commit_sha = "8888888"
commit_date = "2026-07-28T00:00:00Z"
async def make_request(self, method, url, **kwargs):
if "/releases" in url:
return True, [
{
"tag_name": "v3.0.0",
"body": "- Feature A\n- Feature B",
"published_at": "2026-07-20T00:00:00Z",
}
]
if "/compare/" in url:
return True, {"behind_by": 3}
return True, {
"sha": self.commit_sha + "0" * 33,
"commit": {
"message": "latest commit",
"committer": {"date": self.commit_date},
},
}
@pytest.mark.asyncio
async def test_check_updates_nightly_response_includes_behind_and_date(monkeypatch, tmp_path):
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(_CheckUpdatesDownloader()))
monkeypatch.setattr(
update_routes.UpdateRoutes,
"_get_local_version",
staticmethod(lambda: "v1.0.0"),
)
monkeypatch.setattr(
update_routes.UpdateRoutes,
"_get_git_info",
staticmethod(lambda: {
"commit_hash": "abc1234",
"short_hash": "abc1234",
"branch": "main",
"commit_date": "2026-01-01",
}),
)
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
routes_file.parent.mkdir(parents=True)
routes_file.write_text("")
monkeypatch.setattr(update_routes, "__file__", str(routes_file))
(tmp_path / ".git").mkdir()
req = _fake_request(query_params={"nightly": "true"})
resp = await update_routes.UpdateRoutes.check_updates(req)
data = _raw_body(resp)
assert data["success"] is True
assert data["nightly"] is True
assert data["has_git"] is True
assert data["behind_by"] == 3
assert data["commit_date"] == "2026-07-28"
assert data["latest_version"] == "main-8888888"
assert isinstance(data["releases"], list)
assert len(data["releases"]) == 1
assert data["releases"][0]["version"] == "v3.0.0"
def _raw_body(response):
import json
return json.loads(response._body.decode())
+67 -1
View File
@@ -183,7 +183,7 @@ class FakeCache:
def __init__(self, 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":
data = sorted(self.items, key=lambda x: x["model_name"].lower())
if order == "desc":
@@ -1252,3 +1252,69 @@ async def test_get_model_civitai_url_falls_back_when_host_setting_is_not_a_strin
"model_id": "123",
"version_id": "456",
}
class TestHfGroupKey:
"""Tests for _extract_hf_group_key and _extract_group_key."""
# --- _extract_hf_group_key ---
def test_hf_group_key_valid_url(self):
"""Standard HF URL returns hf:user/repo."""
item = {"hf_url": "https://huggingface.co/unsloth/qwen-edit"}
assert BaseModelService._extract_hf_group_key(item) == "hf:unsloth/qwen-edit"
def test_hf_group_key_url_with_subpath(self):
"""URL with subpath still extracts just owner/repo."""
item = {"hf_url": "https://huggingface.co/user/repo/resolve/main/file.safetensors"}
assert BaseModelService._extract_hf_group_key(item) == "hf:user/repo"
def test_hf_group_key_empty_url(self):
"""Empty hf_url returns None."""
assert BaseModelService._extract_hf_group_key({"hf_url": ""}) is None
def test_hf_group_key_no_url(self):
"""Missing hf_url key returns None."""
assert BaseModelService._extract_hf_group_key({}) is None
def test_hf_group_key_none_url(self):
"""None hf_url returns None."""
assert BaseModelService._extract_hf_group_key({"hf_url": None}) is None
def test_hf_group_key_invalid_url(self):
"""Malformed HF URL returns None."""
assert BaseModelService._extract_hf_group_key({"hf_url": "not-a-url"}) is None
assert BaseModelService._extract_hf_group_key({"hf_url": "https://example.com"}) is None
# --- _extract_group_key ---
def test_group_key_civitai_only(self):
"""CivitAI modelId returned as int."""
item = {"civitai": {"modelId": 123}}
assert BaseModelService._extract_group_key(item) == 123
def test_group_key_hf_only(self):
"""HF-only item returns hf:user/repo string."""
item = {"hf_url": "https://huggingface.co/user/repo"}
assert BaseModelService._extract_group_key(item) == "hf:user/repo"
def test_group_key_civitai_preferred(self):
"""CivitAI modelId takes precedence over hf_url."""
item = {
"civitai": {"modelId": 456},
"hf_url": "https://huggingface.co/other/repo",
}
assert BaseModelService._extract_group_key(item) == 456
def test_group_key_neither(self):
"""No CivitAI or HF returns None."""
assert BaseModelService._extract_group_key({}) is None
assert BaseModelService._extract_group_key({"some": "data"}) is None
def test_group_key_civitai_none_model_id(self):
"""civitai.modelId=None falls through to HF."""
item = {
"civitai": {"modelId": None},
"hf_url": "https://huggingface.co/user/repo",
}
assert BaseModelService._extract_group_key(item) == "hf:user/repo"
+142
View File
@@ -363,6 +363,148 @@ async def test_check_pending_models_handles_corrupted_progress_file(
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
def settings_manager():
return get_settings_manager()
+161
View File
@@ -35,9 +35,11 @@ class DummyDownloader:
def reset_singletons():
CivitaiClient._instance = None
ModelMetadataProviderManager._instance = None
civitai_client_module._creator_model_count_cache.clear()
yield
CivitaiClient._instance = None
ModelMetadataProviderManager._instance = None
civitai_client_module._creator_model_count_cache.clear()
@pytest.fixture
@@ -622,3 +624,162 @@ async def test_get_image_info_handles_invalid_id(monkeypatch, downloader, caplog
assert result is None
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:
self._cache = SimpleNamespace(raw_data=models)
self.sync_calls: list[tuple[str, dict]] = []
async def get_cached_data(self):
return self._cache
@@ -38,6 +39,14 @@ class StubScanner:
break
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:
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.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")
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 manager._progress["failed_models"] == {model_hash}
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"]
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())
assert files == ["image_0.png", "image_1.png"]
assert (model_dir / "image_0.png").read_bytes() == b"first"
assert files == ["image_1.png"]
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
def 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"
)
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._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"
)
await scanner._cache.resort()
scanner._cache._last_sort = ("name", "asc")
scanner._cache._last_sort = ("name", "asc", None)
scanner._tags_count = {"alpha": 1}
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()
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):
lora_root = tmp_path / "loras"
lora_root.mkdir()
+119
View File
@@ -139,3 +139,122 @@ def test_contains_dynamic_syntax_detects_wildcards_and_options():
assert contains_dynamic_syntax("__flower__") is True
assert contains_dynamic_syntax("{red|blue}") is True
assert contains_dynamic_syntax("{2$$, $$red|blue|green}") is True
# ---------------------------------------------------------------------------
# _pick_weighted_or_plain
# ---------------------------------------------------------------------------
def test_pick_weighted_or_plain_plain_values(monkeypatch, tmp_path):
"""Plain values without :: are picked via rng.choice (fast path)."""
service, _ = _make_service(monkeypatch, tmp_path)
import random
rng = random.Random(42)
result = service._pick_weighted_or_plain(["red", "green", "blue"], rng)
assert result in {"red", "green", "blue"}
assert "::" not in result
def test_pick_weighted_or_plain_deterministic_with_seed(monkeypatch, tmp_path):
"""Same seed produces the same result for plain values."""
service, _ = _make_service(monkeypatch, tmp_path)
import random
first = service._pick_weighted_or_plain(["a", "b", "c"], random.Random(99))
second = service._pick_weighted_or_plain(["a", "b", "c"], random.Random(99))
assert first == second
def test_pick_weighted_or_plain_weighted_values(monkeypatch, tmp_path):
"""Weighted values use weighted selection and strip the N:: prefix."""
service, _ = _make_service(monkeypatch, tmp_path)
import random
values = ["3::apple", "1::banana"]
results = {"apple": 0, "banana": 0}
for seed in range(4000):
result = service._pick_weighted_or_plain(values, random.Random(seed))
assert result in results, f"Unexpected result: {result!r}"
assert "::" not in result
results[result] += 1
total = results["apple"] + results["banana"]
# 3:1 weight → apple ≈ 75%, banana ≈ 25%
assert 2700 < results["apple"] < 3300, f"apple count out of range: {results['apple']}"
assert 700 < results["banana"] < 1300, f"banana count out of range: {results['banana']}"
def test_pick_weighted_or_plain_weight_one_values(monkeypatch, tmp_path):
"""Values with explicit 1:: prefix have prefix stripped but are not weighted."""
service, _ = _make_service(monkeypatch, tmp_path)
import random
# All weights are 1.0 → no actual weighting, but :: prefix is stripped
values = ["1::foo", "1::bar"]
rng = random.Random(42)
results = {service._pick_weighted_or_plain(values, rng) for _ in range(200)}
assert results == {"foo", "bar"}
# Ensure the prefix is always stripped
for result in results:
assert "::" not in result
def test_pick_weighted_or_plain_mixed_weighted_and_plain(monkeypatch, tmp_path):
"""Mixed list with some weighted and some unweighted values."""
service, _ = _make_service(monkeypatch, tmp_path)
import random
values = ["5::x", "y", "z"] # x has weight 5, y/z have default weight 1
results = {"x": 0, "y": 0, "z": 0}
for seed in range(4000):
result = service._pick_weighted_or_plain(values, random.Random(seed))
assert result in results
assert "::" not in result
results[result] += 1
# x (5) vs combined y+z (1+1=2) → ~71% / ~29%
x_pct = results["x"] / sum(results.values())
assert 0.65 < x_pct < 0.78, f"x proportion out of range: {x_pct:.3f}"
def test_pick_weighted_or_plain_invalid_weight_prefix(monkeypatch, tmp_path):
"""Invalid numeric prefix (e.g. 1.2.3) is NOT treated as a weight and
the prefix is NOT stripped, matching the updated strict regex."""
service, _ = _make_service(monkeypatch, tmp_path)
import random
rng = random.Random(42)
# "1.2.3::a" is not a valid number → treated as plain text value
result = service._pick_weighted_or_plain(["1.2.3::a", "b"], rng)
# It should keep the full text including :: because the prefix isn't a
# valid numeric weight according to the strict regex
assert result == "1.2.3::a" or result == "b"
def test_pick_weighted_or_plain_glob_aggregation(monkeypatch, tmp_path):
"""Weighted wildcard resolution through glob aggregation (__*__)."""
service, wildcards_dir = _make_service(monkeypatch, tmp_path)
wildcards_dir.mkdir()
(wildcards_dir / "animals").mkdir()
(wildcards_dir / "animals" / "cat.txt").write_text("3::tabby\n1::persian\n", encoding="utf-8")
(wildcards_dir / "animals" / "dog.txt").write_text("retriever\npoodle\n", encoding="utf-8")
# __animals/*__ aggregates all values across both files
# Weighted values should have :: stripped
results = {"tabby": 0, "persian": 0, "retriever": 0, "poodle": 0}
for seed in range(4000):
expanded = service.expand_text("__animals/*__", seed=seed)
assert expanded in results, f"Unexpected result: {expanded!r}"
assert "::" not in expanded
results[expanded] += 1
# tabby (3) vs persian (1) → ~75% / ~25% within the cat subset
cat_total = results["tabby"] + results["persian"]
if cat_total > 0:
tabby_pct = results["tabby"] / cat_total
assert 0.65 < tabby_pct < 0.85, f"tabby proportion out of range: {tabby_pct:.3f}"
@@ -63,7 +63,7 @@ async def test_start_download_bootstraps_progress_and_task(
release = asyncio.Event()
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()
await release.wait()
@@ -93,6 +93,44 @@ async def test_start_download_bootstraps_progress_and_task(
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:
settings_manager = get_settings_manager()
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:
self.cache = SimpleNamespace(raw_data=cache_items)
self.updates: List[Tuple[str, str, Dict[str, Any]]] = []
self.sync_updates: List[Tuple[str, Dict[str, Any]]] = []
async def get_cached_data(self):
return self.cache
@@ -23,6 +24,10 @@ class StubScanner:
self.updates.append((old_path, new_path, metadata))
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)
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 Path(patch_metadata_manager[0][0]) == model_file
assert scanner.updates
assert scanner.sync_updates
@pytest.mark.asyncio
@@ -151,8 +156,8 @@ async def test_update_metadata_after_import_preserves_existing_metadata(
assert saved_payload["civitai"]["trainedWords"] == ["foo"]
assert {entry["id"] for entry in saved_payload["civitai"]["customImages"]} == {"existing-id", "new-id"}
assert scanner.updates
updated_metadata = scanner.updates[-1][2]
assert scanner.sync_updates
updated_metadata = scanner.sync_updates[-1][1]
assert updated_metadata["civitai"]["images"] == existing_payload["civitai"]["images"]
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"
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:
def __init__(self, models: list[Dict[str, Any]]) -> None:
self._cache = SimpleNamespace(raw_data=models)
@@ -45,7 +45,7 @@ export interface AutocompleteTextWidgetInterface {
const props = defineProps<{
widget: AutocompleteTextWidgetInterface
node: { id: number }
modelType?: 'loras' | 'embeddings' | 'custom_words' | 'prompt'
modelType?: 'loras' | 'prompt'
placeholder?: string
showPreview?: boolean
spellcheck?: boolean
@@ -98,7 +98,7 @@ interface LoraInfoWidget {
onSetValue?: (v: unknown) => void
callback?: unknown
options?: {
getValue?: () => LoraInfoWidgetValue
getValue?: () => unknown
setValue?: (v: unknown) => void
}
node?: { widgets?: Array<{ id?: string }>; widgets_values?: Array<unknown> }
@@ -299,8 +299,12 @@ onMounted(() => {
// ComponentWidgetImpl.value getter/setter delegates to options.getValue/options.setValue.
// These must be set for workflow JSON persistence (LGraphNode.serialize/configure) to work.
props.widget.options.getValue = buildValue
props.widget.options.setValue = applyValue
if (props.widget.options) {
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)
props.widget.serializeValue = async () => buildValue()
@@ -3,7 +3,7 @@ import { ref, onMounted, onUnmounted, type Ref } from 'vue'
// Dynamic import type for AutoComplete class
type AutoCompleteClass = new (
inputElement: HTMLTextAreaElement,
modelType: 'loras' | 'embeddings' | 'custom_words' | 'prompt',
modelType: 'loras' | 'prompt',
options?: AutocompleteOptions
) => AutoCompleteInstance
@@ -29,7 +29,7 @@ export interface UseAutocompleteOptions {
export function useAutocomplete(
textareaRef: Ref<HTMLTextAreaElement | null>,
modelType: 'loras' | 'embeddings' | 'custom_words' | 'prompt' = 'loras',
modelType: 'loras' | 'prompt' = 'loras',
options: UseAutocompleteOptions = {}
) {
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 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
import { app } from '../../../scripts/app.js'
// @ts-ignore - ComfyUI external module
@@ -718,7 +721,7 @@ function createLoraInfoWidget(node: any) {
function createAutocompleteTextWidgetFactory(
node: any,
widgetName: string,
modelType: 'loras' | 'embeddings' | 'prompt',
modelType: 'loras' | 'prompt',
inputOptions: { placeholder?: string } = {}
) {
const metadataWidgetName = `__lm_autocomplete_meta_${widgetName}`
@@ -835,7 +838,7 @@ function createAutocompleteTextWidgetFactory(
applyAutocompleteTextLayoutFix(
widget,
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`) || {}
return createAutocompleteTextWidgetFactory(node, 'text', 'loras', options)
},
// Autocomplete text widget for embeddings (used by Prompt node)
// @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)
// Autocomplete text widget for prompt (used by Prompt and Text nodes)
// @ts-ignore
AUTOCOMPLETE_TEXT_PROMPT(node) {
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."
});
}
}
});

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