Compare commits

..

61 Commits

Author SHA1 Message Date
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
Will Miao f49b4ba4db fix(metadata-overwrite): rename 'checkpoint' input to 'model' 2026-07-26 10:59:38 +08:00
Will Miao 84e708328b fix: correct return_types propagation to GenericNodeExtractor
Two bugs prevented type-signature-based fallback from working:

- metadata_hook.py used getattr(obj.__class__, 'RETURN_TYPES')
  which fails when _async_map_node_over_list is called with
  a class (not instance) — obj.__class__ is the metaclass
  'type', which has no RETURN_TYPES. Fixed: getattr(obj, ...).

- metadata_registry.py used type(extractor) is GenericNodeExtractor
  to dispatch return_types. NODE_EXTRACTORS stores class
  references, not instances; type(Class) is always 'type',
  never the class. Fixed: extractor is GenericNodeExtractor.
2026-07-26 10:34:20 +08:00
Will Miao 125bed3f09 feat: add Metadata Overwrite node for manual generation params override 2026-07-26 08:50:21 +08:00
Will Miao 077e70169d feat: add type-signature-based fallback for unregistered nodes
GenericNodeExtractor (previously a no-op) now inspects
RETURN_TYPES to detect MODEL loaders and CONDITIONING
encoders in nodes not registered in NODE_EXTRACTORS.

- Propagate return_types from the hook layer through the
  registry to GenericNodeExtractor.extract() and update().
- MODEL detection: scan ckpt_name/unet_name/model_path/
  model_name/gguf_name fields, validate by extension.
- CONDITIONING detection: scan text/clip_l/t5xxl/prompt
  fields, store prompt text and conditioning tensor.
- _fill_missing_metadata also checks node_cache, so
  GenericNodeExtractor-handled nodes survive cache.
2026-07-25 22:14:51 +08:00
Will Miao e6dc169a05 feat: add meta hints user marks for metadata heuristic override
Users can now right-click nodes and assign meta hints
(primary_model, primary_sampler, positive_prompt,
negative_prompt) to override the metadata processor's
heuristic inference.

- Store extra_data from the API request so workflow node
  properties (including lm_marker_role) are accessible
  during metadata processing.
- _get_user_marks scans extra_data.extra_pnginfo.workflow
  for meta_* marks, falling back to prompt.original_prompt.
- extract_generation_params checks user marks before
  heuristic inference for sampler, model, and prompts.
- Warn on duplicate marks or invalid marked nodes.
2026-07-25 22:13:52 +08:00
Will Miao f34c02756d fix(recipes): eliminate O(n) fuzzy search fallback over 42k+ recipes
Drop the SequenceMatcher-based fuzzy_match fallback that froze the server
when FTS returned empty results. FTS now returns empty set for zero results
(no fallback), and when the index is not yet ready, search returns empty
rather than scanning all items in Python.
2026-07-25 17:34:37 +08:00
Will Miao 1e4c315481 fix(ModelModal): respect civitai_host setting for creator profile link 2026-07-25 07:15:12 +08:00
Will Miao a8283a0d00 fix(SaveImageLM): clarify embed_workflow tooltip — explains drag-and-drop workflow restoration
The previous tooltip was misleading: users thought workflow embedding was
automatic. New wording explains this opt-in flag stores the complete
workflow inside images, allowing one-click restoration via drag-and-drop.
PNG and WebP only.
2026-07-24 19:53:59 +08:00
Will Miao 55896669fc feat(SaveImageLM): expose webp_method and jpeg_subsampling as conditional node inputs
Add two new optional parameters to the Save Image node:

- webp_method (INT, 0-6, default 6): Controls WebP compression level.
  0=fastest/largest, 6=slowest/smallest. Previously hardcoded to 0.
- jpeg_subsampling (INT, 0-2, default 0): Controls JPEG chroma
  subsampling. 0=4:4:4 (best quality), 1=4:2:2, 2=4:2:0.

Frontend JS extension hides/disables each parameter when the
selected file_format doesn't apply (e.g., webp_method is hidden
when saving as PNG or JPEG). 7 new tests cover parameter plumbing
and default consistency across INPUT_TYPES, save_images(), and
process_image().
2026-07-24 19:32:51 +08:00
Will Miao e341e0b9d2 fix(test): update parameters assertion to include Version: ComfyUI after metadata format upgrade 2026-07-24 18:29:07 +08:00
Will Miao e6538c83bb fix(metadata): restore sha256 after hydrate_model_data to prevent KeyError in CivitAI fetch
hydrate_model_data replaces model_data with .metadata.json content which
may lack sha256 (corrupted file, concurrent write, etc.). Restore the
cached sha256 after hydration and persist the fix back to disk so
subsequent lookups don't hit the same error.

Also improve error log to include file_path for debugging.
2026-07-24 12:07:18 +08:00
Will Miao 92e1285ea5 feat(SaveImageLM): upgrade metadata output to A1111/Civitai-compatible format
- Replace plain-text Lora hashes with Hashes JSON dict matching A1111 convention
- Add Civitai resources JSON array with AIR URNs for direct model version linking
- Add Clip skip, Version: ComfyUI fields to generation params line
- Build AIR strings from local scanner cache (no API calls needed)
- Add complete sampler name mapping (CIVITAI_SAMPLER_MAP) and base model → AIR slug mapping (BASE_MODEL_AIR_SLUG) sourced from civitai ecosystem constants
- Remove lora text prepending from prompt line; LoRA info now in structured JSON sections
2026-07-24 06:20:28 +08:00
Will Miao 2aabd1d90e fix(ai): use json_schema instead of json_object for broader provider compatibility (#1033)
LM Studio and some other OpenAI-compatible servers reject
response_format=json_object but accept json_schema. Switch to the
equivalent json_schema format and add a fallback that retries
without response_format when the provider rejects the format type.
2026-07-23 09:17:29 +08:00
Will Miao 7b8b778f83 fix(widget): restore strength drag on lora entries and header
widget.value is a getter/setter that returns a new array on every read,
so handleStrengthDrag with updateWidget=false mutated a discarded copy.
Introduce __dragActive flag to suppress renderLoras in setValue during
drag, allowing mutations to persist through widget.value without
destroying the DOM. Use try-finally to guarantee flag cleanup.
2026-07-23 08:31:34 +08:00
Will Miao 7c8dc57d55 fix(security): use abspath instead of realpath in containment checks to support symlinks (#1028) 2026-07-23 07:06:41 +08:00
Will Miao fe95fae5f2 fix(workflow): include Create Hook LoRA in lora_code_update handler 2026-07-22 11:40:56 +08:00
Will Miao ce8a95abf7 chore(release): bump version to v1.1.9 2026-07-21 22:22:39 +08:00
Will Miao c8e7e543d6 fix(api): remove overstrict model type validation in getApiEndpoints
The validation in getApiEndpoints threw for page types not in
MODEL_TYPES (e.g. 'recipes'), crashing the recipes page initialization
when FilterManager calls it via createBaseModelTags(). The throw was
synchronous and outside the fetch().catch() chain, causing an uncaught
promise rejection that aborted the entire app initialization.

getApiEndpoints is a URL builder -- validation belongs to callers that
need strict type checking (they already use isValidModelType()). For
non-model-type pages like recipes, the generated URLs are correct
(the backend does have /api/lm/recipes/* routes).

Fixes regression from f53f859a (feat(filter): add debounced tag search).
2026-07-21 22:09:30 +08:00
Will Miao a9dbb15ffa fix(create_hook_lora): lazy import comfy.hooks/comfy.utils to fix CI pipeline (#744) 2026-07-21 18:44:04 +08:00
Will Miao cf64043f7d fix(security): add library root containment check for delete/move/rename operations (#1028) 2026-07-21 15:23:38 +08:00
Will Miao ccaff92c18 fix(nodes): register Create Hook LoRA node in workflow target registries 2026-07-21 14:56:39 +08:00
Will Miao 585b5c922a feat(nodes): add Create Hook LoRA (LoraManager) node for multi-LoRA hook pipelines 2026-07-21 09:44:28 +08:00
Will Miao ea80c2224c fix(download): prevent path traversal in download template resolution (#1028) 2026-07-20 21:08:43 +08:00
willmiao 8b0f56c1a6 docs: auto-update supporters list in README 2026-07-20 12:42:20 +00:00
117 changed files with 27086 additions and 20821 deletions
+7 -1
View File
@@ -137,7 +137,13 @@ npm run test:coverage # Generate coverage report
- Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json) - Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json)
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"` - Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
- Run `python scripts/sync_translation_keys.py` after adding UI strings to `locales/en.json` - Run `python scripts/sync_translation_keys.py` after adding UI strings to `locales/en.json`
- Symlinks require normalized paths - Symlinks require normalized paths.
**Business paths vs real paths**: All stored paths and operation routing use the
original paths as they appear under configured model roots — symlinks are NOT
resolved. `os.path.realpath` is only for scanner dedup and the symlink cache.
Any path passed to `os.remove`/`os.rename`/`shutil.move` or validated by a
containment check MUST use the business path (i.e. `os.path.abspath`, not
`realpath`).
## Git / Commit Messages ## Git / Commit Messages
+2 -2
View File
File diff suppressed because one or more lines are too long
+10
View File
@@ -17,6 +17,8 @@ try: # pragma: no cover - import fallback for pytest collection
from .py.nodes.lora_cycler import LoraCyclerLM from .py.nodes.lora_cycler import LoraCyclerLM
from .py.nodes.lora_info import LoraInfoLM from .py.nodes.lora_info import LoraInfoLM
from .py.nodes.lora_syntax_to_path import LoraSyntaxToPath from .py.nodes.lora_syntax_to_path import LoraSyntaxToPath
from .py.nodes.create_hook_lora import CreateHookLoraLM
from .py.nodes.metadata_overwrite import MetadataOverwriteLM
from .py.metadata_collector import init as init_metadata_collector from .py.metadata_collector import init as init_metadata_collector
except ( except (
ImportError ImportError
@@ -62,6 +64,12 @@ except (
LoraSyntaxToPath = importlib.import_module( LoraSyntaxToPath = importlib.import_module(
"py.nodes.lora_syntax_to_path" "py.nodes.lora_syntax_to_path"
).LoraSyntaxToPath ).LoraSyntaxToPath
CreateHookLoraLM = importlib.import_module(
"py.nodes.create_hook_lora"
).CreateHookLoraLM
MetadataOverwriteLM = importlib.import_module(
"py.nodes.metadata_overwrite"
).MetadataOverwriteLM
init_metadata_collector = importlib.import_module("py.metadata_collector").init init_metadata_collector = importlib.import_module("py.metadata_collector").init
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
@@ -83,6 +91,8 @@ NODE_CLASS_MAPPINGS = {
LoraCyclerLM.NAME: LoraCyclerLM, LoraCyclerLM.NAME: LoraCyclerLM,
LoraInfoLM.NAME: LoraInfoLM, LoraInfoLM.NAME: LoraInfoLM,
LoraSyntaxToPath.NAME: LoraSyntaxToPath, LoraSyntaxToPath.NAME: LoraSyntaxToPath,
CreateHookLoraLM.NAME: CreateHookLoraLM,
MetadataOverwriteLM.NAME: MetadataOverwriteLM,
} }
WEB_DIRECTORY = "./web/comfyui" WEB_DIRECTORY = "./web/comfyui"
+313 -291
View File
File diff suppressed because it is too large Load Diff
+2227 -2205
View File
File diff suppressed because it is too large Load Diff
+23 -1
View File
@@ -714,7 +714,9 @@
"versionsCount": "Local Versions", "versionsCount": "Local Versions",
"versionsCountDesc": "Most versions first", "versionsCountDesc": "Most versions first",
"versionsCountAsc": "Fewest versions first", "versionsCountAsc": "Fewest versions first",
"versionIdDesc": "Newest version first" "versionIdDesc": "Newest version first",
"random": "Random",
"randomAction": "Randomize (shuffle)"
}, },
"refresh": { "refresh": {
"title": "Refresh model list", "title": "Refresh model list",
@@ -771,6 +773,8 @@
"deleteAll": "Delete Selected", "deleteAll": "Delete Selected",
"downloadMissingLoras": "Download Missing LoRAs", "downloadMissingLoras": "Download Missing LoRAs",
"downloadExamples": "Download Example Images", "downloadExamples": "Download Example Images",
"downloadMissingExamples": "Download Missing",
"reprocessExamples": "Re-process All",
"clear": "Clear Selection", "clear": "Clear Selection",
"skipMetadataRefreshCount": "Skip ({count} models)", "skipMetadataRefreshCount": "Skip ({count} models)",
"resumeMetadataRefreshCount": "Resume ({count} models)", "resumeMetadataRefreshCount": "Resume ({count} models)",
@@ -806,6 +810,8 @@
"sendToWorkflowReplace": "Send to Workflow (Replace)", "sendToWorkflowReplace": "Send to Workflow (Replace)",
"openExamples": "Open Examples Folder", "openExamples": "Open Examples Folder",
"downloadExamples": "Download Example Images", "downloadExamples": "Download Example Images",
"downloadMissingExamples": "Download Missing",
"reprocessExamples": "Re-process All",
"replacePreview": "Replace Preview", "replacePreview": "Replace Preview",
"setContentRating": "Set Content Rating", "setContentRating": "Set Content Rating",
"moveToFolder": "Move to Folder", "moveToFolder": "Move to Folder",
@@ -1548,6 +1554,7 @@
"empty": "No version history available for this model yet.", "empty": "No version history available for this model yet.",
"error": "Failed to load versions.", "error": "Failed to load versions.",
"missingModelId": "This model is missing a Civitai model id.", "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": { "confirm": {
"delete": "Delete this version from your library?" "delete": "Delete this version from your library?"
}, },
@@ -1751,6 +1758,12 @@
"checkingMessage": "Please wait while we check for the latest version.", "checkingMessage": "Please wait while we check for the latest version.",
"showNotifications": "Show update notifications", "showNotifications": "Show update notifications",
"latestBadge": "Latest", "latestBadge": "Latest",
"latestMain": "Latest main",
"channel": "Update Channel",
"channels": {
"release": "Release",
"nightly": "Nightly"
},
"updateProgress": { "updateProgress": {
"preparing": "Preparing update...", "preparing": "Preparing update...",
"installing": "Installing update...", "installing": "Installing update...",
@@ -1771,6 +1784,15 @@
"warning": "Warning: Nightly builds may contain experimental features and could be unstable.", "warning": "Warning: Nightly builds may contain experimental features and could be unstable.",
"enable": "Enable Nightly Updates" "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": { "banners": {
"recent": "Recent messages", "recent": "Recent messages",
"empty": "No recent banners yet.", "empty": "No recent banners yet.",
+2227 -2205
View File
File diff suppressed because it is too large Load Diff
+2227 -2205
View File
File diff suppressed because it is too large Load Diff
+2227 -2205
View File
File diff suppressed because it is too large Load Diff
+2227 -2205
View File
File diff suppressed because it is too large Load Diff
+2227 -2205
View File
File diff suppressed because it is too large Load Diff
+2227 -2205
View File
File diff suppressed because it is too large Load Diff
+2227 -2205
View File
File diff suppressed because it is too large Load Diff
+2227 -2205
View File
File diff suppressed because it is too large Load Diff
+15 -1
View File
@@ -1,5 +1,11 @@
"""Constants used by the metadata collector""" """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 # Metadata categories
MODELS = "models" MODELS = "models"
PROMPTS = "prompts" PROMPTS = "prompts"
@@ -9,6 +15,14 @@ EMBEDDINGS = "embeddings"
SIZE = "size" SIZE = "size"
IMAGES = "images" IMAGES = "images"
IS_SAMPLER = "is_sampler" # New constant to mark sampler nodes IS_SAMPLER = "is_sampler" # New constant to mark sampler nodes
OVERWRITE = "overwrite" # Manual metadata overwrite from MetadataOverwriteLM node
# Field names that the MetadataOverwriteLM node and its extractor share
METADATA_OVERWRITE_FIELDS = (
"prompt", "negative_prompt", "seed", "steps", "cfg_scale",
"sampler", "scheduler", "model", "loras", "size",
"clip_skip", "additional_data",
)
# Complete list of categories to track # Complete list of categories to track
METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES] METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES, OVERWRITE]
+14 -4
View File
@@ -83,7 +83,8 @@ class MetadataHook:
# Record inputs before execution # Record inputs before execution
if node_id is not None: if node_id is not None:
registry.record_node_execution(node_id, class_type, input_data_all, None) return_types = getattr(obj, 'RETURN_TYPES', None)
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
except Exception as e: except Exception as e:
logger.error(f"Error collecting metadata (pre-execution): {str(e)}") logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
@@ -114,7 +115,8 @@ class MetadataHook:
# Record outputs after execution # Record outputs after execution
if node_id is not None: if node_id is not None:
registry.update_node_execution(node_id, class_type, results) return_types = getattr(obj, 'RETURN_TYPES', None)
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
except Exception as e: except Exception as e:
logger.error(f"Error collecting metadata (post-execution): {str(e)}") logger.error(f"Error collecting metadata (post-execution): {str(e)}")
@@ -136,6 +138,9 @@ class MetadataHook:
if hasattr(prompt, 'original_prompt'): if hasattr(prompt, 'original_prompt'):
registry.set_current_prompt(prompt) registry.set_current_prompt(prompt)
# Store extra_data for accessing full workflow node properties
registry.set_extra_data(extra_data)
# Execute the original function # Execute the original function
return original_execute(*args, **kwargs) return original_execute(*args, **kwargs)
@@ -163,7 +168,8 @@ class MetadataHook:
class_type = obj.__class__.__name__ class_type = obj.__class__.__name__
node_id = unique_id node_id = unique_id
if node_id is not None: if node_id is not None:
registry.record_node_execution(node_id, class_type, input_data_all, None) return_types = getattr(obj, 'RETURN_TYPES', None)
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
except Exception as e: except Exception as e:
logger.error(f"Error collecting metadata (pre-execution): {str(e)}") logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
@@ -180,7 +186,8 @@ class MetadataHook:
class_type = obj.__class__.__name__ class_type = obj.__class__.__name__
node_id = unique_id node_id = unique_id
if node_id is not None: if node_id is not None:
registry.update_node_execution(node_id, class_type, results) return_types = getattr(obj, 'RETURN_TYPES', None)
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
except Exception as e: except Exception as e:
logger.error(f"Error collecting metadata (post-execution): {str(e)}") logger.error(f"Error collecting metadata (post-execution): {str(e)}")
@@ -202,6 +209,9 @@ class MetadataHook:
if hasattr(prompt, 'original_prompt'): if hasattr(prompt, 'original_prompt'):
registry.set_current_prompt(prompt) registry.set_current_prompt(prompt)
# Store extra_data for accessing full workflow node properties
registry.set_extra_data(extra_data)
# Execute the original function # Execute the original function
return await original_execute(*args, **kwargs) return await original_execute(*args, **kwargs)
+133 -9
View File
@@ -1,15 +1,68 @@
import json import json
import logging
import os import os
from .constants import IMAGES from .constants import IMAGES
# Check if running in standalone mode # Check if running in standalone mode
standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1" or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0" standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1" or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER, OVERWRITE
from .node_extractors import NODE_EXTRACTORS
logger = logging.getLogger(__name__)
# Keys that identify metadata hint marks stored in node.properties.lm_marker_role
_META_MARK_PREFIX = "meta_"
_MARK_PRIMARY_MODEL = "primary_model"
_MARK_PRIMARY_SAMPLER = "primary_sampler"
_MARK_POSITIVE_PROMPT = "positive_prompt"
_MARK_NEGATIVE_PROMPT = "negative_prompt"
class MetadataProcessor: class MetadataProcessor:
"""Process and format collected metadata""" """Process and format collected metadata"""
@staticmethod
def _get_user_marks(metadata):
"""Scan workflow nodes (from extra_data.extra_pnginfo.workflow) for user-assigned
metadata hint marks stored in node.properties.lm_marker_role.
Returns a dict mapping mark type keys to node IDs.
Example: {'primary_model': '42', 'primary_sampler': '17'}
"""
marks: dict[str, str] = {}
# Primary source: extra_data.extra_pnginfo.workflow.nodes (has full properties)
extra_data = metadata.get("extra_data")
if extra_data and isinstance(extra_data, dict):
extra_pnginfo = extra_data.get("extra_pnginfo", {})
if isinstance(extra_pnginfo, dict):
workflow = extra_pnginfo.get("workflow", {})
nodes = workflow.get("nodes", [])
for node in nodes:
node_id = str(node.get("id", ""))
role = node.get("properties", {}).get("lm_marker_role", "")
if role.startswith(_META_MARK_PREFIX):
mark_type = role[len(_META_MARK_PREFIX):]
if mark_type in marks:
logger.warning(
"Duplicate meta hint '%s': node %s (previous: %s), "
"last match wins",
mark_type, node_id, marks[mark_type],
)
marks[mark_type] = node_id
# Fallback: try prompt.original_prompt (API-only submissions may not have workflow)
if not marks:
prompt = metadata.get("current_prompt")
if prompt and getattr(prompt, "original_prompt", None):
for node_id, node_data in prompt.original_prompt.items():
role = node_data.get("properties", {}).get("lm_marker_role", "")
if role.startswith(_META_MARK_PREFIX):
mark_type = role[len(_META_MARK_PREFIX):]
marks[mark_type] = node_id
return marks
@staticmethod @staticmethod
def find_primary_sampler(metadata, downstream_id=None): def find_primary_sampler(metadata, downstream_id=None):
""" """
@@ -471,20 +524,57 @@ class MetadataProcessor:
"checkpoint": None, "checkpoint": None,
"loras": "", "loras": "",
"size": None, "size": None,
"clip_skip": None "clip_skip": None,
"additional_data": "",
} }
# Get the prompt object for node relationship tracing # Get the prompt object for node relationship tracing
prompt = metadata.get("current_prompt") prompt = metadata.get("current_prompt")
# Find the primary KSampler node # ---- User marks: override heuristic inference with user-assigned hints ----
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id) user_marks = MetadataProcessor._get_user_marks(metadata)
# Directly get checkpoint from metadata instead of tracing # Find the primary KSampler node (user mark takes priority)
# Pass primary_sampler_id to avoid redundant calculation primary_sampler_id = None
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id) primary_sampler = None
if checkpoint: if _MARK_PRIMARY_SAMPLER in user_marks:
params["checkpoint"] = checkpoint marked_id = user_marks[_MARK_PRIMARY_SAMPLER]
sampler_data = metadata.get(SAMPLING, {}).get(marked_id)
if sampler_data and sampler_data.get(IS_SAMPLER):
primary_sampler_id = marked_id
primary_sampler = sampler_data
else:
logger.warning(
"User-marked primary sampler %s has no runtime metadata, "
"falling back to heuristic",
marked_id,
)
if primary_sampler is None:
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id)
# Resolve checkpoint / model (user mark takes priority)
if _MARK_PRIMARY_MODEL in user_marks:
marked_id = user_marks[_MARK_PRIMARY_MODEL]
if marked_id in metadata.get(MODELS, {}):
params["checkpoint"] = metadata[MODELS][marked_id].get("name")
else:
extra_data = metadata.get("extra_data")
extra_pnginfo = extra_data.get("extra_pnginfo", {}) if extra_data and isinstance(extra_data, dict) else {}
workflow = extra_pnginfo.get("workflow", {}) if isinstance(extra_pnginfo, dict) else {}
node_type = "unknown"
for n in workflow.get("nodes", []):
if str(n.get("id", "")) == marked_id:
node_type = n.get("type", "unknown")
break
logger.warning(
"User-marked primary model %s (type=%s, registered=%s) has no runtime metadata, "
"falling back to heuristic",
marked_id, node_type, node_type in NODE_EXTRACTORS,
)
if params["checkpoint"] is None:
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id)
if checkpoint:
params["checkpoint"] = checkpoint
# Check if guidance parameter exists in any sampling node # Check if guidance parameter exists in any sampling node
for node_id, sampler_info in metadata.get(SAMPLING, {}).items(): for node_id, sampler_info in metadata.get(SAMPLING, {}).items():
@@ -540,6 +630,21 @@ class MetadataProcessor:
# For SamplerCustom, handle any additional parameters # For SamplerCustom, handle any additional parameters
MetadataProcessor.handle_custom_advanced_sampler(metadata, prompt, primary_sampler_id, params) MetadataProcessor.handle_custom_advanced_sampler(metadata, prompt, primary_sampler_id, params)
# ---- User marks: override prompts with explicitly tagged nodes ----
prompts_data = metadata.get(PROMPTS, {})
if _MARK_POSITIVE_PROMPT in user_marks:
pos_id = user_marks[_MARK_POSITIVE_PROMPT]
if pos_id in prompts_data:
prompt_text = prompts_data[pos_id].get("text") or prompts_data[pos_id].get("positive_text")
if prompt_text:
params["prompt"] = prompt_text
if _MARK_NEGATIVE_PROMPT in user_marks:
neg_id = user_marks[_MARK_NEGATIVE_PROMPT]
if neg_id in prompts_data:
prompt_text = prompts_data[neg_id].get("text") or prompts_data[neg_id].get("negative_text")
if prompt_text:
params["negative_prompt"] = prompt_text
# Size extraction is same for all sampler types # Size extraction is same for all sampler types
# Check if the sampler itself has size information (from latent_image) # Check if the sampler itself has size information (from latent_image)
if primary_sampler_id in metadata.get(SIZE, {}): if primary_sampler_id in metadata.get(SIZE, {}):
@@ -569,6 +674,25 @@ class MetadataProcessor:
if params["clip_skip"] is None: if params["clip_skip"] is None:
params["clip_skip"] = "1" params["clip_skip"] = "1"
# ---- Apply manual metadata overwrites ----
for overwrite_info in metadata.get(OVERWRITE, {}).values():
overwrite_params = overwrite_info.get("parameters", {})
for key, value in overwrite_params.items():
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),
# but the internal pipeline key remains "checkpoint" for backward compatibility
# with A1111 metadata format and downstream consumers.
if params.get("model"):
params["checkpoint"] = params["model"]
del params["model"]
return params return params
@staticmethod @staticmethod
+36 -13
View File
@@ -1,7 +1,7 @@
import time import time
from nodes import NODE_CLASS_MAPPINGS # type: ignore from nodes import NODE_CLASS_MAPPINGS # type: ignore
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
from .constants import METADATA_CATEGORIES, IMAGES from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE
class MetadataRegistry: class MetadataRegistry:
@@ -61,6 +61,7 @@ class MetadataRegistry:
{ {
"execution_order": [], "execution_order": [],
"current_prompt": None, # Will store the prompt object "current_prompt": None, # Will store the prompt object
"extra_data": None, # Will store the API extra_data for workflow metadata
"timestamp": time.time(), "timestamp": time.time(),
} }
) )
@@ -75,6 +76,11 @@ class MetadataRegistry:
# Store the prompt in the metadata for later relationship tracing # Store the prompt in the metadata for later relationship tracing
self.prompt_metadata[self.current_prompt_id]["current_prompt"] = prompt self.prompt_metadata[self.current_prompt_id]["current_prompt"] = prompt
def set_extra_data(self, extra_data):
"""Store the API extra_data (contains extra_pnginfo.workflow with node properties)"""
if self.current_prompt_id and self.current_prompt_id in self.prompt_metadata:
self.prompt_metadata[self.current_prompt_id]["extra_data"] = extra_data
def get_metadata(self, prompt_id=None): def get_metadata(self, prompt_id=None):
"""Get collected metadata for a prompt""" """Get collected metadata for a prompt"""
key = prompt_id if prompt_id is not None else self.current_prompt_id key = prompt_id if prompt_id is not None else self.current_prompt_id
@@ -122,20 +128,28 @@ class MetadataRegistry:
cache_key = f"{node_id}:{class_type}" cache_key = f"{node_id}:{class_type}"
# Check if this node type is relevant for metadata collection # Check if this node type is relevant for metadata collection
if class_type in NODE_EXTRACTORS: if class_type in NODE_EXTRACTORS or cache_key in self.node_cache:
# Check if we have cached metadata for this node # Check if we have cached metadata for this node
if cache_key in self.node_cache: if cache_key in self.node_cache:
cached_data = self.node_cache[cache_key] cached_data = self.node_cache[cache_key]
# Detect bypass (mode=4) / mute (mode=2) — these nodes
# were intentionally disabled and should not contribute
# overwrite values from a previous execution's cache.
node_mode = node_data.get("mode", 0)
node_is_disabled = node_mode in (2, 4)
# Apply cached metadata to the current metadata # Apply cached metadata to the current metadata
for category in self.metadata_categories: for category in self.metadata_categories:
if category == OVERWRITE and node_is_disabled:
continue
if category in cached_data and node_id in cached_data[category]: if category in cached_data and node_id in cached_data[category]:
if node_id not in metadata[category]: if node_id not in metadata[category]:
metadata[category][node_id] = cached_data[category][ metadata[category][node_id] = cached_data[category][
node_id node_id
] ]
def record_node_execution(self, node_id, class_type, inputs, outputs): def record_node_execution(self, node_id, class_type, inputs, outputs, return_types=None):
"""Record information about a node's execution""" """Record information about a node's execution"""
if not self.current_prompt_id: if not self.current_prompt_id:
return return
@@ -158,17 +172,18 @@ class MetadataRegistry:
# Extract node-specific metadata # Extract node-specific metadata
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor) extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
extractor.extract( if extractor is GenericNodeExtractor:
node_id, extractor.extract(node_id, processed_inputs, outputs,
processed_inputs, self.prompt_metadata[self.current_prompt_id],
outputs, return_types=return_types)
self.prompt_metadata[self.current_prompt_id], else:
) extractor.extract(node_id, processed_inputs, outputs,
self.prompt_metadata[self.current_prompt_id])
# Cache this node's metadata # Cache this node's metadata
self._cache_node_metadata(node_id, class_type) self._cache_node_metadata(node_id, class_type)
def update_node_execution(self, node_id, class_type, outputs): def update_node_execution(self, node_id, class_type, outputs, return_types=None):
"""Update node metadata with output information""" """Update node metadata with output information"""
if not self.current_prompt_id: if not self.current_prompt_id:
return return
@@ -179,9 +194,17 @@ class MetadataRegistry:
# Use the same extractor to update with outputs # Use the same extractor to update with outputs
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor) extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
if hasattr(extractor, "update"): if hasattr(extractor, "update"):
extractor.update( if extractor is GenericNodeExtractor:
node_id, processed_outputs, self.prompt_metadata[self.current_prompt_id] extractor.update(
) node_id, processed_outputs,
self.prompt_metadata[self.current_prompt_id],
return_types=return_types,
)
else:
extractor.update(
node_id, processed_outputs,
self.prompt_metadata[self.current_prompt_id],
)
# Update the cached metadata for this node # Update the cached metadata for this node
self._cache_node_metadata(node_id, class_type) self._cache_node_metadata(node_id, class_type)
+96 -4
View File
@@ -2,7 +2,8 @@ import json
import os import os
import re import re
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE
from .overwrite_utils import collect_overwrite_params
def _store_checkpoint_metadata(metadata, node_id, model_name): def _store_checkpoint_metadata(metadata, node_id, model_name):
@@ -31,10 +32,77 @@ class NodeMetadataExtractor:
pass pass
class GenericNodeExtractor(NodeMetadataExtractor): class GenericNodeExtractor(NodeMetadataExtractor):
"""Default extractor for nodes without specific handling""" """Fallback extractor with type-signature-based detection.
When a node is not in the NODE_EXTRACTORS registry, the hook layer
passes ``return_types`` from ``obj.RETURN_TYPES``:
* ``MODEL`` output: common input fields (ckpt_name, unet_name, etc.)
are checked for a model file name and stored as checkpoint metadata.
* ``CONDITIONING`` output: common text input fields are checked for
prompt text and stored as prompt metadata.
"""
# Input field names that carry a model path in loader-style nodes.
_MODEL_NAME_FIELDS = (
"ckpt_name", "unet_name", "model_path", "model_name", "gguf_name",
)
# Extensions used by checkpoint_scanner.py — only record values that look
# like real model filenames to avoid capturing unrelated string fields.
_MODEL_EXTENSIONS = {
".ckpt", ".pt", ".pt2", ".bin", ".pth", ".safetensors", ".pkl", ".sft", ".gguf",
}
# Input field names that may carry prompt text in encoder-style nodes.
_TEXT_FIELDS = ("text", "clip_l", "t5xxl", "prompt", "positive", "negative")
@staticmethod @staticmethod
def extract(node_id, inputs, outputs, metadata): def extract(node_id, inputs, outputs, metadata, return_types=None):
pass if return_types is None:
return
# — MODEL loader detection (checkpoint / UNET / GGUF) —
if "MODEL" in return_types or any("MODEL" in str(t) for t in return_types):
for field in GenericNodeExtractor._MODEL_NAME_FIELDS:
val = inputs.get(field)
if val and isinstance(val, str) and val.strip():
name = val.strip()
if not any(name.lower().endswith(ext) for ext in GenericNodeExtractor._MODEL_EXTENSIONS):
continue
_store_checkpoint_metadata(metadata, node_id, name)
return
# — CONDITIONING encoder detection (CLIPTextEncode, Flux, custom) —
if "CONDITIONING" in return_types or any("CONDITIONING" in str(t) for t in return_types):
text = None
for field in GenericNodeExtractor._TEXT_FIELDS:
val = inputs.get(field)
if val and isinstance(val, str) and val.strip():
text = val.strip()
break
if text:
prompt_data = metadata.setdefault(PROMPTS, {})
prompt_data[node_id] = {
"text": text,
"node_id": node_id,
}
@staticmethod
def update(node_id, outputs, metadata, return_types=None):
if return_types is None:
return
if "CONDITIONING" not in return_types and not any(
"CONDITIONING" in str(t) for t in return_types
):
return
if node_id not in metadata.get(PROMPTS, {}):
return
if outputs and isinstance(outputs, list) and len(outputs) > 0:
if isinstance(outputs[0], tuple) and len(outputs[0]) > 0:
cond = outputs[0][0]
if cond is not None:
metadata[PROMPTS][node_id]["conditioning"] = cond
class CheckpointLoaderExtractor(NodeMetadataExtractor): class CheckpointLoaderExtractor(NodeMetadataExtractor):
@staticmethod @staticmethod
@@ -1154,6 +1222,28 @@ class CR_ApplyControlNetStackExtractor(NodeMetadataExtractor):
metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive
metadata[PROMPTS][node_id]["negative_encoded"] = transformed_negative metadata[PROMPTS][node_id]["negative_encoded"] = transformed_negative
class MetadataOverwriteExtractor(NodeMetadataExtractor):
"""Extract manually specified metadata from MetadataOverwriteLM node.
Stores truthy input values under the OVERWRITE category so that
extract_generation_params can merge them over the inferred params.
"""
@staticmethod
def extract(node_id, inputs, outputs, metadata):
if not inputs:
return
overwrite_params = collect_overwrite_params(inputs)
if overwrite_params:
metadata.setdefault(OVERWRITE, {})
metadata[OVERWRITE][node_id] = {
"parameters": overwrite_params,
"node_id": node_id,
}
# Registry of node-specific extractors # Registry of node-specific extractors
# Keys are node class names # Keys are node class names
NODE_EXTRACTORS = { NODE_EXTRACTORS = {
@@ -1221,5 +1311,7 @@ NODE_EXTRACTORS = {
"CFGGuider": CFGGuiderExtractor, # Add CFGGuider "CFGGuider": CFGGuiderExtractor, # Add CFGGuider
# Image # Image
"VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor "VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor
# Metadata overwrite
"MetadataOverwriteLM": MetadataOverwriteExtractor,
# Add other nodes as needed # Add other nodes as needed
} }
+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
+117
View File
@@ -0,0 +1,117 @@
"""Create Hook LoRA (LoraManager) — multi-LoRA hook node compatible with ComfyUI's built-in hook pipeline.
Produces ``("HOOKS",)`` output that chains seamlessly with downstream hook consumers
(ConditioningSetProperties, SetHookKeyframes, CombineHooks, SetClipHooks, etc.).
"""
from __future__ import annotations
import logging
import os
from ..utils.utils import get_lora_info_absolute
from .utils import (
FlexibleOptionalInputType,
any_type,
apply_lora_syntax_format,
get_loras_list,
)
logger = logging.getLogger(__name__)
class CreateHookLoraLM:
NAME = "Create Hook LoRA (LoraManager)"
CATEGORY = "Lora Manager/hooks"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": (
"AUTOCOMPLETE_TEXT_LORAS",
{
"placeholder": "Search LoRAs to add...",
"tooltip": (
"Search and select LoRAs. Each LoRA gets its own "
"model/clip strength. Hooks chain with prev_hooks."
),
},
),
},
"optional": FlexibleOptionalInputType(any_type),
}
RETURN_TYPES = ("HOOKS", "STRING", "STRING")
RETURN_NAMES = ("HOOKS", "trigger_words", "active_loras")
FUNCTION = "create_hook"
def create_hook(self, text: str, **kwargs):
"""Create a HookGroup from the selected LoRAs, chained with prev_hooks.
Each active LoRA from the widget is loaded and wrapped in a WeightHook
via :func:`comfy.hooks.create_hook_lora`. All hooks are combined into a
single group and returned alongside trigger words and a human-readable
summary of the active LoRAs.
"""
del text # used by the frontend widget only
# Lazy imports: comfy is not available in CI/test environment at module level
import comfy.hooks # type: ignore # noqa: C0415
import comfy.utils # type: ignore # noqa: C0415
prev_hooks: comfy.hooks.HookGroup | None = kwargs.get("prev_hooks")
hook_group = prev_hooks.clone() if prev_hooks is not None else comfy.hooks.HookGroup()
all_trigger_words: list[str] = []
active_loras: list[tuple[str, float, float]] = []
for lora in get_loras_list(kwargs):
if not lora.get("active", False):
continue
lora_name = apply_lora_syntax_format(lora["name"])
model_strength = float(lora["strength"])
clip_strength = float(lora.get("clipStrength", model_strength))
# Skip useless no-op entries (both strengths are zero)
if model_strength == 0.0 and clip_strength == 0.0:
continue
lora_path, trigger_words = get_lora_info_absolute(lora_name)
if not lora_path or not os.path.isfile(lora_path):
logger.warning("LoRA '%s' not found — skipping", lora_name)
continue
try:
lora_weights = comfy.utils.load_torch_file(lora_path, safe_load=True)
lora_hooks = comfy.hooks.create_hook_lora(
lora=lora_weights,
strength_model=model_strength,
strength_clip=clip_strength,
)
except Exception:
logger.exception("Failed to load LoRA '%s' — skipping", lora_name)
continue
hook_group = hook_group.clone_and_combine(lora_hooks)
active_loras.append((lora_name, model_strength, clip_strength))
all_trigger_words.extend(trigger_words)
# Format trigger words (group mode separator)
trigger_words_text = ",, ".join(all_trigger_words) if all_trigger_words else ""
# Format active LoRAs summary
formatted_loras = []
for name, model_s, clip_s in active_loras:
if abs(model_s - clip_s) > 0.001:
formatted_loras.append(
f"<lora:{name}:{model_s}:{clip_s}>"
)
else:
formatted_loras.append(f"<lora:{name}:{model_s}>")
active_loras_text = " ".join(formatted_loras)
return (hook_group, trigger_words_text, active_loras_text)
+86 -10
View File
@@ -1,26 +1,102 @@
from __future__ import annotations
import inspect
import re
from typing import Any
_STACK_INPUT_PATTERN = re.compile(r"^lora_stack(?:_([ab])|(\d+))$")
def _is_stack_input(name: str) -> bool:
return bool(_STACK_INPUT_PATTERN.match(name))
def _stack_slot_number(name: str) -> int:
"""Numeric slot used to order stack inputs; legacy a/b map to 1/2."""
match = _STACK_INPUT_PATTERN.match(name)
if not match:
return -1
letter, digits = match.group(1), match.group(2)
if digits is not None:
return int(digits)
return 1 if letter == "a" else 2
class _LoraStackOptionalInputs:
"""Lookup that preserves explicit optional inputs and dynamic lora_stack slots."""
def __init__(self, explicit_inputs: dict[str, tuple[str, dict[str, Any]]]) -> None:
self._explicit_inputs = explicit_inputs
def __contains__(self, item: object) -> bool:
if not isinstance(item, str):
return False
return item in self._explicit_inputs or _is_stack_input(item)
def __getitem__(self, key: str) -> tuple[str, dict[str, Any]]:
if key in self._explicit_inputs:
return self._explicit_inputs[key]
if _is_stack_input(key):
return (
"LORA_STACK",
{
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
},
)
raise KeyError(key)
class LoraStackCombinerLM: class LoraStackCombinerLM:
NAME = "Lora Stack Combiner (LoraManager)" NAME = "Lora Stack Combiner (LoraManager)"
CATEGORY = "Lora Manager/stackers" CATEGORY = "Lora Manager/stackers"
DESCRIPTION = (
"Combines multiple LoRA stacks into a single stack. "
"Supports dynamic inputs: connect a stack to add more inputs."
)
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
optional_inputs: dict[str, tuple[str, dict[str, Any]]] = {
"lora_stack1": (
"LORA_STACK",
{
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
},
),
"lora_stack2": (
"LORA_STACK",
{
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
},
),
}
stack = inspect.stack()
if len(stack) > 2 and stack[2].function == "get_input_info":
optional_inputs = _LoraStackOptionalInputs(optional_inputs) # type: ignore[assignment]
return { return {
"required": { "required": {},
"lora_stack_a": ("LORA_STACK",), "optional": optional_inputs,
"lora_stack_b": ("LORA_STACK",),
},
} }
RETURN_TYPES = ("LORA_STACK",) RETURN_TYPES = ("LORA_STACK",)
RETURN_NAMES = ("LORA_STACK",) RETURN_NAMES = ("LORA_STACK",)
FUNCTION = "combine_stacks" FUNCTION = "combine_stacks"
def combine_stacks(self, lora_stack_a, lora_stack_b): def combine_stacks(self, lora_stack1=None, lora_stack2=None, **kwargs):
combined_stack = [] stacks = {
"lora_stack1": lora_stack1,
"lora_stack2": lora_stack2,
}
for key, value in kwargs.items():
if _is_stack_input(key) and value is not None:
stacks[key] = value
if lora_stack_a: combined_stack = []
combined_stack.extend(lora_stack_a) for key in sorted(stacks, key=_stack_slot_number):
if lora_stack_b: stack = stacks[key]
combined_stack.extend(lora_stack_b) if stack:
combined_stack.extend(stack)
return (combined_stack,) return (combined_stack,)
+169
View File
@@ -0,0 +1,169 @@
"""Metadata Overwrite node — allows users to manually specify generation parameters
that override the automatically collected/inferred 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 CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL
from ..metadata_collector.overwrite_utils import collect_overwrite_params
class MetadataOverwriteLM:
NAME = "Metadata Overwrite (LoraManager)"
CATEGORY = "Lora Manager/utils"
DESCRIPTION = (
"Manually specify generation parameters to override automatically collected "
"metadata. Only filled/connected inputs will take effect — empty defaults "
"are ignored."
)
@classmethod
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"optional": {
"prompt": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": "Positive prompt. Only overwrites when non-empty.",
},
),
"negative_prompt": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": "Negative prompt. Only overwrites when non-empty.",
},
),
"seed": (
"INT",
{
"default": 0,
"min": 0,
"max": 0xFFFFFFFFFFFFFFFF,
"control_after_generate": False,
"tooltip": "Seed value. Only overwrites when > 0.",
},
),
"steps": (
"INT",
{
"default": 0,
"min": 0,
"max": 10000,
"tooltip": "Number of steps. Only overwrites when > 0.",
},
),
"cfg_scale": (
"FLOAT",
{
"default": 0.0,
"min": 0.0,
"max": 100.0,
"tooltip": "CFG scale. Only overwrites when > 0.",
},
),
"sampler": (
"STRING",
{
"default": "",
"tooltip": "Sampler name. Only overwrites when non-empty.",
},
),
"scheduler": (
"STRING",
{
"default": "",
"tooltip": "Scheduler name. Only overwrites when non-empty.",
},
),
"model": (
"STRING,MODEL",
{
"default": "",
"widgetType": "STRING",
"tooltip": (
"The checkpoint or diffusion model (UNet) used "
"for generation. Fill in the name manually or "
"connect a MODEL output — the model name is then "
"extracted automatically. Only overwrites when "
"non-empty."
),
},
),
"loras": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": (
"LoRA syntax, e.g. <lora:name:strength> "
"or <lora:name:model_strength:clip_strength>, "
"separated by spaces. Only overwrites when non-empty."
),
},
),
"size": (
"STRING",
{
"default": "",
"tooltip": (
"Image size in WIDTHxHEIGHT format (e.g. 512x768). "
"Only overwrites when non-empty."
),
},
),
"clip_skip": (
"INT",
{
"default": _CLIP_SKIP_SENTINEL,
"min": -25,
"max": 24,
"tooltip": (
"Clip skip (ComfyUI: -24..-1, A1111: 1+). "
"Default -25 means not set — any other value "
"overwrites."
),
},
),
"additional_data": (
"STRING",
{
"default": "",
"multiline": True,
"tooltip": (
"Additional data to embed in the image metadata. "
"Inserted between Clip skip and Model hash in the "
"A1111-compatible parameters string. "
'Example: "Copyright": "Some license info"'
),
},
),
},
}
RETURN_TYPES = ("METADATA",)
RETURN_NAMES = ("metadata",)
FUNCTION = "collect_metadata"
OUTPUT_NODE = True
def collect_metadata(self, **kwargs: Any) -> tuple[dict[str, Any]]:
"""Collect non-default input values into a metadata dict.
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.
"""
return (collect_overwrite_params(kwargs),)
+346 -127
View File
@@ -16,6 +16,156 @@ from PIL import Image, PngImagePlugin
import piexif import piexif
import logging import logging
# Civitai-compatible sampler name mapping: ComfyUI internal → A1111 display name
CIVITAI_SAMPLER_MAP = {
"euler": "Euler",
"euler_ancestral": "Euler a",
"lms": "LMS",
"heun": "Heun",
"dpm_2": "DPM2",
"dpm_2_ancestral": "DPM2 a",
"dpmpp_2s_ancestral": "DPM++ 2S a",
"dpmpp_2m": "DPM++ 2M",
"dpmpp_sde": "DPM++ SDE",
"dpmpp_sde_gpu": "DPM++ SDE",
"dpmpp_2m_sde": "DPM++ 2M SDE",
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
"dpmpp_3m_sde": "DPM++ 3M SDE",
"dpm_fast": "DPM fast",
"dpm_adaptive": "DPM adaptive",
"ddim": "DDIM",
"plms": "PLMS",
"uni_pc_bh2": "UniPC",
"uni_pc": "UniPC",
"lcm": "LCM",
}
# Base model display name → AIR URN slug
# Sourced from civitai source: src/shared/constants/basemodel.constants.ts
BASE_MODEL_AIR_SLUG = {
# Stable Diffusion family
"SD 1.4": "sd1",
"SD 1.5": "sd1",
"SD 1.5 LCM": "sd1",
"SD 1.5 Hyper": "sd1",
"SD 2.0": "sd2",
"SD 2.0 768": "sd2",
"SD 2.1": "sd2",
"SD 2.1 768": "sd2",
"SD 2.1 Unclip": "sd2",
"SD 3.0": "sd3",
"SD 3.5": "sd35",
"SD 3.5 Large": "sd35",
"SD 3.5 Large Turbo": "sd35",
"SD 3.5 Medium": "sd35",
"SDXL 0.9": "sdxl",
"SDXL 1.0": "sdxl",
"SDXL 1.0 LCM": "sdxl",
"SDXL Lightning": "sdxl",
"SDXL Hyper": "sdxl",
"SDXL Turbo": "sdxl",
"SDXL Distilled": "sdxldistilled",
"Stable Cascade": "scascade",
"Stable Video Diffusion": "svd",
"SVD": "svd",
"SVD XT": "svdxt",
# SDXL community fine-tunes
"Pony": "pony",
"Pony Diffusion": "pony",
"Illustrious": "illustrious",
"NoobAI": "noobai",
"Animagine": "illustrious",
# Flux family
"Flux.1": "flux1",
"Flux.1 D": "flux1",
"Flux.1 S": "flux1",
"Flux.1 Krea": "fluxkrea",
"Flux.1 Kontext": "flux1kontext",
"Flux.2": "flux2",
"Flux.2 D": "flux2",
"Flux.2 Klein 9B": "flux2klein_9b",
"Flux.2 Klein 9B Base": "flux2klein_9b_base",
"Flux.2 Klein 4B": "flux2klein_4b",
"Flux.2 Klein 4B Base": "flux2klein_4b_base",
# Other image models (sorted alphabetically)
"AuraFlow": "auraflow",
"Chroma": "chroma",
"HiDream": "hidream",
"HiDream-O1": "hidream-o1",
"Hunyuan DiT": "hydit1",
"Hunyuan Video": "hyv1",
"Kolors": "kolors",
"Lumina": "lumina",
"Mochi": "mochi",
"ODOR": "odor",
"PixArt Alpha": "pixarta",
"PixArt Sigma": "pixarte",
"Playground v2": "playgroundv2",
"Playground v2.5": "playgroundv2",
"Pony Diffusion V7": "ponyv7",
# Video models
"CogVideoX": "cogvideox",
"LTX Video": "ltxv",
"LTX Video 2": "ltxv2",
"LTX Video 2.3": "ltxv23",
"Wan Video": "wanvideo",
"Wan Video 1.3B T2V": "wanvideo_13b_t2v",
"Wan Video 14B T2V": "wanvideo_14b_t2v",
"Wan Video 14B I2V 480p": "wanvideo_14b_i2v_480p",
"Wan Video 14B I2V 720p": "wanvideo_14b_i2v_720p",
# Third-party / proprietary image models
"Boogu": "boogu",
"Ernie": "ernie",
"Grok": "grok",
"HappyHorse": "happyhorse",
"Ideogram": "ideogram",
"Ideogram 4.0": "ideogram",
"Imagen": "imagen4",
"Imagen 4": "imagen4",
"Krea": "krea2",
"Krea 2": "krea2",
"Lens": "lens",
"MAI": "mai",
"Nano Banana": "nanobanana",
"OpenAI": "openai",
"Reve": "reve",
"Reve 2": "reve",
"Reve 2.1": "reve",
"Seedream": "seedream",
"Sora": "sora2",
"Sora 2": "sora2",
"Veo": "veo3",
"Veo 2": "veo3",
"Veo 3": "veo3",
"ZImageTurbo": "zimageturbo",
"ZImageBase": "zimagebase",
"ZImage": "zimagebase",
# Third-party video models
"Hailuo by MiniMax": "minimax",
"Haiper": "haiper",
"Kling": "kling",
"Lightricks": "lightricks",
"Seedance": "seedance",
"Vidu": "vidu",
# Qwen family
"Qwen": "qwen",
"Qwen 2": "qwen2",
# Anima
"Anima": "anima",
# Special
"Upscaler": "upscaler",
"Other": "other",
}
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -70,11 +220,29 @@ class SaveImageLM:
"tooltip": "Compression quality for JPEG and lossy WebP formats (1-100). Higher values mean better quality but larger files.", "tooltip": "Compression quality for JPEG and lossy WebP formats (1-100). Higher values mean better quality but larger files.",
}, },
), ),
"webp_method": (
"INT",
{
"default": 6,
"min": 0,
"max": 6,
"tooltip": "WebP compression method (0-6). 0=fastest/largest, 6=slowest/smallest. Only applies when file_format is 'webp'.",
},
),
"jpeg_subsampling": (
"INT",
{
"default": 0,
"min": 0,
"max": 2,
"tooltip": "JPEG chroma subsampling level. 0=4:4:4 (best quality), 1=4:2:2, 2=4:2:0 (smallest files). Only applies when file_format is 'jpeg'.",
},
),
"embed_workflow": ( "embed_workflow": (
"BOOLEAN", "BOOLEAN",
{ {
"default": False, "default": False,
"tooltip": "Embeds the complete workflow data into the image metadata. Only works with PNG and WebP formats.", "tooltip": "When enabled, saved images store the complete workflow. Drag the image back into ComfyUI to restore the original node graph. PNG and WebP only.",
}, },
), ),
"save_with_metadata": ( "save_with_metadata": (
@@ -142,148 +310,194 @@ class SaveImageLM:
return None return None
def format_metadata(self, metadata_dict): def _resolve_model_cache_entry(self, scanner_type: str, name: str):
"""Format metadata in the requested format similar to userComment example""" """Resolve model hash, civitai metadata, and base_model from scanner cache.
if not metadata_dict: Returns (hash_str, civitai_dict, base_model_str). All values are empty defaults when not found."""
return "" scanner = ServiceRegistry.get_service_sync(scanner_type)
if scanner is None or not name:
return "", {}, ""
# Helper function to only add parameter if value is not None entry = self._get_cached_model_by_name(scanner, name)
def add_param_if_not_none(param_list, label, value): if entry is None:
if value is not None: basename = os.path.splitext(os.path.basename(name))[0]
param_list.append(f"{label}: {value}") hash_val = scanner.get_hash_by_filename(basename)
return (hash_val or "").lower(), {}, ""
hash_val = (entry.get("sha256") or "").lower()
civitai = entry.get("civitai") or {}
base_model = entry.get("base_model") or ""
return hash_val, civitai, base_model
@staticmethod
def _get_civitai_sampler_name(sampler_name: str, scheduler: str) -> str:
if sampler_name in CIVITAI_SAMPLER_MAP:
civitai_name = CIVITAI_SAMPLER_MAP[sampler_name]
if scheduler == "karras":
civitai_name += " Karras"
elif scheduler == "exponential":
civitai_name += " Exponential"
return civitai_name
else:
if scheduler and scheduler != "normal":
return f"{sampler_name}_{scheduler}"
return sampler_name
@staticmethod
def _build_air_string(base_model: str, model_type: str, model_id: int, version_id: int) -> str:
slug = BASE_MODEL_AIR_SLUG.get(base_model, "other")
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:
"""Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources."""
if not metadata_dict: return ""
# Extract the prompt and negative prompt
prompt = metadata_dict.get("prompt", "") prompt = metadata_dict.get("prompt", "")
negative_prompt = metadata_dict.get("negative_prompt", "") negative_prompt = metadata_dict.get("negative_prompt", "")
steps = metadata_dict.get("steps")
# Extract loras from the prompt if present cfg = metadata_dict.get("guidance")
if cfg is None:
cfg = metadata_dict.get("cfg_scale")
if cfg is None:
cfg = metadata_dict.get("cfg")
seed = metadata_dict.get("seed")
size = metadata_dict.get("size")
sampler = metadata_dict.get("sampler") or ""
scheduler = metadata_dict.get("scheduler") or "normal"
checkpoint = metadata_dict.get("checkpoint") or ""
loras_text = metadata_dict.get("loras", "") loras_text = metadata_dict.get("loras", "")
lora_hashes = {} clip_skip = metadata_dict.get("clip_skip")
# If loras are found, add them on a new line after the prompt # Parse LoRA entries from <lora:name:strength> format
lora_entries: list[tuple[str, float]] = []
if loras_text: if loras_text:
prompt_with_loras = f"{prompt}\n{loras_text}" for match in re.findall(r"<lora:([^:]+):([^>]+)>", loras_text):
lora_name, strength_str = match
try:
strength = float(strength_str)
except (ValueError, TypeError):
strength = 1.0
lora_entries.append((lora_name, strength))
# Extract lora names from the format <lora:name:strength> # Resolve checkpoint hash and Civitai data from local cache
lora_matches = re.findall(r"<lora:([^:]+):([^>]+)>", loras_text) ckpt_hash, ckpt_civitai, ckpt_base_model = "", {}, ""
ckpt_display_name = ""
if checkpoint:
ckpt_hash, ckpt_civitai, ckpt_base_model = self._resolve_model_cache_entry(
"checkpoint_scanner", checkpoint
)
ckpt_display_name = os.path.splitext(os.path.basename(checkpoint))[0]
# Get hash for each lora # Resolve LoRA hash and Civitai data from local cache
for lora_name, strength in lora_matches: loras_data: list[dict] = []
hash_value = self.get_lora_hash(lora_name) for lora_name, strength in lora_entries:
if hash_value: lora_hash, lora_civitai, lora_base_model = self._resolve_model_cache_entry(
lora_hashes[lora_name] = hash_value "lora_scanner", lora_name
else: )
prompt_with_loras = prompt loras_data.append({
"name": lora_name,
"strength": strength,
"hash": lora_hash,
"civitai": lora_civitai,
"base_model": lora_base_model,
})
# Format the first part (prompt and loras) # Build Hashes JSON (A1111 / Civitai standard format)
metadata_parts = [prompt_with_loras] hashes: dict[str, str] = {}
if ckpt_hash:
hashes["model"] = ckpt_hash[:10].upper()
for lora in loras_data:
if lora["hash"]:
hashes[f"LORA:{lora['name']}"] = lora["hash"][:10].upper()
# Add negative prompt # Build Civitai resources JSON array
civitai_resources: list[dict] = []
if ckpt_civitai.get("id", 0) > 0:
ckpt_resource: dict = {}
ckpt_type = (ckpt_civitai.get("model") or {}).get("type", "Checkpoint")
model_id = ckpt_civitai.get("modelId", 0)
version_id = ckpt_civitai.get("id", 0)
if model_id and version_id:
ckpt_resource["air"] = self._build_air_string(
ckpt_base_model, ckpt_type, int(model_id), int(version_id)
)
elif version_id:
ckpt_resource["modelVersionId"] = int(version_id)
if ckpt_civitai.get("name"):
ckpt_resource["versionName"] = ckpt_civitai["name"]
if ckpt_resource:
civitai_resources.append(ckpt_resource)
for lora in loras_data:
lora_civitai = lora["civitai"]
if not lora_civitai or lora_civitai.get("id", 0) <= 0:
continue
lora_resource: dict = {"weight": lora["strength"]}
lora_type = (lora_civitai.get("model") or {}).get("type", "LORA")
model_id = lora_civitai.get("modelId", 0)
version_id = lora_civitai.get("id", 0)
if model_id and version_id:
lora_resource["air"] = self._build_air_string(
lora["base_model"], lora_type, int(model_id), int(version_id)
)
elif version_id:
lora_resource["modelVersionId"] = int(version_id)
if lora_civitai.get("name"):
lora_resource["versionName"] = lora_civitai["name"]
civitai_resources.append(lora_resource)
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 [""]
if negative_prompt: if negative_prompt:
metadata_parts.append(f"Negative prompt: {negative_prompt}") lines.append(f"Negative prompt: {negative_prompt}")
# Format the second part (generation parameters) params: list[str] = []
params = [] if steps is not None:
params.append(f"Steps: {steps}")
# Add standard parameters in the correct order
if "steps" in metadata_dict:
add_param_if_not_none(params, "Steps", metadata_dict.get("steps"))
# Combine sampler and scheduler information
sampler_name = None
scheduler_name = None
if "sampler" in metadata_dict:
sampler = metadata_dict.get("sampler")
# Convert ComfyUI sampler names to user-friendly names
sampler_mapping = {
"euler": "Euler",
"euler_ancestral": "Euler a",
"dpm_2": "DPM2",
"dpm_2_ancestral": "DPM2 a",
"heun": "Heun",
"dpm_fast": "DPM fast",
"dpm_adaptive": "DPM adaptive",
"lms": "LMS",
"dpmpp_2s_ancestral": "DPM++ 2S a",
"dpmpp_sde": "DPM++ SDE",
"dpmpp_sde_gpu": "DPM++ SDE",
"dpmpp_2m": "DPM++ 2M",
"dpmpp_2m_sde": "DPM++ 2M SDE",
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
"ddim": "DDIM",
}
sampler_name = sampler_mapping.get(sampler, sampler)
if "scheduler" in metadata_dict:
scheduler = metadata_dict.get("scheduler")
scheduler_mapping = {
"normal": "Simple",
"karras": "Karras",
"exponential": "Exponential",
"sgm_uniform": "SGM Uniform",
"sgm_quadratic": "SGM Quadratic",
}
scheduler_name = scheduler_mapping.get(scheduler, scheduler)
# Add combined sampler and scheduler information
if sampler_name: if sampler_name:
if scheduler_name: if scheduler_name:
params.append(f"Sampler: {sampler_name} {scheduler_name}") params.append(f"Sampler: {sampler_name} {scheduler_name}")
else: else:
params.append(f"Sampler: {sampler_name}") 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 is not None:
try:
params.append(f"Clip skip: {abs(int(clip_skip))}")
except (ValueError, TypeError):
pass
additional_data = metadata_dict.get("additional_data", "")
if additional_data:
params.append(additional_data)
if ckpt_hash:
params.append(f"Model hash: {ckpt_hash[:10].upper()}")
if ckpt_display_name:
params.append(f"Model: {ckpt_display_name}")
if hashes:
params.append(f"Hashes: {json.dumps(hashes, separators=(',', ':'))}")
params.append("Version: ComfyUI")
if civitai_resources:
params.append(
f"Civitai resources: {json.dumps(civitai_resources, separators=(',', ':'))}"
)
# CFG scale (Use guidance if available, otherwise fall back to cfg_scale or cfg) lines.append(", ".join(params))
if "guidance" in metadata_dict: return "\n".join(lines)
add_param_if_not_none(params, "CFG scale", metadata_dict.get("guidance"))
elif "cfg_scale" in metadata_dict:
add_param_if_not_none(params, "CFG scale", metadata_dict.get("cfg_scale"))
elif "cfg" in metadata_dict:
add_param_if_not_none(params, "CFG scale", metadata_dict.get("cfg"))
# Seed
if "seed" in metadata_dict:
add_param_if_not_none(params, "Seed", metadata_dict.get("seed"))
# Size
if "size" in metadata_dict:
add_param_if_not_none(params, "Size", metadata_dict.get("size"))
# Model info
if "checkpoint" in metadata_dict:
# Ensure checkpoint is a string before processing
checkpoint = metadata_dict.get("checkpoint")
if checkpoint is not None:
# Get model hash
model_hash = self.get_checkpoint_hash(checkpoint)
# Extract basename without path
checkpoint_name = os.path.basename(checkpoint)
# Remove extension if present
checkpoint_name = os.path.splitext(checkpoint_name)[0]
# Add model hash if available
if model_hash:
params.append(
f"Model hash: {model_hash[:10]}, Model: {checkpoint_name}"
)
else:
params.append(f"Model: {checkpoint_name}")
# Add LoRA hashes if available
if lora_hashes:
lora_hash_parts = []
for lora_name, hash_value in lora_hashes.items():
lora_hash_parts.append(f"{lora_name}: {hash_value[:10]}")
if lora_hash_parts:
params.append(f'Lora hashes: "{", ".join(lora_hash_parts)}"')
# Combine all parameters with commas
metadata_parts.append(", ".join(params))
# Join all parts with a new line
return "\n".join(metadata_parts)
# credit to nkchocoai # credit to nkchocoai
# Add format_filename method to handle pattern substitution # Add format_filename method to handle pattern substitution
@@ -573,6 +787,8 @@ class SaveImageLM:
extra_pnginfo=None, extra_pnginfo=None,
lossless_webp=True, lossless_webp=True,
quality=100, quality=100,
webp_method=6,
jpeg_subsampling=0,
embed_workflow=False, embed_workflow=False,
save_with_metadata=True, save_with_metadata=True,
add_counter_to_filename=True, add_counter_to_filename=True,
@@ -627,15 +843,14 @@ class SaveImageLM:
elif file_format == "jpeg": elif file_format == "jpeg":
file = base_filename + ".jpg" file = base_filename + ".jpg"
file_extension = ".jpg" file_extension = ".jpg"
save_kwargs = {"quality": quality, "optimize": True} save_kwargs = {"quality": quality, "optimize": True, "subsampling": jpeg_subsampling}
elif file_format == "webp": elif file_format == "webp":
file = base_filename + ".webp" file = base_filename + ".webp"
file_extension = ".webp" file_extension = ".webp"
# Add optimization param to control performance
save_kwargs = { save_kwargs = {
"quality": quality, "quality": quality,
"lossless": lossless_webp, "lossless": lossless_webp,
"method": 0, "method": webp_method,
} }
else: else:
raise ValueError(f"Unsupported file format: {file_format}") raise ValueError(f"Unsupported file format: {file_format}")
@@ -722,6 +937,8 @@ class SaveImageLM:
extra_pnginfo=None, extra_pnginfo=None,
lossless_webp=True, lossless_webp=True,
quality=100, quality=100,
webp_method=6,
jpeg_subsampling=0,
embed_workflow=False, embed_workflow=False,
save_with_metadata=True, save_with_metadata=True,
add_counter_to_filename=True, add_counter_to_filename=True,
@@ -751,6 +968,8 @@ class SaveImageLM:
extra_pnginfo, extra_pnginfo,
lossless_webp, lossless_webp,
quality, quality,
webp_method,
jpeg_subsampling,
embed_workflow, embed_workflow,
save_with_metadata, save_with_metadata,
add_counter_to_filename, add_counter_to_filename,
+21
View File
@@ -7,6 +7,21 @@ from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_c
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def _reload_gguf_unet(
unet_path: str, weight_dtype: str, disable_dynamic: bool = False
) -> object:
"""Reload a GGUF diffusion model from disk (cached_patcher_init factory).
Mirrors the GGUF branch of UNETLoaderLM.load_unet so ModelPatcher
deepclone/dynamic machinery can rebuild GGUF models with the correct
GGMLOps. ``disable_dynamic`` is accepted for signature compatibility
with core ComfyUI loaders.
"""
loader = UNETLoaderLM()
model, = loader._load_gguf_unet(unet_path, unet_path, weight_dtype)
return model
class UNETLoaderLM: class UNETLoaderLM:
"""UNET Loader with support for extra folder paths """UNET Loader with support for extra folder paths
@@ -196,6 +211,12 @@ class UNETLoaderLM:
# Wrap with GGUFModelPatcher # Wrap with GGUFModelPatcher
model = GGUFModelPatcher.clone(model) model = GGUFModelPatcher.clone(model)
# Register a reload factory so the MODEL carries its source path
# (cached_patcher_init) like core ComfyUI loaders do — required
# for model-name extraction downstream and for ModelPatcher
# deepclone/dynamic machinery.
model.cached_patcher_init = (_reload_gguf_unet, (unet_path, weight_dtype))
return (model,) return (model,)
except Exception as e: except Exception as e:
+42 -3
View File
@@ -1562,6 +1562,11 @@ class SettingsHandler:
{"success": False, "error": validation_error} {"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 ( if value == "__DELETE__" and key in (
"proxy_username", "proxy_username",
"proxy_password", "proxy_password",
@@ -2585,6 +2590,8 @@ class ModelLibraryHandler:
status=400, status=400,
) )
cursor = request.query.get("cursor")
metadata_provider = await self._metadata_provider_factory() metadata_provider = await self._metadata_provider_factory()
if not metadata_provider: if not metadata_provider:
return web.json_response( return web.json_response(
@@ -2593,7 +2600,7 @@ class ModelLibraryHandler:
) )
try: try:
models = await metadata_provider.get_user_models(username) result = await metadata_provider.get_user_models(username, cursor)
except NotImplementedError: except NotImplementedError:
return web.json_response( return web.json_response(
{ {
@@ -2603,14 +2610,35 @@ class ModelLibraryHandler:
status=501, status=501,
) )
if models is None: if result is None:
return web.json_response( return web.json_response(
{"success": False, "error": "Failed to fetch user models"}, {"success": False, "error": "Failed to fetch user models"},
status=502, status=502,
) )
if isinstance(result, dict):
models = result.get("items")
next_cursor = result.get("nextCursor")
else:
# Defensive: tolerate providers that still return a raw list
models = result
next_cursor = None
if not isinstance(models, list): if not isinstance(models, list):
models = [] models = []
if next_cursor is not None and not isinstance(next_cursor, str):
next_cursor = str(next_cursor)
estimated_total = None
if cursor is None:
get_count = getattr(metadata_provider, "get_creator_model_count", None)
if get_count is not None:
try:
estimated_total = await get_count(username)
except Exception: # best-effort only
estimated_total = None
if not isinstance(estimated_total, int):
estimated_total = None
lora_scanner = await self._service_registry.get_lora_scanner() lora_scanner = await self._service_registry.get_lora_scanner()
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner() checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
@@ -2630,6 +2658,7 @@ class ModelLibraryHandler:
versions: list[dict] = [] versions: list[dict] = []
history_service = await self._get_download_history_service() history_service = await self._get_download_history_service()
model_ids: list[int] = [] model_ids: list[int] = []
model_count = 0
for model in models: for model in models:
try: try:
model_ids.append(int(model.get("id"))) model_ids.append(int(model.get("id")))
@@ -2663,6 +2692,8 @@ class ModelLibraryHandler:
if model_type not in normalized_allowed_types: if model_type not in normalized_allowed_types:
continue continue
model_count += 1
scanner = type_scanner_map.get(model_type) scanner = type_scanner_map.get(model_type)
if scanner is None: if scanner is None:
return web.json_response( return web.json_response(
@@ -2728,7 +2759,15 @@ class ModelLibraryHandler:
) )
return web.json_response( return web.json_response(
{"success": True, "username": username, "versions": versions} {
"success": True,
"username": username,
"versions": versions,
"modelCount": model_count,
"nextCursor": next_cursor,
"hasMore": next_cursor is not None,
"estimatedTotal": estimated_total,
}
) )
except Exception as exc: # pragma: no cover - defensive logging except Exception as exc: # pragma: no cover - defensive logging
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True) logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
+36 -2
View File
@@ -394,12 +394,14 @@ class ModelListingHandler:
) )
# View-local-versions filter: show all local versions of a specific model # 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") civitai_model_id = request.query.get("civitai_model_id")
if civitai_model_id is not None: if civitai_model_id is not None:
try: try:
civitai_model_id = int(civitai_model_id) civitai_model_id = int(civitai_model_id)
except (TypeError, ValueError): except (TypeError, ValueError):
civitai_model_id = None # Keep as string — could be an HF group key (e.g. "hf:user/repo")
pass
return { return {
"page": page, "page": page,
@@ -537,6 +539,7 @@ class ModelManagementHandler:
# Update model_data with new hash # Update model_data with new hash
model_data["sha256"] = sha256 model_data["sha256"] = sha256
model_data["hash_status"] = "completed" model_data["hash_status"] = "completed"
hash_status = "completed"
else: else:
return web.json_response( return web.json_response(
{"success": False, "error": "No SHA256 hash found"}, status=400 {"success": False, "error": "No SHA256 hash found"}, status=400
@@ -544,6 +547,32 @@ class ModelManagementHandler:
await MetadataManager.hydrate_model_data(model_data) await MetadataManager.hydrate_model_data(model_data)
# hydrate_model_data replaces model_data with .metadata.json content,
# which may lack sha256. Restore from cache and persist the fix.
if not model_data.get("sha256"):
if sha256:
model_data["sha256"] = sha256
model_data["hash_status"] = model_data.get("hash_status", hash_status)
data_to_save = model_data.copy()
data_to_save.pop("folder", None)
await MetadataManager.save_metadata(file_path, data_to_save)
else:
sha256 = await calculate_sha256(file_path)
if sha256:
model_data["sha256"] = sha256.lower()
model_data["hash_status"] = "completed"
data_to_save = model_data.copy()
data_to_save.pop("folder", None)
await MetadataManager.save_metadata(file_path, data_to_save)
else:
return web.json_response(
{
"success": False,
"error": "Failed to compute SHA256 hash for model",
},
status=500,
)
success, error = await self._metadata_sync.fetch_and_update_model( success, error = await self._metadata_sync.fetch_and_update_model(
sha256=model_data["sha256"], sha256=model_data["sha256"],
file_path=file_path, file_path=file_path,
@@ -566,7 +595,12 @@ class ModelManagementHandler:
{"success": False, "error": OFFLINE_FRIENDLY_MESSAGE}, {"success": False, "error": OFFLINE_FRIENDLY_MESSAGE},
status=503, status=503,
) )
self._logger.error("Error fetching from CivitAI: %s", exc, exc_info=True) self._logger.error(
"Error fetching from CivitAI for %s: %s",
locals().get("file_path", "unknown"),
exc,
exc_info=True,
)
return web.json_response({"success": False, "error": str(exc)}, status=500) return web.json_response({"success": False, "error": str(exc)}, status=500)
async def relink_civitai(self, request: web.Request) -> web.Response: async def relink_civitai(self, request: web.Request) -> web.Response:
+305 -37
View File
@@ -38,6 +38,84 @@ def _clean_excludes() -> List[str]:
return excludes 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: class UpdateRoutes:
"""Routes for handling plugin update checks""" """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/check-updates', UpdateRoutes.check_updates)
app.router.add_get('/api/lm/version-info', UpdateRoutes.get_version_info) 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/perform-update', UpdateRoutes.perform_update)
app.router.add_post('/api/lm/switch-channel', UpdateRoutes.switch_channel)
@staticmethod @staticmethod
async def check_updates(request): async def check_updates(request):
@@ -65,10 +144,17 @@ class UpdateRoutes:
# Fetch remote version from GitHub # Fetch remote version from GitHub
if nightly: if nightly:
remote_version, changelog = await UpdateRoutes._get_nightly_version() local_hash = git_info.get('short_hash', '')
releases = None 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: else:
remote_version, changelog, releases = await UpdateRoutes._get_remote_version() remote_version, changelog, releases = await UpdateRoutes._get_remote_version()
behind_by = 0
commit_date = ''
# Compare versions # Compare versions
if nightly: if nightly:
@@ -81,6 +167,10 @@ class UpdateRoutes:
remote_version.replace('v', '') 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 = { response_data = {
'success': True, 'success': True,
'current_version': local_version, 'current_version': local_version,
@@ -88,13 +178,13 @@ class UpdateRoutes:
'update_available': update_available, 'update_available': update_available,
'changelog': changelog, 'changelog': changelog,
'git_info': git_info, '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) return web.json_response(response_data)
except NETWORK_EXCEPTIONS as e: except NETWORK_EXCEPTIONS as e:
@@ -126,9 +216,14 @@ class UpdateRoutes:
# Format: version-short_hash # Format: version-short_hash
version_string = f"{local_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({ return web.json_response({
'success': True, 'success': True,
'version': version_string 'version': version_string,
'has_git': has_git
}) })
except Exception as e: except Exception as e:
@@ -156,20 +251,22 @@ class UpdateRoutes:
if os.path.exists(settings_path): if os.path.exists(settings_path):
with open(settings_path, 'r', encoding='utf-8') as f: with open(settings_path, 'r', encoding='utf-8') as f:
settings_backup = f.read() 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') staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
if os.path.exists(git_folder): try:
# Git update git_folder = os.path.join(plugin_root, '.git')
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly) if os.path.exists(git_folder):
else: success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
# Fallback: Download ZIP and replace files else:
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root) 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: if settings_backup and success:
with open(settings_path, 'w', encoding='utf-8') as f: with open(settings_path, 'w', encoding='utf-8') as f:
f.write(settings_backup) f.write(settings_backup)
logger.info("Restored settings.json") logger.debug("Restored settings.json content (%d bytes)", len(settings_backup))
if success: if success:
return web.json_response({ return web.json_response({
@@ -190,6 +287,164 @@ class UpdateRoutes:
'error': str(e) '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 @staticmethod
async def _download_and_replace_zip(plugin_root: str) -> tuple[bool, str]: async def _download_and_replace_zip(plugin_root: str) -> tuple[bool, str]:
""" """
@@ -244,8 +499,7 @@ class UpdateRoutes:
except Exception: except Exception:
logger.debug("Could not close downloaded-version history database", exc_info=True) 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=list(_PRESERVE_DIRS))
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=['settings.json', 'civitai', 'model_cache', 'cache', 'wildcards', 'backups', 'stats'])
# Extract ZIP to temp dir # Extract ZIP to temp dir
with tempfile.TemporaryDirectory() as tmp_dir: with tempfile.TemporaryDirectory() as tmp_dir:
@@ -255,7 +509,7 @@ class UpdateRoutes:
extracted_root = next(os.scandir(tmp_dir)).path extracted_root = next(os.scandir(tmp_dir)).path
# Copy files, skipping user data that should be preserved # 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): for item in os.listdir(extracted_root):
if item in skip_items: if item in skip_items:
continue continue
@@ -272,7 +526,7 @@ class UpdateRoutes:
# for ComfyUI Manager to work properly # for ComfyUI Manager to work properly
tracking_info_file = os.path.join(plugin_root, '.tracking') tracking_info_file = os.path.join(plugin_root, '.tracking')
tracking_files = [] tracking_files = []
skip_tracked = {'civitai', 'wildcards', 'backups', 'stats'} skip_tracked = set(_PRESERVE_DIRS) - {'settings.json'}
for root, dirs, files in os.walk(extracted_root): for root, dirs, files in os.walk(extracted_root):
# Skip user data directories and their contents # Skip user data directories and their contents
rel_root = os.path.relpath(root, extracted_root) rel_root = os.path.relpath(root, extracted_root)
@@ -296,6 +550,7 @@ class UpdateRoutes:
logger.error(f"ZIP update failed: {e}", exc_info=True) logger.error(f"ZIP update failed: {e}", exc_info=True)
return False, "" return False, ""
@staticmethod
def _clean_plugin_folder(plugin_root, skip_files=None): def _clean_plugin_folder(plugin_root, skip_files=None):
skip_files = skip_files or [] skip_files = skip_files or []
for item in os.listdir(plugin_root): for item in os.listdir(plugin_root):
@@ -308,41 +563,54 @@ class UpdateRoutes:
os.remove(path) os.remove(path)
@staticmethod @staticmethod
async def _get_nightly_version() -> tuple[str, List[str]]: async def _get_nightly_version(local_hash: str = "") -> tuple[str, List[str], int, str]:
"""
Fetch latest commit from main branch
"""
repo_owner = "willmiao" repo_owner = "willmiao"
repo_name = "ComfyUI-Lora-Manager" 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" github_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/commits/main"
try: try:
downloader = await get_downloader() 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: if not success:
logger.warning(f"Failed to fetch GitHub commit: {data}") logger.warning("Failed to fetch GitHub commit: %s", data)
return "main", [] return "main", [], 0, ""
commit_sha = data.get('sha', '')[:7] # Short hash commit_sha = data.get('sha', '')[:7]
commit_message = data.get('commit', {}).get('message', '') commit_message = data.get('commit', {}).get('message', '')
commit_date = data.get('commit', {}).get('committer', {}).get('date', '')[:10]
# Format as "main-{short_hash}"
version = f"main-{commit_sha}" version = f"main-{commit_sha}"
# Use commit message as changelog
changelog = [commit_message] if commit_message else [] 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: except NETWORK_EXCEPTIONS as e:
logger.warning("Unable to reach GitHub for nightly version: %s", e) logger.warning("Unable to reach GitHub for nightly version: %s", e)
return "main", [] return "main", [], 0, ""
except Exception as e: except Exception as e:
logger.error(f"Error fetching nightly version: {e}", exc_info=True) logger.error("Error fetching nightly version: %s", e, exc_info=True)
return "main", [] return "main", [], 0, ""
@staticmethod @staticmethod
def _compare_nightly_versions(local_git_info: Dict[str, str], remote_version: str) -> bool: 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 from abc import ABC, abstractmethod
import asyncio import asyncio
import re 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 logging
import os import os
import time import time
@@ -109,12 +110,15 @@ class BaseModelService(ABC):
if civitai_model_id is not None: if civitai_model_id is not None:
sorted_data = [ sorted_data = [
item for item in 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), # VLM mode: always sort by version ID descending (newest version first),
# regardless of the current sort_by preference. # regardless of the current sort_by preference.
# Fall back to modified timestamp for non-CivitAI sources.
sorted_data.sort( 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, reverse=True,
) )
@@ -129,18 +133,21 @@ class BaseModelService(ABC):
ufs = self.settings.get("version_grouping", "same_base") ufs = self.settings.get("version_grouping", "same_base")
group_by_base = ufs == "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 version_counter = {} # same-key -> count
standalone = [] standalone = []
for item in sorted_data: for item in sorted_data:
mid = self._extract_model_id(item) mid = self._extract_group_key(item)
if mid is None: if mid is None:
standalone.append(item) standalone.append(item)
continue continue
key = (mid, item.get("base_model") or "") if group_by_base else mid key = (mid, item.get("base_model") or "") if group_by_base else mid
# Count all versions per key # Count all versions per key
version_counter[key] = version_counter.get(key, 0) + 1 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]: if key not in dedup_map or vid > dedup_map[key][1]:
dedup_map[key] = (item, vid) dedup_map[key] = (item, vid)
# Attach version_count to each surviving grouped item (shallow copy # Attach version_count to each surviving grouped item (shallow copy
@@ -174,16 +181,19 @@ class BaseModelService(ABC):
model_groups: Dict[Any, List[Dict]] = {} model_groups: Dict[Any, List[Dict]] = {}
ungrouped_standalone: List[Dict] = [] ungrouped_standalone: List[Dict] = []
for item in sorted_data: for item in sorted_data:
mid = self._extract_model_id(item) mid = self._extract_group_key(item)
if mid is None: if mid is None:
ungrouped_standalone.append(item) ungrouped_standalone.append(item)
continue continue
key = (mid, item.get("base_model") or "") if group_by_base else mid key = (mid, item.get("base_model") or "") if group_by_base else mid
model_groups.setdefault(key, []).append(item) 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(): for items in model_groups.values():
items.sort( 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, reverse=True,
) )
# Sort groups by version count # 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("model_name") or item.get("file_name") or "").lower(),
item.get("file_path", "").lower(), item.get("file_path", "").lower(),
) )
elif key_name == "random":
# Seeded random shuffle: same seed -> same order (stable pagination)
rng = random.Random(sort_params.seed or "random")
result = list(data)
rng.shuffle(result)
return result
elif key_name == "size": elif key_name == "size":
key_fn = lambda item: ( key_fn = lambda item: (
int(item.get("size", 0) or 0), int(item.get("size", 0) or 0),
@@ -697,6 +713,33 @@ class BaseModelService(ABC):
return annotated 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 @staticmethod
def _extract_model_id(item: Dict) -> Optional[int]: def _extract_model_id(item: Dict) -> Optional[int]:
civitai = item.get("civitai") if isinstance(item, dict) else None civitai = item.get("civitai") if isinstance(item, dict) else None
+88 -5
View File
@@ -2,6 +2,7 @@ import asyncio
import copy import copy
import logging import logging
import os import os
import time
from collections import OrderedDict from collections import OrderedDict
from typing import Any, Optional, Dict, Tuple, List, Sequence from typing import Any, Optional, Dict, Tuple, List, Sequence
from .connectivity_guard import ( from .connectivity_guard import (
@@ -19,6 +20,12 @@ from ..utils.civitai_utils import resolve_license_payload
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Best-effort cache for creator model counts, keyed by lowercase username.
# Values are (monotonic timestamp, count or None); None results are cached
# too so repeated failures don't hammer the API.
_CREATOR_COUNT_CACHE_TTL_SECONDS = 600
_creator_model_count_cache: Dict[str, Tuple[float, Optional[int]]] = {}
class CivitaiClient: class CivitaiClient:
_instance = None _instance = None
@@ -743,17 +750,34 @@ class CivitaiClient:
return all_versions if all_versions else None return all_versions if all_versions else None
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(
"""Fetch all models for a specific Civitai user.""" self, username: str, cursor: Optional[str] = None
) -> Optional[Dict[str, Any]]:
"""Fetch one page (up to 100 models) for a specific Civitai user.
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
or None on failure. Pass ``cursor`` (from a previous response's
``nextCursor``) to fetch subsequent pages.
"""
if not username: if not username:
return None return None
params: Dict[str, Any] = {
"username": username,
"nsfw": "true",
"limit": 100,
"sort": "Newest",
"period": "AllTime",
}
if cursor:
params["cursor"] = cursor
try: try:
success, result = await self._make_request( success, result = await self._make_request(
"GET", "GET",
f"{self.base_url}/models", f"{self.base_url}/models",
use_auth=True, use_auth=True,
params={"username": username, "nsfw": "true"}, params=params,
) )
if not success: if not success:
@@ -765,7 +789,7 @@ class CivitaiClient:
items = result.get("items") if isinstance(result, dict) else None items = result.get("items") if isinstance(result, dict) else None
if not isinstance(items, list): if not isinstance(items, list):
return [] items = []
for model in items: for model in items:
versions = model.get("modelVersions") versions = model.get("modelVersions")
@@ -774,9 +798,68 @@ class CivitaiClient:
for version in versions: for version in versions:
self._remove_comfy_metadata(version) self._remove_comfy_metadata(version)
return items next_cursor: Optional[str] = None
metadata = result.get("metadata") if isinstance(result, dict) else None
if isinstance(metadata, dict):
raw_cursor = metadata.get("nextCursor")
if raw_cursor is not None:
next_cursor = str(raw_cursor)
return {"items": items, "nextCursor": next_cursor}
except RateLimitError: except RateLimitError:
raise raise
except Exception as exc: # pragma: no cover - defensive logging except Exception as exc: # pragma: no cover - defensive logging
logger.error("Error fetching models for %s: %s", username, exc) logger.error("Error fetching models for %s: %s", username, exc)
return None return None
async def get_creator_model_count(self, username: str) -> Optional[int]:
"""Best-effort lookup of a creator's published model count.
Uses the ``/creators`` endpoint (a contains-match query), picking the
entry whose username matches exactly (case-insensitive). Returns None
on any failure; never raises. Results (including None) are cached
for ``_CREATOR_COUNT_CACHE_TTL_SECONDS``.
"""
if not username:
return None
cache_key = username.lower()
cached = _creator_model_count_cache.get(cache_key)
if cached is not None:
cached_at, cached_count = cached
if time.monotonic() - cached_at < _CREATOR_COUNT_CACHE_TTL_SECONDS:
return cached_count
count: Optional[int] = None
try:
success, result = await self._make_request(
"GET",
f"{self.base_url}/creators",
use_auth=True,
params={"query": username, "limit": 10},
)
if success and isinstance(result, dict):
creators = result.get("items")
if isinstance(creators, list):
for creator in creators:
if not isinstance(creator, dict):
continue
creator_name = creator.get("username")
if not isinstance(creator_name, str):
continue
if creator_name.lower() != cache_key:
continue
model_count = creator.get("modelCount")
if isinstance(model_count, (int, float)) and not isinstance(
model_count, bool
):
count = int(model_count)
break
except Exception as exc: # best-effort only, never propagate
logger.debug(
"Failed to fetch creator model count for %s: %s", username, exc
)
_creator_model_count_cache[cache_key] = (time.monotonic(), count)
return count
+22
View File
@@ -1389,7 +1389,17 @@ class DownloadManager:
# Update save directory with relative path if provided # Update save directory with relative path if provided
if relative_path: if relative_path:
base_save_dir = save_dir
save_dir = os.path.join(save_dir, relative_path) save_dir = os.path.join(save_dir, relative_path)
# Security: validate path containment after joining
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
if not resolved_dir.startswith(base_dir + os.sep) and resolved_dir != base_dir:
logger.warning(
"Path traversal detected: %s escapes %s",
resolved_dir, base_dir,
)
return {"success": False, "error": "Download path is outside allowed directory"}
# Create directory if it doesn't exist # Create directory if it doesn't exist
os.makedirs(save_dir, exist_ok=True) os.makedirs(save_dir, exist_ok=True)
@@ -1827,6 +1837,9 @@ class DownloadManager:
model_tags, model_type model_tags, model_type
) )
if not first_tag:
first_tag = "no tags" # Default if no tags available
# Format the template with available data # Format the template with available data
formatted_path = path_template formatted_path = path_template
formatted_path = formatted_path.replace("{base_model}", mapped_base_model) formatted_path = formatted_path.replace("{base_model}", mapped_base_model)
@@ -1842,6 +1855,15 @@ class DownloadManager:
if model_type == "embedding": if model_type == "embedding":
formatted_path = formatted_path.replace(" ", "_") formatted_path = formatted_path.replace(" ", "_")
# Sanitize the resolved path to prevent path traversal:
# - Strip leading slashes (prevents os.path.join from treating path as absolute)
# - Collapse double slashes from empty placeholder substitutions
# - Strip trailing slashes for cleanliness
formatted_path = formatted_path.lstrip("/")
while "//" in formatted_path:
formatted_path = formatted_path.replace("//", "/")
formatted_path = formatted_path.rstrip("/")
return formatted_path return formatted_path
async def _execute_download( async def _execute_download(
+34 -2
View File
@@ -31,7 +31,7 @@ class DownloadQueueService:
_instance: Optional[DownloadQueueService] = None _instance: Optional[DownloadQueueService] = None
_class_lock: asyncio.Lock = asyncio.Lock() _class_lock: asyncio.Lock = asyncio.Lock()
_SCHEMA = """ _SCHEMA_TABLES = """
CREATE TABLE IF NOT EXISTS download_queue ( CREATE TABLE IF NOT EXISTS download_queue (
download_id TEXT PRIMARY KEY, download_id TEXT PRIMARY KEY,
model_id INTEGER, 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_completed ON download_history(completed_at DESC);
CREATE INDEX IF NOT EXISTS idx_dh_status ON download_history(status); 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 CREATE UNIQUE INDEX IF NOT EXISTS idx_dh_download_id
ON download_history(download_id) WHERE download_id IS NOT NULL; ON download_history(download_id) WHERE download_id IS NOT NULL;
""" """
@@ -115,10 +118,39 @@ class DownloadQueueService:
if self._schema_initialized: if self._schema_initialized:
return return
with self._connect() as conn: 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() conn.commit()
self._schema_initialized = True 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: def get_database_path(self) -> str:
"""Return the resolved database file path.""" """Return the resolved database file path."""
return self._db_path return self._db_path
+42 -8
View File
@@ -566,18 +566,52 @@ class LLMService:
if effective_max is None: if effective_max is None:
effective_max = 4096 effective_max = 4096
result = await self.chat_completion( # Use json_schema (not json_object) for broader provider compatibility:
messages=messages, # LM Studio and some other OpenAI-compatible servers reject
model=model, # json_object but accept json_schema. {"type": "object"} is
temperature=temperature, # functionally equivalent — it accepts any JSON object without
response_format={"type": "json_object"}, # constraining specific fields.
max_tokens=effective_max, response_format = {
) "type": "json_schema",
"json_schema": {
"name": "metadata",
"schema": {"type": "object"},
},
}
try:
result = await self.chat_completion(
messages=messages,
model=model,
temperature=temperature,
response_format=response_format,
max_tokens=effective_max,
)
except LLMResponseError as e:
# Only fall back when the provider rejects the response_format
# type value (e.g. "'response_format.type' must be..."). Avoid
# catching unrelated 400 errors whose body happens to mention
# "response_format" (e.g. "model does not support
# response_format restrictions on this endpoint").
if "'response_format.type'" not in str(e).lower():
raise
logger.info(
"Provider rejected response_format, retrying without it. "
"Falling back to prompt-only JSON mode. Error: %s",
e,
)
result = await self.chat_completion(
messages=messages,
model=model,
temperature=temperature,
response_format=None,
max_tokens=effective_max,
)
content = result.get("content", "") or "" content = result.get("content", "") or ""
if not content: if not content:
raise LLMResponseError( raise LLMResponseError(
"LLM returned empty content in json_object mode. " "LLM returned empty content. "
f"Raw response: {json.dumps(result)[:500]}" f"Raw response: {json.dumps(result)[:500]}"
) )
+21 -12
View File
@@ -1,6 +1,7 @@
import asyncio import asyncio
import time import time
import logging import logging
import random
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from typing import Any, Dict, List, Optional, Tuple from typing import Any, Dict, List, Optional, Tuple
@@ -38,8 +39,8 @@ class ModelCache:
def __post_init__(self): def __post_init__(self):
self._lock = asyncio.Lock() self._lock = asyncio.Lock()
# Cache for last sort: (sort_key, order) -> sorted list # Cache for last sort: (sort_key, order, seed) -> sorted list
self._last_sort: Tuple[str, str] = (None, None) self._last_sort: Tuple[Optional[str], str, Optional[str]] = (None, "asc", None)
self._last_sorted_data: List[Dict] = [] self._last_sorted_data: List[Dict] = []
self._normalize_raw_data() self._normalize_raw_data()
self.name_display_mode = self._normalize_display_mode(self.name_display_mode) self.name_display_mode = self._normalize_display_mode(self.name_display_mode)
@@ -203,9 +204,9 @@ class ModelCache:
async def resort(self): async def resort(self):
"""Resort cached data according to last sort mode if set""" """Resort cached data according to last sort mode if set"""
async with self._lock: async with self._lock:
if self._last_sort != (None, None): if self._last_sort[0] is not None:
sort_key, order = self._last_sort sort_key, order, seed = self._last_sort
sorted_data = self._sort_data(self.raw_data, sort_key, order) sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
self._last_sorted_data = sorted_data self._last_sorted_data = sorted_data
# Update folder list # Update folder list
# else: do nothing # else: do nothing
@@ -218,7 +219,7 @@ class ModelCache:
self.folders = sorted(list(all_folders), key=lambda x: x.lower()) self.folders = sorted(list(all_folders), key=lambda x: x.lower())
self.rebuild_version_index() self.rebuild_version_index()
def _sort_data(self, data: List[Dict], sort_key: str, order: str) -> List[Dict]: def _sort_data(self, data: List[Dict], sort_key: str, order: str, seed: Optional[str] = None) -> List[Dict]:
"""Sort data by sort_key and order""" """Sort data by sort_key and order"""
start_time = time.perf_counter() start_time = time.perf_counter()
reverse = (order == 'desc') reverse = (order == 'desc')
@@ -265,6 +266,13 @@ class ModelCache:
), ),
reverse=reverse reverse=reverse
) )
elif sort_key == 'random':
# Random shuffle seeded for stable pagination: the same seed
# always yields the same order, so successive page requests
# stay consistent while browsing.
rng = random.Random(seed or 'random')
result = list(data)
rng.shuffle(result)
elif sort_key == 'versions_count': elif sort_key == 'versions_count':
# Pre-dedup sort: fall back to name sort. # Pre-dedup sort: fall back to name sort.
# Actual re-sort by version_count happens in get_paginated_data after dedup. # Actual re-sort by version_count happens in get_paginated_data after dedup.
@@ -285,15 +293,16 @@ class ModelCache:
logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration) logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration)
return result return result
async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc') -> List[Dict]: async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc', seed: Optional[str] = None) -> List[Dict]:
"""Get sorted data by sort_key and order, using cache if possible""" """Get sorted data by sort_key and order, using cache if possible"""
async with self._lock: async with self._lock:
if (sort_key, order) == self._last_sort: cache_key = (sort_key, order, seed)
if cache_key == self._last_sort:
return self._last_sorted_data return self._last_sorted_data
start_time = time.perf_counter() start_time = time.perf_counter()
sorted_data = self._sort_data(self.raw_data, sort_key, order) sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
self._last_sort = (sort_key, order) self._last_sort = cache_key
self._last_sorted_data = sorted_data self._last_sorted_data = sorted_data
duration = time.perf_counter() - start_time duration = time.perf_counter() - start_time
@@ -313,8 +322,8 @@ class ModelCache:
self.name_display_mode = normalized self.name_display_mode = normalized
if self._last_sort[0] == 'name': if self._last_sort[0] == 'name':
sort_key, order = self._last_sort sort_key, order, seed = self._last_sort
self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order) self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
async def update_preview_url(self, file_path: str, preview_url: str, preview_nsfw_level: int) -> bool: async def update_preview_url(self, file_path: str, preview_url: str, preview_nsfw_level: int) -> bool:
"""Update preview_url for a specific model in all cached data """Update preview_url for a specific model in all cached data
+4
View File
@@ -8,6 +8,7 @@ from abc import ABC, abstractmethod
from ..utils.utils import calculate_relative_path_for_model, remove_empty_dirs from ..utils.utils import calculate_relative_path_for_model, remove_empty_dirs
from ..utils.constants import AUTO_ORGANIZE_BATCH_SIZE from ..utils.constants import AUTO_ORGANIZE_BATCH_SIZE
from ..services.settings_manager import get_settings_manager from ..services.settings_manager import get_settings_manager
from ..services.model_lifecycle_service import _require_path_in_library_roots
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -493,6 +494,9 @@ class ModelMoveService:
Dictionary with move result Dictionary with move result
""" """
try: try:
_require_path_in_library_roots(file_path, self.scanner, label="Source path")
_require_path_in_library_roots(target_path, self.scanner, label="Target path")
if use_default_paths: if use_default_paths:
# Find the model in cache to get metadata # Find the model in cache to get metadata
cache = await self.scanner.get_cached_data() cache = await self.scanner.get_cached_data()
+41
View File
@@ -48,6 +48,36 @@ async def delete_model_artifacts(
return deleted return deleted
def _require_path_in_library_roots(file_path: str, scanner, *, label: str = "path") -> None:
"""Raise ``ValueError`` if *file_path* is not inside a configured model root.
Uses ``os.path.abspath()`` (NOT ``realpath``) to resolve ``..`` and ``.``
while preserving symlinks this keeps the check in business-path space.
Skips when the scanner does not expose ``get_model_roots`` or the list
is empty.
"""
roots = None
if hasattr(scanner, "get_model_roots"):
try:
roots = scanner.get_model_roots()
except NotImplementedError:
roots = None
if not roots:
return
resolved = os.path.abspath(os.path.normpath(file_path))
for root in roots:
root_resolved = os.path.abspath(os.path.normpath(root))
if resolved == root_resolved or resolved.startswith(root_resolved + os.sep):
return
raise ValueError(
f"{label} '{file_path}' is outside configured library directories"
)
class ModelLifecycleService: class ModelLifecycleService:
"""Co-ordinate destructive and mutating model operations.""" """Co-ordinate destructive and mutating model operations."""
@@ -74,6 +104,8 @@ class ModelLifecycleService:
if not file_path: if not file_path:
raise ValueError("Model path is required") raise ValueError("Model path is required")
_require_path_in_library_roots(file_path, self._scanner, label="File path")
cache = await self._scanner.get_cached_data() cache = await self._scanner.get_cached_data()
cached_entry = None cached_entry = None
@@ -182,6 +214,8 @@ class ModelLifecycleService:
if not file_path: if not file_path:
raise ValueError("Model path is required") raise ValueError("Model path is required")
_require_path_in_library_roots(file_path, self._scanner, label="File path")
metadata_path = os.path.splitext(file_path)[0] + ".metadata.json" metadata_path = os.path.splitext(file_path)[0] + ".metadata.json"
metadata = await self._metadata_loader(metadata_path) metadata = await self._metadata_loader(metadata_path)
metadata["exclude"] = True metadata["exclude"] = True
@@ -229,6 +263,8 @@ class ModelLifecycleService:
if not file_path: if not file_path:
raise ValueError("Model path is required") raise ValueError("Model path is required")
_require_path_in_library_roots(file_path, self._scanner, label="File path")
if not os.path.exists(file_path): if not os.path.exists(file_path):
raise ValueError("Model file does not exist") raise ValueError("Model file does not exist")
@@ -270,6 +306,9 @@ class ModelLifecycleService:
if not file_paths: if not file_paths:
raise ValueError("No file paths provided for deletion") raise ValueError("No file paths provided for deletion")
for path in file_paths:
_require_path_in_library_roots(path, self._scanner, label="File path")
return await self._scanner.bulk_delete_models(file_paths) return await self._scanner.bulk_delete_models(file_paths)
async def rename_model( async def rename_model(
@@ -280,6 +319,8 @@ class ModelLifecycleService:
if not file_path or not new_file_name: if not file_path or not new_file_name:
raise ValueError("File path and new file name are required") raise ValueError("File path and new file name are required")
_require_path_in_library_roots(file_path, self._scanner, label="File path")
invalid_chars = {"/", "\\", ":", "*", "?", '"', "<", ">", "|"} invalid_chars = {"/", "\\", ":", "*", "?", '"', "<", ">", "|"}
if any(char in new_file_name for char in invalid_chars): if any(char in new_file_name for char in invalid_chars):
raise ValueError("Invalid characters in file name") raise ValueError("Invalid characters in file name")
+50 -11
View File
@@ -143,10 +143,18 @@ class ModelMetadataProvider(ABC):
pass pass
@abstractmethod @abstractmethod
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
"""Fetch models owned by the specified user""" """Fetch one page of models owned by the specified user.
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
or None when unsupported/failed. ``cursor`` continues a previous page.
"""
pass pass
async def get_creator_model_count(self, username: str) -> Optional[int]:
"""Published model count for the user; None when unsupported."""
return None
class CivitaiModelMetadataProvider(ModelMetadataProvider): class CivitaiModelMetadataProvider(ModelMetadataProvider):
"""Provider that uses Civitai API for metadata""" """Provider that uses Civitai API for metadata"""
@@ -175,8 +183,11 @@ class CivitaiModelMetadataProvider(ModelMetadataProvider):
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]: async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
return await self.client.get_model_version_info(version_id) return await self.client.get_model_version_info(version_id)
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
return await self.client.get_user_models(username) return await self.client.get_user_models(username, cursor)
async def get_creator_model_count(self, username: str) -> Optional[int]:
return await self.client.get_creator_model_count(username)
class CivArchiveModelMetadataProvider(ModelMetadataProvider): class CivArchiveModelMetadataProvider(ModelMetadataProvider):
"""Provider that uses CivArchive API for metadata""" """Provider that uses CivArchive API for metadata"""
@@ -196,7 +207,7 @@ class CivArchiveModelMetadataProvider(ModelMetadataProvider):
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]: async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
return await self.client.get_model_version_info(version_id) return await self.client.get_model_version_info(version_id)
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
"""Not supported by CivArchive provider""" """Not supported by CivArchive provider"""
return None return None
@@ -347,7 +358,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
version_data = await self._get_version_with_model_data(db, model_id, version_id) version_data = await self._get_version_with_model_data(db, model_id, version_id)
return version_data, None return version_data, None
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
"""Listing models by username is not supported for archive database""" """Listing models by username is not supported for archive database"""
return None return None
@@ -602,13 +613,14 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue continue
return None return None
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
for provider, label in self._iter_providers(): for provider, label in self._iter_providers():
try: try:
result = await self._call_with_rate_limit( result = await self._call_with_rate_limit(
label, label,
provider.get_user_models, provider.get_user_models,
username, username,
cursor=cursor,
) )
if result is not None: if result is not None:
return result return result
@@ -624,6 +636,19 @@ class FallbackMetadataProvider(ModelMetadataProvider):
continue continue
return None return None
async def get_creator_model_count(self, username: str) -> Optional[int]:
for provider, label in self._iter_providers():
try:
result = await provider.get_creator_model_count(username)
if result is not None:
return result
except Exception as e:
logger.debug(
"Provider %s failed for get_creator_model_count: %s", label, e
)
continue
return None
def _iter_providers(self): def _iter_providers(self):
return zip(self.providers, self._provider_labels) return zip(self.providers, self._provider_labels)
@@ -704,13 +729,17 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
version_id, version_id,
) )
async def get_user_models(self, username: str) -> Optional[List[Dict]]: async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
return await self._rate_limit_helper.run( return await self._rate_limit_helper.run(
self._label, self._label,
self._provider.get_user_models, self._provider.get_user_models,
username, username,
cursor=cursor,
) )
async def get_creator_model_count(self, username: str) -> Optional[int]:
return await self._provider.get_creator_model_count(username)
class ModelMetadataProviderManager: class ModelMetadataProviderManager:
"""Manager for selecting and using model metadata providers""" """Manager for selecting and using model metadata providers"""
@@ -776,10 +805,20 @@ class ModelMetadataProviderManager:
except NotImplementedError: except NotImplementedError:
return None return None
async def get_user_models(self, username: str, provider_name: str = None) -> Optional[List[Dict]]: async def get_user_models(
"""Fetch models owned by the specified user""" self,
username: str,
provider_name: str = None,
cursor: Optional[str] = None,
) -> Optional[Dict]:
"""Fetch one page of models owned by the specified user"""
provider = self._get_provider(provider_name) provider = self._get_provider(provider_name)
return await provider.get_user_models(username) return await provider.get_user_models(username, cursor)
async def get_creator_model_count(self, username: str, provider_name: str = None) -> Optional[int]:
"""Best-effort published model count for the specified user"""
provider = self._get_provider(provider_name)
return await provider.get_creator_model_count(username)
def _get_provider(self, provider_name: str = None) -> ModelMetadataProvider: def _get_provider(self, provider_name: str = None) -> ModelMetadataProvider:
"""Get provider by name or default provider""" """Get provider by name or default provider"""
+11 -3
View File
@@ -85,6 +85,7 @@ class SortParams:
key: str key: str
order: str order: str
seed: Optional[str] = None
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -116,7 +117,7 @@ class ModelCacheRepository:
async def fetch_sorted(self, params: SortParams) -> List[Dict[str, Any]]: async def fetch_sorted(self, params: SortParams) -> List[Dict[str, Any]]:
"""Fetch cached data pre-sorted according to ``params``.""" """Fetch cached data pre-sorted according to ``params``."""
cache = await self.get_cache() cache = await self.get_cache()
return await cache.get_sorted_data(params.key, params.order) return await cache.get_sorted_data(params.key, params.order, params.seed)
@staticmethod @staticmethod
def parse_sort(sort_by: str) -> SortParams: def parse_sort(sort_by: str) -> SortParams:
@@ -132,10 +133,17 @@ class ModelCacheRepository:
sort_key = sort_by.strip().lower() or "name" sort_key = sort_by.strip().lower() or "name"
order = "asc" order = "asc"
if order not in ("asc", "desc"): seed = None
if sort_key == "random":
# Random sort: the portion after ':' is the shuffle seed.
# A stable seed keeps paginated requests consistent; order is
# meaningless for a random shuffle.
seed = order if order and order not in ("asc", "desc") else None
order = "asc"
elif order not in ("asc", "desc"):
order = "asc" order = "asc"
return SortParams(key=sort_key, order=order) return SortParams(key=sort_key, order=order, seed=seed)
class ModelFilterSet: class ModelFilterSet:
+41 -10
View File
@@ -14,7 +14,7 @@ from ..utils.metadata_manager import MetadataManager
from ..utils.civitai_utils import resolve_license_info from ..utils.civitai_utils import resolve_license_info
from .model_cache import ModelCache from .model_cache import ModelCache
from .model_hash_index import ModelHashIndex from .model_hash_index import ModelHashIndex
from .model_lifecycle_service import delete_model_artifacts from .model_lifecycle_service import delete_model_artifacts, _require_path_in_library_roots
from .service_registry import ServiceRegistry from .service_registry import ServiceRegistry
from .websocket_manager import ws_manager from .websocket_manager import ws_manager
from .persistent_model_cache import get_persistent_cache from .persistent_model_cache import get_persistent_cache
@@ -927,6 +927,25 @@ class ModelScanner:
# Update cache data # Update cache data
self._cache.raw_data = [item for item in self._cache.raw_data if item['file_path'] not in missing_files] 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 # Resort cache if changes were made
if total_added > 0 or total_removed > 0: if total_added > 0 or total_removed > 0:
# Update folders list # Update folders list
@@ -1352,18 +1371,25 @@ class ModelScanner:
# Update folder in metadata # Update folder in metadata
metadata_dict['folder'] = folder metadata_dict['folder'] = folder
# Add to cache 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) self._cache.raw_data.append(metadata_dict)
self._cache.add_to_version_index(metadata_dict)
# Resort cache data
await self._cache.resort() 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 # Update the hash index
self._hash_index.add_entry(metadata_dict['sha256'], metadata_dict['file_path']) self._hash_index.add_entry(metadata_dict['sha256'], metadata_dict['file_path'])
await self._persist_current_cache() await self._persist_current_cache()
@@ -1395,6 +1421,9 @@ class ModelScanner:
base_name = os.path.splitext(os.path.basename(source_path))[0] base_name = os.path.splitext(os.path.basename(source_path))[0]
source_dir = os.path.dirname(source_path) source_dir = os.path.dirname(source_path)
_require_path_in_library_roots(source_path, self, label="Source path")
_require_path_in_library_roots(target_path, self, label="Target path")
os.makedirs(target_path, exist_ok=True) os.makedirs(target_path, exist_ok=True)
def get_source_hash(): def get_source_hash():
@@ -1723,7 +1752,7 @@ class ModelScanner:
# ---- Conditional resort (only when sort-key fields changed) ---- # ---- Conditional resort (only when sort-key fields changed) ----
need_resort = False need_resort = False
_last = cache._last_sort _last = cache._last_sort
sort_key: Optional[str] = _last[0] if _last != (None, None) else None sort_key: Optional[str] = _last[0] if _last[0] is not None else None
if sort_key == "name": if sort_key == "name":
if ( if (
old_model_name != desired_entry.get("model_name", "") old_model_name != desired_entry.get("model_name", "")
@@ -1971,6 +2000,8 @@ class ModelScanner:
break break
try: try:
_require_path_in_library_roots(file_path, self, label="File path")
target_dir = os.path.dirname(file_path) target_dir = os.path.dirname(file_path)
base_name = os.path.basename(file_path) base_name = os.path.basename(file_path)
file_name, main_extension = os.path.splitext(base_name) file_name, main_extension = os.path.splitext(base_name)
+16 -48
View File
@@ -21,7 +21,7 @@ from .checkpoint_scanner import CheckpointScanner
from .settings_manager import get_settings_manager from .settings_manager import get_settings_manager
from .recipes.errors import RecipeNotFoundError from .recipes.errors import RecipeNotFoundError
from ..utils.civitai_utils import extract_civitai_image_id from ..utils.civitai_utils import extract_civitai_image_id
from ..utils.utils import calculate_recipe_fingerprint, fuzzy_match from ..utils.utils import calculate_recipe_fingerprint
from natsort import natsorted from natsort import natsorted
import sys import sys
import re import re
@@ -1020,13 +1020,16 @@ class RecipeScanner:
try: try:
result = self._fts_index.search(search, fields) result = self._fts_index.search(search, fields)
# Return None if empty to trigger fuzzy fallback # Return empty set for empty FTS results — do NOT fall back to
# Empty FTS results may indicate query syntax issues or need for fuzzy matching # Python fuzzy matching, which freezes the server with 10k+ recipes.
# FTS5 prefix matching with unicode61 tokenizer correctly handles
# compound tokens (e.g. "illustrious" matches "path/illustrious/model").
# If FTS returns nothing, there are genuinely no matching recipes.
if not result: if not result:
return None return set()
return result return result
except Exception as exc: except Exception as exc:
logger.debug("FTS search failed, falling back to fuzzy search: %s", exc) logger.debug("FTS search failed, falling back to title-only search: %s", exc)
return None return None
def _update_fts_index_for_recipe( def _update_fts_index_for_recipe(
@@ -2079,49 +2082,14 @@ class RecipeScanner:
if str(item.get("id", "")) in fts_matching_ids if str(item.get("id", "")) in fts_matching_ids
] ]
else: else:
# Fallback to fuzzy_match (slower but always available) # FTS index not yet built — return empty rather than
# Build the search predicate based on search options # scanning 42k+ items in Python. The FTS background build
def matches_search(item): # finishes in seconds; by the time a user navigates here
# Search in title if enabled # and types a search, it is already available.
if search_options.get("title", True): logger.debug(
if fuzzy_match(str(item.get("title", "")), search): "FTS index not ready — search '%s' returning empty", search
return True )
filtered_data = []
# Search in tags if enabled
if search_options.get("tags", True) and "tags" in item:
for tag in item["tags"]:
if fuzzy_match(tag, search):
return True
# Search in lora file names if enabled
if search_options.get("lora_name", True) and "loras" in item:
for lora in item["loras"]:
if fuzzy_match(str(lora.get("file_name", "")), search):
return True
# Search in lora model names if enabled
if search_options.get("lora_model", True) and "loras" in item:
for lora in item["loras"]:
if fuzzy_match(str(lora.get("modelName", "")), search):
return True
# Search in prompt and negative_prompt if enabled
if search_options.get("prompt", True) and "gen_params" in item:
gen_params = item["gen_params"]
if fuzzy_match(str(gen_params.get("prompt", "")), search):
return True
if fuzzy_match(
str(gen_params.get("negative_prompt", "")), search
):
return True
# No match found
return False
# Filter the data using the search predicate
filtered_data = [
item for item in filtered_data if matches_search(item)
]
# Apply additional filters # Apply additional filters
if filters: if filters:
@@ -126,6 +126,7 @@ class BulkMetadataRefreshUseCase:
if sha256: if sha256:
model["sha256"] = sha256 model["sha256"] = sha256
model["hash_status"] = "completed" model["hash_status"] = "completed"
hash_status = "completed"
else: else:
self._logger.error(f"Failed to calculate hash for {file_path}") self._logger.error(f"Failed to calculate hash for {file_path}")
failures.append({"name": model.get("model_name", file_path or "Unknown"), "error": "Failed to calculate hash"}) failures.append({"name": model.get("model_name", file_path or "Unknown"), "error": "Failed to calculate hash"})
@@ -148,6 +149,16 @@ class BulkMetadataRefreshUseCase:
continue continue
await MetadataManager.hydrate_model_data(model) await MetadataManager.hydrate_model_data(model)
# hydrate_model_data replaces model with .metadata.json content,
# which may lack sha256. Restore from cache and persist the fix.
if not model.get("sha256"):
model["sha256"] = sha256
model["hash_status"] = model.get("hash_status", hash_status)
data_to_save = model.copy()
data_to_save.pop("folder", None)
await MetadataManager.save_metadata(file_path, data_to_save)
result, error_msg = await self._metadata_sync.fetch_and_update_model( result, error_msg = await self._metadata_sync.fetch_and_update_model(
sha256=model["sha256"], sha256=model["sha256"],
file_path=model["file_path"], file_path=model["file_path"],
+36 -3
View File
@@ -19,7 +19,7 @@ logger = logging.getLogger(__name__)
_WILDCARD_PATTERN = re.compile(r"__([\w\s.\-+/*\\]+?)__") _WILDCARD_PATTERN = re.compile(r"__([\w\s.\-+/*\\]+?)__")
_OPTION_PATTERN = re.compile(r"{([^{}]*?)}") _OPTION_PATTERN = re.compile(r"{([^{}]*?)}")
_TRIGGER_WORD_PATTERN = re.compile(r"^trigger_words\d+$") _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+)?$") _NUMERIC_PATTERN = re.compile(r"^-?\d+(\.\d+)?$")
@@ -390,7 +390,7 @@ class WildcardService:
) -> str | None: ) -> str | None:
keyword = _normalize_wildcard_key(raw_key) keyword = _normalize_wildcard_key(raw_key)
if keyword in wildcard_dict: if keyword in wildcard_dict:
return rng.choice(wildcard_dict[keyword]) return self._pick_weighted_or_plain(wildcard_dict[keyword], rng)
if "*" in keyword: if "*" in keyword:
regex_pattern = keyword.replace("*", ".*").replace("+", r"\+") regex_pattern = keyword.replace("*", ".*").replace("+", r"\+")
@@ -400,7 +400,7 @@ class WildcardService:
if compiled.match(key): if compiled.match(key):
aggregated.extend(values) aggregated.extend(values)
if aggregated: if aggregated:
return rng.choice(aggregated) return self._pick_weighted_or_plain(aggregated, rng)
if "/" not in keyword: if "/" not in keyword:
fallback_keyword = _normalize_wildcard_key(f"*/{keyword}") fallback_keyword = _normalize_wildcard_key(f"*/{keyword}")
@@ -409,6 +409,39 @@ class WildcardService:
return None 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: def is_trigger_words_input(name: str) -> bool:
return bool(_TRIGGER_WORD_PATTERN.match(name)) return bool(_TRIGGER_WORD_PATTERN.match(name))
+1
View File
@@ -12,6 +12,7 @@ NODE_TYPES = {
"Lora Loader (LoraManager)": 1, "Lora Loader (LoraManager)": 1,
"Lora Stacker (LoraManager)": 2, "Lora Stacker (LoraManager)": 2,
"WanVideo Lora Select (LoraManager)": 3, "WanVideo Lora Select (LoraManager)": 3,
"Create Hook LoRA (LoraManager)": 4,
} }
# Default ComfyUI node color when bgcolor is null # Default ComfyUI node color when bgcolor is null
+132 -33
View File
@@ -14,11 +14,16 @@ from ..services.service_registry import ServiceRegistry
from ..utils.example_images_paths import ( from ..utils.example_images_paths import (
ExampleImagePathResolver, ExampleImagePathResolver,
ensure_library_root_exists, ensure_library_root_exists,
get_example_images_root,
is_hash_folder,
uses_library_scoped_folders, uses_library_scoped_folders,
) )
from ..utils.metadata_manager import MetadataManager from ..utils.metadata_manager import MetadataManager
from .example_images_processor import ExampleImagesProcessor from .example_images_processor import ExampleImagesProcessor
from .example_images_metadata import MetadataUpdater from .example_images_metadata import (
MetadataUpdater,
update_cache_from_metadata,
)
from ..services.downloader import get_downloader from ..services.downloader import get_downloader
from ..services.settings_manager import get_settings_manager from ..services.settings_manager import get_settings_manager
@@ -87,6 +92,13 @@ class _DownloadProgress(dict):
return snapshot return snapshot
# When fewer candidates than this remain in check_pending_models, probe each
# model folder directly (preserving legacy-folder migration semantics). Above
# it, build a folder index with a single directory scan so libraries with
# 100k+ models do not pay one syscall per candidate.
_BULK_LOOKUP_THRESHOLD = 1000
def _model_directory_has_files(path: str) -> bool: def _model_directory_has_files(path: str) -> bool:
"""Return True when the provided directory exists and contains entries.""" """Return True when the provided directory exists and contains entries."""
@@ -103,6 +115,36 @@ def _model_directory_has_files(path: str) -> bool:
return False return False
def _build_example_folder_index(output_dir: str) -> dict[str, bool]:
"""Build a ``{hash: has_files}`` index for a library's example-image folders.
A single directory scan over the library root replaces ``O(candidates)``
per-folder ``os.scandir`` calls, which is required for libraries with
100k+ models. Each hash folder is classified by whether it contains any
entries, matching the semantics of ``_model_directory_has_files``.
"""
index: dict[str, bool] = {}
if not output_dir or not os.path.isdir(output_dir):
return index
try:
with os.scandir(output_dir) as entries:
for entry in entries:
name = entry.name
if not entry.is_dir() or not is_hash_folder(name):
continue
try:
with os.scandir(entry.path) as subentries:
index[name.lower()] = any(subentries)
except OSError:
index[name.lower()] = False
except OSError:
pass
return index
class DownloadManager: class DownloadManager:
"""Manages downloading example images for models.""" """Manages downloading example images for models."""
@@ -130,6 +172,7 @@ class DownloadManager:
model_types = data.get("model_types", ["lora", "checkpoint"]) model_types = data.get("model_types", ["lora", "checkpoint"])
delay = float(data.get("delay", 0.2)) delay = float(data.get("delay", 0.2))
force = data.get("force", False) force = data.get("force", False)
model_hashes = data.get("model_hashes", [])
# Step 2: Validate configuration (fast lookup) # Step 2: Validate configuration (fast lookup)
settings_manager = get_settings_manager() settings_manager = get_settings_manager()
@@ -199,6 +242,7 @@ class DownloadManager:
delay, delay,
active_library, active_library,
force, force,
model_hashes,
) )
) )
@@ -410,14 +454,49 @@ class DownloadManager:
# Calculate pending count: check which models actually need processing. # Calculate pending count: check which models actually need processing.
# A model is pending if it has a hash, is not already processed or known-failed, # A model is pending if it has a hash, is not already processed or known-failed,
# and its folder doesn't exist or is empty. # and its folder doesn't exist or is empty.
pending_hashes = set() candidate_hashes = [
for model_hash, model_name in all_models_with_hash: model_hash
if model_hash not in processed_models and model_hash not in failed_models: for model_hash, _ in all_models_with_hash
if model_hash not in processed_models
and model_hash not in failed_models
]
pending_hashes: set[str] = set()
# For small candidate counts the existing per-folder check is fine
# and handles legacy folder migration.
# For large libraries, scan the library root once and do set lookups.
if len(candidate_hashes) <= _BULK_LOOKUP_THRESHOLD or not output_dir:
for model_hash in candidate_hashes:
model_dir = ExampleImagePathResolver.get_model_folder( model_dir = ExampleImagePathResolver.get_model_folder(
model_hash, active_library model_hash, active_library
) )
if not _model_directory_has_files(model_dir): if not _model_directory_has_files(model_dir):
pending_hashes.add(model_hash) pending_hashes.add(model_hash)
else:
folder_index = await asyncio.get_event_loop().run_in_executor(
None, _build_example_folder_index, output_dir
)
# In multi-library mode, folders that have not been consolidated
# into the library root yet (startup migration skipped, failed
# move, or created at the legacy path afterwards) still live at
# the legacy root/<hash> location. Only scan that root when at
# least one candidate is missing from the library-root index, so
# the fully-consolidated case does not pay an extra directory
# pass on every call.
if uses_library_scoped_folders() and any(
not folder_index.get(model_hash, False)
for model_hash in candidate_hashes
):
legacy_root = get_example_images_root()
if legacy_root and legacy_root != output_dir:
legacy_index = await asyncio.get_event_loop().run_in_executor(
None, _build_example_folder_index, legacy_root
)
for hash_key, has_files in legacy_index.items():
folder_index.setdefault(hash_key, has_files)
for model_hash in candidate_hashes:
if not folder_index.get(model_hash, False):
pending_hashes.add(model_hash)
pending_count = len(pending_hashes) pending_count = len(pending_hashes)
@@ -500,8 +579,9 @@ class DownloadManager:
delay, delay,
library_name, library_name,
force: bool = False, force: bool = False,
model_hashes: list[str] | None = None,
): ):
"""Download example images for all models.""" """Download example images for all models (or only the given hashes)."""
downloader = await get_downloader() downloader = await get_downloader()
@@ -529,6 +609,18 @@ class DownloadManager:
if model.get("sha256"): if model.get("sha256"):
all_models.append((scanner_type, model, scanner)) all_models.append((scanner_type, model, scanner))
# Restrict to the requested hashes when provided (empty = all models).
# Explicit targets are a directed user request, so previously failed
# models are retried instead of skipped.
explicit_targets = bool(model_hashes)
if model_hashes:
hash_set = {h.lower() for h in model_hashes}
all_models = [
(scanner_type, model, scanner)
for scanner_type, model, scanner in all_models
if model.get("sha256", "").lower() in hash_set
]
# Update total count # Update total count
self._progress["total"] = len(all_models) self._progress["total"] = len(all_models)
logger.debug(f"Found {self._progress['total']} models to process") logger.debug(f"Found {self._progress['total']} models to process")
@@ -552,6 +644,7 @@ class DownloadManager:
downloader, downloader,
library_name, library_name,
force, force,
explicit_targets,
) )
# Update progress # Update progress
@@ -648,6 +741,7 @@ class DownloadManager:
downloader, downloader,
library_name, library_name,
force: bool = False, force: bool = False,
explicit_targets: bool = False,
): ):
"""Process a single model download.""" """Process a single model download."""
@@ -670,8 +764,9 @@ class DownloadManager:
self._progress["current_model"] = f"{model_name} ({model_hash[:8]})" self._progress["current_model"] = f"{model_name} ({model_hash[:8]})"
await self._broadcast_progress(status="running") await self._broadcast_progress(status="running")
# Skip if already in failed models (unless force mode is enabled) # Skip if already in failed models (unless force mode is enabled or
if not force and model_hash in self._progress["failed_models"]: # the model was explicitly targeted by hash)
if not force and not explicit_targets and model_hash in self._progress["failed_models"]:
logger.debug(f"Skipping known failed model: {model_name}") logger.debug(f"Skipping known failed model: {model_name}")
return False return False
@@ -680,30 +775,34 @@ class DownloadManager:
) )
existing_files = _model_directory_has_files(model_dir) existing_files = _model_directory_has_files(model_dir)
# Skip if already processed AND directory exists with files # Model-level guard: a populated folder counts as done. Explicitly
if model_hash in self._progress["processed_models"]: # targeted models bypass it so the per-image existence pre-check can
if existing_files: # fill individual gaps without re-fetching existing files.
logger.debug(f"Skipping already processed model: {model_name}") if not explicit_targets:
# Skip if already processed AND directory exists with files
if model_hash in self._progress["processed_models"]:
if existing_files:
logger.debug(f"Skipping already processed model: {model_name}")
return False
logger.debug(
"Model %s (%s) marked as processed but folder empty or missing, reprocessing triggered",
model_name,
model_hash,
)
# Track that we are reprocessing this model for summary logging
self._progress["reprocessed_models"].add(model_hash)
# Remove from processed models since we need to reprocess
self._progress["processed_models"].discard(model_hash)
if existing_files and model_hash not in self._progress["processed_models"]:
logger.debug(
"Model folder already populated for %s, marking as processed without download",
model_name,
)
self._progress["processed_models"].add(model_hash)
return False return False
logger.debug(
"Model %s (%s) marked as processed but folder empty or missing, reprocessing triggered",
model_name,
model_hash,
)
# Track that we are reprocessing this model for summary logging
self._progress["reprocessed_models"].add(model_hash)
# Remove from processed models since we need to reprocess
self._progress["processed_models"].discard(model_hash)
if existing_files and model_hash not in self._progress["processed_models"]:
logger.debug(
"Model folder already populated for %s, marking as processed without download",
model_name,
)
self._progress["processed_models"].add(model_hash)
return False
if not model_dir: if not model_dir:
logger.warning( logger.warning(
"Unable to resolve example images folder for model %s (%s)", "Unable to resolve example images folder for model %s (%s)",
@@ -807,7 +906,7 @@ class DownloadManager:
model_name, model_name,
) )
# Clear failed_models so non-force runs can retry # Clear failed_models so non-force runs can retry
if force and model_hash in self._progress["failed_models"]: if (force or explicit_targets) and model_hash in self._progress["failed_models"]:
self._progress["failed_models"].discard(model_hash) self._progress["failed_models"].discard(model_hash)
logger.info( logger.info(
f"Removed {model_name} from failed_models after force retry with rate-limited images" f"Removed {model_name} from failed_models after force retry with rate-limited images"
@@ -827,7 +926,7 @@ class DownloadManager:
) )
elif success: elif success:
self._progress["processed_models"].add(model_hash) self._progress["processed_models"].add(model_hash)
if force and model_hash in self._progress["failed_models"]: if (force or explicit_targets) and model_hash in self._progress["failed_models"]:
self._progress["failed_models"].discard(model_hash) self._progress["failed_models"].discard(model_hash)
logger.info( logger.info(
f"Removed {model_name} from failed_models after successful force retry" f"Removed {model_name} from failed_models after successful force retry"
@@ -1343,8 +1442,8 @@ class DownloadManager:
await MetadataManager.save_metadata(file_path, model_copy) await MetadataManager.save_metadata(file_path, model_copy)
try: try:
await scanner.update_single_model_cache( await update_cache_from_metadata(
file_path, file_path, model_data scanner, file_path, model_copy
) )
except AttributeError: except AttributeError:
logger.debug( logger.debug(
+57 -41
View File
@@ -1,3 +1,4 @@
import inspect
import logging import logging
import os import os
import re import re
@@ -28,6 +29,31 @@ if TYPE_CHECKING: # pragma: no cover - import for type checkers only
from ..services.settings_manager import SettingsManager from ..services.settings_manager import SettingsManager
async def update_cache_from_metadata(
scanner: Any, file_path: str, metadata: Dict[str, Any]
) -> bool:
"""Update the scanner cache from a metadata dict using the in-place sync path.
``sync_cache_from_metadata`` patches the existing cache entry incrementally
(tag/hash/version indexes, targeted single-row SQL update) and only resorts
when a sort-key field changed. This avoids the ``O(n)`` full-list resort and
full cache rewrite that ``update_single_model_cache`` performs on every call,
which is critical for libraries with 100k+ models.
Falls back to the legacy full update when the scanner does not expose an
async ``sync_cache_from_metadata`` method.
Returns:
``True`` if the cache entry was updated, ``False`` otherwise.
"""
sync_method = getattr(scanner, "sync_cache_from_metadata", None)
if inspect.iscoroutinefunction(sync_method):
return await sync_method(file_path, metadata)
return await scanner.update_single_model_cache(file_path, file_path, metadata)
def _build_metadata_sync_service(settings_manager: "SettingsManager") -> MetadataSyncService: def _build_metadata_sync_service(settings_manager: "SettingsManager") -> MetadataSyncService:
"""Construct a metadata sync service bound to the provided settings.""" """Construct a metadata sync service bound to the provided settings."""
@@ -103,7 +129,7 @@ class MetadataUpdater:
progress['refreshed_models'].add(model_hash) progress['refreshed_models'].add(model_hash)
async def update_cache_func(old_path, new_path, metadata): async def update_cache_func(old_path, new_path, metadata):
return await scanner.update_single_model_cache(old_path, new_path, metadata) return await update_cache_from_metadata(scanner, new_path, metadata)
await MetadataManager.hydrate_model_data(model_data) await MetadataManager.hydrate_model_data(model_data)
success, error = await _get_metadata_sync_service().fetch_and_update_model( success, error = await _get_metadata_sync_service().fetch_and_update_model(
@@ -234,6 +260,7 @@ class MetadataUpdater:
# Save metadata to .metadata.json file # Save metadata to .metadata.json file
file_path = model.get('file_path') file_path = model.get('file_path')
model_copy: Optional[Dict[str, Any]] = None
try: try:
model_copy = model.copy() model_copy = model.copy()
model_copy.pop('folder', None) model_copy.pop('folder', None)
@@ -242,13 +269,17 @@ class MetadataUpdater:
except Exception as e: except Exception as e:
logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}") logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}")
# Save updated metadata to scanner cache # Save updated metadata to scanner cache. sync_cache_from_metadata
success = await scanner.update_single_model_cache(file_path, file_path, model) # returns False both for "already in sync" and for actual failures,
if success: # so the cache sync result is deliberately not treated as an error;
# the return value reflects whether the metadata was persisted.
if file_path and model_copy is not None:
await update_cache_from_metadata(scanner, file_path, model_copy)
logger.info(f"Successfully updated metadata for {model.get('model_name')} with {len(images)} local examples") logger.info(f"Successfully updated metadata for {model.get('model_name')} with {len(images)} local examples")
return True return True
else:
logger.warning(f"Failed to update metadata for {model.get('model_name')}") logger.warning(f"Failed to update metadata for {model.get('model_name')}")
return False
return False return False
except Exception as e: except Exception as e:
@@ -336,6 +367,7 @@ class MetadataUpdater:
# Save metadata to .metadata.json file # Save metadata to .metadata.json file
file_path = model_data.get('file_path') file_path = model_data.get('file_path')
model_copy: Optional[Dict[str, Any]] = None
if file_path: if file_path:
try: try:
model_copy = model_data.copy() model_copy = model_data.copy()
@@ -346,8 +378,8 @@ class MetadataUpdater:
logger.error(f"Failed to save metadata: {str(e)}") logger.error(f"Failed to save metadata: {str(e)}")
# Save updated metadata to scanner cache # Save updated metadata to scanner cache
if file_path: if file_path and model_copy is not None:
await scanner.update_single_model_cache(file_path, file_path, model_data) await update_cache_from_metadata(scanner, file_path, model_copy)
# Get regular images array (might be None) # Get regular images array (might be None)
regular_images = civitai_data.get('images', []) regular_images = civitai_data.get('images', [])
@@ -475,13 +507,19 @@ class MetadataUpdater:
return False return False
model_folder = get_model_folder(model_hash) model_folder = get_model_folder(model_hash)
if not model_folder: if not model_folder or not os.path.isdir(model_folder):
return False return False
civitai = getattr(metadata, "civitai", None) civitai = getattr(metadata, "civitai", None)
if not isinstance(civitai, dict): if not isinstance(civitai, dict):
return False 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 has_changes = False
custom_images = civitai.get("customImages") custom_images = civitai.get("customImages")
@@ -493,24 +531,15 @@ class MetadataUpdater:
if not img_id: if not img_id:
continue 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) 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: if stale:
for idx in reversed(stale): for idx in reversed(stale):
@@ -532,22 +561,9 @@ class MetadataUpdater:
# is gone. # is gone.
continue 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) 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: if stale:
for idx in reversed(stale): for idx in reversed(stale):
+98 -2
View File
@@ -3,11 +3,19 @@ import logging
import os import os
import re import re
import json import json
import shutil
from ..services.settings_manager import get_settings_manager from ..services.settings_manager import get_settings_manager
from ..services.service_registry import ServiceRegistry 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.metadata_manager import MetadataManager
from ..utils.example_images_processor import ExampleImagesProcessor from ..utils.example_images_processor import ExampleImagesProcessor
from ..utils.example_images_metadata import update_cache_from_metadata
from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -36,6 +44,90 @@ settings = _SettingsProxy()
class ExampleImagesMigration: class ExampleImagesMigration:
"""Handles migrations for example images naming conventions""" """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 @staticmethod
async def check_and_run_migrations(): async def check_and_run_migrations():
"""Check if migrations are needed and run them in background""" """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") logger.debug("No example images path configured or path doesn't exist, skipping migrations")
return 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(): for library_name, library_path in iter_library_roots():
if not library_path or not os.path.exists(library_path): if not library_path or not os.path.exists(library_path):
continue continue
@@ -326,7 +422,7 @@ class ExampleImagesMigration:
await MetadataManager.save_metadata(file_path, model_copy) await MetadataManager.save_metadata(file_path, model_copy)
# Update scanner cache # Update scanner cache
await scanner.update_single_model_cache(file_path, file_path, model_metadata) await update_cache_from_metadata(scanner, file_path, model_copy)
updated_models += 1 updated_models += 1
except Exception as e: except Exception as e:
+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: 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: if not model_hash:
return "" return ""
@@ -113,35 +118,6 @@ def get_model_folder(model_hash: str, library_name: Optional[str] = None) -> str
exc, exc,
) )
return legacy_folder 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 return resolved_folder
+33 -3
View File
@@ -9,7 +9,7 @@ from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
from ..services.service_registry import ServiceRegistry from ..services.service_registry import ServiceRegistry
from ..services.settings_manager import get_settings_manager from ..services.settings_manager import get_settings_manager
from ..utils.example_images_paths import get_model_folder, get_model_relative_path from ..utils.example_images_paths import get_model_folder, get_model_relative_path
from .example_images_metadata import MetadataUpdater from .example_images_metadata import MetadataUpdater, update_cache_from_metadata
from ..utils.metadata_manager import MetadataManager from ..utils.metadata_manager import MetadataManager
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -113,6 +113,26 @@ class ExampleImagesProcessor:
message = str(error).lower() message = str(error).lower()
return '404' in message or 'file not found' in message return '404' in message or 'file not found' in message
@staticmethod
def _example_image_file_exists(model_dir: str, index: int, media_type_hint: str | None = None) -> bool:
"""Return True when the file that would be written for a media index already exists.
The final filename (``image_{index}{extension}``) depends on the downloaded
content, so the extension cannot be known ahead of time. The post-download
check skips the write when the exact target file exists; this pre-check
approximates that with the candidate extensions for the media type (videos
only when the metadata hints at a video) so the network request is avoided
for files that already exist on disk.
"""
if media_type_hint == "video":
extensions = SUPPORTED_MEDIA_EXTENSIONS['videos']
else:
extensions = SUPPORTED_MEDIA_EXTENSIONS['images']
return any(
os.path.exists(os.path.join(model_dir, f"image_{index}{ext}"))
for ext in extensions
)
@staticmethod @staticmethod
async def download_model_images(model_hash, model_name, model_images, model_dir, optimize, downloader): async def download_model_images(model_hash, model_name, model_images, model_dir, optimize, downloader):
"""Download images for a single model """Download images for a single model
@@ -140,6 +160,11 @@ class ExampleImagesProcessor:
if optimize and 'civitai.com' in image_url: if optimize and 'civitai.com' in image_url:
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url) image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
# Skip the download when the file already exists on disk
if ExampleImagesProcessor._example_image_file_exists(model_dir, i, image.get("type")):
logger.debug("File already exists, skipping download for %s", image_url)
continue
# Download the file first to determine the actual file type # Download the file first to determine the actual file type
try: try:
logger.debug(f"Downloading media file {i} for {model_name}") logger.debug(f"Downloading media file {i} for {model_name}")
@@ -229,6 +254,11 @@ class ExampleImagesProcessor:
if optimize and 'civitai.com' in image_url: if optimize and 'civitai.com' in image_url:
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url) image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
# Skip the download when the file already exists on disk
if ExampleImagesProcessor._example_image_file_exists(model_dir, i, image.get("type")):
logger.debug("File already exists, skipping download for %s", image_url)
continue
async def _attempt_download() -> tuple: async def _attempt_download() -> tuple:
logger.debug("Downloading media file %s for %s", i, model_name) logger.debug("Downloading media file %s for %s", i, model_name)
return await downloader.download_to_memory( return await downloader.download_to_memory(
@@ -644,7 +674,7 @@ class ExampleImagesProcessor:
}, status=500) }, status=500)
# Update cache # Update cache
await scanner.update_single_model_cache(file_path, file_path, model_data) await update_cache_from_metadata(scanner, file_path, model_data)
# Get regular images array (might be None) # Get regular images array (might be None)
regular_images = civitai_data.get('images', []) regular_images = civitai_data.get('images', [])
@@ -759,7 +789,7 @@ class ExampleImagesProcessor:
model_copy = model_data.copy() model_copy = model_data.copy()
model_copy.pop('folder', None) model_copy.pop('folder', None)
await MetadataManager.save_metadata(file_path, model_copy) await MetadataManager.save_metadata(file_path, model_copy)
await scanner.update_single_model_cache(file_path, file_path, model_data) await update_cache_from_metadata(scanner, file_path, model_copy)
return web.json_response({ return web.json_response({
'success': True, 'success': True,
+54 -1
View File
@@ -1,7 +1,7 @@
from difflib import SequenceMatcher from difflib import SequenceMatcher
import os import os
import re import re
from typing import Dict from typing import Any, Dict, List, Optional
from ..services.service_registry import ServiceRegistry from ..services.service_registry import ServiceRegistry
from ..config import config from ..config import config
from ..services.settings_manager import get_settings_manager from ..services.settings_manager import get_settings_manager
@@ -294,6 +294,53 @@ def _format_model_name_for_comfyui(file_path: str, model_roots: list) -> str:
return os.path.basename(file_path) return os.path.basename(file_path)
def model_patcher_to_name(model_patcher: Any) -> Optional[str]:
"""Extract a ComfyUI-style model name from a MODEL (ModelPatcher) object.
Core ComfyUI loaders record the absolute weight file path on the patcher's
``cached_patcher_init`` attribute:
- load_checkpoint_guess_config -> (fn, (ckpt_path, ...), index)
- load_diffusion_model -> (fn, (unet_path, model_options))
Patcher clones (LoRA loaders, model merges, ...) preserve the attribute,
so the name is recoverable anywhere downstream of a core loader including
from LoRA Manager's own loaders (CheckpointLoaderLM / UNETLoaderLM), which
call the same core load functions.
The absolute path is converted to the ComfyUI-style relative name used by
the metadata pipeline (covering standard ComfyUI roots and LoRA Manager
extra folder paths).
Returns None when the path cannot be recovered (e.g. third-party loaders
that never set ``cached_patcher_init``).
"""
init = getattr(model_patcher, "cached_patcher_init", None)
if not isinstance(init, (tuple, list)) or len(init) < 2:
return None
args = init[1]
abs_path = args[0] if args else None
if not isinstance(abs_path, str) or not abs_path:
return None
return _abs_model_path_to_name(abs_path)
def _abs_model_path_to_name(abs_path: str) -> str:
"""Convert an absolute model path to a ComfyUI-style relative name.
Tries standard ComfyUI model roots plus LoRA Manager extra folder paths;
falls back to the bare filename.
"""
try:
roots: List[str] = list(config.base_models_roots or [])
roots.extend(config.extra_checkpoints_roots or [])
roots.extend(config.extra_unet_roots or [])
formatted = _format_model_name_for_comfyui(abs_path, roots)
if formatted:
return formatted
except Exception:
pass
return os.path.basename(abs_path)
def fuzzy_match(text: str, pattern: str, threshold: float = 0.85) -> bool: def fuzzy_match(text: str, pattern: str, threshold: float = 0.85) -> bool:
""" """
Check if text matches pattern using fuzzy matching. Check if text matches pattern using fuzzy matching.
@@ -488,6 +535,12 @@ def calculate_relative_path_for_model(
if model_type == "embedding": if model_type == "embedding":
formatted_path = formatted_path.replace(" ", "_") formatted_path = formatted_path.replace(" ", "_")
# Sanitize the resolved path to prevent path traversal
formatted_path = formatted_path.lstrip("/")
while "//" in formatted_path:
formatted_path = formatted_path.replace("//", "/")
formatted_path = formatted_path.rstrip("/")
return formatted_path return formatted_path
+1 -1
View File
@@ -1,7 +1,7 @@
[project] [project]
name = "comfyui-lora-manager" name = "comfyui-lora-manager"
description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!" description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!"
version = "1.1.8" version = "1.2.0"
license = {file = "LICENSE"} license = {file = "LICENSE"}
dependencies = [ dependencies = [
"aiohttp", "aiohttp",
+1
View File
@@ -151,6 +151,7 @@ body.modal-open {
.support-section, .support-section,
.changelog-section, .changelog-section,
.update-info, .update-info,
.update-channels,
.info-item, .info-item,
.path-preview { .path-preview {
background: var(--surface-subtle); background: var(--surface-subtle);
+126 -9
View File
@@ -93,15 +93,13 @@
.update-content { .update-content {
display: flex; display: flex;
flex-direction: column; flex-direction: column;
gap: var(--space-3); gap: var(--space-2);
} }
.update-info { .update-info {
display: flex; display: flex;
justify-content: space-between; justify-content: space-between;
align-items: center; align-items: center;
border-radius: var(--border-radius-sm);
padding: var(--space-3);
} }
.update-info .version-info { .update-info .version-info {
@@ -175,7 +173,6 @@
border: 1px solid var(--lora-border); border: 1px solid var(--lora-border);
border-radius: var(--border-radius-sm); border-radius: var(--border-radius-sm);
padding: var(--space-2); padding: var(--space-2);
margin: var(--space-2) 0;
} }
[data-theme="dark"] .update-progress { [data-theme="dark"] .update-progress {
@@ -233,11 +230,6 @@
} }
/* Changelog section */ /* Changelog section */
.changelog-section {
border-radius: var(--border-radius-sm);
padding: var(--space-3);
}
.changelog-section h3 { .changelog-section h3 {
margin-top: 0; margin-top: 0;
margin-bottom: var(--space-2); margin-bottom: var(--space-2);
@@ -349,6 +341,131 @@
text-decoration: underline; 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 section */
.update-preferences { .update-preferences {
border-top: 1px solid var(--lora-border); border-top: 1px solid var(--lora-border);
+2 -5
View File
@@ -49,10 +49,6 @@ export const MODEL_CONFIG = {
* @returns {Object} Object containing all API endpoints for the model type * @returns {Object} Object containing all API endpoints for the model type
*/ */
export function getApiEndpoints(modelType) { export function getApiEndpoints(modelType) {
if (!Object.values(MODEL_TYPES).includes(modelType)) {
throw new Error(`Invalid model type: ${modelType}`);
}
return { return {
// Base CRUD operations // Base CRUD operations
list: `/api/lm/${modelType}/list`, list: `/api/lm/${modelType}/list`,
@@ -188,7 +184,8 @@ export const DOWNLOAD_ENDPOINTS = {
downloadGet: '/api/lm/download-model-get', downloadGet: '/api/lm/download-model-get',
cancelGet: '/api/lm/cancel-download-get', cancelGet: '/api/lm/cancel-download-get',
progress: '/api/lm/download-progress', progress: '/api/lm/download-progress',
exampleImages: '/api/lm/force-download-example-images' // New endpoint for downloading example images exampleImages: '/api/lm/force-download-example-images', // Re-process example images ignoring previous status
exampleImagesMissing: '/api/lm/download-example-images' // Download only missing example images
}; };
// Hugging Face API endpoints // Hugging Face API endpoints
+8 -2
View File
@@ -1641,7 +1641,7 @@ export class BaseModelApiClient {
} }
} }
async downloadExampleImages(modelHashes, modelTypes = null) { async downloadExampleImages(modelHashes, modelTypes = null, { force = true } = {}) {
let ws = null; let ws = null;
await state.loadingManager.showWithProgress(async (loading) => { await state.loadingManager.showWithProgress(async (loading) => {
@@ -1700,8 +1700,13 @@ export class BaseModelApiClient {
// Determine optimize setting // Determine optimize setting
const optimize = state.global?.settings?.optimize_example_images ?? true; const optimize = state.global?.settings?.optimize_example_images ?? true;
// force=false routes to the regular endpoint, which skips already-processed models
const endpoint = force
? DOWNLOAD_ENDPOINTS.exampleImages
: DOWNLOAD_ENDPOINTS.exampleImagesMissing;
// Make the API request to start the download process // Make the API request to start the download process
const response = await fetch(DOWNLOAD_ENDPOINTS.exampleImages, { const response = await fetch(endpoint, {
method: 'POST', method: 'POST',
headers: { headers: {
'Content-Type': 'application/json' 'Content-Type': 'application/json'
@@ -1710,6 +1715,7 @@ export class BaseModelApiClient {
model_hashes: modelHashes, model_hashes: modelHashes,
output_dir: outputDir, output_dir: outputDir,
optimize: optimize, optimize: optimize,
force: force,
model_types: modelTypes || [this.apiConfig.config.singularName] model_types: modelTypes || [this.apiConfig.config.singularName]
}) })
}); });
@@ -137,11 +137,10 @@ export class BulkContextMenu extends BaseContextMenu {
downloadMissingLorasItem.style.display = currentModelType === 'recipes' ? 'flex' : 'none'; downloadMissingLorasItem.style.display = currentModelType === 'recipes' ? 'flex' : 'none';
} }
const downloadExampleImagesItem = this.menu.querySelector('[data-action="download-example-images"]'); const downloadExampleImagesSubmenu = this.menu.querySelector('[data-has-submenu="download-example-images"]');
if (downloadExampleImagesItem) { if (downloadExampleImagesSubmenu) {
// Show on model pages (loras, checkpoints, embeddings), hide on recipes // Show on model pages (loras, checkpoints, embeddings), hide on recipes
const modelPages = ['loras', 'checkpoints', 'embeddings']; downloadExampleImagesSubmenu.style.display = ['loras', 'checkpoints', 'embeddings'].includes(currentModelType) ? 'flex' : 'none';
downloadExampleImagesItem.style.display = modelPages.includes(currentModelType) ? 'flex' : 'none';
} }
const skipMetadataRefreshItem = this.menu.querySelector('[data-action="skip-metadata-refresh"]'); const skipMetadataRefreshItem = this.menu.querySelector('[data-action="skip-metadata-refresh"]');
@@ -294,8 +293,11 @@ export class BulkContextMenu extends BaseContextMenu {
case 'download-missing-loras': case 'download-missing-loras':
this.handleDownloadMissingLoras(); this.handleDownloadMissingLoras();
break; break;
case 'download-missing-example-images':
this.handleDownloadExampleImages({ force: false });
break;
case 'download-example-images': case 'download-example-images':
this.handleDownloadExampleImages(); this.handleDownloadExampleImages({ force: true });
break; break;
case 'clear': case 'clear':
bulkManager.clearSelection(); bulkManager.clearSelection();
@@ -340,7 +342,7 @@ export class BulkContextMenu extends BaseContextMenu {
await bulkMissingLoraDownloadManager.downloadMissingLoras(selectedRecipes); await bulkMissingLoraDownloadManager.downloadMissingLoras(selectedRecipes);
} }
async handleDownloadExampleImages() { async handleDownloadExampleImages({ force = true } = {}) {
if (state.selectedModels.size === 0) { if (state.selectedModels.size === 0) {
return; return;
} }
@@ -361,7 +363,7 @@ export class BulkContextMenu extends BaseContextMenu {
try { try {
const apiClient = getModelApiClient(); const apiClient = getModelApiClient();
await apiClient.downloadExampleImages([...hashes]); await apiClient.downloadExampleImages([...hashes], null, { force });
} catch (error) { } catch (error) {
console.error('Bulk download example images failed:', error); console.error('Bulk download example images failed:', error);
} }
@@ -347,7 +347,10 @@ export const ModelContextMenuMixin = {
openExampleImagesFolder(this.currentCard.dataset.sha256); openExampleImagesFolder(this.currentCard.dataset.sha256);
return true; return true;
case 'download-examples': case 'download-examples':
this.downloadExampleImages(); this.downloadExampleImages(false);
return true;
case 'download-examples-force':
this.downloadExampleImages(true);
return true; return true;
case 'civitai': case 'civitai':
if (this.currentCard.dataset.from_civitai === 'true') { if (this.currentCard.dataset.from_civitai === 'true') {
@@ -378,7 +381,7 @@ export const ModelContextMenuMixin = {
}, },
// Download example images method // Download example images method
async downloadExampleImages() { async downloadExampleImages(force = false) {
const modelHash = this.currentCard.dataset.sha256; const modelHash = this.currentCard.dataset.sha256;
if (!modelHash) { if (!modelHash) {
showToast('toast.contextMenu.missingHash', {}, 'error'); showToast('toast.contextMenu.missingHash', {}, 'error');
@@ -387,7 +390,7 @@ export const ModelContextMenuMixin = {
try { try {
const apiClient = getModelApiClient(); const apiClient = getModelApiClient();
await apiClient.downloadExampleImages([modelHash]); await apiClient.downloadExampleImages([modelHash], null, { force });
} catch (error) { } catch (error) {
console.error('Error downloading example images:', error); console.error('Error downloading example images:', error);
} }
@@ -260,8 +260,9 @@ export class RecipeContextMenu extends BaseContextMenu {
strength: lora.strength || 1.0, strength: lora.strength || 1.0,
// Model identifiers // Model identifiers
modelId: lora.modelId || lora.model_id || civitaiInfo.modelId,
hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash, hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash,
modelVersionId: civitaiInfo.id || lora.modelVersionId, id: civitaiInfo.id || lora.modelVersionId,
// Metadata // Metadata
thumbnailUrl: civitaiInfo.images?.[0]?.url || '', thumbnailUrl: civitaiInfo.images?.[0]?.url || '',
+1
View File
@@ -1421,6 +1421,7 @@ class RecipeModal {
strength: lora.strength || 1.0, strength: lora.strength || 1.0,
// Model identifiers // Model identifiers
modelId: lora.modelId || lora.model_id || civitaiInfo.modelId,
hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash, hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash,
id: civitaiInfo.id || lora.modelVersionId, id: civitaiInfo.id || lora.modelVersionId,
+56 -17
View File
@@ -108,10 +108,20 @@ export class PageControls {
const sortSelect = document.getElementById('sortSelect'); const sortSelect = document.getElementById('sortSelect');
if (sortSelect) { if (sortSelect) {
initSortDropdown(sortSelect); initSortDropdown(sortSelect);
sortSelect.value = this.pageState.sortBy; this.applySortToSelect(this.pageState.sortBy);
sortSelect.addEventListener('change', async (e) => { sortSelect.addEventListener('change', async (e) => {
this.pageState.sortBy = e.target.value; let value = e.target.value;
this.saveSortPreference(e.target.value); if (value.startsWith('random')) {
// Every pick of Random reshuffles the list: generate a
// fresh seed so the backend keeps a stable order across
// paginated requests.
value = this._randomizeSortValue();
}
this.pageState.sortBy = value;
this.saveSortPreference(value);
// Reset the seeded Random option when switching away from
// Random, or re-apply the fresh seed when picking it again.
this.applySortToSelect(value);
await this.resetAndReload(); await this.resetAndReload();
}); });
} }
@@ -312,6 +322,44 @@ export class PageControls {
} }
} }
/**
* Apply a sort value to the native sort <select>, keeping the Random
* option's value in sync when the persisted value carries a seed
* (e.g. "random:abc123"). Must be used instead of assigning
* sortSelect.value directly whenever the value may be a seeded random
* sort, otherwise the native select has no matching option.
* @param {string} sortValue - Sort value like "name:asc" or "random:<seed>"
*/
applySortToSelect(sortValue) {
const sortSelect = document.getElementById('sortSelect');
if (!sortSelect) return;
const randomOpt = sortSelect.querySelector('option[value="random"], option[value^="random:"]');
if (randomOpt) {
randomOpt.value = String(sortValue).startsWith('random') ? sortValue : 'random';
}
sortSelect.value = sortValue;
}
/**
* Generate a fresh seeded random sort value ("random:<seed>") and keep
* the native <select> in sync so its value matches the persisted sort
* string and the dropdown shows the selected label.
* @returns {string} The new sort value, e.g. "random:abc123xyz"
*/
_randomizeSortValue() {
const seed = Math.random().toString(36).slice(2, 12);
const value = `random:${seed}`;
const sortSelect = document.getElementById('sortSelect');
if (sortSelect) {
const randomOpt = sortSelect.querySelector('option[value="random"], option[value^="random:"]');
if (randomOpt) {
randomOpt.value = value;
}
sortSelect.value = value;
}
return value;
}
/** /**
* Load sort preference from storage * Load sort preference from storage
*/ */
@@ -326,10 +374,7 @@ export class PageControls {
// Handle legacy format conversion // Handle legacy format conversion
const convertedSort = this.convertLegacySortFormat(savedSort); const convertedSort = this.convertLegacySortFormat(savedSort);
this.pageState.sortBy = convertedSort; this.pageState.sortBy = convertedSort;
const sortSelect = document.getElementById('sortSelect'); this.applySortToSelect(convertedSort);
if (sortSelect) {
sortSelect.value = convertedSort;
}
} }
} }
@@ -523,9 +568,9 @@ export class PageControls {
this.pageState.sortBy = restoredSort; this.pageState.sortBy = restoredSort;
this.saveSortPreference(restoredSort); this.saveSortPreference(restoredSort);
this._removeVlmSortOption(); this._removeVlmSortOption();
this.applySortToSelect(restoredSort);
const sortSelect = document.getElementById('sortSelect'); const sortSelect = document.getElementById('sortSelect');
if (sortSelect) { if (sortSelect) {
sortSelect.value = restoredSort;
sortSelect.disabled = false; sortSelect.disabled = false;
} }
} }
@@ -575,10 +620,7 @@ export class PageControls {
const savedGroupedSort = getStorageItem(groupedKey); const savedGroupedSort = getStorageItem(groupedKey);
if (savedGroupedSort) { if (savedGroupedSort) {
this.pageState.sortBy = savedGroupedSort; this.pageState.sortBy = savedGroupedSort;
const sortSelect = document.getElementById('sortSelect'); this.applySortToSelect(savedGroupedSort);
if (sortSelect) {
sortSelect.value = savedGroupedSort;
}
} }
} else { } else {
// Leaving group mode: persist current sort for next time, restore non-group sort // Leaving group mode: persist current sort for next time, restore non-group sort
@@ -586,10 +628,7 @@ export class PageControls {
const savedNormalSort = getStorageItem(`${this.pageType}_sort`); const savedNormalSort = getStorageItem(`${this.pageType}_sort`);
if (savedNormalSort) { if (savedNormalSort) {
this.pageState.sortBy = savedNormalSort; this.pageState.sortBy = savedNormalSort;
const sortSelect = document.getElementById('sortSelect'); this.applySortToSelect(savedNormalSort);
if (sortSelect) {
sortSelect.value = savedNormalSort;
}
} }
} }
} }
@@ -874,7 +913,7 @@ export class PageControls {
} }
if (sortSelect) { if (sortSelect) {
sortSelect.value = this.pageState.sortBy; this.applySortToSelect(this.pageState.sortBy);
} }
if (searchInput) { if (searchInput) {
searchInput.value = this.pageState.filters?.search || ''; searchInput.value = this.pageState.filters?.search || '';
+13 -3
View File
@@ -96,7 +96,16 @@ export function initSortDropdown(select) {
}; };
const choose = (value) => { const choose = (value) => {
if (select.value === value) return; if (select.value === value) {
// Re-picking the already-selected option is normally a no-op,
// matching native <select> behavior. The seeded Random sort is
// the exception: clicking it again should reshuffle, so let the
// change handler (PageControls) generate a fresh seed.
if (String(value).startsWith('random')) {
select.dispatchEvent(new Event('change', { bubbles: true }));
}
return;
}
select.value = value; select.value = value;
select.dispatchEvent(new Event('change', { bubbles: true })); select.dispatchEvent(new Event('change', { bubbles: true }));
}; };
@@ -277,9 +286,10 @@ export function initSortDropdown(select) {
} }
// Rebuild the menu when <option>s change (VLM adds/removes a temporary // Rebuild the menu when <option>s change (VLM adds/removes a temporary
// option at runtime). // option at runtime, and the seeded Random sort option gets a new value
// attribute each time it is picked).
const observer = new MutationObserver(() => buildMenu()); const observer = new MutationObserver(() => buildMenu());
observer.observe(select, { childList: true }); observer.observe(select, { childList: true, subtree: true, attributes: true, attributeFilter: ['value'] });
buildMenu(); buildMenu();
group.dataset.sortReady = '1'; group.dataset.sortReady = '1';
+6
View File
@@ -489,6 +489,12 @@ export function createModelCard(model, modelType) {
const modelId = civitaiData?.modelId ?? civitaiData?.model_id; const modelId = civitaiData?.modelId ?? civitaiData?.model_id;
if (modelId !== undefined && modelId !== null && modelId !== '') { if (modelId !== undefined && modelId !== null && modelId !== '') {
card.dataset.modelId = 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 // LoRA specific data
+10 -2
View File
@@ -473,7 +473,14 @@ export async function showModelModal(model, modelType) {
const loadingExamplesText = translate('modals.model.loading.examples', {}, 'Loading examples...'); const loadingExamplesText = translate('modals.model.loading.examples', {}, 'Loading examples...');
const loadingVersionsText = translate('modals.model.loading.versions', {}, 'Loading versions...'); 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 civitaiVersionId = modelWithFullData.civitai?.id || '';
const navAriaLabel = translate('modals.model.navigation.label', {}, 'Model navigation'); const navAriaLabel = translate('modals.model.navigation.label', {}, 'Model navigation');
const previousTitle = translate('modals.model.navigation.previousWithShortcut', {}, 'Previous model (←)'); const previousTitle = translate('modals.model.navigation.previousWithShortcut', {}, 'Previous model (←)');
@@ -885,7 +892,8 @@ function setupEventHandlers(filePath, modelType) {
case 'view-creator': case 'view-creator':
const username = target.dataset.username; const username = target.dataset.username;
if (username) { if (username) {
window.open(`https://civitai.com/user/${username}`, '_blank'); const host = state.global.settings.civitai_host || 'civitai.com';
window.open(`https://${host}/user/${username}`, '_blank');
} }
break; break;
case 'open-file-location': case 'open-file-location':
@@ -950,6 +950,26 @@ export function initVersionsTab({
renderErrorState(container, translate('modals.model.versions.missingModelId', {}, 'This model is missing a Civitai model id.')); renderErrorState(container, translate('modals.model.versions.missingModelId', {}, 'This model is missing a Civitai model id.'));
return; 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) { if (controller.hasLoaded && !forceRefresh) {
return; return;
} }
+48 -4
View File
@@ -27,6 +27,8 @@ export class BulkManager {
// Drag detection properties // Drag detection properties
this.dragThreshold = 5; // Pixels to move before considering it a drag this.dragThreshold = 5; // Pixels to move before considering it a drag
this.dragDelayMs = 100; // Minimum hold time before a drag is treated as a marquee
this.minMarqueeSize = 10; // Minimum drag box (px) before a marquee counts as a selection
this.mouseDownTime = 0; this.mouseDownTime = 0;
this.mouseDownPosition = { x: 0, y: 0 }; this.mouseDownPosition = { x: 0, y: 0 };
@@ -88,7 +90,7 @@ export class BulkManager {
moveAll: true, moveAll: true,
autoOrganize: false, autoOrganize: false,
deleteAll: true, deleteAll: true,
setContentRating: false, setContentRating: true,
skipMetadataRefresh: false, skipMetadataRefresh: false,
setFavorite: true, setFavorite: true,
unfavorite: true, unfavorite: true,
@@ -173,6 +175,19 @@ export class BulkManager {
}); });
eventManager.addHandler('mousemove', 'bulkManager-marquee-move', (e) => { eventManager.addHandler('mousemove', 'bulkManager-marquee-move', (e) => {
// Only track marquee/drag while the left button is physically held.
// mouseup can be missed (release outside the window, focus loss, driver quirks),
// so mousemove must verify the button state itself instead of relying on it.
if (!(e.buttons & 1)) {
if (this.isMarqueeActive) {
this.endMarqueeSelection(e);
} else {
this.mouseDownTime = 0;
this.isDragging = false;
}
return false;
}
if (this.isMarqueeActive) { if (this.isMarqueeActive) {
this.lastClientX = e.clientX; this.lastClientX = e.clientX;
this.lastClientY = e.clientY; this.lastClientY = e.clientY;
@@ -184,7 +199,10 @@ export class BulkManager {
const dy = e.clientY - this.mouseDownPosition.y; const dy = e.clientY - this.mouseDownPosition.y;
const distance = Math.sqrt(dx * dx + dy * dy); const distance = Math.sqrt(dx * dx + dy * dy);
if (distance >= this.dragThreshold) { // Require both enough movement AND enough hold time so quick
// click jitter from micro-movement input devices is not a marquee.
const heldTime = Date.now() - this.mouseDownTime;
if (heldTime >= this.dragDelayMs && distance >= this.dragThreshold) {
this.isDragging = true; this.isDragging = true;
this.startMarqueeSelection(e, true); this.startMarqueeSelection(e, true);
} }
@@ -1510,14 +1528,18 @@ export class BulkManager {
let failureCount = 0; let failureCount = 0;
try { try {
const apiClient = getModelApiClient(); const isRecipesPage = state.currentPageType === 'recipes';
for (const filePath of targets) { for (const filePath of targets) {
if (cancelled) { if (cancelled) {
showToast('toast.api.operationCancelled', {}, 'info'); showToast('toast.api.operationCancelled', {}, 'info');
break; break;
} }
try { try {
await apiClient.saveModelMetadata(filePath, { preview_nsfw_level: level }); if (isRecipesPage) {
await updateRecipeMetadata(filePath, { preview_nsfw_level: level });
} else {
await getModelApiClient().saveModelMetadata(filePath, { preview_nsfw_level: level });
}
successCount++; successCount++;
} catch (error) { } catch (error) {
failureCount++; failureCount++;
@@ -1958,9 +1980,31 @@ export class BulkManager {
// Remove visual feedback class // Remove visual feedback class
document.body.classList.remove('marquee-selecting'); document.body.classList.remove('marquee-selecting');
// Compute the actual drag box size in document coordinates, matching how
// updateMarqueeSelectionFromPosition tracks the rectangle. Client-space
// size would wrongly flag auto-scroll marquees (tiny pointer movement,
// large document-space box) as accidental clicks.
const container = document.querySelector('.page-content');
const scrollX = container?.scrollLeft || 0;
const scrollY = container?.scrollTop || 0;
const dragWidth = Math.abs((e.clientX + scrollX) - this.marqueeStartDoc.x);
const dragHeight = Math.abs((e.clientY + scrollY) - this.marqueeStartDoc.y);
const isTinyMarquee = dragWidth < this.minMarqueeSize && dragHeight < this.minMarqueeSize;
// Get selection count // Get selection count
const selectionCount = state.selectedModels.size; const selectionCount = state.selectedModels.size;
// A tiny box (e.g. click jitter that happened to graze a card) is treated
// as an accidental click: undo any selection and leave bulk mode.
if (isTinyMarquee) {
this.clearSelection();
if (state.bulkMode) {
this.toggleBulkMode();
}
this.initialSelectedModels.clear();
return;
}
// If no models were selected, exit bulk mode // If no models were selected, exit bulk mode
if (selectionCount === 0) { if (selectionCount === 0) {
if (state.bulkMode) { if (state.bulkMode) {
+27 -8
View File
@@ -158,6 +158,7 @@ export class DownloadManager {
this.modelVersionId = null; this.modelVersionId = null;
this.source = null; this.source = null;
this.selectedFile = null; this.selectedFile = null;
this._isDiffusionModel = false;
this.selectedFolder = ''; this.selectedFolder = '';
this.batchModels = []; this.batchModels = [];
@@ -787,24 +788,40 @@ export class DownloadManager {
async proceedToLocationContent() { async proceedToLocationContent() {
try { try {
// Fetch model roots const _isDiffusionModel = this.selectedFile
const rootsData = await this.apiClient.fetchModelRoots(); ? (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'); const modelRoot = document.getElementById('modelRoot');
modelRoot.innerHTML = rootsData.roots.map(root => modelRoot.innerHTML = rootsData.roots.map(root =>
`<option value="${root}">${root}</option>` `<option value="${root}">${root}</option>`
).join(''); ).join('');
// Set default root if available const singularType = this._isDiffusionModel
const singularType = this.apiClient.modelType.replace(/s$/, ''); ? 'unet'
: this.apiClient.modelType.replace(/s$/, '');
const defaultRootKey = `default_${singularType}_root`; const defaultRootKey = `default_${singularType}_root`;
const defaultRoot = state.global.settings[defaultRootKey]; 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); console.log('Available roots:', rootsData.roots);
if (defaultRoot && rootsData.roots.includes(defaultRoot)) { if (defaultRoot && rootsData.roots.includes(defaultRoot)) {
console.log(`Setting default root: ${defaultRoot}`); console.log(`Setting default root: ${defaultRoot}`);
modelRoot.value = 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 // Set autocomplete="off" on folderPath input
const folderPathInput = document.getElementById('folderPath'); const folderPathInput = document.getElementById('folderPath');
if (folderPathInput) { if (folderPathInput) {
@@ -1776,13 +1793,15 @@ export class DownloadManager {
const modelRoot = document.getElementById('modelRoot').value; const modelRoot = document.getElementById('modelRoot').value;
const config = this.apiClient.apiConfig.config; 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 (modelRoot) {
if (this.useDefaultPath) { if (this.useDefaultPath) {
// Show actual template path
try { 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 templates = state.global.settings.download_path_templates;
const template = templates[singularType]; const template = templates[singularType];
fullPath += `/${template}`; fullPath += `/${template}`;
+6 -2
View File
@@ -729,10 +729,12 @@ export class FilterManager {
const pageState = getCurrentPageState(); const pageState = getCurrentPageState();
const storageKey = `${this.currentPage}_filters`; 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(); const filtersSnapshot = this.cloneFilters();
// Don't persist EMPTY_WILDCARD_MARKER - it's a runtime-only marker // Don't persist EMPTY_WILDCARD_MARKER - it's a runtime-only marker
filtersSnapshot.baseModel = filtersSnapshot.baseModel.filter(m => m !== EMPTY_WILDCARD_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); setStorageItem(storageKey, filtersSnapshot);
// Update state with current filters // Update state with current filters
@@ -984,6 +986,7 @@ export class FilterManager {
} }
cloneFilters() { cloneFilters() {
const pageState = getCurrentPageState();
return { return {
...this.filters, ...this.filters,
baseModel: [...(this.filters.baseModel || [])], baseModel: [...(this.filters.baseModel || [])],
@@ -991,7 +994,8 @@ export class FilterManager {
autoTags: { ...(this.filters.autoTags || {}) }, autoTags: { ...(this.filters.autoTags || {}) },
license: { ...(this.filters.license || {}) }, license: { ...(this.filters.license || {}) },
modelTypes: [...(this.filters.modelTypes || [])], 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; return data;
} }
async loadLoraRoots() { showNoRootsPlaceholder(select) {
try { select.innerHTML = '';
const defaultLoraRootSelect = document.getElementById('defaultLoraRoot'); const option = document.createElement('option');
if (!defaultLoraRootSelect) return; 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 // Fetch lora roots
const response = await fetch('/api/lm/loras/roots'); const response = await fetch('/api/lm/loras/roots');
if (!response.ok) { if (!response.ok) {
@@ -1530,10 +1539,12 @@ export class SettingsManager {
const data = await response.json(); const data = await response.json();
if (!data.roots || data.roots.length === 0) { if (!data.roots || data.roots.length === 0) {
throw new Error('No LoRA roots found'); this.showNoRootsPlaceholder(defaultLoraRootSelect);
return;
} }
defaultLoraRootSelect.innerHTML = ''; defaultLoraRootSelect.innerHTML = '';
defaultLoraRootSelect.disabled = false;
// Add options for each root // Add options for each root
data.roots.forEach(root => { data.roots.forEach(root => {
@@ -1548,15 +1559,16 @@ export class SettingsManager {
} catch (error) { } catch (error) {
console.error('Error loading LoRA roots:', error); console.error('Error loading LoRA roots:', error);
this.showNoRootsPlaceholder(defaultLoraRootSelect);
showToast('toast.settings.loraRootsFailed', { message: error.message }, 'error'); showToast('toast.settings.loraRootsFailed', { message: error.message }, 'error');
} }
} }
async loadCheckpointRoots() { async loadCheckpointRoots() {
try { const defaultCheckpointRootSelect = document.getElementById('defaultCheckpointRoot');
const defaultCheckpointRootSelect = document.getElementById('defaultCheckpointRoot'); if (!defaultCheckpointRootSelect) return;
if (!defaultCheckpointRootSelect) return;
try {
// Fetch checkpoint roots (checkpoint paths only, not unet) // Fetch checkpoint roots (checkpoint paths only, not unet)
const response = await fetch('/api/lm/checkpoints/checkpoints_roots'); const response = await fetch('/api/lm/checkpoints/checkpoints_roots');
if (!response.ok) { if (!response.ok) {
@@ -1565,10 +1577,12 @@ export class SettingsManager {
const data = await response.json(); const data = await response.json();
if (!data.roots || data.roots.length === 0) { if (!data.roots || data.roots.length === 0) {
throw new Error('No checkpoint roots found'); this.showNoRootsPlaceholder(defaultCheckpointRootSelect);
return;
} }
defaultCheckpointRootSelect.innerHTML = ''; defaultCheckpointRootSelect.innerHTML = '';
defaultCheckpointRootSelect.disabled = false;
// Add options for each root // Add options for each root
data.roots.forEach(root => { data.roots.forEach(root => {
@@ -1583,15 +1597,16 @@ export class SettingsManager {
} catch (error) { } catch (error) {
console.error('Error loading checkpoint roots:', error); console.error('Error loading checkpoint roots:', error);
this.showNoRootsPlaceholder(defaultCheckpointRootSelect);
showToast('toast.settings.checkpointRootsFailed', { message: error.message }, 'error'); showToast('toast.settings.checkpointRootsFailed', { message: error.message }, 'error');
} }
} }
async loadUnetRoots() { async loadUnetRoots() {
try { const defaultUnetRootSelect = document.getElementById('defaultUnetRoot');
const defaultUnetRootSelect = document.getElementById('defaultUnetRoot'); if (!defaultUnetRootSelect) return;
if (!defaultUnetRootSelect) return;
try {
// Fetch unet roots (diffusion model paths only) // Fetch unet roots (diffusion model paths only)
const response = await fetch('/api/lm/checkpoints/unet_roots'); const response = await fetch('/api/lm/checkpoints/unet_roots');
if (!response.ok) { if (!response.ok) {
@@ -1600,10 +1615,12 @@ export class SettingsManager {
const data = await response.json(); const data = await response.json();
if (!data.roots || data.roots.length === 0) { if (!data.roots || data.roots.length === 0) {
throw new Error('No diffusion model roots found'); this.showNoRootsPlaceholder(defaultUnetRootSelect);
return;
} }
defaultUnetRootSelect.innerHTML = ''; defaultUnetRootSelect.innerHTML = '';
defaultUnetRootSelect.disabled = false;
// Add options for each root // Add options for each root
data.roots.forEach(root => { data.roots.forEach(root => {
@@ -1618,15 +1635,16 @@ export class SettingsManager {
} catch (error) { } catch (error) {
console.error('Error loading diffusion model roots:', error); console.error('Error loading diffusion model roots:', error);
this.showNoRootsPlaceholder(defaultUnetRootSelect);
showToast('toast.settings.unetRootsFailed', { message: error.message }, 'error'); showToast('toast.settings.unetRootsFailed', { message: error.message }, 'error');
} }
} }
async loadEmbeddingRoots() { async loadEmbeddingRoots() {
try { const defaultEmbeddingRootSelect = document.getElementById('defaultEmbeddingRoot');
const defaultEmbeddingRootSelect = document.getElementById('defaultEmbeddingRoot'); if (!defaultEmbeddingRootSelect) return;
if (!defaultEmbeddingRootSelect) return;
try {
// Fetch embedding roots // Fetch embedding roots
const response = await fetch('/api/lm/embeddings/roots'); const response = await fetch('/api/lm/embeddings/roots');
if (!response.ok) { if (!response.ok) {
@@ -1635,10 +1653,12 @@ export class SettingsManager {
const data = await response.json(); const data = await response.json();
if (!data.roots || data.roots.length === 0) { if (!data.roots || data.roots.length === 0) {
throw new Error('No embedding roots found'); this.showNoRootsPlaceholder(defaultEmbeddingRootSelect);
return;
} }
defaultEmbeddingRootSelect.innerHTML = ''; defaultEmbeddingRootSelect.innerHTML = '';
defaultEmbeddingRootSelect.disabled = false;
// Add options for each root // Add options for each root
data.roots.forEach(root => { data.roots.forEach(root => {
@@ -1653,6 +1673,7 @@ export class SettingsManager {
} catch (error) { } catch (error) {
console.error('Error loading embedding roots:', error); console.error('Error loading embedding roots:', error);
this.showNoRootsPlaceholder(defaultEmbeddingRootSelect);
showToast('toast.settings.embeddingRootsFailed', { message: error.message }, 'error'); showToast('toast.settings.embeddingRootsFailed', { message: error.message }, 'error');
} }
} }
+282 -39
View File
@@ -6,6 +6,7 @@ import {
setStoredVersionInfo, setStoredVersionInfo,
isVersionMatch isVersionMatch
} from '../utils/storageHelpers.js'; } from '../utils/storageHelpers.js';
import { state } from '../state/index.js';
import { bannerService } from './BannerService.js'; import { bannerService } from './BannerService.js';
import { translate } from '../utils/i18nHelpers.js'; import { translate } from '../utils/i18nHelpers.js';
@@ -24,7 +25,11 @@ export class UpdateService {
this.updateNotificationsEnabled = getStorageItem('show_update_notifications', true); this.updateNotificationsEnabled = getStorageItem('show_update_notifications', true);
this.lastCheckTime = parseInt(getStorageItem('last_update_check') || '0'); this.lastCheckTime = parseInt(getStorageItem('last_update_check') || '0');
this.isUpdating = false; 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.currentVersionInfo = null;
this.versionMismatch = false; this.versionMismatch = false;
this.activeNotificationTab = 'updates'; this.activeNotificationTab = 'updates';
@@ -49,43 +54,180 @@ export class UpdateService {
updateBtn.addEventListener('click', () => this.performUpdate()); updateBtn.addEventListener('click', () => this.performUpdate());
} }
// Register event listener for nightly update toggle this.wireChannelButtons();
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.setupNotificationCenter(); this.setupNotificationCenter();
window.addEventListener('lm:banner-history-updated', this.handleBannerHistoryUpdated); window.addEventListener('lm:banner-history-updated', this.handleBannerHistoryUpdated);
this.updateTabBadges(); this.updateTabBadges();
// Perform update check if needed // Perform update check if needed
this.checkForUpdates().then(() => { this.checkVersionInfo().then(() => {
// Ensure badges are updated after checking this.checkForUpdates().then(() => {
this.updateBadgeVisibility(); this.updateBadgeVisibility();
});
}); });
// Immediately update modal content with current values (even if from default)
this.updateModalContent(); this.updateModalContent();
// Check version info for mismatch after loading basic info
this.checkVersionInfo();
} }
updateNightlyWarning() { wireChannelButtons() {
const warning = document.getElementById('nightlyWarning'); const releaseBtn = document.getElementById('channelRelease');
if (warning) { const nightlyBtn = document.getElementById('channelNightly');
warning.style.display = this.nightlyMode ? 'flex' : 'none'; 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() { setupNotificationCenter() {
@@ -355,6 +497,18 @@ export class UpdateService {
} }
async checkForUpdates({ force = false } = {}) { 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) { if (!force && !this.updateNotificationsEnabled) {
return; return;
} }
@@ -373,7 +527,8 @@ export class UpdateService {
try { try {
// Call backend API to check for updates with nightly flag // 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(); const data = await response.json();
if (data.success) { if (data.success) {
@@ -381,17 +536,35 @@ export class UpdateService {
this.latestVersion = data.latest_version || "v0.0.0"; this.latestVersion = data.latest_version || "v0.0.0";
this.updateInfo = data; this.updateInfo = data;
this.gitInfo = data.git_info || this.gitInfo; this.gitInfo = data.git_info || this.gitInfo;
this.hasGit = data.has_git || false;
// Explicitly set update availability based on version comparison if (needsMigration || this.channelMode === null) {
this.updateAvailable = this.isNewerVersion(this.latestVersion, this.currentVersion); 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();
}
// Update last check time
this.lastCheckTime = now; this.lastCheckTime = now;
setStorageItem('last_update_check', now.toString()); setStorageItem('last_update_check', now.toString());
// Update UI
this.updateBadgeVisibility(); this.updateBadgeVisibility();
this.updateModalContent(); this.updateModalContent();
this.updateChannelUI();
console.log("Update check complete:", { console.log("Update check complete:", {
currentVersion: this.currentVersion, currentVersion: this.currentVersion,
@@ -436,6 +609,28 @@ export class UpdateService {
return false; return false;
} }
_getTodayKey() {
const now = new Date();
const month = String(now.getMonth() + 1).padStart(2, '0');
const day = String(now.getDate()).padStart(2, '0');
return `${now.getFullYear()}-${month}-${day}`;
}
_isNightlyBadgeAllowed() {
if (this.channelMode !== 'nightly') {
return true;
}
// Keep the badge visible for the rest of the session once shown, but do
// not show it again on later sessions within the same calendar day.
return this.nightlyNotifyDate !== this._getTodayKey() || this.nightlyBadgeShown;
}
_markNightlyNotified() {
this.nightlyNotifyDate = this._getTodayKey();
this.nightlyBadgeShown = true;
setStorageItem('nightly_notify_date', this.nightlyNotifyDate);
}
updateBadgeVisibility() { updateBadgeVisibility() {
const updateToggle = document.querySelector('.update-toggle'); const updateToggle = document.querySelector('.update-toggle');
const updateBadge = document.querySelector('.update-toggle .update-badge'); const updateBadge = document.querySelector('.update-toggle .update-badge');
@@ -443,9 +638,12 @@ export class UpdateService {
? bannerService.getUnreadBannerCount() ? bannerService.getUnreadBannerCount()
: 0; : 0;
// Force updating badges visibility based on current state
const shouldShowUpdate = this.updateNotificationsEnabled && this.updateAvailable && this._isNightlyBadgeAllowed();
if (updateToggle) { if (updateToggle) {
let tooltipKey = 'header.actions.notifications'; let tooltipKey = 'header.actions.notifications';
if (this.updateNotificationsEnabled && this.updateAvailable) { if (shouldShowUpdate) {
tooltipKey = 'update.updateAvailable'; tooltipKey = 'update.updateAvailable';
} else if (unreadBanners > 0) { } else if (unreadBanners > 0) {
tooltipKey = 'update.tabs.messages'; tooltipKey = 'update.tabs.messages';
@@ -453,8 +651,6 @@ export class UpdateService {
updateToggle.title = translate(tooltipKey); updateToggle.title = translate(tooltipKey);
} }
// Force updating badges visibility based on current state
const shouldShowUpdate = this.updateNotificationsEnabled && this.updateAvailable;
const shouldShow = shouldShowUpdate || unreadBanners > 0; const shouldShow = shouldShowUpdate || unreadBanners > 0;
if (updateBadge) { if (updateBadge) {
@@ -482,8 +678,31 @@ export class UpdateService {
if (currentVersionEl) currentVersionEl.textContent = this.currentVersion; 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) { 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 // Update update button state
@@ -599,8 +818,12 @@ export class UpdateService {
// Update GitHub link to point to the specific release if available // Update GitHub link to point to the specific release if available
const githubLink = modal.querySelector('.update-link'); const githubLink = modal.querySelector('.update-link');
if (githubLink && this.latestVersion) { if (githubLink && this.latestVersion) {
const versionTag = this.latestVersion.replace(/^v/, ''); if (this.updateInfo?.nightly) {
githubLink.href = `https://github.com/willmiao/ComfyUI-Lora-Manager/releases/tag/v${versionTag}`; 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' 'Content-Type': 'application/json'
}, },
body: JSON.stringify({ body: JSON.stringify({
nightly: this.nightlyMode nightly: this.channelMode === 'nightly'
}) })
}); });
@@ -699,6 +922,25 @@ export class UpdateService {
} }
} }
_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) { showUpdateCompleteMessage(newVersion) {
const modal = document.getElementById('updateModal'); const modal = document.getElementById('updateModal');
if (!modal) return; if (!modal) return;
@@ -771,6 +1013,7 @@ export class UpdateService {
// Update the modal content immediately with current data // Update the modal content immediately with current data
this.updateModalContent(); this.updateModalContent();
this.updateChannelUI();
this.renderRecentBanners(); this.renderRecentBanners();
// Show the modal with current data // Show the modal with current data
@@ -801,8 +1044,8 @@ export class UpdateService {
if (data.success) { if (data.success) {
this.currentVersionInfo = data.version; this.currentVersionInfo = data.version;
this.hasGit = data.has_git || false;
// Check if version matches stored version
this.versionMismatch = !isVersionMatch(this.currentVersionInfo); this.versionMismatch = !isVersionMatch(this.currentVersionInfo);
if (this.versionMismatch) { if (this.versionMismatch) {
+18 -3
View File
@@ -3,6 +3,7 @@ import { translate } from '../../utils/i18nHelpers.js';
import { getModelApiClient } from '../../api/modelApiFactory.js'; import { getModelApiClient } from '../../api/modelApiFactory.js';
import { MODEL_TYPES } from '../../api/apiConfig.js'; import { MODEL_TYPES } from '../../api/apiConfig.js';
import { getStorageItem } from '../../utils/storageHelpers.js'; import { getStorageItem } from '../../utils/storageHelpers.js';
import { state } from '../../state/index.js';
export class DownloadManager { export class DownloadManager {
constructor(importManager) { constructor(importManager) {
@@ -125,11 +126,25 @@ export class DownloadManager {
showToast('toast.recipes.nameSaved', { name: this.importManager.recipeName }, 'success'); showToast('toast.recipes.nameSaved', { name: this.importManager.recipeName }, 'success');
} }
// Close modal
modalManager.closeModal('importModal'); modalManager.closeModal('importModal');
// Refresh the recipe if (isDownloadOnly && state.virtualScroller) {
window.recipeManager.loadRecipes(true); 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) { } catch (error) {
console.error('Error:', error); console.error('Error:', error);
+7 -3
View File
@@ -333,6 +333,7 @@ export const PATH_TEMPLATE_PLACEHOLDERS = [
export const DEFAULT_PATH_TEMPLATES = { export const DEFAULT_PATH_TEMPLATES = {
lora: '{base_model}/{first_tag}', lora: '{base_model}/{first_tag}',
checkpoint: '{base_model}', checkpoint: '{base_model}',
unet: '{base_model}',
embedding: '{first_tag}' embedding: '{first_tag}'
}; };
@@ -369,21 +370,24 @@ export function getMatureBlurThreshold(settings = {}) {
export const NODE_TYPES = { export const NODE_TYPES = {
LORA_LOADER: 1, LORA_LOADER: 1,
LORA_STACKER: 2, LORA_STACKER: 2,
WAN_VIDEO_LORA_SELECT: 3 WAN_VIDEO_LORA_SELECT: 3,
HOOK_LORA: 4
}; };
// Node type names to IDs mapping // Node type names to IDs mapping
export const NODE_TYPE_NAMES = { export const NODE_TYPE_NAMES = {
"Lora Loader (LoraManager)": NODE_TYPES.LORA_LOADER, "Lora Loader (LoraManager)": NODE_TYPES.LORA_LOADER,
"Lora Stacker (LoraManager)": NODE_TYPES.LORA_STACKER, "Lora Stacker (LoraManager)": NODE_TYPES.LORA_STACKER,
"WanVideo Lora Select (LoraManager)": NODE_TYPES.WAN_VIDEO_LORA_SELECT "WanVideo Lora Select (LoraManager)": NODE_TYPES.WAN_VIDEO_LORA_SELECT,
"Create Hook LoRA (LoraManager)": NODE_TYPES.HOOK_LORA
}; };
// Node type icons // Node type icons
export const NODE_TYPE_ICONS = { export const NODE_TYPE_ICONS = {
[NODE_TYPES.LORA_LOADER]: "fas fa-l", [NODE_TYPES.LORA_LOADER]: "fas fa-l",
[NODE_TYPES.LORA_STACKER]: "fas fa-s", [NODE_TYPES.LORA_STACKER]: "fas fa-s",
[NODE_TYPES.WAN_VIDEO_LORA_SELECT]: "fas fa-w" [NODE_TYPES.WAN_VIDEO_LORA_SELECT]: "fas fa-w",
[NODE_TYPES.HOOK_LORA]: "fas fa-h"
}; };
// Default ComfyUI node color when bgcolor is null // Default ComfyUI node color when bgcolor is null
+6 -1
View File
@@ -32,7 +32,12 @@
<div class="context-menu-separator menu-section-break"></div> <div class="context-menu-separator menu-section-break"></div>
<!-- Media / Preview --> <!-- Media / Preview -->
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div> <div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }}</div> <div class="context-menu-item has-submenu" data-has-submenu="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }} <i class="fas fa-chevron-right submenu-arrow"></i>
<div class="context-submenu">
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadMissingExamples') }}</div>
<div class="context-menu-item" data-action="download-examples-force"><i class="fas fa-redo-alt"></i> {{ t('loras.contextMenu.reprocessExamples') }}</div>
</div>
</div>
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div> <div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
<div class="context-menu-separator menu-section-break"></div> <div class="context-menu-separator menu-section-break"></div>
<!-- Attributes --> <!-- Attributes -->
+24 -4
View File
@@ -44,8 +44,18 @@
<div class="context-menu-item" data-action="preview"> <div class="context-menu-item" data-action="preview">
<i class="fas fa-folder-open"></i> <span>{{ t('loras.contextMenu.openExamples') }}</span> <i class="fas fa-folder-open"></i> <span>{{ t('loras.contextMenu.openExamples') }}</span>
</div> </div>
<div class="context-menu-item" data-action="download-examples"> <div class="context-menu-item has-submenu" data-has-submenu="download-examples">
<i class="fas fa-download"></i> <span>{{ t('loras.contextMenu.downloadExamples') }}</span> <i class="fas fa-download"></i>
<span>{{ t('loras.contextMenu.downloadExamples') }}</span>
<i class="fas fa-chevron-right submenu-arrow"></i>
<div class="context-submenu">
<div class="context-menu-item" data-action="download-examples">
<i class="fas fa-download"></i> <span>{{ t('loras.contextMenu.downloadMissingExamples') }}</span>
</div>
<div class="context-menu-item" data-action="download-examples-force">
<i class="fas fa-redo-alt"></i> <span>{{ t('loras.contextMenu.reprocessExamples') }}</span>
</div>
</div>
</div> </div>
<div class="context-menu-item" data-action="replace-preview"> <div class="context-menu-item" data-action="replace-preview">
<i class="fas fa-image"></i> <span>{{ t('loras.contextMenu.replacePreview') }}</span> <i class="fas fa-image"></i> <span>{{ t('loras.contextMenu.replacePreview') }}</span>
@@ -136,8 +146,18 @@
</div> </div>
<div class="context-menu-section" data-section="download"> <div class="context-menu-section" data-section="download">
<div class="context-menu-section-header">{{ t('loras.bulkOperations.sections.download') }}</div> <div class="context-menu-section-header">{{ t('loras.bulkOperations.sections.download') }}</div>
<div class="context-menu-item" data-action="download-example-images"> <div class="context-menu-item has-submenu" data-has-submenu="download-example-images">
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadExamples') }}</span> <i class="fas fa-download"></i>
<span>{{ t('loras.bulkOperations.downloadExamples') }}</span>
<i class="fas fa-chevron-right submenu-arrow"></i>
<div class="context-submenu">
<div class="context-menu-item" data-action="download-missing-example-images">
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingExamples') }}</span>
</div>
<div class="context-menu-item" data-action="download-example-images">
<i class="fas fa-redo-alt"></i> <span>{{ t('loras.bulkOperations.reprocessExamples') }}</span>
</div>
</div>
</div> </div>
<div class="context-menu-item" data-action="download-missing-loras"> <div class="context-menu-item" data-action="download-missing-loras">
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingLoras') }}</span> <i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingLoras') }}</span>
+5
View File
@@ -48,6 +48,11 @@
<option value="versions_count:asc">{{ t('loras.controls.sort.versionsCountAsc', default='Fewest versions first') }}</option> <option value="versions_count:asc">{{ t('loras.controls.sort.versionsCountAsc', default='Fewest versions first') }}</option>
</optgroup> </optgroup>
{% endif %} {% endif %}
{% if page_id != 'recipes' %}
<optgroup label="{{ t('loras.controls.sort.random', default='Random') }}">
<option value="random">{{ t('loras.controls.sort.randomAction', default='Randomize (shuffle)') }}</option>
</optgroup>
{% endif %}
{% if page_id == 'recipes' %} {% if page_id == 'recipes' %}
<optgroup label="{{ t('recipes.controls.sort.lorasCount') }}"> <optgroup label="{{ t('recipes.controls.sort.lorasCount') }}">
<option value="loras_count:desc">{{ t('recipes.controls.sort.lorasCountDesc') }}</option> <option value="loras_count:desc">{{ t('recipes.controls.sort.lorasCountDesc') }}</option>
@@ -19,6 +19,20 @@
<div class="notification-panels"> <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="notification-panel active" id="updatesPanel" role="tabpanel" aria-labelledby="updatesTab" aria-hidden="false" tabindex="0" data-notification-panel="updates">
<div class="update-content"> <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="update-info">
<div class="version-info"> <div class="version-info">
<div class="current-version"> <div class="current-version">
+6 -1
View File
@@ -32,7 +32,12 @@
<div class="context-menu-separator menu-section-break"></div> <div class="context-menu-separator menu-section-break"></div>
<!-- Media / Preview --> <!-- Media / Preview -->
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div> <div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }}</div> <div class="context-menu-item has-submenu" data-has-submenu="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }} <i class="fas fa-chevron-right submenu-arrow"></i>
<div class="context-submenu">
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadMissingExamples') }}</div>
<div class="context-menu-item" data-action="download-examples-force"><i class="fas fa-redo-alt"></i> {{ t('loras.contextMenu.reprocessExamples') }}</div>
</div>
</div>
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div> <div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
<div class="context-menu-separator menu-section-break"></div> <div class="context-menu-separator menu-section-break"></div>
<!-- Attributes --> <!-- Attributes -->
+1
View File
@@ -85,6 +85,7 @@ sys.modules['comfy.utils'] = comfy_mock.utils
sys.modules['comfy.sd'] = comfy_mock.sd sys.modules['comfy.sd'] = comfy_mock.sd
sys.modules['comfy.model_management'] = comfy_mock.model_management sys.modules['comfy.model_management'] = comfy_mock.model_management
sys.modules['comfy.comfy_types'] = comfy_mock.comfy_types sys.modules['comfy.comfy_types'] = comfy_mock.comfy_types
sys.modules['comfy.hooks'] = MockModule("comfy.hooks")
execution_mock = MockModule("execution") execution_mock = MockModule("execution")
execution_mock.PromptExecutor = mock.MagicMock() execution_mock.PromptExecutor = mock.MagicMock()
@@ -2155,4 +2155,35 @@ describe('Interaction-level regression coverage', () => {
excludedItem.dispatchEvent(new Event('click', { bubbles: true })); excludedItem.dispatchEvent(new Event('click', { bubbles: true }));
expect(window.pageControls.enterExcludedView).toHaveBeenCalledTimes(1); expect(window.pageControls.enterExcludedView).toHaveBeenCalledTimes(1);
}); });
it('routes single-model example downloads to missing-only and force paths', async () => {
document.body.innerHTML = `
<div id="loraContextMenu" class="context-menu">
<div class="context-menu-item has-submenu" data-has-submenu="download-examples">
<div class="context-submenu">
<div class="context-menu-item" data-action="download-examples"></div>
<div class="context-menu-item" data-action="download-examples-force"></div>
</div>
</div>
</div>
`;
const { LoraContextMenu } = await import('../../../static/js/components/ContextMenu/LoraContextMenu.js');
const contextMenu = new LoraContextMenu();
const card = document.createElement('div');
card.className = 'model-card';
card.dataset.filepath = '/models/test.safetensors';
card.dataset.sha256 = 'abc123hash';
document.body.appendChild(card);
contextMenu.showMenu(100, 100, card);
document.querySelector('[data-action="download-examples"]').dispatchEvent(new Event('click', { bubbles: true }));
expect(downloadExampleImagesApiMock).toHaveBeenCalledWith(['abc123hash'], null, { force: false });
contextMenu.showMenu(100, 100, card);
document.querySelector('[data-action="download-examples-force"]').dispatchEvent(new Event('click', { bubbles: true }));
expect(downloadExampleImagesApiMock).toHaveBeenCalledWith(['abc123hash'], null, { force: true });
});
}); });
@@ -0,0 +1,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);
});
});
@@ -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', () => { describe('SettingsManager library controls', () => {
it('loads libraries and populates the select', async () => { it('loads libraries and populates the select', async () => {
const manager = createManager(); const manager = createManager();
+123 -2
View File
@@ -1,12 +1,26 @@
import { describe, beforeEach, afterEach, expect, it, vi } from 'vitest'; import { describe, beforeEach, afterEach, expect, it, vi } from 'vitest';
import { UpdateService } from '../../../static/js/managers/UpdateService.js'; import { UpdateService } from '../../../static/js/managers/UpdateService.js';
import { state } from '../../../static/js/state/index.js';
function createFetchResponse(payload) { function createFetchResponse(payload) {
return { 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', () => { describe('UpdateService passive checks', () => {
let service; let service;
let fetchMock; let fetchMock;
@@ -16,10 +30,13 @@ describe('UpdateService passive checks', () => {
success: true, success: true,
current_version: 'v1.0.0', current_version: 'v1.0.0',
latest_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; global.fetch = fetchMock;
stubSettingsUpdateChannel('release');
service = new UpdateService(); service = new UpdateService();
service.updateNotificationsEnabled = false; service.updateNotificationsEnabled = false;
service.lastCheckTime = 0; service.lastCheckTime = 0;
@@ -28,6 +45,7 @@ describe('UpdateService passive checks', () => {
afterEach(() => { afterEach(() => {
delete global.fetch; delete global.fetch;
clearSettingsUpdateChannel();
}); });
it('skips passive update checks when notifications are disabled', async () => { 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'); 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);
});
});
@@ -30,10 +30,10 @@ def test_metadata_hook_installs_and_traces_execution(monkeypatch, metadata_regis
calls = [] calls = []
def record_stub(self, node_id, class_type, inputs, outputs): def record_stub(self, node_id, class_type, inputs, outputs, return_types=None):
calls.append(("record", node_id, class_type, inputs)) calls.append(("record", node_id, class_type, inputs))
def update_stub(self, node_id, class_type, outputs): def update_stub(self, node_id, class_type, outputs, return_types=None):
calls.append(("update", node_id, class_type, outputs)) calls.append(("update", node_id, class_type, outputs))
monkeypatch.setattr(MetadataRegistry, "record_node_execution", record_stub) monkeypatch.setattr(MetadataRegistry, "record_node_execution", record_stub)
@@ -820,3 +820,227 @@ def test_lora_manager_checkpoint_and_unet_loaders_extract_models(metadata_regist
"type": "checkpoint", "type": "checkpoint",
"node_id": "unet_node", "node_id": "unet_node",
} }
# ---------------------------------------------------------------------------
# MetadataOverwriteExtractor & overwrite merge tests
# ---------------------------------------------------------------------------
from py.metadata_collector.constants import OVERWRITE, METADATA_OVERWRITE_FIELDS
from py.metadata_collector.node_extractors import MetadataOverwriteExtractor
def test_metadata_overwrite_extractor_stores_truthy_values(metadata_registry):
"""Extractor should store truthy inputs under the OVERWRITE category."""
metadata_registry.start_collection("prompt-ow")
metadata = metadata_registry.prompt_metadata["prompt-ow"]
inputs = {
"prompt": "a beautiful landscape",
"negative_prompt": "",
"seed": 42,
"steps": 0,
"cfg_scale": 7.5,
"sampler": "",
"scheduler": "",
"model": "myModel.safetensors",
"loras": "<lora:detail:0.8>",
"size": "1024x768",
"clip_skip": 0,
"additional_data": '{"Copyright": "CC0"}',
}
MetadataOverwriteExtractor.extract("ow-1", inputs, None, metadata)
assert OVERWRITE in metadata
assert "ow-1" in metadata[OVERWRITE]
params = metadata[OVERWRITE]["ow-1"]["parameters"]
# Truthy values stored
assert params["prompt"] == "a beautiful landscape"
assert params["seed"] == 42
assert params["cfg_scale"] == 7.5
assert params["model"] == "myModel.safetensors"
assert params["loras"] == "<lora:detail:0.8>"
assert params["size"] == "1024x768"
assert params["additional_data"] == '{"Copyright": "CC0"}'
# Falsy values NOT stored
assert "negative_prompt" not in params
assert "steps" not in params
assert "sampler" not in params
assert "scheduler" 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()
def test_metadata_overwrite_extractor_empty_inputs(metadata_registry):
"""Extractor with all-falsy inputs should NOT create OVERWRITE category."""
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": CLIP_SKIP_SENTINEL})
MetadataOverwriteExtractor.extract("ow-2", inputs, None, metadata)
# start_collection pre-creates empty dicts for all categories,
# but no node should have populated OVERWRITE with any data
assert not metadata[OVERWRITE]
metadata_registry.clear_metadata()
def test_extract_generation_params_applies_overwrite(metadata_registry, populated_registry, monkeypatch):
"""overwrite values should replace inferred params in extract_generation_params."""
import py.metadata_collector.metadata_processor as mp
monkeypatch.setattr(mp, "standalone_mode", False)
metadata = populated_registry["metadata"]
registry_obj = populated_registry["registry"]
# Simulate the MetadataOverwriteLM node having been executed with overwrite values
registry_obj.start_collection("promptA")
# Re-populate with the same data (start_collection resets)
registry_obj.set_current_prompt(populated_registry["prompt"])
metadata2 = registry_obj.prompt_metadata["promptA"]
# Inject overwrite data into metadata
metadata2[OVERWRITE] = {
"ow-1": {
"parameters": {
"seed": 777,
"additional_data": '{"AuthorURL": "https://civitai.com/user/foo"}',
},
"node_id": "ow-1",
}
}
# Copy other categories from original populated metadata
for cat in ("models", "prompts", "sampling", "loras", "size", "images"):
if cat in metadata:
metadata2[cat] = metadata[cat]
metadata2["execution_order"] = metadata["execution_order"]
params = MetadataProcessor.extract_generation_params(metadata2, id="vae")
# Overwritten values
assert params["seed"] == 777
assert params["additional_data"] == '{"AuthorURL": "https://civitai.com/user/foo"}'
# Inferred values still present (not overwritten)
assert params["prompt"] == "A castle on a hill"
assert params["cfg_scale"] == 7.5
assert params["checkpoint"] == "model.safetensors"
registry_obj.clear_metadata()
def test_extract_generation_params_overwrite_falsy_skipped(metadata_registry, populated_registry, monkeypatch):
"""Overwrite entries with falsy values should NOT replace inferred params."""
import py.metadata_collector.metadata_processor as mp
monkeypatch.setattr(mp, "standalone_mode", False)
metadata = populated_registry["metadata"]
registry_obj = populated_registry["registry"]
registry_obj.start_collection("promptA")
registry_obj.set_current_prompt(populated_registry["prompt"])
metadata2 = registry_obj.prompt_metadata["promptA"]
# 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": {
"seed": 0,
"steps": 0,
"cfg_scale": 0.0,
"prompt": "",
"clip_skip": 0,
},
"node_id": "ow-1",
}
}
for cat in ("models", "prompts", "sampling", "loras", "size", "images"):
if cat in metadata:
metadata2[cat] = metadata[cat]
metadata2["execution_order"] = metadata["execution_order"]
params = MetadataProcessor.extract_generation_params(metadata2, id="vae")
# Falsy overwrites should NOT have replaced inferred values
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()
def test_fill_missing_metadata_skips_overwrite_for_bypassed_node(metadata_registry):
"""Bypassed (mode=4) node should not have OVERWRITE filled from cache."""
metadata_registry.start_collection("prompt-bypass")
# Simulate a previous execution that cached overwrite data
metadata_registry.record_node_execution(
"ow-1",
"MetadataOverwriteLM",
{"seed": 99, "prompt": "test", "steps": 0, "cfg_scale": 0.0,
"negative_prompt": "", "sampler": "", "scheduler": "", "model": "",
"loras": "", "size": "", "clip_skip": 0, "additional_data": ""},
None,
)
# Now start a new prompt where the node is bypassed (mode=4)
metadata_registry.start_collection("prompt-bypass-2")
original_prompt = {
"ow-1": {"class_type": "MetadataOverwriteLM", "inputs": {}, "mode": 4},
}
metadata_registry.set_current_prompt(
SimpleNamespace(original_prompt=original_prompt)
)
metadata = metadata_registry.get_metadata("prompt-bypass-2")
# The overwrite data should NOT be present (node was bypassed, not
# a cache hit — it should not inherit previous execution's overwrite)
assert "ow-1" not in metadata.get(OVERWRITE, {})
metadata_registry.clear_metadata()
def test_fill_missing_metadata_fills_overwrite_for_muted_node(metadata_registry):
"""Muted (mode=2) node should also not have OVERWRITE filled from cache."""
metadata_registry.start_collection("prompt-mute")
# Simulate a previous execution that cached overwrite data
metadata_registry.record_node_execution(
"ow-1",
"MetadataOverwriteLM",
{"seed": 88, "prompt": "test2", "steps": 0, "cfg_scale": 0.0,
"negative_prompt": "", "sampler": "", "scheduler": "", "model": "",
"loras": "", "size": "", "clip_skip": 0, "additional_data": ""},
None,
)
# Start a new prompt where the node is muted (mode=2)
metadata_registry.start_collection("prompt-mute-2")
original_prompt = {
"ow-1": {"class_type": "MetadataOverwriteLM", "inputs": {}, "mode": 2},
}
metadata_registry.set_current_prompt(
SimpleNamespace(original_prompt=original_prompt)
)
metadata = metadata_registry.get_metadata("prompt-mute-2")
assert "ow-1" not in metadata.get(OVERWRITE, {})
metadata_registry.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(): def test_combine_stacks_preserves_order():
@@ -49,3 +56,109 @@ def test_combine_stacks_allows_duplicate_entries():
(combined_stack,) = node.combine_stacks([duplicate_entry], [duplicate_entry]) (combined_stack,) = node.combine_stacks([duplicate_entry], [duplicate_entry])
assert combined_stack == [duplicate_entry, duplicate_entry] assert combined_stack == [duplicate_entry, duplicate_entry]
def test_combine_stacks_returns_empty_when_both_unconnected():
node = LoraStackCombinerLM()
(combined_stack,) = node.combine_stacks()
assert combined_stack == []
def test_combine_stacks_returns_other_when_one_unconnected():
node = LoraStackCombinerLM()
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
(combined_stack_a,) = node.combine_stacks(lora_stack1=stack_a)
(combined_stack_b,) = node.combine_stacks(lora_stack2=stack_a)
assert combined_stack_a == stack_a
assert combined_stack_b == stack_a
def test_combine_stacks_with_dynamic_third_slot():
node = LoraStackCombinerLM()
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
(combined_stack,) = node.combine_stacks(
lora_stack1=stack_a, lora_stack2=stack_b, lora_stack3=stack_c
)
assert combined_stack == stack_a + stack_b + stack_c
def test_combine_stacks_orders_by_slot_number_not_call_order():
node = LoraStackCombinerLM()
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
(combined_stack,) = node.combine_stacks(
lora_stack3=stack_c, lora_stack2=stack_b, lora_stack1=stack_a
)
assert combined_stack == stack_a + stack_b + stack_c
def test_combine_stacks_accepts_only_dynamic_slot():
node = LoraStackCombinerLM()
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
(combined_stack,) = node.combine_stacks(lora_stack3=stack_c)
assert combined_stack == stack_c
def test_combine_stacks_handles_legacy_input_names():
node = LoraStackCombinerLM()
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
(combined_stack,) = node.combine_stacks(lora_stack_a=stack_a, lora_stack_b=stack_b)
assert combined_stack == stack_a + stack_b
def test_input_types_exposes_two_default_slots():
input_types = LoraStackCombinerLM.INPUT_TYPES()
assert set(input_types["optional"]) == {"lora_stack1", "lora_stack2"}
assert input_types["optional"]["lora_stack1"][0] == "LORA_STACK"
assert input_types["optional"]["lora_stack2"][0] == "LORA_STACK"
def test_input_types_recognizes_dynamic_slots_from_get_input_info(monkeypatch):
frames = [None, None, types.SimpleNamespace(function="get_input_info")]
monkeypatch.setattr(
"py.nodes.lora_stack_combiner.inspect.stack", lambda: frames
)
input_types = LoraStackCombinerLM.INPUT_TYPES()
optional = input_types["optional"]
assert "lora_stack3" in optional
assert optional["lora_stack3"][0] == "LORA_STACK"
assert "lora_stack25" in optional
assert optional["lora_stack25"][0] == "LORA_STACK"
def test_lora_stack_optional_inputs_proxy():
proxy = _LoraStackOptionalInputs({"lora_stack1": ("LORA_STACK", {})})
assert "lora_stack1" in proxy
assert "lora_stack2" in proxy
assert "lora_stack10" in proxy
assert "lora_stack_a" in proxy
assert "lora_stack" not in proxy
assert "lora_stacka" not in proxy
assert "lora_stack_1" not in proxy
assert "text" not in proxy
assert proxy["lora_stack1"][0] == "LORA_STACK"
assert proxy["lora_stack5"][0] == "LORA_STACK"
with pytest.raises(KeyError):
proxy["not_a_stack"]
+100 -1
View File
@@ -59,7 +59,7 @@ def test_save_image_defaults_to_writing_png_metadata(monkeypatch, tmp_path):
image_path = tmp_path / "sample_00001_.png" image_path = tmp_path / "sample_00001_.png"
with Image.open(image_path) as img: with Image.open(image_path) as img:
assert img.info["parameters"] == "prompt text\nSeed: 123" assert img.info["parameters"] == "prompt text\nSeed: 123, Version: ComfyUI"
def test_save_image_skips_png_parameters_when_metadata_disabled_and_keeps_workflow( def test_save_image_skips_png_parameters_when_metadata_disabled_and_keeps_workflow(
@@ -363,3 +363,102 @@ def test_save_image_as_recipe_writes_recipe_without_async_scanner_calls(
assert recipe["gen_params"] == {"prompt": "prompt text", "seed": 123} assert recipe["gen_params"] == {"prompt": "prompt text", "seed": 123}
assert scanner._json_path_map[recipe["id"]] == os.path.normpath(str(recipe_files[0])) assert scanner._json_path_map[recipe["id"]] == os.path.normpath(str(recipe_files[0]))
assert scanner.fts_updates == [(recipe["id"], "add")] assert scanner.fts_updates == [(recipe["id"], "add")]
# ---------------------------------------------------------------------------
# Tests for webp_method and jpeg_subsampling parameters
# ---------------------------------------------------------------------------
def _capture_save_kwargs(monkeypatch):
"""Monkeypatch Image.Image.save to capture kwargs while still saving to disk."""
real_save = Image.Image.save
captured_kwargs = {}
def _fake_save(self, fp, *args, **kwargs):
captured_kwargs.update(kwargs)
return real_save(self, fp, *args, **kwargs)
monkeypatch.setattr(Image.Image, "save", _fake_save)
return captured_kwargs
def test_webp_method_default_passed_to_pillow_save(monkeypatch, tmp_path):
_configure_save_paths(monkeypatch, tmp_path)
_configure_metadata(monkeypatch, {"prompt": "test", "seed": 1})
captured = _capture_save_kwargs(monkeypatch)
node = SaveImageLM()
node.save_images([_make_image()], "ComfyUI", "webp", id="node-1")
assert "method" in captured
assert captured["method"] == 6
def test_webp_method_custom_value_passed_to_pillow_save(monkeypatch, tmp_path):
_configure_save_paths(monkeypatch, tmp_path)
_configure_metadata(monkeypatch, {"prompt": "test", "seed": 1})
captured = _capture_save_kwargs(monkeypatch)
node = SaveImageLM()
node.save_images(
[_make_image()], "ComfyUI", "webp", id="node-1", webp_method=3
)
assert captured["method"] == 3
def test_jpeg_subsampling_default_passed_to_pillow_save(monkeypatch, tmp_path):
_configure_save_paths(monkeypatch, tmp_path)
_configure_metadata(monkeypatch, {"prompt": "test", "seed": 1})
captured = _capture_save_kwargs(monkeypatch)
node = SaveImageLM()
node.save_images([_make_image()], "ComfyUI", "jpeg", id="node-1")
assert "subsampling" in captured
assert captured["subsampling"] == 0
def test_jpeg_subsampling_custom_value_passed_to_pillow_save(monkeypatch, tmp_path):
_configure_save_paths(monkeypatch, tmp_path)
_configure_metadata(monkeypatch, {"prompt": "test", "seed": 1})
captured = _capture_save_kwargs(monkeypatch)
node = SaveImageLM()
node.save_images(
[_make_image()], "ComfyUI", "jpeg", id="node-1", jpeg_subsampling=1
)
assert captured["subsampling"] == 1
class TestParameterDefaultConsistency:
"""Verify defaults match across INPUT_TYPES, save_images(), and process_image()."""
def test_webp_method_defaults_are_consistent(self):
input_types = SaveImageLM.INPUT_TYPES()
optional = input_types["optional"]
assert optional["webp_method"][1]["default"] == 6
assert SaveImageLM.save_images.__defaults__[4] == 6 # positional: webp_method=6 is at index 4
assert SaveImageLM.process_image.__defaults__[6] == 6
def test_jpeg_subsampling_defaults_are_consistent(self):
input_types = SaveImageLM.INPUT_TYPES()
optional = input_types["optional"]
assert optional["jpeg_subsampling"][1]["default"] == 0
assert SaveImageLM.save_images.__defaults__[5] == 0
assert SaveImageLM.process_image.__defaults__[7] == 0
def test_png_does_not_pass_webp_method_or_jpeg_subsampling(monkeypatch, tmp_path):
_configure_save_paths(monkeypatch, tmp_path)
_configure_metadata(monkeypatch, {"prompt": "test", "seed": 1})
captured = _capture_save_kwargs(monkeypatch)
node = SaveImageLM()
node.save_images([_make_image()], "ComfyUI", "png", id="node-1")
assert "method" not in captured
assert "subsampling" not in captured
+97 -5
View File
@@ -900,18 +900,28 @@ class FakeMetadataProvider:
async def get_model_versions(self, _model_id): async def get_model_versions(self, _model_id):
return {"modelVersions": [], "name": "", "type": "lora"} return {"modelVersions": [], "name": "", "type": "lora"}
async def get_user_models(self, _username): async def get_user_models(self, _username, cursor=None):
return [] return {"items": [], "nextCursor": None}
async def get_creator_model_count(self, _username):
return None
class FakeUserModelsProvider(FakeMetadataProvider): class FakeUserModelsProvider(FakeMetadataProvider):
def __init__(self, models): def __init__(self, models, next_cursor=None, estimated_total=None):
self.models = models self.models = models
self.next_cursor = next_cursor
self.estimated_total = estimated_total
self.received_usernames: list[str] = [] self.received_usernames: list[str] = []
self.received_cursors: list = []
async def get_user_models(self, username): async def get_user_models(self, username, cursor=None):
self.received_usernames.append(username) self.received_usernames.append(username)
return self.models self.received_cursors.append(cursor)
return {"items": self.models, "nextCursor": self.next_cursor}
async def get_creator_model_count(self, _username):
return self.estimated_total
async def fake_metadata_provider_factory(): async def fake_metadata_provider_factory():
@@ -1286,6 +1296,88 @@ async def test_get_civitai_user_models_requires_username():
assert "username" in payload["error"].lower() assert "username" in payload["error"].lower()
@pytest.mark.asyncio
async def test_get_civitai_user_models_returns_pagination_fields():
models = [
{
"id": 1,
"name": "Model A",
"type": "LORA",
"tags": [],
"modelVersions": [
{"id": 100, "name": "v1", "images": [{"url": "http://example.com/a.jpg"}]},
],
},
{
"id": 2,
"name": "Unsupported",
"type": "Other",
"modelVersions": [{"id": 200, "name": "v1"}],
},
]
provider = FakeUserModelsProvider(models, next_cursor="cursor-token", estimated_total=2140)
async def provider_factory():
return provider
handler = ModelLibraryHandler(
ServiceRegistryAdapter(
get_lora_scanner=fake_scanner_factory,
get_checkpoint_scanner=fake_scanner_factory,
get_embedding_scanner=fake_scanner_factory,
get_downloaded_version_history_service=fake_download_history_service_factory,
),
metadata_provider_factory=provider_factory,
)
response = await handler.get_civitai_user_models(
FakeRequest(query={"username": "pixel"})
)
payload = json.loads(response.text)
assert response.status == 200
assert payload["success"] is True
# modelCount only counts models surviving the type filter
assert payload["modelCount"] == 1
assert payload["nextCursor"] == "cursor-token"
assert payload["hasMore"] is True
# first page includes the estimated total
assert payload["estimatedTotal"] == 2140
assert provider.received_cursors == [None]
@pytest.mark.asyncio
async def test_get_civitai_user_models_passes_cursor_and_omits_estimate():
provider = FakeUserModelsProvider([], next_cursor=None, estimated_total=999)
async def provider_factory():
return provider
handler = ModelLibraryHandler(
ServiceRegistryAdapter(
get_lora_scanner=fake_scanner_factory,
get_checkpoint_scanner=fake_scanner_factory,
get_embedding_scanner=fake_scanner_factory,
get_downloaded_version_history_service=fake_download_history_service_factory,
),
metadata_provider_factory=provider_factory,
)
response = await handler.get_civitai_user_models(
FakeRequest(query={"username": "pixel", "cursor": "opaque-token"})
)
payload = json.loads(response.text)
assert response.status == 200
assert payload["success"] is True
assert payload["nextCursor"] is None
assert payload["hasMore"] is False
# cursor requests must not include the estimated total
assert payload["estimatedTotal"] is None
assert provider.received_cursors == ["opaque-token"]
def test_ensure_handler_mapping_caches_result(): def test_ensure_handler_mapping_caches_result():
call_records = [] call_records = []
+328 -1
View File
@@ -1,10 +1,33 @@
import logging import logging
import os
import shutil
from aiohttp import ClientError from aiohttp import ClientError
from aiohttp import web
import pytest import pytest
from py.routes import update_routes 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: class OfflineDownloader:
async def make_request(self, *_, **__): async def make_request(self, *_, **__):
return False, "Cannot connect to host" 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) caplog.set_level(logging.WARNING)
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(RaisingDownloader())) 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 version == "main"
assert changelog == [] assert changelog == []
assert behind_by == 0
assert commit_date == ""
assert "Unable to reach GitHub for nightly version" in caplog.text assert "Unable to reach GitHub for nightly version" in caplog.text
assert "Traceback" not 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] clean_args = clean_calls[0][1]
for name in update_routes._PRESERVE_DIRS: for name in update_routes._PRESERVE_DIRS:
assert name in clean_args, f"{name} missing from git clean excludes (stable)" 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): def __init__(self, items):
self.items = list(items) self.items = list(items)
async def get_sorted_data(self, sort_key, order): async def get_sorted_data(self, sort_key, order, seed=None):
if sort_key == "name": if sort_key == "name":
data = sorted(self.items, key=lambda x: x["model_name"].lower()) data = sorted(self.items, key=lambda x: x["model_name"].lower())
if order == "desc": if order == "desc":
@@ -1252,3 +1252,69 @@ async def test_get_model_civitai_url_falls_back_when_host_setting_is_not_a_strin
"model_id": "123", "model_id": "123",
"version_id": "456", "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 assert result["pending_count"] == 1
@pytest.mark.asyncio
@pytest.mark.usefixtures("tmp_path")
async def test_check_pending_models_uses_bulk_folder_index_for_large_libraries(
monkeypatch: pytest.MonkeyPatch,
tmp_path,
settings_manager,
):
"""For >1000 candidates the pre-check scans the library root once instead of
probing every folder individually."""
ws_manager = RecordingWebSocketManager()
manager = download_module.DownloadManager(ws_manager=ws_manager)
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
# 1500 unprocessed models triggers the bulk lookup path
models = [
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
for i in range(1500)
]
# Create folders with files for the first 500 models
for i in range(500):
model_dir = tmp_path / f"{i:064x}"
model_dir.mkdir()
(model_dir / "image_0.png").write_text("data")
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
per_model_checks = 0
def counting_model_directory_has_files(path: str) -> bool:
nonlocal per_model_checks
per_model_checks += 1
return False
monkeypatch.setattr(
download_module,
"_model_directory_has_files",
counting_model_directory_has_files,
)
result = await manager.check_pending_models(["lora"])
assert result["success"] is True
assert result["total_models"] == 1500
assert result["pending_count"] == 1000
assert result["needs_download"] is True
# The per-folder check should not be used once we cross the threshold.
assert per_model_checks == 0
@pytest.mark.asyncio
@pytest.mark.usefixtures("tmp_path")
async def test_check_pending_models_uses_per_folder_check_for_small_candidate_sets(
monkeypatch: pytest.MonkeyPatch,
tmp_path,
settings_manager,
):
"""For <=1000 candidates the pre-check keeps the accurate per-folder path."""
ws_manager = RecordingWebSocketManager()
manager = download_module.DownloadManager(ws_manager=ws_manager)
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
models = [
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
for i in range(500)
]
# Create folders with files for the first 200 models
for i in range(200):
model_dir = tmp_path / f"{i:064x}"
model_dir.mkdir()
(model_dir / "image_0.png").write_text("data")
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
per_model_checks = 0
original_has_files = download_module._model_directory_has_files
def counting_model_directory_has_files(path: str) -> bool:
nonlocal per_model_checks
per_model_checks += 1
return original_has_files(path)
monkeypatch.setattr(
download_module,
"_model_directory_has_files",
counting_model_directory_has_files,
)
result = await manager.check_pending_models(["lora"])
assert result["success"] is True
assert result["total_models"] == 500
assert result["pending_count"] == 300
assert result["needs_download"] is True
# Per-folder path should run once per candidate.
assert per_model_checks == 500
@pytest.mark.asyncio
@pytest.mark.usefixtures("tmp_path")
async def test_check_pending_models_bulk_index_includes_legacy_folders(
monkeypatch: pytest.MonkeyPatch,
tmp_path,
settings_manager,
):
"""In multi-library mode the bulk index also scans the legacy root so models
whose folders have not been consolidated yet are not reported pending."""
ws_manager = RecordingWebSocketManager()
manager = download_module.DownloadManager(ws_manager=ws_manager)
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
monkeypatch.setitem(settings_manager.settings, "libraries", {"default": {}, "extra": {}})
monkeypatch.setitem(settings_manager.settings, "active_library", "extra")
# 1500 unprocessed models triggers the bulk lookup path
models = [
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
for i in range(1500)
]
# Folders live at the LEGACY root/<hash> path (not yet consolidated)
for i in range(500):
model_dir = tmp_path / f"{i:064x}"
model_dir.mkdir()
(model_dir / "image_0.png").write_text("data")
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
result = await manager.check_pending_models(["lora"])
assert result["success"] is True
assert result["total_models"] == 1500
assert result["pending_count"] == 1000
assert result["needs_download"] is True
@pytest.fixture @pytest.fixture
def settings_manager(): def settings_manager():
return get_settings_manager() return get_settings_manager()
+161
View File
@@ -35,9 +35,11 @@ class DummyDownloader:
def reset_singletons(): def reset_singletons():
CivitaiClient._instance = None CivitaiClient._instance = None
ModelMetadataProviderManager._instance = None ModelMetadataProviderManager._instance = None
civitai_client_module._creator_model_count_cache.clear()
yield yield
CivitaiClient._instance = None CivitaiClient._instance = None
ModelMetadataProviderManager._instance = None ModelMetadataProviderManager._instance = None
civitai_client_module._creator_model_count_cache.clear()
@pytest.fixture @pytest.fixture
@@ -622,3 +624,162 @@ async def test_get_image_info_handles_invalid_id(monkeypatch, downloader, caplog
assert result is None assert result is None
assert "Invalid image ID format" in caplog.text assert "Invalid image ID format" in caplog.text
async def test_get_user_models_requests_first_page_with_stable_params(downloader):
request_calls = []
async def fake_make_request(method, url, use_auth=True, **kwargs):
request_calls.append({"method": method, "url": url, "kwargs": kwargs})
return True, {
"items": [
{
"id": 1,
"modelVersions": [
{"id": 100, "images": [{"meta": {"comfy": {"x": 1}}}]}
],
}
],
"metadata": {"nextCursor": "next-token"},
}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
result = await client.get_user_models("pixel")
assert result is not None
assert result["nextCursor"] == "next-token"
assert len(result["items"]) == 1
# comfy metadata is still stripped
assert "comfy" not in result["items"][0]["modelVersions"][0]["images"][0]["meta"]
call = request_calls[0]
assert call["method"] == "GET"
assert call["url"] == "https://civitai.red/api/v1/models"
params = call["kwargs"]["params"]
assert params["username"] == "pixel"
assert params["nsfw"] == "true"
assert params["limit"] == 100
assert params["sort"] == "Newest"
assert params["period"] == "AllTime"
assert "cursor" not in params
async def test_get_user_models_passes_cursor_and_stringifies_next_cursor(downloader):
request_calls = []
async def fake_make_request(method, url, use_auth=True, **kwargs):
request_calls.append(kwargs)
return True, {"items": [], "metadata": {"nextCursor": 12345}}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
result = await client.get_user_models("pixel", cursor="opaque-token")
assert request_calls[0]["params"]["cursor"] == "opaque-token"
assert result == {"items": [], "nextCursor": "12345"}
async def test_get_user_models_without_next_cursor_returns_none_cursor(downloader):
async def fake_make_request(method, url, use_auth=True, **kwargs):
return True, {"items": [{"id": 1, "modelVersions": []}], "metadata": {}}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
result = await client.get_user_models("pixel")
assert result == {"items": [{"id": 1, "modelVersions": []}], "nextCursor": None}
async def test_get_user_models_failure_returns_none(downloader):
async def fake_make_request(method, url, use_auth=True, **kwargs):
return False, "500 server error"
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
result = await client.get_user_models("pixel")
assert result is None
async def test_get_creator_model_count_matches_exact_username(downloader):
request_calls = []
async def fake_make_request(method, url, use_auth=True, **kwargs):
request_calls.append({"url": url, "kwargs": kwargs})
return True, {
"items": [
{"username": "pixelart", "modelCount": 5},
{"username": "Pixel", "modelCount": 2140},
]
}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
count = await client.get_creator_model_count("pixel")
assert count == 2140
assert request_calls[0]["url"] == "https://civitai.red/api/v1/creators"
assert request_calls[0]["kwargs"]["params"] == {"query": "pixel", "limit": 10}
async def test_get_creator_model_count_without_exact_match_returns_none(downloader):
async def fake_make_request(method, url, use_auth=True, **kwargs):
return True, {"items": [{"username": "pixelart", "modelCount": 5}]}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
count = await client.get_creator_model_count("pixel")
assert count is None
async def test_get_creator_model_count_caches_results(downloader):
request_count = 0
async def fake_make_request(method, url, use_auth=True, **kwargs):
nonlocal request_count
request_count += 1
return True, {"items": [{"username": "pixel", "modelCount": 42}]}
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
assert await client.get_creator_model_count("pixel") == 42
# case-insensitive cache key, second call served from cache
assert await client.get_creator_model_count("Pixel") == 42
assert request_count == 1
async def test_get_creator_model_count_caches_failures(downloader):
request_count = 0
async def fake_make_request(method, url, use_auth=True, **kwargs):
nonlocal request_count
request_count += 1
return False, "500 server error"
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
assert await client.get_creator_model_count("pixel") is None
assert await client.get_creator_model_count("pixel") is None
assert request_count == 1
async def test_get_creator_model_count_never_raises(downloader):
async def fake_make_request(method, url, use_auth=True, **kwargs):
return True, "unexpected non-dict payload"
downloader.make_request = fake_make_request
client = await CivitaiClient.get_instance()
assert await client.get_creator_model_count("pixel") is None
@@ -1189,6 +1189,109 @@ def test_relative_path_sanitizes_model_and_version_placeholders():
assert relative_path == "Fancy_Model/Version_One" assert relative_path == "Fancy_Model/Version_One"
def test_relative_path_empty_first_tag_fallback():
"""Test that empty first_tag falls back to 'no tags'."""
manager = DownloadManager()
settings_manager = get_settings_manager()
settings_manager.settings["download_path_templates"]["lora"] = (
"{base_model}/{first_tag}"
)
version_info = {
"baseModel": "SDXL",
"model": {"name": "Test Model", "tags": []},
"creator": {"username": "Author"},
}
relative_path = manager._calculate_relative_path(version_info, "lora")
assert relative_path == "SDXL/no tags"
def test_relative_path_empty_base_model_and_first_tag():
"""Test that empty base_model + empty first_tag does NOT produce a leading slash."""
manager = DownloadManager()
settings_manager = get_settings_manager()
settings_manager.settings["download_path_templates"]["lora"] = (
"{base_model}/{first_tag}"
)
version_info = {
"baseModel": "",
"model": {"name": "Test Model", "tags": []},
"creator": {"username": "Author"},
}
relative_path = manager._calculate_relative_path(version_info, "lora")
assert not relative_path.startswith("/")
assert relative_path == "no tags"
def test_relative_path_sanitizes_double_slashes():
"""Test that empty placeholder substitutions don't produce double slashes."""
manager = DownloadManager()
settings_manager = get_settings_manager()
settings_manager.settings["download_path_templates"]["lora"] = (
"{base_model}/{first_tag}/{author}"
)
version_info = {
"baseModel": "SDXL",
"model": {"name": "Test Model", "tags": []},
"creator": {"username": "Author"},
}
relative_path = manager._calculate_relative_path(version_info, "lora")
assert "//" not in relative_path
assert relative_path == "SDXL/no tags/Author"
def test_download_containment_accepts_symlink_save_dir(tmp_path):
"""Verify the download path containment check (download_manager.py:1395-1397)
accepts save directories reached through user-created symlinks inside the
library root reproducing the symlink scenario from issue #1028."""
# Library root with a symlink subdirectory pointing to an external drive
lora_root = tmp_path / "loras"
lora_root.mkdir()
external_drive = tmp_path / "external" / "models"
external_drive.mkdir(parents=True)
symlink = lora_root / "Krea 2"
symlink.symlink_to(str(external_drive))
# Simulate a download: base_save_dir = library root,
# relative_path = "Krea 2/concept/NewModel"
base_save_dir = str(lora_root)
save_dir = os.path.join(base_save_dir, "Krea 2", "concept", "NewModel")
# Replicate the exact containment check from download_manager.py
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
# Must NOT be rejected — symlinks are legitimate business paths
assert resolved_dir.startswith(base_dir + os.sep)
def test_download_containment_rejects_dot_dot_traversal(tmp_path):
"""Verify the download path containment check still blocks ``..`` traversal
after the realpath abspath change."""
lora_root = tmp_path / "loras"
lora_root.mkdir()
base_save_dir = str(lora_root)
save_dir = os.path.join(base_save_dir, "..", "..", "etc", "passwd")
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
# Must be rejected — dot-dot escapes the library root
assert not resolved_dir.startswith(base_dir + os.sep)
assert resolved_dir != base_dir
def test_distribute_preview_to_entries_moves_and_copies(tmp_path): def test_distribute_preview_to_entries_moves_and_copies(tmp_path):
"""Test that preview distribution moves file to first entry and copies to others.""" """Test that preview distribution moves file to first entry and copies to others."""
manager = DownloadManager() manager = DownloadManager()
@@ -26,6 +26,7 @@ class StubScanner:
def __init__(self, models: list[dict]) -> None: def __init__(self, models: list[dict]) -> None:
self._cache = SimpleNamespace(raw_data=models) self._cache = SimpleNamespace(raw_data=models)
self.sync_calls: list[tuple[str, dict]] = []
async def get_cached_data(self): async def get_cached_data(self):
return self._cache return self._cache
@@ -38,6 +39,14 @@ class StubScanner:
break break
return True return True
async def sync_cache_from_metadata(self, file_path: str, metadata: dict) -> bool:
self.sync_calls.append((file_path, metadata))
for index, model in enumerate(self._cache.raw_data):
if model.get("file_path") == metadata.get("file_path"):
self._cache.raw_data[index] = metadata
break
return True
def _patch_scanner(monkeypatch: pytest.MonkeyPatch, scanner: StubScanner) -> None: def _patch_scanner(monkeypatch: pytest.MonkeyPatch, scanner: StubScanner) -> None:
async def _get_lora_scanner(cls): async def _get_lora_scanner(cls):
@@ -520,7 +529,8 @@ async def test_not_found_example_images_are_cleaned(
model_dir = images_root / model_hash model_dir = images_root / model_hash
model_dir.mkdir(parents=True, exist_ok=True) model_dir.mkdir(parents=True, exist_ok=True)
(model_dir / "image_0.png").write_bytes(b"first") # Pre-existing file collides with the valid image index (1) so the
# pre-download existence check must skip it without a network request
(model_dir / "image_1.png").write_bytes(b"second") (model_dir / "image_1.png").write_bytes(b"second")
async def fake_process_local_examples(*_args, **_kwargs): async def fake_process_local_examples(*_args, **_kwargs):
@@ -588,6 +598,9 @@ async def test_not_found_example_images_are_cleaned(
assert missing_url in downloader.calls assert missing_url in downloader.calls
assert manager._progress["failed_models"] == {model_hash} assert manager._progress["failed_models"] == {model_hash}
assert model_hash in manager._progress["processed_models"] assert model_hash in manager._progress["processed_models"]
assert scanner.sync_calls
assert len(scanner.sync_calls) == 1
assert scanner.sync_calls[0][0] == str(model_path)
remaining_images = model_metadata["civitai"]["images"] remaining_images = model_metadata["civitai"]["images"]
assert remaining_images == [ assert remaining_images == [
@@ -596,11 +609,188 @@ async def test_not_found_example_images_are_cleaned(
] ]
files = sorted(p.name for p in model_dir.iterdir()) files = sorted(p.name for p in model_dir.iterdir())
assert files == ["image_0.png", "image_1.png"] assert files == ["image_1.png"]
assert (model_dir / "image_0.png").read_bytes() == b"first"
assert (model_dir / "image_1.png").read_bytes() == b"second" assert (model_dir / "image_1.png").read_bytes() == b"second"
async def test_failed_models_retried_when_explicitly_targeted(
monkeypatch: pytest.MonkeyPatch,
tmp_path,
settings_manager,
):
ws_manager = RecordingWebSocketManager()
manager = download_module.DownloadManager(ws_manager=ws_manager)
images_root = tmp_path / "examples"
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(images_root))
model_hash = "a" * 64
model_path = tmp_path / "model.safetensors"
model_path.write_text("data", encoding="utf-8")
model_metadata = {
"sha256": model_hash,
"model_name": "Failed Example",
"file_path": str(model_path),
"file_name": "model.safetensors",
"civitai": {"images": [{"url": "https://example.com/valid.png"}]},
}
scanner = StubScanner([model_metadata.copy()])
_patch_scanner(monkeypatch, scanner)
# Persist a previous failure so the skip path is exercised
images_root.mkdir(parents=True, exist_ok=True)
(images_root / ".download_progress.json").write_text(
json.dumps(
{
"failed_models": [model_hash],
"processed_models": [],
"rate_limited_models": [],
}
),
encoding="utf-8",
)
async def fake_process_local_examples(*_args, **_kwargs):
return False
async def fake_get_updated_model(model_hash_arg, _scanner):
return model_metadata
class DownloaderStub:
def __init__(self):
self.calls: list[str] = []
async def download_to_memory(self, url, *_args, **_kwargs):
self.calls.append(url)
return True, b"\x89PNG\r\n\x1a\n", {"content-type": "image/png"}
downloader = DownloaderStub()
async def fake_get_downloader():
return downloader
monkeypatch.setattr(
download_module.ExampleImagesProcessor,
"process_local_examples",
staticmethod(fake_process_local_examples),
)
monkeypatch.setattr(
download_module.MetadataUpdater,
"get_updated_model",
staticmethod(fake_get_updated_model),
)
monkeypatch.setattr(download_module, "get_downloader", fake_get_downloader)
# Without explicit hashes the previously failed model is skipped
skipped_manager = download_module.DownloadManager(ws_manager=RecordingWebSocketManager())
result = await skipped_manager.start_download({"model_types": ["lora"], "delay": 0})
assert result["success"] is True
if skipped_manager._download_task is not None:
await asyncio.wait_for(skipped_manager._download_task, timeout=1)
assert downloader.calls == []
# With explicit hashes the previously failed model is retried and cleared
result = await manager.start_download(
{"model_types": ["lora"], "delay": 0, "model_hashes": [model_hash]}
)
assert result["success"] is True
if manager._download_task is not None:
await asyncio.wait_for(manager._download_task, timeout=1)
assert downloader.calls == ["https://example.com/valid.png"]
assert manager._progress["failed_models"] == set()
assert model_hash in manager._progress["processed_models"]
async def test_explicit_targets_fill_partial_example_gaps(
monkeypatch: pytest.MonkeyPatch,
tmp_path,
settings_manager,
):
ws_manager = RecordingWebSocketManager()
images_root = tmp_path / "examples"
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(images_root))
model_hash = "b" * 64
model_path = tmp_path / "model.safetensors"
model_path.write_text("data", encoding="utf-8")
model_metadata = {
"sha256": model_hash,
"model_name": "Partial Example",
"file_path": str(model_path),
"file_name": "model.safetensors",
"civitai": {
"images": [
{"url": "https://example.com/first.png"},
{"url": "https://example.com/second.png"},
]
},
}
scanner = StubScanner([model_metadata.copy()])
_patch_scanner(monkeypatch, scanner)
# Simulate a partially populated folder: index 0 already downloaded
model_dir = images_root / model_hash
model_dir.mkdir(parents=True, exist_ok=True)
(model_dir / "image_0.png").write_bytes(b"existing")
async def fake_process_local_examples(*_args, **_kwargs):
return False
async def fake_get_updated_model(model_hash_arg, _scanner):
return model_metadata
class DownloaderStub:
def __init__(self):
self.calls: list[str] = []
async def download_to_memory(self, url, *_args, **_kwargs):
self.calls.append(url)
return True, b"\x89PNG\r\n\x1a\n", {"content-type": "image/png"}
downloader = DownloaderStub()
async def fake_get_downloader():
return downloader
monkeypatch.setattr(
download_module.ExampleImagesProcessor,
"process_local_examples",
staticmethod(fake_process_local_examples),
)
monkeypatch.setattr(
download_module.MetadataUpdater,
"get_updated_model",
staticmethod(fake_get_updated_model),
)
monkeypatch.setattr(download_module, "get_downloader", fake_get_downloader)
# Untargeted run treats the populated folder as done
untargeted = download_module.DownloadManager(ws_manager=RecordingWebSocketManager())
result = await untargeted.start_download({"model_types": ["lora"], "delay": 0})
assert result["success"] is True
if untargeted._download_task is not None:
await asyncio.wait_for(untargeted._download_task, timeout=1)
assert downloader.calls == []
# Explicitly targeted run fills only the missing index, skipping the
# existing file without a network request
targeted = download_module.DownloadManager(ws_manager=ws_manager)
result = await targeted.start_download(
{"model_types": ["lora"], "delay": 0, "model_hashes": [model_hash]}
)
assert result["success"] is True
if targeted._download_task is not None:
await asyncio.wait_for(targeted._download_task, timeout=1)
assert downloader.calls == ["https://example.com/second.png"]
assert (model_dir / "image_1.png").exists()
assert (model_dir / "image_0.png").read_bytes() == b"existing"
@pytest.fixture @pytest.fixture
def settings_manager(): def settings_manager():
return get_settings_manager() return get_settings_manager()
+50
View File
@@ -243,6 +243,56 @@ class TestLLMServiceChatCompletionJson:
assert result == {"key": "value"} assert result == {"key": "value"}
@pytest.mark.asyncio
async def test_chat_completion_json_falls_back_on_response_format_rejection(
self, llm_service,
):
"""Retry without response_format when provider rejects it (HTTP 400)."""
error_response = MockResponse(
400,
text_data=(
'{"error":"\'response_format.type\' must be '
'\'json_schema\' or \'text\'"}'
),
)
success_response = MockResponse(
200,
json_data={
"choices": [{"message": {"content": '{"key": "value"}'}}],
"usage": {},
"model": "local-model",
},
)
call_index = 0
class FallbackMockSession:
def __init__(self):
self.last_url = None
self.last_json = None
def post(self, url, json=None, headers=None):
nonlocal call_index
self.last_url = url
self.last_json = json
call_index += 1
return error_response if call_index == 1 else success_response
async def __aenter__(self):
return self
async def __aexit__(self, *args):
pass
with mock.patch("aiohttp.ClientSession", return_value=FallbackMockSession()):
result = await llm_service.chat_completion_json(
system_prompt="You are helpful.",
user_prompt="Return JSON.",
)
assert result == {"key": "value"}
assert call_index == 2
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_chat_completion_json_raises_on_non_json(self, llm_service): async def test_chat_completion_json_raises_on_non_json(self, llm_service):
# Non-JSON content raises LLMResponseError (salvage also fails) # Non-JSON content raises LLMResponseError (salvage also fails)
+169 -1
View File
@@ -1,13 +1,181 @@
import json import json
import os
from pathlib import Path from pathlib import Path
import pytest import pytest
from py.services.model_lifecycle_service import ModelLifecycleService from py.services.model_lifecycle_service import ModelLifecycleService, _require_path_in_library_roots
from py.utils.metadata_manager import MetadataManager from py.utils.metadata_manager import MetadataManager
from py.utils.models import LoraMetadata from py.utils.models import LoraMetadata
class ScannerWithRoots:
def __init__(self, roots):
self._roots = list(roots)
def get_model_roots(self):
return self._roots
class TestRequirePathInLibraryRoots:
def test_accepts_path_within_root(self, tmp_path):
root = tmp_path / "loras"
root.mkdir()
model = root / "model.safetensors"
model.write_text("")
scanner = ScannerWithRoots([str(root)])
_require_path_in_library_roots(str(model), scanner)
def test_rejects_path_outside_roots(self, tmp_path):
root = tmp_path / "loras"
root.mkdir()
outside = tmp_path / "outside" / "model.safetensors"
outside.parent.mkdir(parents=True)
outside.write_text("")
scanner = ScannerWithRoots([str(root)])
with pytest.raises(ValueError, match="outside configured library"):
_require_path_in_library_roots(str(outside), scanner)
def test_passes_when_no_roots_configured(self, tmp_path):
f = tmp_path / "model.safetensors"
f.write_text("")
scanner = ScannerWithRoots([])
_require_path_in_library_roots(str(f), scanner)
def test_accepts_path_matching_root_exactly(self, tmp_path):
root = tmp_path / "loras"
root.mkdir()
scanner = ScannerWithRoots([str(root)])
_require_path_in_library_roots(str(root), scanner)
def test_accepts_symlink_within_root(self, tmp_path):
"""Symlinks under a configured root are legitimate business paths
and should be accepted containment works on business-path space,
not resolved physical paths."""
root = tmp_path / "loras"
root.mkdir()
outside_dir = tmp_path / "outside"
outside_dir.mkdir()
outside_file = outside_dir / "escaped.safetensors"
outside_file.write_text("")
symlink = root / "link.safetensors"
symlink.symlink_to(outside_file)
scanner = ScannerWithRoots([str(root)])
# Symlink path is under root in business-path space → accepted
_require_path_in_library_roots(str(symlink), scanner)
def test_rejects_dot_dot_traversal(self, tmp_path):
"""Verify that ``..`` components are still resolved and blocked —
``abspath`` normalises dot-dot but does not resolve symlinks."""
root = tmp_path / "loras"
root.mkdir()
# A path that traverses up out of the root via ..
escaped = os.path.join(str(root), "..", "..", "etc", "passwd")
scanner = ScannerWithRoots([str(root)])
with pytest.raises(ValueError, match="outside configured library"):
_require_path_in_library_roots(escaped, scanner)
class ScannerForDelete:
def __init__(self, raw_data, roots, model_type="lora"):
self.model_type = model_type
self.cache = DummyCache(raw_data)
self._hash_index = DummyHashIndex()
self._roots = list(roots)
self._persist_calls = []
def get_model_roots(self):
return self._roots
async def get_cached_data(self):
return self.cache
async def _persist_current_cache(self):
self._persist_calls.append(True)
@pytest.mark.asyncio
async def test_delete_model_rejects_path_outside_roots(tmp_path: Path):
root = tmp_path / "loras"
root.mkdir()
model = root / "model.safetensors"
model.write_bytes(b"data")
scanner = ScannerForDelete(
raw_data=[{"file_path": str(model)}],
roots=[str(root)],
)
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({"civitai": {"modelId": 1}}),
metadata_loader=lambda x: {},
)
# Path within root should work (model file exists)
result = await service.delete_model(str(model))
assert result["success"] is True
# Path outside root should be rejected
outside = tmp_path / "outside.safetensors"
outside.write_bytes(b"data")
scanner2 = ScannerForDelete(
raw_data=[],
roots=[str(root)],
)
service2 = ModelLifecycleService(
scanner=scanner2,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
)
with pytest.raises(ValueError, match="outside configured library"):
await service2.delete_model(str(outside))
@pytest.mark.asyncio
async def test_rename_model_rejects_path_outside_roots(tmp_path: Path):
root = tmp_path / "loras"
root.mkdir()
scanner = ScannerWithRoots([str(root)])
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
)
outside = tmp_path / "outside.safetensors"
outside.write_bytes(b"data")
with pytest.raises(ValueError, match="outside configured library"):
await service.rename_model(file_path=str(outside), new_file_name="new_name")
@pytest.mark.asyncio
async def test_bulk_delete_rejects_any_path_outside_roots(tmp_path: Path):
root = tmp_path / "loras"
root.mkdir()
model_ok = root / "model.safetensors"
model_ok.write_bytes(b"data")
outside = tmp_path / "outside.safetensors"
outside.write_bytes(b"data")
scanner = ScannerWithRoots([str(root)])
service = ModelLifecycleService(
scanner=scanner,
metadata_manager=DummyMetadataManager({}),
metadata_loader=lambda x: {},
)
with pytest.raises(ValueError, match="outside configured library"):
await service.bulk_delete_models([str(model_ok), str(outside)])
class DummyCache: class DummyCache:
def __init__(self, raw_data): def __init__(self, raw_data):
self.raw_data = raw_data self.raw_data = raw_data
+2 -2
View File
@@ -884,7 +884,7 @@ async def test_sync_cache_conditional_resort_skipped(tmp_path: Path, monkeypatch
raw_data=[dict(entry)], folders=[], name_display_mode="model_name" raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
) )
await scanner._cache.resort() await scanner._cache.resort()
scanner._cache._last_sort = ("name", "asc") # name sort is active scanner._cache._last_sort = ("name", "asc", None) # name sort is active
scanner._tags_count = {"alpha": 1} scanner._tags_count = {"alpha": 1}
scanner._hash_index.add_entry("abc123", "/m/a.safetensors") scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
@@ -935,7 +935,7 @@ async def test_sync_cache_conditional_resort_triggered(tmp_path: Path, monkeypat
raw_data=[dict(entry)], folders=[], name_display_mode="model_name" raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
) )
await scanner._cache.resort() await scanner._cache.resort()
scanner._cache._last_sort = ("name", "asc") scanner._cache._last_sort = ("name", "asc", None)
scanner._tags_count = {"alpha": 1} scanner._tags_count = {"alpha": 1}
scanner._hash_index.add_entry("abc123", "/m/a.safetensors") scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
+97
View File
@@ -0,0 +1,97 @@
"""Tests for sort parsing and the seeded random sort mode."""
import asyncio
import pytest
from py.services.model_cache import ModelCache
from py.services.model_query import ModelCacheRepository, SortParams
def _make_cache(items):
return ModelCache(
raw_data=[
{
"file_path": f"/models/{name}.safetensors",
"file_name": f"{name}.safetensors",
"model_name": name,
"folder": "",
"size": 100,
"modified": 0.0,
}
for name in items
],
folders=[],
)
class TestParseSort:
def test_random_with_seed(self):
params = ModelCacheRepository.parse_sort("random:abc123")
assert params == SortParams(key="random", order="asc", seed="abc123")
def test_random_without_seed(self):
params = ModelCacheRepository.parse_sort("random")
assert params == SortParams(key="random", order="asc", seed=None)
def test_random_empty_seed_falls_back_to_none(self):
params = ModelCacheRepository.parse_sort("random:")
assert params.seed is None
def test_regular_sorts_unaffected(self):
params = ModelCacheRepository.parse_sort("name:desc")
assert params == SortParams(key="name", order="desc", seed=None)
class TestRandomShuffle:
@pytest.mark.asyncio
async def test_same_seed_yields_same_order(self):
cache = _make_cache(["a", "b", "c", "d", "e"])
await asyncio.sleep(0) # allow background resort task to run
first = await cache.get_sorted_data("random", "asc", "seed1")
second = await cache.get_sorted_data("random", "asc", "seed1")
assert [item["model_name"] for item in first] == [
item["model_name"] for item in second
]
@pytest.mark.asyncio
async def test_different_seeds_yield_different_orders(self):
cache = _make_cache([f"m{i}" for i in range(20)])
await asyncio.sleep(0)
first = await cache.get_sorted_data("random", "asc", "seed-a")
second = await cache.get_sorted_data("random", "asc", "seed-b")
assert [item["model_name"] for item in first] != [
item["model_name"] for item in second
]
@pytest.mark.asyncio
async def test_shuffle_is_a_permutation(self):
cache = _make_cache(["a", "b", "c", "d", "e"])
await asyncio.sleep(0)
shuffled = await cache.get_sorted_data("random", "asc", "seed")
assert sorted(item["model_name"] for item in shuffled) == [
"a",
"b",
"c",
"d",
"e",
]
assert len({item["file_path"] for item in shuffled}) == 5
@pytest.mark.asyncio
async def test_missing_seed_is_stable(self):
cache = _make_cache(["a", "b", "c", "d", "e"])
await asyncio.sleep(0)
first = await cache.get_sorted_data("random", "asc")
second = await cache.get_sorted_data("random", "asc")
assert [item["model_name"] for item in first] == [
item["model_name"] for item in second
]
+9 -1
View File
@@ -77,7 +77,15 @@ def recipe_scanner(tmp_path: Path, monkeypatch):
monkeypatch.setattr(config, "loras_roots", [str(tmp_path)]) monkeypatch.setattr(config, "loras_roots", [str(tmp_path)])
stub = StubLoraScanner() stub = StubLoraScanner()
scanner = RecipeScanner(lora_scanner=stub) scanner = RecipeScanner(lora_scanner=stub)
asyncio.run(scanner.refresh_cache(force=True))
async def _init():
await scanner.refresh_cache(force=True)
# Wait for FTS index build to finish — asyncio.run()
# cancels background tasks on return, so we must await it here.
if scanner._fts_index_task:
await scanner._fts_index_task
asyncio.run(_init())
yield scanner, stub yield scanner, stub
RecipeScanner._instance = None RecipeScanner._instance = None
settings_manager_module.reset_settings_manager() settings_manager_module.reset_settings_manager()
+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("__flower__") is True
assert contains_dynamic_syntax("{red|blue}") is True assert contains_dynamic_syntax("{red|blue}") is True
assert contains_dynamic_syntax("{2$$, $$red|blue|green}") 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}"

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