Compare commits

..

69 Commits

Author SHA1 Message Date
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
Will Miao 8022d12f03 chore(release): bump version to v1.1.8 2026-07-20 20:42:04 +08:00
Will Miao 3939f7f91b chore: Update lora manager basic example workflow 2026-07-20 20:20:35 +08:00
Will Miao aebf2e37dd fix(filter): apply preset on full tile click and suppress i18n double-translate warnings
- Move preset apply handler from span.preset-name to div.filter-preset so
  clicking anywhere on the tile triggers the preset, not just the label text.
- Add whitespace heuristic in showToast() to skip translate() for plain
  messages that are already translated at the call site. This prevents
  i18next from logging 'Translation key not found' for pre-translated
  strings like 'Preset "name" applied'.
2026-07-20 17:57:14 +08:00
Will Miao f53f859a71 feat(filter): add debounced tag search with backend search-tags endpoint 2026-07-20 17:37:47 +08:00
Will Miao d916375abe fix(checkpoint): populate hash index from pre-computed metadata to prevent repeated hash re-calculation (#1002) 2026-07-20 12:24:54 +08:00
Will Miao 57983df4bd fix(recipe): resolve recipe metadata update bugs in cache sort, allowed fields, and bulk API routing
- Use safe .get() in RecipeCache._resort_locked instead of itemgetter to prevent KeyError when recipe missing created_date; align sort key with _sort_cache_sync (prefer modified, fallback created_date, fallback 0)
- Add base_model to allowed_fields in persistence_service.update_recipe() so the field passes validation
- Route bulk base model updates through updateRecipeMetadata() on recipes page instead of generic saveModelMetadata(), matching existing isRecipesPage pattern used in setBulkFavorites and saveBulkTags
2026-07-20 11:20:06 +08:00
Will Miao c68d7559a0 fix(widget): correct reorder drop indicator position when container is scrolled
The drop indicator top position was calculated using only
getBoundingClientRect() offsets (post-CSS-transform viewport space)
without accounting for container.scrollTop (pre-transform layout space).
This caused the indicator to drift upward as the user scrolled down,
eventually disappearing entirely.

Fixed by adding container.scrollTop to the position calculation and
only dividing the GBCR visual-diff portion by scale, since scrollTop
is already in pre-transform coordinate space.
2026-07-19 22:40:07 +08:00
Will Miao 9a8f5bf2d6 fix(ui): reposition download settings before AI provider section 2026-07-19 17:57:02 +08:00
Will Miao a8d742b031 feat(metadata): add CivArchive API toggle and provider fallback order settings
- Add enable_civarchive_api toggle (default on) to allow disabling
  CivArchive to avoid its rate-limit windows entirely
- Add metadata_provider_order dropdown with two presets:
  CivitAI → CivArchive → Archive DB (default) and
  CivitAI → Archive DB → CivArchive
- Wire both settings through backend (metadata_service, settings_manager,
  misc_handlers) and frontend (SettingsManager, state, settings modal)
- Reorder Metadata section in settings modal: toggles → status/management
  → fallback order, for natural top-down workflow
- Make update_metadata_providers() log the effective provider chain
  using actually-registered providers rather than settings assumptions
- Add 5 test cases covering all provider-combination paths
- Complete i18n translations for 6 new keys across all 9 non-English locales
2026-07-19 17:51:50 +08:00
Will Miao c27e4d1bfc feat(cache): opportunistic cache sync on metadata read with in-place update
- Add PersistentModelCache.update_single_model() for lightweight targeted
  SQL update (single row + incremental tag/hash deltas, no full table scan)
- Add ModelScanner.sync_cache_from_metadata() with compare-first logic:
  skips entirely when cache is already in sync; when stale, updates the
  entry in-place (O(1) instead of O(n) remove+append), incrementally
  adjusts tag counts/hash index/version index, and resorts only when
  sort-relevant fields changed
- Wire sync_cache_from_metadata() into BaseModelService.get_model_metadata()
  via fire-and-forget asyncio.create_task — disk I/O is already paid for
- Include identity re-validation guard against concurrent cache replacement
- Add 16 tests covering _cache_entries_differ, sync_cache_from_metadata
  (no-change, in-place, fallback, conditional resort), and
  update_single_model (insert, tag delta, hash delta)
2026-07-19 08:32:21 +08:00
Will Miao d15a8aa9a2 fix(workflow): accept non-string widget values and support GlobalSeed node in gen-params (#1026) 2026-07-18 22:57:20 +08:00
Will Miao 74a7d12ca4 fix(test): add missing options mock in LoraInfoWidget test 2026-07-18 22:14:03 +08:00
Will Miao 2f94a9773e feat(workflow): redirect gen-params updates to connected Primitive nodes (#1026)
When a KSampler marked as 'Send Gen Params Target' has widget inputs
wired to Primitive nodes (PrimitiveNode, PrimitiveInt, PrimitiveFloat,
etc.), sending gen params from the Lora Manager UI now updates the
Primitive node's value instead of the KSampler widget. This is
necessary because ComfyUI's execution engine reads from the connected
input, ignoring the widget value when a wire is present.

Also fix two minor issues found during review:
- Remove unnecessary String() wrapping on numeric gen params (seed,
  steps, cfg) to preserve native types through the JSON/WS path
- Correct misleading isNodeEnabled comment: LGraphEventMode values
  are 0=Always, 2=Never, 4=Bypass (not 'Normal/Enabled')
2026-07-18 22:10:48 +08:00
Will Miao 37bdfa21ea fix(standalone): ensure sys.path includes script dir for python_embeded compatibility (#1025) 2026-07-18 21:22:25 +08:00
Will Miao f0bf2728c9 fix(downloads): accept download_id in history delete/retry endpoints, add unique index 2026-07-18 21:10:42 +08:00
Will Miao dc715aa273 fix(download): fallback to downloadUrl when all mirrors are deleted
When Civitai returns 404 for /models/{id} (e.g. due to Civitai API bug
where un-deleted models still get 404), the fallback to CivArchive
provides metadata.  However CivArchive may return mirrors with every
entry marked deletedAt, while the file's downloadUrl is still valid.

Before this fix, _build_download_urls_from_file_info used an if/else
that skipped the downloadUrl fallback whenever the mirrors array was
non-empty, even when all mirrors were filtered out.  Now downloadUrl
is always tried when no usable mirror remains.

Also deduplicated the inline mirror-processing code at the second call
site by replacing it with a call to the shared helper.
2026-07-18 18:28:32 +08:00
Will Miao 7ee2361e87 fix(config): remove stale 'default' library entry and consolidate example images on startup 2026-07-18 17:25:36 +08:00
Will Miao e04c22f83f fix(widgets): allow text selection in LoraInfoWidget description tab 2026-07-17 18:33:59 +08:00
Will Miao 681cc13e90 fix(widgets): persist LoRA entry selection and active tab across save/load 2026-07-17 18:27:34 +08:00
Will Miao 090e0297d4 fix(downloader): hold session lock in retry paths to prevent session close race
Refactor _create_session() to make-before-break: snapshot old session,
assign new one first, then close old.  Previously, concurrent download
retries called _create_session() without the session lock (violating its
docstring contract) and closed the old session while other coroutines
held active references — causing aiohttp to raise "NoneType has no
attribute connect" when dereferencing the torn-down connector.

Also wrap the two _create_session() calls in the integrity-retry and
network-retry paths with self._session_lock to match the locking
discipline used by the session property and refresh_session().
2026-07-17 17:21:05 +08:00
Will Miao 6f71335be4 feat(widgets): add Description tab to LoraInfoWidget with dual-mode rendering support
- Add Notes/Description tab switching with tab state persistence in widget value
- Lazy-load model description and version description from /lm/loras/metadata
- Render CivitAI HTML descriptions inline via v-html
- Auto-fetch description when LoRA selection changes while on Description tab
- Fix Vue mode height containment via contain:layout size (lm-vue-node class)
- Fix scroll wheel isolation: widget scroll vs canvas zoom in both render modes
- Add docs/comfyui-dual-mode-widgets.md with widget rendering patterns
2026-07-17 15:04:47 +08:00
Will Miao 7f51812c1e feat(nodes): add LoRA Syntax → Path node (#1015) 2026-07-16 19:57:29 +08:00
Will Miao a9dc4d7b9d fix(widgets): reuse orphaned DOM containers after undo/redo in Vue render mode
In ComfyUI Vue render mode, WidgetDOM.vue reuses its component instance
during undo/redo without re-calling mountWidgetElement(), leaving newly
created widget containers detached from the DOM.

- AutocompleteTextWidget: scan for empty containers by ID prefix and reuse
- Loras widget: scan for empty .lm-loras-container elements and reuse
- Prevent duplicate event listeners by guarding listener setup on new
  containers only
- Keep container in DOM on cleanup (clearChildren instead of remove)
  so it can be found and reused by the next factory invocation
2026-07-16 18:54:00 +08:00
Will Miao 5d50ddb5d4 fix(ui): exit bulk mode after send-to-workflow completes 2026-07-16 18:54:00 +08:00
Will Miao f86198d234 fix(loras): include folder prefix in context menu and bulk send-to-workflow
When using full path lora syntax, the context menu (single/bulk)
and bulk copy actions were passing only the file basename to
buildLoraSyntax(), ignoring the folder prefix. This caused the
output to look like legacy A1111 format even when full path mode
was enabled.

Aligns all entry points with ModelCard.handleSendToWorkflow(),
which correctly includes the folder prefix.

Also fixes selectAllVisibleModels() to cache the folder field,
preventing missing prefix on select-all-then-send flows.
2026-07-16 18:54:00 +08:00
Will Miao ffe65d983c feat(api): add GET endpoints for update-lora-code and update-node-widget
Add GET variants of the two POST endpoints used by the send-to-workflow
feature. Parameters are read from query string instead of JSON body,
supporting both simple repeated node_id params and JSON-encoded node_ids
for complex graph references.
2026-07-15 21:49:22 +08:00
Will Miao b0b5be913c fix(downloads): reject re-insertion of download_ids already in history
In add_to_queue, check download_history before INSERT OR IGNORE.  Without
this check, a fire-and-forget /queue/complete failure on the extension side
would allow the same download_id to be re-inserted after complete_download()
deleted it from the queue — creating phantom queued entries for already-
finished downloads.
2026-07-15 19:12:22 +08:00
Will Miao 01efcbc584 fix(loras): allow toggle deselect on LoRA entry click 2026-07-14 18:23:06 +08:00
Will Miao 02c249917a fix(recipe): ensure custom recipes_path is added to preview allowed roots on startup 2026-07-14 18:15:20 +08:00
Will Miao 419bbc90b2 feat(lora-info): add Lora Info display node
Add a pure frontend node that shows filename and editable notes for
a selected LoRA. Connect any output from a LoRA Loader/Stacker/Randomizer/
WanVideoSelect to the lora_source input — selecting a LoRA in the source
widget updates the info display automatically.

- Python node (LoraInfoLM): display-only, no workflow execution
- Vue widget: filename label, auto-sizing notes textarea, save button
  with ComfyUI toast feedback on save
- Frontend extension: wire-based selection propagation with stale-response
  race guard; clears display on wire disconnect
- Backend: get-notes endpoint now returns file_path alongside notes;
  matching supports full-path lora syntax; fix NoneType crash in
  trigger words endpoint; document cache file_name invariant
- Wired into all four lora widget nodes (Loader, Stacker, Randomizer,
  WanVideoSelect)
2026-07-14 18:00:31 +08:00
willmiao b0c4510fdb docs: auto-update supporters list in README 2026-07-13 14:18:36 +00:00
120 changed files with 29022 additions and 20947 deletions
+8 -1
View File
@@ -102,6 +102,7 @@ npm run test:coverage # Generate coverage report
- ComfyUI: `app.registerExtension()`, `node.addDOMWidget(name, type, element, options)` - ComfyUI: `app.registerExtension()`, `node.addDOMWidget(name, type, element, options)`
- Event handlers via `addEventListener` or widget callbacks - Event handlers via `addEventListener` or widget callbacks
- Shared utilities: `web/comfyui/utils.js` - Shared utilities: `web/comfyui/utils.js`
- Dual-mode rendering patterns (canvas vs Vue): see `docs/comfyui-dual-mode-widgets.md`
### Vue Composables Pattern ### Vue Composables Pattern
@@ -136,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
+18
View File
@@ -15,6 +15,10 @@ try: # pragma: no cover - import fallback for pytest collection
from .py.nodes.lora_pool import LoraPoolLM from .py.nodes.lora_pool import LoraPoolLM
from .py.nodes.lora_randomizer import LoraRandomizerLM from .py.nodes.lora_randomizer import LoraRandomizerLM
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_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
@@ -56,6 +60,16 @@ except (
"py.nodes.lora_randomizer" "py.nodes.lora_randomizer"
).LoraRandomizerLM ).LoraRandomizerLM
LoraCyclerLM = importlib.import_module("py.nodes.lora_cycler").LoraCyclerLM LoraCyclerLM = importlib.import_module("py.nodes.lora_cycler").LoraCyclerLM
LoraInfoLM = importlib.import_module("py.nodes.lora_info").LoraInfoLM
LoraSyntaxToPath = importlib.import_module(
"py.nodes.lora_syntax_to_path"
).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 = {
@@ -75,6 +89,10 @@ NODE_CLASS_MAPPINGS = {
LoraPoolLM.NAME: LoraPoolLM, LoraPoolLM.NAME: LoraPoolLM,
LoraRandomizerLM.NAME: LoraRandomizerLM, LoraRandomizerLM.NAME: LoraRandomizerLM,
LoraCyclerLM.NAME: LoraCyclerLM, LoraCyclerLM.NAME: LoraCyclerLM,
LoraInfoLM.NAME: LoraInfoLM,
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
CreateHookLoraLM.NAME: CreateHookLoraLM,
MetadataOverwriteLM.NAME: MetadataOverwriteLM,
} }
WEB_DIRECTORY = "./web/comfyui" WEB_DIRECTORY = "./web/comfyui"
+138 -125
View File
@@ -10,13 +10,17 @@
"Brennok", "Brennok",
"2018cfh", "2018cfh",
"Insomnia Art Designs", "Insomnia Art Designs",
"Rob Williams",
"Arlecchino Shion", "Arlecchino Shion",
"Charles Blakemore", "Charles Blakemore",
"Rob Williams",
"$MetaSamsara", "$MetaSamsara",
"W+K+White", "W+K+White",
"stone9k", "stone9k",
"Rosenthal",
"Francisco Tatis",
"Mozzel",
"Gingko Biloba", "Gingko Biloba",
"Birdy",
"Kiba", "Kiba",
"onesecondinosaur", "onesecondinosaur",
"Christian Byrne", "Christian Byrne",
@@ -27,22 +31,25 @@
"Phil", "Phil",
"Carl G.", "Carl G.",
"Dsperado", "Dsperado",
"Rosenthal",
"ClockDaemon", "ClockDaemon",
"Francisco Tatis",
"Tobi_Swagg", "Tobi_Swagg",
"SG",
"jmack",
"Andrew Wilson", "Andrew Wilson",
"Greybush", "Greybush",
"Ricky Carter", "Ricky Carter",
"JongWon Han", "JongWon Han",
"VantAI", "VantAI",
"レプサイ",
"Michael Wong", "Michael Wong",
"Illrigger", "Illrigger",
"Tom Corrigan", "Tom Corrigan",
"JackieWang", "JackieWang",
"FreelancerZ", "FreelancerZ",
"fnkylove", "fnkylove",
"Lilleman",
"Robert Stacey", "Robert Stacey",
"PM",
"Edgar Tejeda", "Edgar Tejeda",
"Liam MacDougal", "Liam MacDougal",
"Polymorphic Indeterminate", "Polymorphic Indeterminate",
@@ -50,9 +57,7 @@
"Dogwalkerbr", "Dogwalkerbr",
"Skalabananen", "Skalabananen",
"Marc Whiffen", "Marc Whiffen",
"Birdy",
"itismyelement", "itismyelement",
"Mozzel",
"quarz", "quarz",
"Reno Lam", "Reno Lam",
"jean jahren", "jean jahren",
@@ -65,40 +70,46 @@
"Jonathan Ross", "Jonathan Ross",
"KD", "KD",
"Omnidex", "Omnidex",
"Nazono_hito", "Nolife_M",
"Melville Parrish",
"daniel dove", "daniel dove",
"Lustre",
"Tyler Trebuchon", "Tyler Trebuchon",
"Release Cabrakan", "Release Cabrakan",
"JW Sin", "JW Sin",
"Alex", "Alex",
"SG", "bh",
"carozzz", "carozzz",
"Marlon Daniels",
"James Dooley", "James Dooley",
"zenbound", "zenbound",
"Buzzard", "Buzzard",
"jmack", "Aaron Bleuer",
"LacesOut!",
"Adam Shaw", "Adam Shaw",
"Mark Corneglio", "Mark Corneglio",
"SarcasticHashtag", "SarcasticHashtag",
"RedrockVP", "RedrockVP",
"James Todd", "James Todd",
"Wicked Choices by ASLPro3D", "Wicked Choices by ASLPro3D",
"FinalyFree",
"Weasyl",
"Steven Pfeiffer", "Steven Pfeiffer",
"レプサイ",
"Timmy", "Timmy",
"Johnny", "Johnny",
"Cory Paza",
"Tak", "Tak",
"Lisster", "Lisster",
"runte3221", "runte3221",
"Big Red", "Big Red",
"whudunit", "whudunit",
"Luc Job",
"dl0901dm", "dl0901dm",
"corde",
"Yushio", "Yushio",
"Vik71it", "Vik71it",
"Bishoujoker", "Bishoujoker",
"Echo", "Echo",
"Lilleman",
"PM",
"Todd Keck", "Todd Keck",
"Briton Heilbrun", "Briton Heilbrun",
"wildnut", "wildnut",
@@ -108,53 +119,51 @@
"BadassArabianMofo", "BadassArabianMofo",
"Pascal Dahle", "Pascal Dahle",
"Greg", "Greg",
"Sangheili460",
"MagnaInsomnia",
"Akira_HentAI",
"Karl P.",
"MiraiKuriyamaSy", "MiraiKuriyamaSy",
"otaku fra", "otaku fra",
"lmsupporter", "lmsupporter",
"andrew.tappan", "andrew.tappan",
"Takkan", "Takkan",
"N/A",
"Greenmoustache",
"zounic", "zounic",
"wfpearl", "wfpearl",
"ElitaSSJ4", "ElitaSSJ4",
"Matt+J", "Matt+J",
"Jack B Nimble", "Jack B Nimble",
"Melville Parrish",
"Lustre",
"bh",
"Jwk0205", "Jwk0205",
"Marlon Daniels",
"Starkselle", "Starkselle",
"Aaron Bleuer", "Olive",
"LacesOut!",
"greebles", "greebles",
"Some Guy Named Barry", "Some Guy Named Barry",
"Resist's Creations - Spicy Edition 🔥", "Resist's Creations - Spicy Edition 🔥",
"M Postkasse", "M Postkasse",
"Wolffen", "Wolffen",
"wamekukyouzin",
"drum matthieu",
"Jacob Hoehler", "Jacob Hoehler",
"FinalyFree", "DogmaR34",
"Matt Wenzel", "Matt Wenzel",
"Weasyl",
"Lex Song", "Lex Song",
"Cory Paza", "Christopher Michel",
"Gonzalo Andre Allendes Lopez", "Gonzalo Andre Allendes Lopez",
"Serge Bekenkamp",
"Jimmy Ledbetter", "Jimmy Ledbetter",
"Luc Job", "LeoZero",
"Philip Hempel", "Philip Hempel",
"corde",
"nwalker94", "nwalker94",
"dan", "dan",
"aai", "aai",
"Tori", "Tori",
"Mouthlessman",
"Ran C", "Ran C",
"ViperC", "ViperC",
"Sangheili460",
"MagnaInsomnia",
"Akira_HentAI",
"Karl P.",
"Adam Taylor", "Adam Taylor",
"Weird_With_A_Beard", "Weird_With_A_Beard",
"N/A",
"The Spawn", "The Spawn",
"graysock", "graysock",
"Pozadine1", "Pozadine1",
@@ -162,7 +171,8 @@
"AIGooner", "AIGooner",
"Luc", "Luc",
"ProtonPrince", "ProtonPrince",
"Greenmoustache", "DiffDuck",
"elu3199",
"fancypants", "fancypants",
"John+Edwards", "John+Edwards",
"Joboshy", "Joboshy",
@@ -172,42 +182,39 @@
"contrite831", "contrite831",
"Dan", "Dan",
"Bro Xie", "Bro Xie",
"yer fey",
"batblue", "batblue",
"carey6409", "carey6409",
"Olive",
"太郎 ゲーム", "太郎 ゲーム",
"Roslynd",
"jinxedx", "jinxedx",
"Neco28",
"David Ortega",
"AELOX", "AELOX",
"Gooohokrbe", "Gooohokrbe",
"Dankin-Pics", "Dankin-Pics",
"Nicfit23", "Nicfit23",
"Cristian Vazquez", "Cristian Vazquez",
"wamekukyouzin",
"OldBones", "OldBones",
"drum matthieu",
"Dogmaster",
"Frank Nitty", "Frank Nitty",
"Magic Noob", "Magic Noob",
"Christopher Michel",
"Zach Gonser", "Zach Gonser",
"Serge Bekenkamp",
"DougPeterson", "DougPeterson",
"LeoZero",
"Antonio Pontes", "Antonio Pontes",
"nahinahi9", "Bruce",
"kushiroK9",
"Kevin John Duck", "Kevin John Duck",
"Dustin Chen", "Dustin Chen",
"Kevin Christopher",
"Blackfish95", "Blackfish95",
"Mouthlessman",
"Paul Kroll", "Paul Kroll",
"Penfore", "Penfore",
"Bas Imagineer", "Bas Imagineer",
"John Statham",
"Gordon Cole", "Gordon Cole",
"AbstractAss", "AbstractAss",
"Dušan Ryban", "Dušan Ryban",
"decoy", "decoy",
"DiffDuck",
"elu3199",
"Hasturkun", "Hasturkun",
"Jon Sandman", "Jon Sandman",
"Ubivis", "Ubivis",
@@ -222,34 +229,34 @@
"MJG", "MJG",
"David LaVallee", "David LaVallee",
"linnfrey", "linnfrey",
"ae",
"Tr4shP4nda",
"Jackthemind", "Jackthemind",
"griffin+dahlberg", "griffin+dahlberg",
"jeaness", "jeaness",
"takyamtom", "takyamtom",
"Brian M",
"Josef Lanzl", "Josef Lanzl",
"Nerezza", "Nerezza",
"yer fey", "sanborondon",
"Error_Rule34_Not_found", "Error_Rule34_Not_found",
"aezin", "aezin",
"jcay015", "jcay015",
"Erik Lopez", "Erik Lopez",
"Roslynd",
"Mateo Curić", "Mateo Curić",
"Geolog", "Geolog",
"Neco28",
"Cosmosis", "Cosmosis",
"Eris3D", "Eris3D",
"David Ortega", "m",
"FloPro4Sho", "FloPro4Sho",
"Jamie Ogletree",
"a _", "a _",
"Jeff", "Jeff",
"Bruce",
"Steven Owens", "Steven Owens",
"James Coleman", "James Coleman",
"Kevin Christopher",
"Chad Idk", "Chad Idk",
"dd", "dd",
"John Statham", "Sam",
"sjon kreutz", "sjon kreutz",
"yuxz69", "yuxz69",
"LarsesFPC", "LarsesFPC",
@@ -257,8 +264,6 @@
"esthe", "esthe",
"AlexDuKaNa", "AlexDuKaNa",
"地獄の禄", "地獄の禄",
"ae",
"Tr4shP4nda",
"Gamalonia", "Gamalonia",
"capn", "capn",
"Joseph", "Joseph",
@@ -272,12 +277,16 @@
"Hailshem", "Hailshem",
"Naomi Hale Danchi", "Naomi Hale Danchi",
"epicgamer0020690", "epicgamer0020690",
"Joshua Porrata",
"SuBu",
"RedPIXel",
"Wind",
"IamAyam", "IamAyam",
"Andrew", "Andrew",
"Brian M",
"Robert Wegemund", "Robert Wegemund",
"sanborondon", "Littlehuggy",
"confiscated Zyra", "Andrew Marshall",
"Brian Buie",
"Taylor Funk", "Taylor Funk",
"Thought2Form", "Thought2Form",
"Gerald Welly", "Gerald Welly",
@@ -285,15 +294,19 @@
"Sadlip", "Sadlip",
"Tee Gee", "Tee Gee",
"tarek helmi", "tarek helmi",
"Joey Callahan",
"Max Marklund", "Max Marklund",
"m", "Mike Simone",
"Pierce McBride", "Pierce McBride",
"Joshua Gray", "Joshua Gray",
"Pronredn", "Pronredn",
"Mikko Hemilä", "Mikko Hemilä",
"Jamie Ogletree", "Jacob McDaniel",
"X",
"Temikus", "Temikus",
"Artokun",
"Michael Taylor", "Michael Taylor",
"Derek Baker",
"lh qwe", "lh qwe",
"Martial", "Martial",
"conner", "conner",
@@ -305,24 +318,21 @@
"Decx _", "Decx _",
"Yuji Kaneko", "Yuji Kaneko",
"Rops Alot", "Rops Alot",
"Sam",
"Ace Ventura", "Ace Ventura",
"四糸凜音", "四糸凜音",
"Xeeosat", "Xeeosat",
"Douglas Gaspar", "Douglas Gaspar",
"Saya",
"George", "George",
"dw", "dw",
"FrxzenSnxw",
"WRL_SPR", "WRL_SPR",
"momokai", "momokai",
"몽타주", "몽타주",
"kudari", "kudari",
"ken", "ken",
"Crocket", "Crocket",
"Joshua Porrata",
"keemun", "keemun",
"SuBu",
"RedPIXel",
"Wind",
"Nexus", "Nexus",
"Ramneek“Guy”Ashok", "Ramneek“Guy”Ashok",
"squid_actually", "squid_actually",
@@ -337,37 +347,36 @@
"KitKatM", "KitKatM",
"socrasteeze", "socrasteeze",
"OrganicArtifact", "OrganicArtifact",
"ResidentDeviant",
"MudkipMedkitz", "MudkipMedkitz",
"deanbrian", "deanbrian",
"Alex Wortman", "Alex Wortman",
"Cody", "Cody",
"emadsultan", "emadsultan",
"InformedViewz",
"CHKeeho80",
"Bubbafett",
"leaf",
"Adam Rinehart",
"Pitpe11",
"TheD1rtyD03",
"gzmzmvp", "gzmzmvp",
"Richard", "Richard",
"奚明 刘", "奚明 刘",
"Littlehuggy",
"Aberr", "Aberr",
"Gregory Kozhemiak", "Gregory Kozhemiak",
"준희 김", "준희 김",
"Brian Buie",
"Eric Whitney", "Eric Whitney",
"Joey Callahan",
"Ivan Tadic", "Ivan Tadic",
"Tomohiro Baba", "Tomohiro Baba",
"Mike Simone",
"Noora", "Noora",
"John J Linehan", "John J Linehan",
"Mattssn",
"Elliot E", "Elliot E",
"Morgandel", "Morgandel",
"Theerat Jiramate", "Theerat Jiramate",
"Noah", "Noah",
"Jacob McDaniel",
"X",
"Sloan Steddy", "Sloan Steddy",
"Artokun",
"hexxish", "hexxish",
"Derek Baker",
"Steam Steam", "Steam Steam",
"NICHOLAS BAXLEY", "NICHOLAS BAXLEY",
"CryptoTraderJK", "CryptoTraderJK",
@@ -378,23 +387,14 @@
"Fotek Design", "Fotek Design",
"Nihongasuki", "Nihongasuki",
"MadSpin", "MadSpin",
"FrxzenSnxw",
"inbijiburu", "inbijiburu",
"Nick “Loadstone” D", "Nick “Loadstone” D",
"starbugx", "starbugx",
"dc7431", "dc7431",
"ResidentDeviant",
"Ginnie", "Ginnie",
"Raku", "Raku",
"InformedViewz",
"CHKeeho80",
"Bubbafett",
"leaf",
"Vir", "Vir",
"Skyfire83", "Skyfire83",
"Adam Rinehart",
"Pitpe11",
"TheD1rtyD03",
"moonpetal", "moonpetal",
"g9p0o", "g9p0o",
"Pkrsky", "Pkrsky",
@@ -403,6 +403,8 @@
"SpringBootisTrash", "SpringBootisTrash",
"carsten", "carsten",
"ikok", "ikok",
"quantenmecha",
"Jason+Nash",
"DarkRoast", "DarkRoast",
"Nasty+Hobbit", "Nasty+Hobbit",
"letzte", "letzte",
@@ -414,12 +416,15 @@
"David Schenck", "David Schenck",
"Wolfe7D1", "Wolfe7D1",
"Draven T", "Draven T",
"Time Valentine",
"elleshar666", "elleshar666",
"ACTUALLY_the_Real_Willem_Dafoe", "ACTUALLY_the_Real_Willem_Dafoe",
"Михал Михалыч", "Михал Михалыч",
"Matt",
"Aquatic Coffee", "Aquatic Coffee",
"Kauffy", "Kauffy",
"ethanfel", "ethanfel",
"SPJ",
"Focuschannel", "Focuschannel",
"Edward Kennedy", "Edward Kennedy",
"Nick Kage", "Nick Kage",
@@ -432,12 +437,13 @@
"notedfakes", "notedfakes",
"Michael Scott", "Michael Scott",
"Pat Hen", "Pat Hen",
"Saya", "Solixer",
"Jordan Shaw", "Jordan Shaw",
"Wes Sims", "Wes Sims",
"Donor4115", "Donor4115",
"g unit", "g unit",
"Jimmy Borup", "Jimmy Borup",
"Manu Thetug",
"Filippo Ferrari", "Filippo Ferrari",
"JC", "JC",
"Prompt Pirate", "Prompt Pirate",
@@ -451,6 +457,11 @@
"SomeDude", "SomeDude",
"nanana", "nanana",
"raf8osz", "raf8osz",
"Bob+Barker",
"D",
"Dark_Pest",
"Eldithor",
"Alex",
"Karru", "Karru",
"ChaChanoKo", "ChaChanoKo",
"redcarrot", "redcarrot",
@@ -467,36 +478,34 @@
"Doug+Rintoul", "Doug+Rintoul",
"Noor", "Noor",
"Yorunai", "Yorunai",
"quantenmecha",
"Jason+Nash",
"cocona", "cocona",
"blikkies", "blikkies",
"JBsuede", "JBsuede",
"Time Valentine",
"Shock Shockor", "Shock Shockor",
"りん あめ", "りん あめ",
"Matt",
"Goldwaters", "Goldwaters",
"Zude", "Zude",
"Joaquin Hierrezuelo",
"Frogmilk", "Frogmilk",
"SPJ", "Sean voets",
"Kyler", "Kyler",
"Kor", "Kor",
"Joseph Hanson",
"John Rednoulf",
"Bryan Rutkowski", "Bryan Rutkowski",
"Justin Blaylock", "Justin Blaylock",
"aRtFuL_DodGeR", "aRtFuL_DodGeR",
"Steven",
"TenaciousD", "TenaciousD",
"Dmitry Ryzhov", "Dmitry Ryzhov",
"Edward Ten Eyck", "Edward Ten Eyck",
"Billy Gladky", "Billy Gladky",
"Probis", "Probis",
"Solixer",
"Pete Pain", "Pete Pain",
"ItsGeneralButtNaked", "ItsGeneralButtNaked",
"RHopkirk", "RHopkirk",
"jinksta187", "jinksta187",
"robin.kok.", "robin.kok.",
"Manu Thetug",
"Maxim", "Maxim",
"Karlanx", "Karlanx",
"Lyavph", "Lyavph",
@@ -504,6 +513,7 @@
"Youguang", "Youguang",
"andrewzpong", "andrewzpong",
"BossGame", "BossGame",
"Marcus thronico",
"lrdchs", "lrdchs",
"Tree Tagger", "Tree Tagger",
"Inversity", "Inversity",
@@ -511,6 +521,15 @@
"Kevinj", "Kevinj",
"Mitchell Robson", "Mitchell Robson",
"POPPIN", "POPPIN",
"PoorStudent",
"Alex+Zaw",
"Supporter",
"ExLightSaber",
"Mobius2020",
"YaboiRay",
"Sildoren",
"Darv",
"Seon+Song",
"2turbo", "2turbo",
"Dmitry+Viznesenskiy", "Dmitry+Viznesenskiy",
"tanjin90", "tanjin90",
@@ -528,11 +547,6 @@
"Inkognito", "Inkognito",
"G", "G",
"Tan+Huynh", "Tan+Huynh",
"Bob+Barker",
"D",
"Dark_Pest",
"Eldithor",
"Alex",
"BillyBoy84", "BillyBoy84",
"Buecyb99", "Buecyb99",
"Welkor", "Welkor",
@@ -545,28 +559,27 @@
"G", "G",
"Ronan Delevacq", "Ronan Delevacq",
"Christian Schäfer", "Christian Schäfer",
"Leslie Andrew Ridings",
"Dave Abraham", "Dave Abraham",
"Joaquin Hierrezuelo",
"Locrospiel", "Locrospiel",
"Sean voets",
"Jarrid Lee", "Jarrid Lee",
"Poophead27 Blyat", "Poophead27 Blyat",
"Joseph Hanson",
"John Rednoulf",
"Kyron Mahan", "Kyron Mahan",
"Mythspire", "Mythspire",
"Boba Smith", "Boba Smith",
"TBitz33", "TBitz33",
"Anonym dkjglfleeoeldldldlkf", "Anonym dkjglfleeoeldldldlkf",
"MR.Bear", "MR.Bear",
"matt",
"somethingtosay8", "somethingtosay8",
"Ezokewn", "Ezokewn",
"Terminuz",
"ivistorm", "ivistorm",
"SendingRavens", "SendingRavens",
"Sauv", "Sauv",
"Steven",
"JackJohnnyJim", "JackJohnnyJim",
"Khánh Đặng", "Khánh Đặng",
"Borte",
"Michael Docherty", "Michael Docherty",
"Ted Cart", "Ted Cart",
"Sage Himeros", "Sage Himeros",
@@ -574,6 +587,7 @@
"Paul Hartsuyker", "Paul Hartsuyker",
"elitassj", "elitassj",
"Tigon", "Tigon",
"SkibidiRizzler",
"Tania Nayelli Fernandez", "Tania Nayelli Fernandez",
"Draconach", "Draconach",
"Jacob Winter", "Jacob Winter",
@@ -581,6 +595,8 @@
"Andrew Wilkinson", "Andrew Wilkinson",
"David", "David",
"Meilo", "Meilo",
"Nacho Ferrando",
"Marcos Tortosa Carmona",
"Dkom22", "Dkom22",
"shinonomeiro", "shinonomeiro",
"Snille", "Snille",
@@ -589,7 +605,6 @@
"xybrightsummer", "xybrightsummer",
"jreedatchison", "jreedatchison",
"PhilW", "PhilW",
"Marcus thronico",
"Janik", "Janik",
"Cruel", "Cruel",
"MRBlack", "MRBlack",
@@ -601,6 +616,13 @@
"Scott", "Scott",
"Muratoraccio", "Muratoraccio",
"D", "D",
"Somebody",
"Celestial+Kitten",
"TequiTequi",
"Homero+Banda",
"bakeliteboy",
"Nick",
"てぃんてぃんひーろー",
"Gold_miner_ego", "Gold_miner_ego",
"IshouI;_;", "IshouI;_;",
"Monix", "Monix",
@@ -620,17 +642,8 @@
"you+halo9", "you+halo9",
"cloudghost", "cloudghost",
"Yongkwan+Lee", "Yongkwan+Lee",
"PoorStudent",
"lucites", "lucites",
"Alex+Zaw",
"Mobius2020",
"ExLightSaber",
"YaboiRay",
"nickname", "nickname",
"Sildoren",
"Darv",
"Seon+Song",
"Somebody",
"Balut+Omelette", "Balut+Omelette",
"eriick", "eriick",
"Lev+Lanevskiy", "Lev+Lanevskiy",
@@ -651,37 +664,35 @@
"Vinarus", "Vinarus",
"Josh Snyder", "Josh Snyder",
"ja s", "ja s",
"Leslie Andrew Ridings",
"Doug Mason", "Doug Mason",
"scoreswazey", "scoreswazey",
"Oliverfish",
"Owen Gwosdz", "Owen Gwosdz",
"Room Light", "Room Light",
"Patryk Serious", "Patryk Serious",
"AZ Party Oasis", "AZ Party Oasis",
"Devil Lude", "Gentle Sartori",
"Snorklebort", "Snorklebort",
"David Murcko", "David Murcko",
"vinter",
"TheFusion", "TheFusion",
"Jack Dole", "Jack Dole",
"matt",
"3zS4QNQ4", "3zS4QNQ4",
"Terminuz",
"max blo", "max blo",
"Matt M.", "Matt M.",
"Ivan Imes", "Ivan Imes",
"J M", "J M",
"Slacks",
"Bouya shaka", "Bouya shaka",
"Jack Lawfield", "Jack Lawfield",
"Borte",
"Maso", "Maso",
"Homero Banda",
"yyuvuvu", "yyuvuvu",
"Eric Ketchum",
"Nomki", "Nomki",
"Kevin Wallace", "Kevin Wallace",
"ChicRic", "ChicRic",
"BastardSama", "BastardSama",
"mercur", "mercur",
"SkibidiRizzler",
"Never_M", "Never_M",
"Kalle Björk", "Kalle Björk",
"Yavizu3d", "Yavizu3d",
@@ -689,9 +700,7 @@
"Teriak47", "Teriak47",
"Just me", "Just me",
"Raf Stahelin", "Raf Stahelin",
"Nacho Ferrando",
"Вячеслав Маринин", "Вячеслав Маринин",
"Marcos Tortosa Carmona",
"Cola Matthew", "Cola Matthew",
"OniNoKen", "OniNoKen",
"Iain Wisely", "Iain Wisely",
@@ -734,6 +743,15 @@
"SelfishMedic", "SelfishMedic",
"adderleighn", "adderleighn",
"EnragedAntelope", "EnragedAntelope",
"Brandon+G",
"fazefour33",
"plonk",
"Kotetsu",
"o",
"Tony+V",
"Anvil+G",
"draganjankovic1975dj528",
"MrSEIGE88",
"yarsev", "yarsev",
"M+Alsulaiti", "M+Alsulaiti",
"Mark+Staaf", "Mark+Staaf",
@@ -751,11 +769,6 @@
"miduzza", "miduzza",
"KB", "KB",
"shw", "shw",
"Celestial+Kitten",
"bakeliteboy",
"TequiTequi",
"Homero+Banda",
"Nick",
"Jim", "Jim",
"JoL", "JoL",
"YoruHime", "YoruHime",
@@ -781,14 +794,14 @@
"han b", "han b",
"Nico", "Nico",
"Maximilian Krischan", "Maximilian Krischan",
"Banana Joe", "socialcat",
"proto merp", "proto merp",
"_ G3n", "_ G3n",
"Brandon Thomas", "Brandon Thomas",
"Donovan Jenkins", "Donovan Jenkins",
"Hans Meier", "Hans Meier",
"Dustin Hendel", "Dustin Hendel",
"sicarius", "jboul",
"Michael Eid", "Michael Eid",
"Liberation", "Liberation",
"Bob barker", "Bob barker",
@@ -805,7 +818,6 @@
"jumpd", "jumpd",
"John C", "John C",
"Rim", "Rim",
"Oliverfish",
"yfx507", "yfx507",
"uruksayshi", "uruksayshi",
"Jairus Knudsen", "Jairus Knudsen",
@@ -814,24 +826,25 @@
"nk8", "nk8",
"lylepaul", "lylepaul",
"Middo", "Middo",
"Gary Chaboya",
"Forbidden Atelier", "Forbidden Atelier",
"Thomas Sankowski", "Thomas Sankowski",
"DrB", "DrB",
"Nimhloth", "Nimhloth",
"Adictedtohumping", "Adictedtohumping",
"Moneymaker412K", "Moneymaker412K",
"vinter", "Tsani Prodanov",
"Towelie", "Towelie",
"Jean-françois SEMA", "Jean-françois SEMA",
"Myrthrac",
"Taylor Dominy",
"Andrew Ly", "Andrew Ly",
"Slacks",
"Glenn Hoetker", "Glenn Hoetker",
"john Greene", "john Greene",
"Faburizu", "Faburizu",
"jimyjomson", "jimyjomson",
"JaeHyun Jang", "JaeHyun Jang",
"Michael Hicks", "Michael Hicks",
"Homero Banda",
"Chase Kwon", "Chase Kwon",
"Bob Ling", "Bob Ling",
"Inyoshu", "Inyoshu",
@@ -858,5 +871,5 @@
"Somebody", "Somebody",
"CK" "CK"
], ],
"totalCount": 855 "totalCount": 868
} }
+65
View File
@@ -0,0 +1,65 @@
# ComfyUI Dual-Mode Widget Rendering
ComfyUI custom node widgets render in one of two modes. Patterns that work in one often fail silently in the other. Test both.
## Mode Detection
```js
typeof LiteGraph !== 'undefined' && LiteGraph.vueNodesMode
```
In Vue SFCs, `window.LiteGraph` is unavailable — pass as a prop from `main.ts`.
## Canvas Mode Layout
Uses `computeLayoutSize()` + `distributeSpace()` to allocate widget height within the node. Widgets with `computeLayoutSize` participate in space distribution; those with `computeSize` have fixed height.
- `getMinHeight()` in `addDOMWidget` options → minimum widget height
- `widget.computeLayoutSize()``{ minHeight, minWidth, maxHeight? }`
- Avoid `getMaxHeight()` unless the widget genuinely needs a fixed cap (prevents user resize)
## Vue Mode Layout
Uses CSS Grid (`grid-template-rows`) + `ResizeObserver`. The ResizeObserver watches the widget's DOM and feeds back into grid row sizing. This creates a feedback loop: content grows → row resizes → more space for content → content reflows/grows → row resizes again.
### Height Containment
The fix: `contain: layout size` on the widget root. This tells the browser the element's intrinsic size is CSS-determined, not driven by descendant content. The ResizeObserver sees a stable size and the loop is broken.
```css
.widget-root.lm-vue-node {
height: 100%;
min-height: var(--comfy-widget-min-height, 200px);
contain: layout size;
}
```
Existing examples: `.lm-loras-container.lm-vue-node` and `.comfy-tags-container.lm-vue-node` in `web/comfyui/lm_styles.css`.
**Do NOT** fix height issues with `maxHeight`, `getMaxHeight()`, or inline `max-height` — these prevent the user from resizing the node.
## Scroll Wheel Isolation
Both modes need to distinguish "user wants to scroll widget content" from "user wants to zoom canvas".
**Canvas mode:** Add `@wheel` on widget root. Check `event.target.closest(selector)` for scrollable sub-areas. If scrollable → `event.stopPropagation()`. Otherwise → `app.canvas.processMouseWheel(event)`.
**Vue mode:** Add CSS class `lm-wheel-scrollable` to scrollable elements. The global capture-phase hook in `web/comfyui/utils.js` (`enableListWheelScroll`) detects wheel events on marked elements and manually scrolls them via `element.scrollTop`, consuming the event before canvas zoom sees it.
## DOM Structure
`main.ts` creates an outer `<div>` container, then `vueApp.mount(container)`. The Vue app renders its own root element inside.
- `container.id` / `container.style.*` → outer element
- Vue scoped `<style>``[data-v-hash]` applies only to Vue root
Classes needed by scoped Vue CSS must go on the Vue root element. Pass data as props and bind with `:class` rather than manipulating the DOM from `main.ts`.
## Serialization
For stateful widgets that need workflow persistence:
- `serialize: true` in `addDOMWidget` options
- `serializeValue()` → state snapshot (called on workflow save)
- `onSetValue(v)` → restore state (called on workflow load)
- Always handle missing keys in restored value for backward compatibility with old workflows
File diff suppressed because one or more lines are too long
+2219 -2195
View File
File diff suppressed because it is too large Load Diff
+26 -2
View File
@@ -233,7 +233,7 @@
"presetNamePlaceholder": "Preset name...", "presetNamePlaceholder": "Preset name...",
"baseModel": "Base Model", "baseModel": "Base Model",
"baseModelSearchPlaceholder": "Search base models...", "baseModelSearchPlaceholder": "Search base models...",
"modelTags": "Tags (Top 20)", "modelTags": "Tags",
"modelTypes": "Model Types", "modelTypes": "Model Types",
"license": "License", "license": "License",
"noCreditRequired": "No Credit Required", "noCreditRequired": "No Credit Required",
@@ -241,6 +241,8 @@
"allowSellingGeneratedContentTooltip": "Allow selling generated images", "allowSellingGeneratedContentTooltip": "Allow selling generated images",
"noCreditRequiredTooltip": "Use the model without crediting the creator", "noCreditRequiredTooltip": "Use the model without crediting the creator",
"noTags": "No tags", "noTags": "No tags",
"tagSearchPlaceholder": "Search tags...",
"noTagMatches": "No tags match the current search.",
"autoTags": "Auto Tags", "autoTags": "Auto Tags",
"noBaseModelMatches": "No base models match the current search.", "noBaseModelMatches": "No base models match the current search.",
"clearAll": "Clear All Filters", "clearAll": "Clear All Filters",
@@ -640,7 +642,13 @@
"preparing": "Preparing download...", "preparing": "Preparing download...",
"connecting": "Connecting to download server...", "connecting": "Connecting to download server...",
"completed": "Completed", "completed": "Completed",
"downloadComplete": "Download completed successfully" "downloadComplete": "Download completed successfully",
"enableCivarchiveApi": "Enable CivArchive API as metadata provider",
"enableCivarchiveApiHelp": "When on, CivArchive API is used as a fallback source for model metadata (e.g. for models deleted from CivitAI). Turn off to avoid CivArchive rate limits entirely.",
"providerOrder": "Metadata provider fallback order",
"providerOrderHelp": "CivitAI API is always tried first. Choose the order of the remaining providers when looking up metadata.",
"providerOrderCivitaiArchiveSqlite": "CivitAI → CivArchive → Archive DB",
"providerOrderCivitaiSqliteArchive": "CivitAI → Archive DB → CivArchive"
}, },
"proxySettings": { "proxySettings": {
"enableProxy": "Enable App-level Proxy", "enableProxy": "Enable App-level Proxy",
@@ -1540,6 +1548,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?"
}, },
@@ -1743,6 +1752,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...",
@@ -1763,6 +1778,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 remove the Git repository and install the latest stable release. Future updates will use stable releases only.",
"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.",
+2219 -2195
View File
File diff suppressed because it is too large Load Diff
+2219 -2195
View File
File diff suppressed because it is too large Load Diff
+2219 -2195
View File
File diff suppressed because it is too large Load Diff
+2219 -2195
View File
File diff suppressed because it is too large Load Diff
+2219 -2195
View File
File diff suppressed because it is too large Load Diff
+2219 -2195
View File
File diff suppressed because it is too large Load Diff
+2219 -2195
View File
File diff suppressed because it is too large Load Diff
+2219 -2195
View File
File diff suppressed because it is too large Load Diff
+47 -4
View File
@@ -208,6 +208,12 @@ class Config:
if not isinstance(library_config, dict): if not isinstance(library_config, dict):
return return
# Always read recipes_path — it is independent of extra folder paths
# and must be set before any early returns below.
recipes_path = library_config.get("recipes_path", "")
if isinstance(recipes_path, str) and recipes_path:
self.recipes_path = recipes_path
extra_folder_paths = library_config.get("extra_folder_paths") extra_folder_paths = library_config.get("extra_folder_paths")
if not isinstance(extra_folder_paths, dict): if not isinstance(extra_folder_paths, dict):
return return
@@ -233,10 +239,6 @@ class Config:
extra_embedding extra_embedding
) )
recipes_path = library_config.get("recipes_path", "")
if isinstance(recipes_path, str) and recipes_path:
self.recipes_path = recipes_path
if self.extra_loras_roots: if self.extra_loras_roots:
logger.info( logger.info(
"Found extra LoRA roots:" "Found extra LoRA roots:"
@@ -357,6 +359,47 @@ class Config:
"Failed to rename legacy 'default' library: %s", rename_error "Failed to rename legacy 'default' library: %s", rename_error
) )
# Clean up a stale "default" library entry that has no meaningful
# paths configured (e.g. leftover bootstrap artifact). This only
# fires when "comfyui" already exists so we never delete the last
# remaining library.
if (
"default" in libraries
and "comfyui" in libraries
and isinstance(default_library, Mapping)
):
default_folder_paths = _normalize_library_folder_paths(
default_library
)
default_extra_paths = default_library.get("extra_folder_paths", {})
has_meaningful_paths = bool(default_folder_paths) or bool(
default_extra_paths
) or any(
default_library.get(key)
for key in (
"default_lora_root",
"default_checkpoint_root",
"default_unet_root",
"default_embedding_root",
"recipes_path",
)
)
if not has_meaningful_paths:
try:
settings_service.delete_library("default")
libraries_changed = True
logger.info(
"Removed stale 'default' library entry "
"with no meaningful paths configured"
)
libraries = settings_service.get_libraries()
comfy_library = libraries.get("comfyui", {})
except Exception as delete_error:
logger.debug(
"Failed to remove stale 'default' library: %s",
delete_error,
)
default_lora_root = _resolve_valid_default_root( default_lora_root = _resolve_valid_default_root(
comfy_library.get("default_lora_root", ""), comfy_library.get("default_lora_root", ""),
list(self.loras_roots or []), list(self.loras_roots or []),
+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]
+16 -6
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)}")
@@ -135,10 +137,13 @@ class MetadataHook:
# Store the dynprompt reference for node lookups # Store the dynprompt reference for node lookups
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)
# Replace the functions # Replace the functions
execution._map_node_over_list = map_node_over_list_with_metadata execution._map_node_over_list = map_node_over_list_with_metadata
execution.execute = execute_with_prompt_tracking execution.execute = execute_with_prompt_tracking
@@ -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)
+138 -14
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():
@@ -539,7 +629,22 @@ 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, {}):
@@ -568,7 +673,26 @@ class MetadataProcessor:
break break
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)
+103 -5
View File
@@ -2,7 +2,7 @@ import json
import os import os
import re import re
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER from .constants import CLIP_SKIP_SENTINEL, MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE, METADATA_OVERWRITE_FIELDS
def _store_checkpoint_metadata(metadata, node_id, model_name): def _store_checkpoint_metadata(metadata, node_id, model_name):
@@ -31,11 +31,78 @@ 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
def extract(node_id, inputs, outputs, metadata): def extract(node_id, inputs, outputs, metadata):
@@ -1154,6 +1221,35 @@ 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 = {}
for key in METADATA_OVERWRITE_FIELDS:
value = inputs.get(key)
if key == "clip_skip":
if value != CLIP_SKIP_SENTINEL:
overwrite_params[key] = value
elif value: # truthy — only overwrite when user provided a real value
overwrite_params[key] = value
if overwrite_params:
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 +1317,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
} }
+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)
+45
View File
@@ -0,0 +1,45 @@
"""Lora Info display node — pure frontend node for showing selected LoRA info.
This node does NOT participate in workflow execution. Its single optional
"lora_source" input exists solely as a wire-connection anchor so that the
frontend can traverse the graph and push selection data to connected info nodes.
"""
from __future__ import annotations
class LoraInfoLM:
"""Display node that shows filename and notes for the selected LoRA."""
NAME = "Lora Info (LoraManager)"
CATEGORY = "Lora Manager/utils"
DESCRIPTION = (
"Displays information (filename, notes) about the currently selected "
"LoRA. Connect any output from a LoRA Loader or Stacker to the "
"lora_source input, then select a LoRA in the source widget — the "
"info updates automatically. Does not affect workflow execution."
)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {},
}
RETURN_TYPES = ()
RETURN_NAMES = ()
OUTPUT_NODE = False
FUNCTION = "noop"
def noop(self, **kwargs):
# This node is display-only — no workflow execution needed.
return ()
NODE_CLASS_MAPPINGS = {
LoraInfoLM.NAME: LoraInfoLM,
}
NODE_DISPLAY_NAME_MAPPINGS = {
LoraInfoLM.NAME: "Lora Info (LoraManager)",
}
+2 -17
View File
@@ -1,6 +1,5 @@
import importlib import importlib
import logging import logging
import re
import comfy.sd # type: ignore import comfy.sd # type: ignore
import comfy.utils # type: ignore import comfy.utils # type: ignore
@@ -14,6 +13,7 @@ from .utils import (
extract_lora_name, extract_lora_name,
get_loras_list, get_loras_list,
nunchaku_load_lora, nunchaku_load_lora,
parse_lora_syntax,
) )
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -189,25 +189,10 @@ class LoraTextLoaderLM:
RETURN_NAMES = ("MODEL", "CLIP", "trigger_words", "loaded_loras") RETURN_NAMES = ("MODEL", "CLIP", "trigger_words", "loaded_loras")
FUNCTION = "load_loras_from_text" FUNCTION = "load_loras_from_text"
def parse_lora_syntax(self, text):
"""Parse LoRA syntax from text input."""
pattern = r"<lora:([^:>]+):([^:>]+)(?::([^:>]+))?>"
matches = re.findall(pattern, text, re.IGNORECASE)
loras = []
for match in matches:
model_strength = float(match[1])
loras.append({
"name": match[0],
"model_strength": model_strength,
"clip_strength": float(match[2]) if match[2] else model_strength,
})
return loras
def load_loras_from_text(self, model, lora_syntax, clip=None, lora_stack=None): def load_loras_from_text(self, model, lora_syntax, clip=None, lora_stack=None):
"""Load LoRAs based on text syntax input.""" """Load LoRAs based on text syntax input."""
lora_entries = _collect_stack_entries(lora_stack) lora_entries = _collect_stack_entries(lora_stack)
for lora in self.parse_lora_syntax(lora_syntax): for lora in parse_lora_syntax(lora_syntax):
lora_path, trigger_words = get_lora_info_absolute(lora["name"]) lora_path, trigger_words = get_lora_info_absolute(lora["name"])
lora_entries.append({ lora_entries.append({
"name": lora["name"], "name": lora["name"],
+62
View File
@@ -0,0 +1,62 @@
"""Node to resolve `<lora:name:strength>` syntax to absolute file system paths.
Takes the loaded_loras / active_loras STRING output from LoraLoaderLM or
LoraStackerLM and resolves each lora name to its absolute path on disk via
the scanner cache. Unknown names are returned as-is.
"""
import logging
from ..utils.utils import get_lora_info_absolute
from .utils import parse_lora_syntax
logger = logging.getLogger(__name__)
class LoraSyntaxToPath:
NAME = "LoRA Syntax → Path (LoraManager)"
CATEGORY = "Lora Manager/utils"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"lora_syntax": (
"STRING",
{
"forceInput": True,
"multiline": True,
"tooltip": (
"<lora:name:strength> formatted text from "
"loaded_loras / active_loras output"
),
},
),
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("paths",)
FUNCTION = "resolve"
def resolve(self, lora_syntax: str) -> tuple[str]:
"""Parse <lora:...> syntax and resolve each name to its absolute path."""
if not lora_syntax or not lora_syntax.strip():
logger.info("Received empty lora_syntax input")
return ("",)
parsed = parse_lora_syntax(lora_syntax)
if not parsed:
logger.info("No valid <lora:...> entries found in input")
return ("",)
paths: list[str] = []
for entry in parsed:
try:
absolute_path, _ = get_lora_info_absolute(entry["name"])
paths.append(absolute_path)
except Exception:
logger.warning("Failed to resolve lora '%s', skipping", entry["name"])
continue
return ("\n".join(paths),)
+170
View File
@@ -0,0 +1,170 @@
"""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,
METADATA_OVERWRITE_FIELDS,
)
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",
{
"default": "",
"tooltip": (
"The checkpoint or diffusion model (UNet) used "
"for generation. 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.
"""
result: dict[str, Any] = {}
for key in METADATA_OVERWRITE_FIELDS:
value = kwargs.get(key)
if key == "clip_skip":
if value != _CLIP_SKIP_SENTINEL:
result[key] = value
elif value:
result[key] = value
return (result,)
+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,
+20
View File
@@ -36,6 +36,7 @@ any_type = AnyType("*")
# Common methods extracted from lora_loader.py and lora_stacker.py # Common methods extracted from lora_loader.py and lora_stacker.py
import os import os
import re
import logging import logging
import copy import copy
import sys import sys
@@ -69,6 +70,25 @@ def extract_lora_name(lora_path):
return apply_lora_syntax_format(name_no_ext) return apply_lora_syntax_format(name_no_ext)
def parse_lora_syntax(text: str) -> list[dict]:
"""Parse <lora:name:strength> syntax from text input into a list of dicts.
Each entry contains: name, model_strength, clip_strength.
Supports both ``<lora:name:strength>`` and ``<lora:name:model_strength:clip_strength>``.
"""
pattern = r"<lora:([^:>]+):([^:>]+)(?::([^:>]+))?>"
matches = re.findall(pattern, text, re.IGNORECASE)
loras = []
for match in matches:
model_strength = float(match[1])
loras.append({
"name": match[0],
"model_strength": model_strength,
"clip_strength": float(match[2]) if match[2] else model_strength,
})
return loras
def get_loras_list(kwargs): def get_loras_list(kwargs):
"""Helper to extract loras list from either old or new kwargs format""" """Helper to extract loras list from either old or new kwargs format"""
if "loras" not in kwargs: if "loras" not in kwargs:
+250 -2
View File
@@ -1570,7 +1570,11 @@ class SettingsHandler:
else: else:
self._settings.set(key, value) self._settings.set(key, value)
if key == "enable_metadata_archive_db": if key in (
"enable_metadata_archive_db",
"enable_civarchive_api",
"metadata_provider_order",
):
await self._metadata_provider_updater() await self._metadata_provider_updater()
if key in self._PROXY_KEYS: if key in self._PROXY_KEYS:
@@ -1784,6 +1788,124 @@ class LoraCodeHandler:
logger.error("Failed to update lora code: %s", exc, exc_info=True) logger.error("Failed to update lora code: %s", 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 get_update_lora_code(self, request: web.Request) -> web.Response:
"""GET version of update_lora_code — reads parameters from query string.
Query params:
lora_code (required) the LoRA syntax to send
mode (optional) "append" (default) or "replace"
node_id (repeatable) target node id(s), e.g. node_id=3&node_id=5
node_ids (optional) JSON-encoded array for complex references with graph_id:
[{"node_id":3,"graph_id":"g1"}, ...]
"""
try:
node_ids_raw = request.query.get("node_ids")
node_id_list = request.query.getall("node_id", [])
lora_code = request.query.get("lora_code", "")
mode = request.query.get("mode", "append")
if not lora_code:
return web.json_response(
{"success": False, "error": "Missing lora_code parameter"},
status=400,
)
node_ids = None
if node_ids_raw:
try:
node_ids = json.loads(node_ids_raw)
except (json.JSONDecodeError, TypeError):
return web.json_response(
{"success": False, "error": "node_ids must be a valid JSON array"},
status=400,
)
if not isinstance(node_ids, list) or not node_ids:
return web.json_response(
{"success": False, "error": "node_ids must be a non-empty JSON array"},
status=400,
)
elif node_id_list:
node_ids = node_id_list
results = []
if node_ids is None:
try:
self._prompt_server.instance.send_sync(
"lora_code_update",
{"id": -1, "lora_code": lora_code, "mode": mode},
)
results.append({"node_id": "broadcast", "success": True})
except Exception as exc: # pragma: no cover - defensive logging
logger.error("Error broadcasting lora code: %s", exc)
results.append(
{"node_id": "broadcast", "success": False, "error": str(exc)}
)
else:
for entry in node_ids:
node_identifier = entry
graph_identifier = None
if isinstance(entry, dict):
node_identifier = entry.get("node_id")
graph_identifier = entry.get("graph_id")
if node_identifier is None:
results.append(
{
"node_id": node_identifier,
"graph_id": graph_identifier,
"success": False,
"error": "Missing node_id parameter",
}
)
continue
try:
parsed_node_id = int(node_identifier)
except (TypeError, ValueError):
parsed_node_id = node_identifier
payload = {
"id": parsed_node_id,
"lora_code": lora_code,
"mode": mode,
}
if graph_identifier is not None:
payload["graph_id"] = str(graph_identifier)
try:
self._prompt_server.instance.send_sync(
"lora_code_update",
payload,
)
results.append(
{
"node_id": parsed_node_id,
"graph_id": payload.get("graph_id"),
"success": True,
}
)
except Exception as exc: # pragma: no cover - defensive logging
logger.error(
"Error sending lora code to node %s (graph %s): %s",
parsed_node_id,
graph_identifier,
exc,
)
results.append(
{
"node_id": parsed_node_id,
"graph_id": payload.get("graph_id"),
"success": False,
"error": str(exc),
}
)
return web.json_response({"success": True, "results": results})
except Exception as exc: # pragma: no cover - defensive logging
logger.error("Failed to update lora code (GET): %s", exc, exc_info=True)
return web.json_response({"success": False, "error": str(exc)}, status=500)
class TrainedWordsHandler: class TrainedWordsHandler:
async def get_trained_words(self, request: web.Request) -> web.Response: async def get_trained_words(self, request: web.Request) -> web.Response:
@@ -3353,7 +3475,7 @@ class NodeRegistryHandler:
status=400, status=400,
) )
if not isinstance(value, str) or not value: if value is None or (isinstance(value, str) and not value):
return web.json_response( return web.json_response(
{"success": False, "error": "Missing value parameter"}, status=400 {"success": False, "error": "Missing value parameter"}, status=400
) )
@@ -3431,6 +3553,130 @@ class NodeRegistryHandler:
logger.error("Failed to update node widget: %s", exc, exc_info=True) logger.error("Failed to update node widget: %s", 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 get_update_node_widget(self, request: web.Request) -> web.Response:
"""GET version of update_node_widget — reads parameters from query string.
Query params:
widget_name (optional) the widget name to update (required unless action is set)
action (optional) alternative action, e.g. "inject_text" (required unless widget_name is set)
value (required) the value to set
mode (optional) "replace" (default) or "append"
node_id (repeatable) target node id(s), e.g. node_id=3&node_id=5
node_ids (optional) JSON-encoded array for complex references:
[{"node_id":3,"graph_id":"g1"}, ...]
"""
try:
widget_name = request.query.get("widget_name")
action = request.query.get("action")
value = request.query.get("value")
mode = request.query.get("mode", "replace")
node_ids_raw = request.query.get("node_ids")
node_id_list = request.query.getall("node_id", [])
if not action and (not isinstance(widget_name, str) or not widget_name):
return web.json_response(
{
"success": False,
"error": "Missing parameter: provide either 'action' or 'widget_name'",
},
status=400,
)
if value is None or (isinstance(value, str) and not value):
return web.json_response(
{"success": False, "error": "Missing value parameter"}, status=400
)
node_ids = None
if node_ids_raw:
try:
node_ids = json.loads(node_ids_raw)
except (json.JSONDecodeError, TypeError):
return web.json_response(
{"success": False, "error": "node_ids must be a valid JSON array"},
status=400,
)
if not isinstance(node_ids, list) or not node_ids:
return web.json_response(
{"success": False, "error": "node_ids must be a non-empty JSON array"},
status=400,
)
elif node_id_list:
node_ids = node_id_list
if not isinstance(node_ids, list) or not node_ids:
return web.json_response(
{"success": False, "error": "node_ids must be a non-empty list"},
status=400,
)
results = []
for entry in node_ids:
node_identifier = entry
graph_identifier = None
if isinstance(entry, dict):
node_identifier = entry.get("node_id")
graph_identifier = entry.get("graph_id")
if node_identifier is None:
results.append(
{
"node_id": node_identifier,
"graph_id": graph_identifier,
"success": False,
"error": "Missing node_id parameter",
}
)
continue
try:
parsed_node_id = int(node_identifier)
except (TypeError, ValueError):
parsed_node_id = node_identifier
payload: dict = {
"id": parsed_node_id,
"value": value,
"mode": mode,
}
if action:
payload["action"] = action
if widget_name:
payload["widget_name"] = widget_name
if graph_identifier is not None:
payload["graph_id"] = str(graph_identifier)
try:
self._prompt_server.instance.send_sync("lm_widget_update", payload)
results.append(
{
"node_id": parsed_node_id,
"graph_id": payload.get("graph_id"),
"success": True,
}
)
except Exception as exc: # pragma: no cover - defensive logging
logger.error(
"Error sending widget update to node %s (graph %s): %s",
parsed_node_id,
graph_identifier,
exc,
)
results.append(
{
"node_id": parsed_node_id,
"graph_id": payload.get("graph_id"),
"success": False,
"error": str(exc),
}
)
return web.json_response({"success": True, "results": results})
except Exception as exc: # pragma: no cover - defensive logging
logger.error("Failed to update node widget (GET): %s", exc, exc_info=True)
return web.json_response({"success": False, "error": str(exc)}, status=500)
class MiscHandlerSet: class MiscHandlerSet:
"""Aggregate handlers into a lookup compatible with the registrar.""" """Aggregate handlers into a lookup compatible with the registrar."""
@@ -3497,10 +3743,12 @@ class MiscHandlerSet:
"update_usage_stats": self.usage_stats.update_usage_stats, "update_usage_stats": self.usage_stats.update_usage_stats,
"get_usage_stats": self.usage_stats.get_usage_stats, "get_usage_stats": self.usage_stats.get_usage_stats,
"update_lora_code": self.lora_code.update_lora_code, "update_lora_code": self.lora_code.update_lora_code,
"get_update_lora_code": self.lora_code.get_update_lora_code,
"get_trained_words": self.trained_words.get_trained_words, "get_trained_words": self.trained_words.get_trained_words,
"get_model_example_files": self.model_examples.get_model_example_files, "get_model_example_files": self.model_examples.get_model_example_files,
"register_nodes": self.node_registry.register_nodes, "register_nodes": self.node_registry.register_nodes,
"update_node_widget": self.node_registry.update_node_widget, "update_node_widget": self.node_registry.update_node_widget,
"get_update_node_widget": self.node_registry.get_update_node_widget,
"get_registry": self.node_registry.get_registry, "get_registry": self.node_registry.get_registry,
"check_model_exists": self.model_library.check_model_exists, "check_model_exists": self.model_library.check_model_exists,
"check_models_exist": self.model_library.check_models_exist, "check_models_exist": self.model_library.check_models_exist,
+82 -13
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:
@@ -973,6 +1007,8 @@ class ModelQueryHandler:
limit = int(request.query.get("limit", "20")) limit = int(request.query.get("limit", "20"))
if limit < 0: if limit < 0:
limit = 20 limit = 20
elif limit > 200:
limit = 20
top_tags = await self._service.get_top_tags(limit) top_tags = await self._service.get_top_tags(limit)
return web.json_response({"success": True, "tags": top_tags}) return web.json_response({"success": True, "tags": top_tags})
except Exception as exc: except Exception as exc:
@@ -981,6 +1017,22 @@ class ModelQueryHandler:
{"success": False, "error": "Internal server error"}, status=500 {"success": False, "error": "Internal server error"}, status=500
) )
async def search_tags(self, request: web.Request) -> web.Response:
try:
query = request.query.get("q", "")
limit = int(request.query.get("limit", "20"))
if limit < 0:
limit = 20
elif limit > 200:
limit = 20
tags = await self._service.search_tags(query, limit)
return web.json_response({"success": True, "tags": tags})
except Exception as exc:
self._logger.error("Error searching tags: %s", exc, exc_info=True)
return web.json_response(
{"success": False, "error": "Internal server error"}, status=500
)
async def get_base_models(self, request: web.Request) -> web.Response: async def get_base_models(self, request: web.Request) -> web.Response:
try: try:
limit = int(request.query.get("limit", "20")) limit = int(request.query.get("limit", "20"))
@@ -1275,9 +1327,13 @@ class ModelQueryHandler:
text=f"{self._service.model_type.capitalize()} file name is required", text=f"{self._service.model_type.capitalize()} file name is required",
status=400, status=400,
) )
notes = await self._service.get_model_notes(model_name) result = await self._service.get_model_notes(model_name)
if notes is not None: if result is not None:
return web.json_response({"success": True, "notes": notes}) return web.json_response({
"success": True,
"notes": result["notes"],
"file_path": result["file_path"],
})
return web.json_response( return web.json_response(
{ {
"success": False, "success": False,
@@ -1783,14 +1839,20 @@ class ModelDownloadHandler:
async def delete_download_history_item(self, request: web.Request) -> web.Response: async def delete_download_history_item(self, request: web.Request) -> web.Response:
try: try:
item_id = int(request.query.get("id", "0")) download_id = request.query.get("download_id")
if not item_id: id_str = request.query.get("id")
item_id = int(id_str) if id_str else None
if not download_id and not item_id:
return web.json_response( return web.json_response(
{"success": False, "error": "id is required"}, status=400 {"success": False, "error": "id or download_id is required"},
status=400,
) )
service = await DownloadQueueService.get_instance() service = await DownloadQueueService.get_instance()
deleted = await service.delete_history_item(item_id) deleted = await service.delete_history_item(
id=item_id, download_id=download_id
)
return web.json_response({"success": deleted}) return web.json_response({"success": deleted})
except Exception as exc: except Exception as exc:
self._logger.error( self._logger.error(
@@ -1800,14 +1862,20 @@ class ModelDownloadHandler:
async def retry_download_from_history(self, request: web.Request) -> web.Response: async def retry_download_from_history(self, request: web.Request) -> web.Response:
try: try:
item_id = int(request.query.get("id", "0")) download_id = request.query.get("download_id")
if not item_id: id_str = request.query.get("id")
item_id = int(id_str) if id_str else None
if not download_id and not item_id:
return web.json_response( return web.json_response(
{"success": False, "error": "id is required"}, status=400 {"success": False, "error": "id or download_id is required"},
status=400,
) )
service = await DownloadQueueService.get_instance() service = await DownloadQueueService.get_instance()
item = await service.retry_from_history(item_id) item = await service.retry_from_history(
item_id=item_id, download_id=download_id
)
if item is None: if item is None:
return web.json_response( return web.json_response(
{"success": False, "error": "History item not found or not retryable"}, {"success": False, "error": "History item not found or not retryable"},
@@ -2931,6 +2999,7 @@ class ModelHandlerSet:
"bulk_delete_models": self.management.bulk_delete_models, "bulk_delete_models": self.management.bulk_delete_models,
"verify_duplicates": self.management.verify_duplicates, "verify_duplicates": self.management.verify_duplicates,
"get_top_tags": self.query.get_top_tags, "get_top_tags": self.query.get_top_tags,
"search_tags": self.query.search_tags,
"get_base_models": self.query.get_base_models, "get_base_models": self.query.get_base_models,
"get_model_types": self.query.get_model_types, "get_model_types": self.query.get_model_types,
"scan_models": self.query.scan_models, "scan_models": self.query.scan_models,
+55 -6
View File
@@ -72,6 +72,7 @@ class RecipeHandlerSet:
"save_recipe": self.management.save_recipe, "save_recipe": self.management.save_recipe,
"delete_recipe": self.management.delete_recipe, "delete_recipe": self.management.delete_recipe,
"get_top_tags": self.query.get_top_tags, "get_top_tags": self.query.get_top_tags,
"search_tags": self.query.search_tags,
"get_base_models": self.query.get_base_models, "get_base_models": self.query.get_base_models,
"get_roots": self.query.get_roots, "get_roots": self.query.get_roots,
"get_folders": self.query.get_folders, "get_folders": self.query.get_folders,
@@ -317,12 +318,11 @@ class RecipeQueryHandler:
raise RuntimeError("Recipe scanner unavailable") raise RuntimeError("Recipe scanner unavailable")
limit = int(request.query.get("limit", "20")) limit = int(request.query.get("limit", "20"))
cache = await recipe_scanner.get_cached_data() if limit < 0:
limit = 20
tag_counts: Dict[str, int] = {} elif limit > 200:
for recipe in getattr(cache, "raw_data", []): limit = 20
for tag in recipe.get("tags", []) or []: tag_counts = await self._get_recipe_tag_counts(recipe_scanner)
tag_counts[tag] = tag_counts.get(tag, 0) + 1
sorted_tags = [ sorted_tags = [
{"tag": tag, "count": count} for tag, count in tag_counts.items() {"tag": tag, "count": count} for tag, count in tag_counts.items()
@@ -333,6 +333,55 @@ class RecipeQueryHandler:
self._logger.error("Error retrieving top tags: %s", exc, exc_info=True) self._logger.error("Error retrieving top tags: %s", 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 search_tags(self, request: web.Request) -> web.Response:
try:
await self._ensure_dependencies_ready()
recipe_scanner = self._recipe_scanner_getter()
if recipe_scanner is None:
raise RuntimeError("Recipe scanner unavailable")
query = request.query.get("q", "")
limit = int(request.query.get("limit", "20"))
if limit < 0:
limit = 20
elif limit > 200:
limit = 20
tag_counts = await self._get_recipe_tag_counts(recipe_scanner)
normalized_query = (query or "").strip().lower()
if not normalized_query:
sorted_tags = [
{"tag": tag, "count": count} for tag, count in tag_counts.items()
]
sorted_tags.sort(key=lambda entry: entry["count"], reverse=True)
return web.json_response(
{"success": True, "tags": sorted_tags[: (limit if limit > 0 else 20)]}
)
matched = [
{"tag": tag, "count": count}
for tag, count in tag_counts.items()
if normalized_query in tag.lower()
]
matched.sort(key=lambda entry: entry["count"], reverse=True)
if limit == 0:
result = matched
else:
result = matched[:limit]
return web.json_response({"success": True, "tags": result})
except Exception as exc:
self._logger.error("Error searching recipe tags: %s", exc, exc_info=True)
return web.json_response({"success": False, "error": str(exc)}, status=500)
async def _get_recipe_tag_counts(self, recipe_scanner) -> Dict[str, int]:
"""Compute tag->count mapping from cached recipe data."""
cache = await recipe_scanner.get_cached_data()
tag_counts: Dict[str, int] = {}
for recipe in getattr(cache, "raw_data", []):
for tag in recipe.get("tags", []) or []:
tag_counts[tag] = tag_counts.get(tag, 0) + 1
return tag_counts
async def get_base_models(self, request: web.Request) -> web.Response: async def get_base_models(self, request: web.Request) -> web.Response:
try: try:
await self._ensure_dependencies_ready() await self._ensure_dependencies_ready()
+2
View File
@@ -39,10 +39,12 @@ MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
RouteDefinition("POST", "/api/lm/update-usage-stats", "update_usage_stats"), RouteDefinition("POST", "/api/lm/update-usage-stats", "update_usage_stats"),
RouteDefinition("GET", "/api/lm/get-usage-stats", "get_usage_stats"), RouteDefinition("GET", "/api/lm/get-usage-stats", "get_usage_stats"),
RouteDefinition("POST", "/api/lm/update-lora-code", "update_lora_code"), RouteDefinition("POST", "/api/lm/update-lora-code", "update_lora_code"),
RouteDefinition("GET", "/api/lm/update-lora-code", "get_update_lora_code"),
RouteDefinition("GET", "/api/lm/trained-words", "get_trained_words"), RouteDefinition("GET", "/api/lm/trained-words", "get_trained_words"),
RouteDefinition("GET", "/api/lm/model-example-files", "get_model_example_files"), RouteDefinition("GET", "/api/lm/model-example-files", "get_model_example_files"),
RouteDefinition("POST", "/api/lm/register-nodes", "register_nodes"), RouteDefinition("POST", "/api/lm/register-nodes", "register_nodes"),
RouteDefinition("POST", "/api/lm/update-node-widget", "update_node_widget"), RouteDefinition("POST", "/api/lm/update-node-widget", "update_node_widget"),
RouteDefinition("GET", "/api/lm/update-node-widget", "get_update_node_widget"),
RouteDefinition("GET", "/api/lm/get-registry", "get_registry"), RouteDefinition("GET", "/api/lm/get-registry", "get_registry"),
RouteDefinition("GET", "/api/lm/check-model-exists", "check_model_exists"), RouteDefinition("GET", "/api/lm/check-model-exists", "check_model_exists"),
RouteDefinition("GET", "/api/lm/check-models-exist", "check_models_exist"), RouteDefinition("GET", "/api/lm/check-models-exist", "check_models_exist"),
+1
View File
@@ -46,6 +46,7 @@ COMMON_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
"GET", "/api/lm/{prefix}/auto-organize-progress", "get_auto_organize_progress" "GET", "/api/lm/{prefix}/auto-organize-progress", "get_auto_organize_progress"
), ),
RouteDefinition("GET", "/api/lm/{prefix}/top-tags", "get_top_tags"), RouteDefinition("GET", "/api/lm/{prefix}/top-tags", "get_top_tags"),
RouteDefinition("GET", "/api/lm/{prefix}/search-tags", "search_tags"),
RouteDefinition("GET", "/api/lm/{prefix}/base-models", "get_base_models"), RouteDefinition("GET", "/api/lm/{prefix}/base-models", "get_base_models"),
RouteDefinition("GET", "/api/lm/{prefix}/model-types", "get_model_types"), RouteDefinition("GET", "/api/lm/{prefix}/model-types", "get_model_types"),
RouteDefinition("GET", "/api/lm/{prefix}/scan", "scan_models"), RouteDefinition("GET", "/api/lm/{prefix}/scan", "scan_models"),
+1
View File
@@ -29,6 +29,7 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
RouteDefinition("POST", "/api/lm/recipes/save", "save_recipe"), RouteDefinition("POST", "/api/lm/recipes/save", "save_recipe"),
RouteDefinition("DELETE", "/api/lm/recipe/{recipe_id}", "delete_recipe"), RouteDefinition("DELETE", "/api/lm/recipe/{recipe_id}", "delete_recipe"),
RouteDefinition("GET", "/api/lm/recipes/top-tags", "get_top_tags"), RouteDefinition("GET", "/api/lm/recipes/top-tags", "get_top_tags"),
RouteDefinition("GET", "/api/lm/recipes/search-tags", "search_tags"),
RouteDefinition("GET", "/api/lm/recipes/base-models", "get_base_models"), RouteDefinition("GET", "/api/lm/recipes/base-models", "get_base_models"),
RouteDefinition("GET", "/api/lm/recipes/roots", "get_roots"), RouteDefinition("GET", "/api/lm/recipes/roots", "get_roots"),
RouteDefinition("GET", "/api/lm/recipes/folders", "get_folders"), RouteDefinition("GET", "/api/lm/recipes/folders", "get_folders"),
+219 -32
View File
@@ -47,6 +47,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 +66,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 +89,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 +100,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 +138,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:
@@ -190,6 +207,162 @@ class UpdateRoutes:
'error': str(e) 'error': str(e)
}) })
@staticmethod
async def switch_channel(request):
"""
Switch between release and nightly update channels.
Release Nightly: Initialize a Git repository (from ZIP/CM stable mode)
Nightly Release: Remove .git, download latest release ZIP, write .tracking
"""
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.info("Backed up settings.json before channel switch")
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:
git_backup = None
if os.path.exists(git_folder):
git_backup = UpdateRoutes._backup_git(git_folder, 'release')
success = False
new_version = ''
try:
if os.path.exists(git_folder):
shutil.rmtree(git_folder)
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:
UpdateRoutes._restore_git(git_backup, git_folder, success, 'release')
if settings_backup and success:
with open(settings_path, 'w', encoding='utf-8') as f:
f.write(settings_backup)
logger.info("Restored settings.json after channel switch")
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]:
""" """
@@ -295,7 +468,8 @@ class UpdateRoutes:
except Exception as e: except Exception as e:
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 +482,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:
+81 -13
View File
@@ -1,7 +1,7 @@
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 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 +109,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 +132,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 +180,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
@@ -697,6 +706,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
@@ -804,6 +840,12 @@ class BaseModelService(ABC):
"""Get top tags sorted by frequency""" """Get top tags sorted by frequency"""
return await self.scanner.get_top_tags(limit) return await self.scanner.get_top_tags(limit)
async def search_tags(
self, query: str, limit: int = 50
) -> List[Dict]:
"""Search tags by substring, sorted by frequency"""
return await self.scanner.search_tags(query, limit)
async def get_base_models(self, limit: int = 20) -> List[Dict]: async def get_base_models(self, limit: int = 20) -> List[Dict]:
"""Get base models sorted by frequency""" """Get base models sorted by frequency"""
return await self.scanner.get_base_models(limit) return await self.scanner.get_base_models(limit)
@@ -955,13 +997,21 @@ class BaseModelService(ABC):
return unified_tree return unified_tree
async def get_model_notes(self, model_name: str) -> Optional[str]: async def get_model_notes(self, model_name: str) -> Optional[dict]:
"""Get notes for a specific model file""" """Get notes and file_path for a specific model file.
Supports both simple names (``OWSMianne_ANIMA_V1``) and full-path
syntax (``Anima/character/OWSMianne_ANIMA_V1``).
"""
cache = await self.scanner.get_cached_data() cache = await self.scanner.get_cached_data()
for model in cache.raw_data: for model in cache.raw_data:
if model["file_name"] == model_name: file_name = model.get("file_name", "")
return model.get("notes", "") if file_name == model_name or model_name.endswith("/" + file_name) or model_name.endswith("\\" + file_name):
return {
"notes": model.get("notes", ""),
"file_path": model.get("file_path", ""),
}
return None return None
@@ -1084,6 +1134,11 @@ class BaseModelService(ABC):
Listing/search endpoints return lightweight cache entries; this method performs Listing/search endpoints return lightweight cache entries; this method performs
a lazy read of the on-disk metadata snapshot when callers need full detail. a lazy read of the on-disk metadata snapshot when callers need full detail.
As a beneficial side effect, the in-memory and persistent caches are
opportunistically synchronised with the on-disk metadata this keeps the
caches fresh even when a ``.metadata.json`` file was edited outside of the
normal save path (e.g. manually or by an external script).
""" """
metadata, should_skip = await MetadataManager.load_metadata( metadata, should_skip = await MetadataManager.load_metadata(
file_path, self.metadata_class file_path, self.metadata_class
@@ -1101,6 +1156,19 @@ class BaseModelService(ABC):
MetadataManager.save_metadata(file_path, metadata) MetadataManager.save_metadata(file_path, metadata)
) )
# Opportunistically sync the in-memory + persistent caches.
# The .metadata.json disk read is already paid for; the sync only
# performs work when the cache is actually stale, and uses targeted,
# in-place operations to minimise overhead even with large model sets.
#
# Fire-and-forget by design: the task is intentionally untracked.
# sync_cache_from_metadata handles its own errors internally.
asyncio.create_task(
self.scanner.sync_cache_from_metadata(
file_path, metadata.to_dict()
)
)
return self.filter_civitai_data(metadata.to_dict().get("civitai", {})) return self.filter_civitai_data(metadata.to_dict().get("civitai", {}))
async def get_model_description(self, file_path: str) -> Optional[str]: async def get_model_description(self, file_path: str) -> Optional[str]:
+25
View File
@@ -114,6 +114,13 @@ class CheckpointScanner(ModelScanner):
and metadata.hash_status == "completed" and metadata.hash_status == "completed"
and metadata.sha256 and metadata.sha256
): ):
# Ensure the in-memory hash index is populated even when
# the hash was already computed and persisted to the metadata
# file. Without this, usage tracking (and any other caller
# that queries get_hash_by_filename first) will miss on every
# lookup and keep calling back into this method, creating a
# tight loop that never populates the index.
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
return metadata.sha256 return metadata.sha256
async with self._hash_calculation_lock: async with self._hash_calculation_lock:
@@ -125,6 +132,7 @@ class CheckpointScanner(ModelScanner):
and metadata.hash_status == "completed" and metadata.hash_status == "completed"
and metadata.sha256 and metadata.sha256
): ):
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
return metadata.sha256 return metadata.sha256
task = self._hash_calculation_tasks.get(real_path) task = self._hash_calculation_tasks.get(real_path)
@@ -175,6 +183,9 @@ class CheckpointScanner(ModelScanner):
# Check if hash is already calculated # Check if hash is already calculated
if metadata.hash_status == "completed" and metadata.sha256: if metadata.hash_status == "completed" and metadata.sha256:
# Populate the in-memory hash index even for pre-computed
# hashes, mirroring the fix in calculate_hash_for_model.
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
return metadata.sha256 return metadata.sha256
# Update status to calculating # Update status to calculating
@@ -193,6 +204,20 @@ class CheckpointScanner(ModelScanner):
# Update hash index # Update hash index
self._hash_index.add_entry(sha256.lower(), file_path) self._hash_index.add_entry(sha256.lower(), file_path)
# Update the in-memory cache entry so that subsequent
# _persist_current_cache / _save_persistent_cache calls
# write the hash back to the SQLite models table. Without
# this the hash only lives in the metadata file and the
# in-memory hash index, both of which are lost across
# restarts, causing the same re-computation loop on the
# next session.
if self._cache is not None and self._cache.raw_data:
for entry in self._cache.raw_data:
if entry.get("file_path") == file_path:
entry["sha256"] = sha256.lower()
entry["hash_status"] = "completed"
break
logger.info(f"Hash calculated for checkpoint: {file_path}") logger.info(f"Hash calculated for checkpoint: {file_path}")
return sha256 return sha256
+27 -29
View File
@@ -682,7 +682,10 @@ class DownloadManager:
u for u in download_urls if not u.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES) u for u in download_urls if not u.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
] ]
download_urls = non_civitai_urls + civitai_urls download_urls = non_civitai_urls + civitai_urls
else:
# Fallback: when mirrors is empty or all mirrors have been deleted,
# use the file's downloadUrl directly (e.g. CivitAI download endpoint).
if not download_urls:
download_url = file_info.get("downloadUrl") download_url = file_info.get("downloadUrl")
if download_url: if download_url:
download_urls.append(normalize_civitai_download_url(download_url)) download_urls.append(normalize_civitai_download_url(download_url))
@@ -1386,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)
@@ -1520,35 +1533,8 @@ class DownloadManager:
if not file_info: if not file_info:
return {"success": False, "error": "No suitable file found in metadata"} return {"success": False, "error": "No suitable file found in metadata"}
mirrors = file_info.get("mirrors") or []
download_urls = []
if mirrors:
for mirror in mirrors:
if mirror.get("deletedAt") is None and mirror.get("url"):
download_urls.append(
normalize_civitai_download_url(mirror["url"])
)
# When source is 'civarchive', prioritize non-Civitai URLs download_urls = self._build_download_urls_from_file_info(file_info, source=source)
# This avoids failed downloads from deleted Civitai models
if source == "civarchive" and len(download_urls) > 1:
civitai_urls = [
u
for u in download_urls
if u.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
]
non_civitai_urls = [
u
for u in download_urls
if not u.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
]
download_urls = non_civitai_urls + civitai_urls
else:
download_url = file_info.get("downloadUrl")
if download_url:
download_urls.append(
normalize_civitai_download_url(download_url)
)
if not download_urls: if not download_urls:
return {"success": False, "error": "No mirror URL found"} return {"success": False, "error": "No mirror URL found"}
@@ -1851,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)
@@ -1866,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(
+88 -21
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,
@@ -76,6 +76,11 @@ class DownloadQueueService:
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
ON download_history(download_id) WHERE download_id IS NOT NULL;
"""
@classmethod @classmethod
async def get_instance(cls) -> DownloadQueueService: async def get_instance(cls) -> DownloadQueueService:
"""Return the singleton instance, creating it if necessary.""" """Return the singleton instance, creating it if necessary."""
@@ -113,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
@@ -154,13 +188,23 @@ class DownloadQueueService:
"""Insert a new download into the queue. """Insert a new download into the queue.
Returns the inserted row as a dict (or an empty dict if the Returns the inserted row as a dict (or an empty dict if the
download_id already exists). download_id already exists in the queue or has a terminal
record in history).
""" """
now = time.time() now = time.time()
file_params_json = json.dumps(file_params) if file_params is not None else None file_params_json = json.dumps(file_params) if file_params is not None else None
async with self._lock: async with self._lock:
conn = self._get_conn() conn = self._get_conn()
# Reject download_ids that already have a terminal record in history.
history_row = conn.execute(
"SELECT 1 FROM download_history WHERE download_id = ? LIMIT 1",
(download_id,),
).fetchone()
if history_row is not None:
return {}
conn.execute( conn.execute(
""" """
INSERT OR IGNORE INTO download_queue ( INSERT OR IGNORE INTO download_queue (
@@ -380,7 +424,7 @@ class DownloadQueueService:
) )
conn.execute( conn.execute(
""" """
INSERT INTO download_history ( INSERT OR IGNORE INTO download_history (
download_id, model_id, model_version_id, model_name, download_id, model_id, model_version_id, model_name,
version_name, thumbnail_url, status, error, file_path, version_name, thumbnail_url, status, error, file_path,
bytes_downloaded, total_bytes, completed_at bytes_downloaded, total_bytes, completed_at
@@ -537,17 +581,27 @@ class DownloadQueueService:
"offset": offset, "offset": offset,
} }
async def delete_history_item(self, id: int) -> bool: async def delete_history_item(
"""Delete a single history entry by its *id*. self, id: Optional[int] = None, download_id: Optional[str] = None
) -> bool:
"""Delete a single history entry by *download_id* (preferred) or *id*.
Returns ``True`` if a row was deleted. Returns ``True`` if a row was deleted.
""" """
async with self._lock: async with self._lock:
conn = self._get_conn() conn = self._get_conn()
cursor = conn.execute( if download_id:
"DELETE FROM download_history WHERE id = ?", cursor = conn.execute(
(id,), "DELETE FROM download_history WHERE download_id = ?",
) (download_id,),
)
elif id is not None:
cursor = conn.execute(
"DELETE FROM download_history WHERE id = ?",
(id,),
)
else:
return False
conn.commit() conn.commit()
return cursor.rowcount > 0 return cursor.rowcount > 0
@@ -604,21 +658,34 @@ class DownloadQueueService:
# Retry # Retry
# ------------------------------------------------------------------ # ------------------------------------------------------------------
async def retry_from_history(self, item_id: int) -> Optional[dict[str, Any]]: async def retry_from_history(
self,
item_id: Optional[int] = None,
download_id: Optional[str] = None,
) -> Optional[dict[str, Any]]:
"""Re-queue a failed or canceled download from history. """Re-queue a failed or canceled download from history.
Looks up the history record by its primary key. If the status is Looks up the history record by *download_id* (preferred) or
``failed`` or ``canceled`` a new queue entry is created with the *item_id*. If the status is ``failed`` or ``canceled`` a new
same model metadata and a fresh download id, and the original queue entry is created with the same model metadata and a fresh
history entry is **deleted** to prevent exponential growth when download id, and the original history entry is **deleted** to
the retried item is later canceled or fails again and re-retried. prevent exponential growth when the retried item is later
canceled or fails again and re-retried.
""" """
async with self._lock: async with self._lock:
conn = self._get_conn() conn = self._get_conn()
row = conn.execute( if download_id:
"SELECT * FROM download_history WHERE id = ?", row = conn.execute(
(item_id,), "SELECT * FROM download_history WHERE download_id = ?",
).fetchone() (download_id,),
).fetchone()
elif item_id is not None:
row = conn.execute(
"SELECT * FROM download_history WHERE id = ?",
(item_id,),
).fetchone()
else:
return None
if row is None: if row is None:
return None return None
status = str(row["status"]) status = str(row["status"])
@@ -650,7 +717,7 @@ class DownloadQueueService:
) )
conn.execute( conn.execute(
"DELETE FROM download_history WHERE id = ?", "DELETE FROM download_history WHERE id = ?",
(item_id,), (row["id"],),
) )
conn.commit() conn.commit()
queued = conn.execute( queued = conn.execute(
+19 -10
View File
@@ -270,14 +270,14 @@ class Downloader:
Note: This is private and caller MUST hold self._session_lock. Note: This is private and caller MUST hold self._session_lock.
""" """
# Close existing session if any # Snapshot and clear old session reference before creating the new
if self._session is not None: # one. This ensures self._session is always valid (or None, which
try: # triggers a fresh creation) and avoids a race where concurrent
await self._session.close() # requests hold a reference to a session whose connector has been
except Exception as e: # pragma: no cover # torn down by a premature close() call — the root cause of the
logger.warning(f"Error closing previous session: {e}") # intermittent "NoneType has no attribute connect" crash.
finally: old_session = self._session
self._session = None self._session = None
# Check for app-level proxy settings # Check for app-level proxy settings
proxy_url = None # http(s) proxy, passed via the per-request `proxy=` kwarg proxy_url = None # http(s) proxy, passed via the per-request `proxy=` kwarg
@@ -372,6 +372,13 @@ class Downloader:
self._proxy_url = proxy_url self._proxy_url = proxy_url
self._session_created_at = datetime.now() self._session_created_at = datetime.now()
# Close the previous session now that the replacement is live.
if old_session is not None:
try:
await old_session.close()
except Exception as e: # pragma: no cover
logger.warning(f"Error closing previous session: {e}")
logger.debug( logger.debug(
"Created new HTTP session with proxy settings. App-level proxy: %s, System-level proxy (trust_env): %s", "Created new HTTP session with proxy settings. App-level proxy: %s, System-level proxy (trust_env): %s",
bool(proxy_url), bool(proxy_url),
@@ -753,7 +760,8 @@ class Downloader:
else: else:
resume_offset = 0 resume_offset = 0
total_size = 0 total_size = 0
await self._create_session() async with self._session_lock:
await self._create_session()
continue continue
return False, integrity_error return False, integrity_error
@@ -843,7 +851,8 @@ class Downloader:
logger.info(f"Will resume from byte {resume_offset}") logger.info(f"Will resume from byte {resume_offset}")
# Refresh session to get new connection # Refresh session to get new connection
await self._create_session() async with self._session_lock:
await self._create_session()
continue continue
else: else:
logger.error(f"Max retries exceeded for download: {e}") logger.error(f"Max retries exceeded for download: {e}")
+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]}"
) )
+7 -3
View File
@@ -271,12 +271,16 @@ class LoraService(BaseModelService):
return letters return letters
async def get_lora_trigger_words(self, lora_name: str) -> List[str]: async def get_lora_trigger_words(self, lora_name: str) -> List[str]:
"""Get trigger words for a specific LoRA file""" """Get trigger words for a specific LoRA file.
Supports both simple names and full-path syntax.
"""
cache = await self.scanner.get_cached_data() cache = await self.scanner.get_cached_data()
for lora in cache.raw_data: for lora in cache.raw_data:
if lora["file_name"] == lora_name: file_name = lora.get("file_name", "")
civitai_data = lora.get("civitai", {}) if file_name == lora_name or lora_name.endswith("/" + file_name) or lora_name.endswith("\\" + file_name):
civitai_data = lora.get("civitai") or {}
return civitai_data.get("trainedWords", []) return civitai_data.get("trainedWords", [])
return [] return []
+69 -16
View File
@@ -15,6 +15,17 @@ from .service_registry import ServiceRegistry
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_PROVIDER_DISPLAY_NAMES = {
"civitai_api": "CivitAI",
"civarchive_api": "CivArchive",
"sqlite": "Archive DB",
}
_PRESET_PROVIDER_ORDERS = {
"civitai_archive_sqlite": ["civitai_api", "civarchive_api", "sqlite"],
"civitai_sqlite_archive": ["civitai_api", "sqlite", "civarchive_api"],
}
async def initialize_metadata_providers(): async def initialize_metadata_providers():
"""Initialize and configure all metadata providers based on settings""" """Initialize and configure all metadata providers based on settings"""
provider_manager = await ModelMetadataProviderManager.get_instance() provider_manager = await ModelMetadataProviderManager.get_instance()
@@ -26,7 +37,9 @@ async def initialize_metadata_providers():
# Get settings # Get settings
settings_manager = get_settings_manager() settings_manager = get_settings_manager()
enable_archive_db = settings_manager.get('enable_metadata_archive_db', False) enable_archive_db = settings_manager.get('enable_metadata_archive_db', False)
enable_civarchive_api = settings_manager.get('enable_civarchive_api', True)
provider_order = settings_manager.get('metadata_provider_order', 'civitai_archive_sqlite')
providers = [] providers = []
# Initialize archive database provider if enabled # Initialize archive database provider if enabled
@@ -59,27 +72,48 @@ async def initialize_metadata_providers():
except Exception as e: except Exception as e:
logger.error(f"Failed to initialize Civitai API metadata provider: {e}") logger.error(f"Failed to initialize Civitai API metadata provider: {e}")
# Register CivArchive provider, and all add to fallback providers # Register CivArchive provider when enabled. Civitai API is always
try: # preferred (better metadata); CivArchive mainly recovers metadata for
civarchive_client = await ServiceRegistry.get_civarchive_client() # models deleted from Civitai, so it can be turned off to avoid its long
civarchive_provider = CivArchiveModelMetadataProvider(civarchive_client) # rate-limit windows entirely.
provider_manager.register_provider('civarchive_api', civarchive_provider) if enable_civarchive_api:
providers.append(('civarchive_api', civarchive_provider)) try:
logger.debug("CivArchive metadata provider registered (also included in fallback)") civarchive_client = await ServiceRegistry.get_civarchive_client()
except Exception as e: civarchive_provider = CivArchiveModelMetadataProvider(civarchive_client)
logger.error(f"Failed to initialize CivArchive metadata provider: {e}") provider_manager.register_provider('civarchive_api', civarchive_provider)
providers.append(('civarchive_api', civarchive_provider))
logger.debug("CivArchive metadata provider registered (also included in fallback)")
except Exception as e:
logger.error(f"Failed to initialize CivArchive metadata provider: {e}")
else:
logger.debug("CivArchive metadata provider disabled by setting 'enable_civarchive_api'")
# Preset fallback orderings (see module-level _PRESET_PROVIDER_ORDERS).
# civitai_api is always first (better metadata); the remaining providers
# are arranged by the configured preset. Providers that are not
# registered (disabled/unavailable) are simply skipped, so each preset
# degrades gracefully.
desired_order = _PRESET_PROVIDER_ORDERS.get(
provider_order, _PRESET_PROVIDER_ORDERS["civitai_archive_sqlite"]
)
# Set up fallback provider based on available providers # Set up fallback provider based on available providers
if len(providers) > 1: if len(providers) > 1:
# Always use Civitai API (it has better metadata), then CivArchive API, then Archive DB
ordered_providers: list[tuple[str, ModelMetadataProvider]] = [] ordered_providers: list[tuple[str, ModelMetadataProvider]] = []
ordered_providers.extend([p for p in providers if p[0] == 'civitai_api']) for name in desired_order:
ordered_providers.extend([p for p in providers if p[0] == 'civarchive_api']) ordered_providers.extend([p for p in providers if p[0] == name])
ordered_providers.extend([p for p in providers if p[0] == 'sqlite']) # Include any provider not covered by the preset (defensive) at the end
for p in providers:
if p not in ordered_providers:
ordered_providers.append(p)
if ordered_providers: if ordered_providers:
fallback_provider = FallbackMetadataProvider(ordered_providers) fallback_provider = FallbackMetadataProvider(ordered_providers)
provider_manager.register_provider('fallback', fallback_provider, is_default=True) provider_manager.register_provider('fallback', fallback_provider, is_default=True)
logger.debug(
"Metadata fallback provider order: %s",
", ".join(name for name, _ in ordered_providers),
)
elif len(providers) == 1: elif len(providers) == 1:
# Only one provider available, set it as default # Only one provider available, set it as default
provider_name, provider = providers[0] provider_name, provider = providers[0]
@@ -96,11 +130,30 @@ async def update_metadata_providers():
# Get current settings # Get current settings
settings_manager = get_settings_manager() settings_manager = get_settings_manager()
enable_archive_db = settings_manager.get('enable_metadata_archive_db', False) enable_archive_db = settings_manager.get('enable_metadata_archive_db', False)
enable_civarchive_api = settings_manager.get('enable_civarchive_api', True)
provider_order = settings_manager.get('metadata_provider_order', 'civitai_archive_sqlite')
# Reinitialize all providers with new settings # Reinitialize all providers with new settings
provider_manager = await initialize_metadata_providers() provider_manager = await initialize_metadata_providers()
logger.info(f"Updated metadata providers, archive_db enabled: {enable_archive_db}") # Build effective provider chain for logging (use actually-registered
# providers, not just settings, so a failed init is reflected correctly)
registered = set(provider_manager.providers.keys())
desired = _PRESET_PROVIDER_ORDERS.get(
provider_order, _PRESET_PROVIDER_ORDERS["civitai_archive_sqlite"]
)
chain = "".join(
_PROVIDER_DISPLAY_NAMES[p]
for p in desired
if p in registered and p in _PROVIDER_DISPLAY_NAMES
)
logger.info(
"Updated metadata providers: archive_db=%s, civarchive_api=%s, chain=%s",
enable_archive_db,
enable_civarchive_api,
chain,
)
return provider_manager return provider_manager
except Exception as e: except Exception as e:
logger.error(f"Failed to update metadata providers: {e}") logger.error(f"Failed to update metadata providers: {e}")
+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")
+284 -11
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
@@ -227,6 +227,11 @@ class ModelScanner:
entry: Dict[str, Any] = { entry: Dict[str, Any] = {
'file_path': normalized_path, 'file_path': normalized_path,
# file_name is always stored WITHOUT extension (e.g. "OWSMianne_ANIMA_V1",
# not "OWSMianne_ANIMA_V1.safetensors"). All upstream population points
# (MetadataManager, from_civitai_info, download manager, etc.) strip the
# extension via os.path.splitext before writing. Code consuming this field
# should match against names that are likewise extension-free.
'file_name': get_value('file_name', '') or '', 'file_name': get_value('file_name', '') or '',
'model_name': get_value('model_name', '') or '', 'model_name': get_value('model_name', '') or '',
'folder': normalized_folder, 'folder': normalized_folder,
@@ -922,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
@@ -1347,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', '')
self._cache.raw_data.append(metadata_dict) if file_path:
self._cache.add_to_version_index(metadata_dict) old_entries = [item for item in self._cache.raw_data if item.get('file_path') == file_path]
for old_entry in old_entries:
for tag in old_entry.get('tags', []):
if tag in self._tags_count:
self._tags_count[tag] = max(0, self._tags_count[tag] - 1)
if self._tags_count[tag] == 0:
del self._tags_count[tag]
self._hash_index.remove_by_path(file_path)
self._cache.raw_data = [item for item in self._cache.raw_data if item.get('file_path') != file_path]
for tag in metadata_dict.get('tags', []):
self._tags_count[tag] = self._tags_count.get(tag, 0) + 1
self._cache.raw_data.append(metadata_dict)
# Resort cache data
await self._cache.resort() 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()
@@ -1389,6 +1420,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)
@@ -1561,6 +1595,218 @@ class ModelScanner:
return cache_entry if metadata else True return cache_entry if metadata else True
async def sync_cache_from_metadata(
self, file_path: str, metadata_dict: Dict[str, Any]
) -> bool:
"""Opportunistically sync in-memory and persistent caches from metadata.
Builds a prospective cache entry from *metadata_dict* (deserialized
``.metadata.json`` content) and compares it against the current cache
entry. When the two are already identical this method returns
``False`` without touching anything avoiding the overhead of
``update_single_model_cache``, which always removes and re-inserts
the entry, triggers a full resort, and persists via the heavyweight
``save_cache()``.
When differences are detected the update is applied **in-place** with
targeted operations:
* The existing ``raw_data`` entry is modified rather than removed and
re-appended (O(1) instead of O(n)).
* Tag counts and the hash index are updated incrementally.
* The version index is rebuilt only for the affected entry.
* ``resort()`` is called **only** when a sort-relevant field changed
(``model_name`` / ``file_name`` for name-sort, ``modified`` for
date-sort, ``size`` for size-sort).
* The persistent (SQLite) cache receives a targeted single-row update
via :meth:`PersistentModelCache.update_single_model` rather than a
full-table ``save_cache()``.
Returns:
``True`` if any cache update was performed, ``False`` if the
caches were already in sync.
.. note::
This is a **best-effort** operation. Failures are logged but
never propagated callers should fire-and-forget via
:func:`asyncio.create_task`.
"""
try:
return await self._sync_cache_from_metadata_impl(
file_path, metadata_dict
)
except Exception:
logger.warning(
"sync_cache_from_metadata failed for %s",
file_path,
exc_info=True,
)
return False
async def _sync_cache_from_metadata_impl(
self, file_path: str, metadata_dict: Dict[str, Any]
) -> bool:
cache = await self.get_cached_data()
# Locate the existing cache entry -----------------------------------
existing_idx: Optional[int] = None
existing_entry: Optional[Dict[str, Any]] = None
for i, item in enumerate(cache.raw_data):
if item.get("file_path") == file_path:
existing_entry = item
existing_idx = i
break
# Build the desired entry from metadata ------------------------------
folder_value = (
existing_entry.get("folder", "")
if existing_entry
else self._calculate_folder(file_path)
)
desired_entry = self._build_cache_entry(
metadata_dict,
folder=folder_value,
file_path_override=file_path,
)
# Ensure sha256 is populated (defensive — metadata should have it)
if (
not desired_entry.get("sha256")
and file_path
and os.path.exists(file_path)
):
try:
sha256 = await calculate_sha256(file_path)
if sha256:
desired_entry["sha256"] = sha256.lower()
except Exception:
pass
# Not in cache at all — delegate to the full update path ------------
if existing_entry is None:
result = await self.update_single_model_cache(
file_path, file_path, metadata_dict
)
return bool(result)
# Compare — skip everything if already in sync -----------------------
if not self._cache_entries_differ(existing_entry, desired_entry):
return False
# Re-validate: the cache may have been replaced concurrently
# (e.g. by _apply_scan_result). Use identity check, not equality,
# so we detect when the raw_data list was swapped out from under us.
if self._cache is None or not any(
item is existing_entry for item in self._cache.raw_data
):
return False
# ---- Differences detected: apply targeted, in-place updates --------
# Snapshot old values for delta computations
old_tags = list(existing_entry.get("tags") or [])
old_sha256: str = existing_entry.get("sha256", "") or ""
old_model_name: str = existing_entry.get("model_name", "") or ""
old_file_name: str = existing_entry.get("file_name", "") or ""
old_modified: float = float(existing_entry.get("modified", 0.0) or 0.0)
old_size: int = int(existing_entry.get("size", 0) or 0)
old_civitai = existing_entry.get("civitai")
# ---- In-place update of the cache entry ----
existing_entry.clear()
existing_entry.update(desired_entry)
# ---- Incremental tag count update ----
new_tags: set = set(desired_entry.get("tags") or [])
old_tag_set: set = set(old_tags)
for tag in old_tag_set - new_tags:
current = self._tags_count.get(tag, 0)
if current <= 1:
self._tags_count.pop(tag, None)
else:
self._tags_count[tag] = current - 1
for tag in new_tags - old_tag_set:
self._tags_count[tag] = self._tags_count.get(tag, 0) + 1
# ---- Incremental hash index update ----
new_sha = (desired_entry.get("sha256", "") or "").lower()
old_sha = (old_sha256 or "").lower()
if new_sha != old_sha:
if old_sha:
self._hash_index.remove_by_path(file_path)
if new_sha:
self._hash_index.add_entry(new_sha, file_path)
# ---- Incremental version index update ----
new_civitai = desired_entry.get("civitai")
if old_civitai != new_civitai:
temp_old = {
"file_path": file_path,
"file_name": old_file_name,
"civitai": old_civitai,
}
cache.remove_from_version_index(temp_old)
cache.add_to_version_index(existing_entry)
# ---- Conditional resort (only when sort-key fields changed) ----
need_resort = False
_last = cache._last_sort
sort_key: Optional[str] = _last[0] if _last != (None, None) else None
if sort_key == "name":
if (
old_model_name != desired_entry.get("model_name", "")
or old_file_name != desired_entry.get("file_name", "")
):
need_resort = True
elif sort_key == "date":
if old_modified != float(desired_entry.get("modified", 0.0) or 0.0):
need_resort = True
elif sort_key == "size":
if old_size != int(desired_entry.get("size", 0) or 0):
need_resort = True
if need_resort:
await cache.resort()
# ---- Targeted SQL update (single row, not full save_cache) ----
persistent = getattr(self, "_persistent_cache", None)
if persistent is not None:
old_item_for_sql: Dict[str, Any] = {
"file_path": file_path,
"tags": old_tags,
"sha256": old_sha256,
}
await asyncio.get_event_loop().run_in_executor(
None,
persistent.update_single_model,
self.model_type,
desired_entry,
old_item_for_sql,
)
return True
@staticmethod
def _cache_entries_differ(a: Dict[str, Any], b: Dict[str, Any]) -> bool:
"""Return ``True`` when two cache-entry dicts differ in any field.
Tag lists are compared order-insensitively; all other keys use
standard equality.
"""
a_tags = sorted(a.get("tags") or [])
b_tags = sorted(b.get("tags") or [])
if a_tags != b_tags:
return True
all_keys = set(a.keys()) | set(b.keys())
for key in all_keys:
if key == "tags":
continue
if a.get(key) != b.get(key):
return True
return False
def has_hash(self, sha256: str) -> bool: def has_hash(self, sha256: str) -> bool:
"""Check if a model with given hash exists""" """Check if a model with given hash exists"""
return self._hash_index.has_hash(sha256.lower()) return self._hash_index.has_hash(sha256.lower())
@@ -1613,7 +1859,32 @@ class ModelScanner:
if limit == 0: if limit == 0:
return sorted_tags return sorted_tags
return sorted_tags[:limit] return sorted_tags[:limit]
async def search_tags(
self, query: str, limit: int = 50
) -> List[Dict[str, any]]:
"""Search tags by case-insensitive substring match, sorted by count.
If query is empty, behaves like get_top_tags (returns top ``limit``
tags). If limit is 0, all matching tags are returned.
"""
await self.get_cached_data()
normalized_query = (query or "").strip().lower()
if not normalized_query:
return await self.get_top_tags(limit if limit > 0 else 20)
matched = [
{"tag": tag, "count": count}
for tag, count in self._tags_count.items()
if normalized_query in tag.lower()
]
matched.sort(key=lambda x: x["count"], reverse=True)
if limit == 0:
return matched
return matched[:limit]
async def get_base_models(self, limit: int = 20) -> List[Dict[str, any]]: async def get_base_models(self, limit: int = 20) -> List[Dict[str, any]]:
"""Get base models sorted by count. If limit is 0, return all.""" """Get base models sorted by count. If limit is 0, return all."""
cache = await self.get_cached_data() cache = await self.get_cached_data()
@@ -1729,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)
+89
View File
@@ -587,6 +587,95 @@ class PersistentModelCache:
placeholders = ", ".join(["?"] * len(self._MODEL_COLUMNS)) placeholders = ", ".join(["?"] * len(self._MODEL_COLUMNS))
return f"INSERT INTO models ({columns}) VALUES ({placeholders})" return f"INSERT INTO models ({columns}) VALUES ({placeholders})"
def update_single_model(
self,
model_type: str,
new_item: Dict,
old_item: Optional[Dict] = None,
) -> None:
"""Update a single model row in the persistent cache.
A lightweight alternative to :meth:`save_cache` that performs a targeted
DELETE + INSERT for the model row and computes incremental tag / hash-index
deltas from *old_item*. When *old_item* is omitted the previous tags and
hash are not cleaned up (callers should only omit it for brand-new entries).
All operations run inside a single transaction so readers see a consistent
view.
"""
if not self.is_enabled():
return
if not self._schema_initialized:
self._initialize_schema()
if not self._schema_initialized:
return
file_path: Optional[str] = new_item.get("file_path")
if not file_path:
return
try:
with self._db_lock:
conn = self._connect()
try:
conn.execute("PRAGMA foreign_keys = ON")
conn.execute("BEGIN")
# --- model row (DELETE + INSERT = upsert) ---
conn.execute(
"DELETE FROM models WHERE model_type = ? AND file_path = ?",
(model_type, file_path),
)
row = self._prepare_model_row(model_type, new_item)
conn.execute(self._insert_model_sql(), row)
# --- tags ---
new_tags: set = set(new_item.get("tags") or [])
old_tags: set = set(old_item.get("tags") or []) if old_item else set()
tags_to_delete = old_tags - new_tags
tags_to_insert = new_tags - old_tags
if tags_to_delete:
conn.executemany(
"DELETE FROM model_tags WHERE model_type = ? AND file_path = ? AND tag = ?",
[(model_type, file_path, t) for t in tags_to_delete],
)
if tags_to_insert:
conn.executemany(
"INSERT INTO model_tags (model_type, file_path, tag) VALUES (?, ?, ?)",
[(model_type, file_path, t) for t in tags_to_insert],
)
# --- hash_index ---
new_sha: Optional[str] = (new_item.get("sha256") or "").lower() or None
old_sha: Optional[str] = (
(old_item.get("sha256") or "").lower() or None
) if old_item else None
if new_sha != old_sha:
if old_sha:
conn.execute(
"DELETE FROM hash_index WHERE model_type = ? AND sha256 = ? AND file_path = ?",
(model_type, old_sha, file_path),
)
if new_sha:
conn.execute(
"INSERT OR IGNORE INTO hash_index (model_type, sha256, file_path) VALUES (?, ?, ?)",
(model_type, new_sha, file_path),
)
conn.execute("COMMIT")
except Exception:
conn.execute("ROLLBACK")
raise
finally:
conn.close()
except Exception as exc:
logger.warning(
"Failed to update single model in persistent cache (%s): %s",
file_path,
exc,
)
def _load_tags(self, conn: sqlite3.Connection, model_type: str) -> Dict[str, List[str]]: def _load_tags(self, conn: sqlite3.Connection, model_type: str) -> Dict[str, List[str]]:
tag_rows = conn.execute( tag_rows = conn.execute(
"SELECT file_path, tag FROM model_tags WHERE model_type = ?", "SELECT file_path, tag FROM model_tags WHERE model_type = ?",
+6 -2
View File
@@ -1,7 +1,6 @@
import asyncio import asyncio
from typing import Iterable, List, Dict, Optional from typing import Iterable, List, Dict, Optional
from dataclasses import dataclass, field from dataclasses import dataclass, field
from operator import itemgetter
from natsort import natsorted from natsort import natsorted
@@ -149,5 +148,10 @@ class RecipeCache:
) )
if not name_only: if not name_only:
self.sorted_by_date = sorted( self.sorted_by_date = sorted(
self.raw_data, key=itemgetter("created_date", "file_path"), reverse=True self.raw_data,
key=lambda x: (
x.get("modified", x.get("created_date", 0)),
x.get("file_path", ""),
),
reverse=True,
) )
+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:
+2 -1
View File
@@ -216,11 +216,12 @@ class RecipePersistenceService:
"preview_nsfw_level", "preview_nsfw_level",
"favorite", "favorite",
"gen_params", "gen_params",
"base_model",
) )
if not any(key in updates for key in allowed_fields): if not any(key in updates for key in allowed_fields):
raise RecipeValidationError( raise RecipeValidationError(
"At least one field to update must be provided (title or tags or source_path or preview_nsfw_level or favorite or gen_params)" "At least one field to update must be provided (title or tags or source_path or preview_nsfw_level or favorite or gen_params or base_model)"
) )
if "gen_params" in updates and not isinstance(updates["gen_params"], dict): if "gen_params" in updates and not isinstance(updates["gen_params"], dict):
+2
View File
@@ -65,6 +65,8 @@ DEFAULT_SETTINGS: Dict[str, Any] = {
"onboarding_completed": False, "onboarding_completed": False,
"dismissed_banners": [], "dismissed_banners": [],
"enable_metadata_archive_db": False, "enable_metadata_archive_db": False,
"enable_civarchive_api": True,
"metadata_provider_order": "civitai_archive_sqlite",
"proxy_enabled": False, "proxy_enabled": False,
"proxy_host": "", "proxy_host": "",
"proxy_port": "", "proxy_port": "",
@@ -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
+29
View File
@@ -113,6 +113,35 @@ 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
+6
View File
@@ -488,6 +488,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.7" version = "1.1.9"
license = {file = "LICENSE"} license = {file = "LICENSE"}
dependencies = [ dependencies = [
"aiohttp", "aiohttp",
+4
View File
@@ -1,6 +1,10 @@
import os import os
import sys import sys
import json import json
# Ensure the script's directory is on sys.path so that py.* imports resolve
# regardless of the current working directory (e.g. when launched via
# ComfyUI's python_embeded from the ComfyUI root directory).
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from py.middleware.cache_middleware import cache_control from py.middleware.cache_middleware import cache_control
from py.middleware.error_middleware import api_json_error from py.middleware.error_middleware import api_json_error
from py.utils.settings_paths import ensure_settings_file from py.utils.settings_paths import ensure_settings_file
+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);
+5
View File
@@ -274,6 +274,11 @@
font-style: italic; font-style: italic;
} }
/* Inline extra tags (selected but not in top-20/appended after API results) */
.filter-tag.extra-tag {
border-style: dashed;
}
/* Ensure solid border and full opacity when active or excluded */ /* Ensure solid border and full opacity when active or excluded */
.filter-tag.special-tag.active, .filter-tag.special-tag.active,
.filter-tag.special-tag.exclude { .filter-tag.special-tag.exclude {
+1 -4
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`,
@@ -93,6 +89,7 @@ export function getApiEndpoints(modelType) {
// Query operations // Query operations
scan: `/api/lm/${modelType}/scan`, scan: `/api/lm/${modelType}/scan`,
topTags: `/api/lm/${modelType}/top-tags`, topTags: `/api/lm/${modelType}/top-tags`,
searchTags: `/api/lm/${modelType}/search-tags`,
baseModels: `/api/lm/${modelType}/base-models`, baseModels: `/api/lm/${modelType}/base-models`,
roots: `/api/lm/${modelType}/roots`, roots: `/api/lm/${modelType}/roots`,
folders: `/api/lm/${modelType}/folders`, folders: `/api/lm/${modelType}/folders`,
@@ -152,7 +152,9 @@ export class LoraContextMenu extends BaseContextMenu {
sendLoraToWorkflow(replaceMode) { sendLoraToWorkflow(replaceMode) {
const card = this.currentCard; const card = this.currentCard;
const usageTips = JSON.parse(card.dataset.usage_tips || '{}'); const usageTips = JSON.parse(card.dataset.usage_tips || '{}');
const loraSyntax = buildLoraSyntax(card.dataset.file_name, usageTips); const folder = card.dataset.folder || '';
const loraName = folder ? `${folder}/${card.dataset.file_name}` : card.dataset.file_name;
const loraSyntax = buildLoraSyntax(loraName, usageTips);
sendLoraToWorkflow(loraSyntax, replaceMode, 'lora'); sendLoraToWorkflow(loraSyntax, replaceMode, 'lora');
} }
@@ -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,
+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;
} }
+17 -5
View File
@@ -397,6 +397,7 @@ export class BulkManager {
const updated = { const updated = {
...existing, ...existing,
fileName: card.dataset.file_name ?? existing.fileName, fileName: card.dataset.file_name ?? existing.fileName,
folder: card.dataset.folder ?? existing.folder,
usageTips: card.dataset.usage_tips ?? existing.usageTips, usageTips: card.dataset.usage_tips ?? existing.usageTips,
modelName: card.dataset.name ?? existing.modelName, modelName: card.dataset.name ?? existing.modelName,
}; };
@@ -494,7 +495,8 @@ export class BulkManager {
if (metadata) { if (metadata) {
const usageTips = JSON.parse(metadata.usageTips || '{}'); const usageTips = JSON.parse(metadata.usageTips || '{}');
loraSyntaxes.push(buildLoraSyntax(metadata.fileName, usageTips)); const loraName = metadata.folder ? `${metadata.folder}/${metadata.fileName}` : metadata.fileName;
loraSyntaxes.push(buildLoraSyntax(loraName, usageTips));
} else { } else {
missingLoras.push(filepath); missingLoras.push(filepath);
} }
@@ -537,7 +539,8 @@ export class BulkManager {
if (metadata) { if (metadata) {
const usageTips = JSON.parse(metadata.usageTips || '{}'); const usageTips = JSON.parse(metadata.usageTips || '{}');
loraSyntaxes.push(buildLoraSyntax(metadata.fileName, usageTips)); const loraName = metadata.folder ? `${metadata.folder}/${metadata.fileName}` : metadata.fileName;
loraSyntaxes.push(buildLoraSyntax(loraName, usageTips));
} else { } else {
missingLoras.push(filepath); missingLoras.push(filepath);
} }
@@ -553,7 +556,8 @@ export class BulkManager {
return; return;
} }
await sendLoraToWorkflow(loraSyntaxes.join(', '), replaceMode, 'lora'); const exitBulkMode = () => { if (state.bulkMode) this.toggleBulkMode(); };
await sendLoraToWorkflow(loraSyntaxes.join(', '), replaceMode, 'lora', exitBulkMode);
} }
async _sendAllEmbeddingsToWorkflow() { async _sendAllEmbeddingsToWorkflow() {
@@ -575,7 +579,8 @@ export class BulkManager {
} }
const joinedCode = embeddingCodes.join(', '); const joinedCode = embeddingCodes.join(', ');
await sendEmbeddingToWorkflow(joinedCode); const exitBulkMode = () => { if (state.bulkMode) this.toggleBulkMode(); };
await sendEmbeddingToWorkflow(joinedCode, exitBulkMode);
} }
showBulkDeleteModal() { showBulkDeleteModal() {
@@ -674,6 +679,7 @@ export class BulkManager {
const modelId = this.parseModelId(item?.civitai?.modelId); const modelId = this.parseModelId(item?.civitai?.modelId);
metadataCache.set(item.file_path, { metadataCache.set(item.file_path, {
fileName: item.file_name, fileName: item.file_name,
folder: item.folder || '',
usageTips: item.usage_tips || '{}', usageTips: item.usage_tips || '{}',
modelName: item.name || item.file_name, modelName: item.name || item.file_name,
...(modelId !== null ? { modelId } : {}) ...(modelId !== null ? { modelId } : {})
@@ -1659,13 +1665,19 @@ export class BulkManager {
cancelled = true; cancelled = true;
}); });
const isRecipesPage = state.currentPageType === 'recipes';
for (const filepath of state.selectedModels) { for (const filepath of state.selectedModels) {
if (cancelled) { if (cancelled) {
showToast('toast.api.operationCancelled', {}, 'info'); showToast('toast.api.operationCancelled', {}, 'info');
break; break;
} }
try { try {
await getModelApiClient().saveModelMetadata(filepath, { base_model: newBaseModel }); if (isRecipesPage) {
await updateRecipeMetadata(filepath, { base_model: newBaseModel });
} else {
await getModelApiClient().saveModelMetadata(filepath, { base_model: newBaseModel });
}
successCount++; successCount++;
} catch (error) { } catch (error) {
errorCount++; errorCount++;
+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}`;
+129 -26
View File
@@ -1,6 +1,7 @@
import { getCurrentPageState } from '../state/index.js'; import { getCurrentPageState } from '../state/index.js';
import { showToast, updatePanelPositions } from '../utils/uiHelpers.js'; import { showToast, updatePanelPositions } from '../utils/uiHelpers.js';
import { getModelApiClient } from '../api/modelApiFactory.js'; import { getModelApiClient } from '../api/modelApiFactory.js';
import { getApiEndpoints } from '../api/apiConfig.js';
import { removeStorageItem, setStorageItem, getStorageItem } from '../utils/storageHelpers.js'; import { removeStorageItem, setStorageItem, getStorageItem } from '../utils/storageHelpers.js';
import { MODEL_TYPE_DISPLAY_NAMES } from '../utils/constants.js'; import { MODEL_TYPE_DISPLAY_NAMES } from '../utils/constants.js';
import { translate } from '../utils/i18nHelpers.js'; import { translate } from '../utils/i18nHelpers.js';
@@ -24,6 +25,12 @@ export class FilterManager {
this.baseModelOptions = []; this.baseModelOptions = [];
this.tagsLoaded = false; this.tagsLoaded = false;
// Tag search state
this.modelTagsSearchInput = document.getElementById('modelTagsSearchInput');
this.tagSearchDebounceTimer = null;
this.tagSearchAbortController = null;
this.tagSearchQuery = '';
// Initialize preset manager // Initialize preset manager
this.presetManager = new FilterPresetManager({ this.presetManager = new FilterPresetManager({
page: this.currentPage, page: this.currentPage,
@@ -123,6 +130,60 @@ export class FilterManager {
this.renderBaseModelTags(); this.renderBaseModelTags();
}); });
} }
if (this.modelTagsSearchInput) {
this.modelTagsSearchInput.addEventListener('input', () => {
clearTimeout(this.tagSearchDebounceTimer);
this.tagSearchDebounceTimer = setTimeout(() => {
this.handleTagSearchInput();
}, 150);
});
}
}
handleTagSearchInput() {
const query = (this.modelTagsSearchInput?.value || '').trim();
const trimmedQuery = query.toLowerCase();
if (trimmedQuery === this.tagSearchQuery) return;
this.tagSearchQuery = trimmedQuery;
if (!trimmedQuery) {
// Empty query: reload top tags (default/common view)
this.loadTopTags();
return;
}
this.searchTags(trimmedQuery);
}
async searchTags(query) {
// Abort any in-flight search request
if (this.tagSearchAbortController) {
this.tagSearchAbortController.abort();
}
this.tagSearchAbortController = new AbortController();
const controller = this.tagSearchAbortController;
try {
const tagsEndpoint = `${getApiEndpoints(this.currentPage).searchTags}?q=${encodeURIComponent(query)}&limit=20`;
const response = await fetch(tagsEndpoint, { signal: controller.signal });
if (!response.ok) throw new Error('Failed to search tags');
const data = await response.json();
if (controller.signal.aborted) return; // stale response
if (data.success && data.tags) {
this.createTagFilterElements(data.tags);
} else {
throw new Error('Invalid response format');
}
} catch (error) {
if (error.name === 'AbortError') return; // expected, ignore
console.error('Error searching tags:', error);
const tagsContainer = document.getElementById('modelTagsFilter');
if (tagsContainer) {
tagsContainer.innerHTML = '<div class="tags-error">Failed to search tags</div>';
}
const emptyState = document.getElementById('modelTagsEmptyState');
if (emptyState) emptyState.hidden = true;
}
} }
getNormalizedSearchQuery(input) { getNormalizedSearchQuery(input) {
@@ -146,15 +207,24 @@ export class FilterManager {
} }
async loadTopTags() { async loadTopTags() {
// Abort any in-flight tag search request
if (this.tagSearchAbortController) {
this.tagSearchAbortController.abort();
this.tagSearchAbortController = null;
}
this.tagSearchQuery = '';
try { try {
// Show loading state // Show loading state
const tagsContainer = document.getElementById('modelTagsFilter'); const tagsContainer = document.getElementById('modelTagsFilter');
const emptyState = document.getElementById('modelTagsEmptyState');
if (!tagsContainer) return; if (!tagsContainer) return;
if (emptyState) emptyState.hidden = true;
tagsContainer.innerHTML = '<div class="tags-loading">Loading tags...</div>'; tagsContainer.innerHTML = '<div class="tags-loading">Loading tags...</div>';
// Determine the API endpoint based on the page type // Determine the API endpoint based on the page type
const tagsEndpoint = `/api/lm/${this.currentPage}/top-tags?limit=20`; const tagsEndpoint = `${getApiEndpoints(this.currentPage).topTags}?limit=20`;
const response = await fetch(tagsEndpoint); const response = await fetch(tagsEndpoint);
if (!response.ok) throw new Error('Failed to fetch tags'); if (!response.ok) throw new Error('Failed to fetch tags');
@@ -179,29 +249,38 @@ export class FilterManager {
createTagFilterElements(tags) { createTagFilterElements(tags) {
const tagsContainer = document.getElementById('modelTagsFilter'); const tagsContainer = document.getElementById('modelTagsFilter');
const emptyState = document.getElementById('modelTagsEmptyState');
if (!tagsContainer) return; if (!tagsContainer) return;
tagsContainer.innerHTML = ''; tagsContainer.innerHTML = '';
if (emptyState) emptyState.hidden = true;
// Collect existing tag names from the API response // Collect existing tag names from the API response
const existingTagNames = new Set(tags.map(t => t.tag)); const existingTagNames = new Set(tags.map(t => t.tag));
// Add any active filter tags that aren't in the top 20 // Collect active filter tags that aren't in the response (excluding __no_tags__)
const missingSelectedTags = [];
if (this.filters.tags) { if (this.filters.tags) {
Object.keys(this.filters.tags).forEach(tagName => { Object.keys(this.filters.tags).forEach(tagName => {
// Skip special tags like __no_tags__
if (tagName.startsWith('__')) return; if (tagName.startsWith('__')) return;
if (!existingTagNames.has(tagName)) { if (!existingTagNames.has(tagName)) {
// Add this tag to the list with count 0 (unknown) missingSelectedTags.push({ tag: tagName, count: 0 });
tags.push({ tag: tagName, count: 0 });
existingTagNames.add(tagName); existingTagNames.add(tagName);
} }
}); });
} }
// Append missing selected tags after the API results so they appear inline
for (const t of missingSelectedTags) {
tags.push(t);
}
if (!tags.length) { if (!tags.length) {
tagsContainer.innerHTML = `<div class="no-tags">No ${this.currentPage === 'recipes' ? 'recipe ' : ''}tags available</div>`; if (this.tagSearchQuery) {
if (emptyState) emptyState.hidden = false;
} else {
tagsContainer.innerHTML = `<div class="no-tags">No ${this.currentPage === 'recipes' ? 'recipe ' : ''}tags available</div>`;
}
return; return;
} }
@@ -209,6 +288,10 @@ export class FilterManager {
const tagEl = document.createElement('div'); const tagEl = document.createElement('div');
tagEl.className = 'filter-tag tag-filter'; tagEl.className = 'filter-tag tag-filter';
const tagName = tag.tag; const tagName = tag.tag;
if (missingSelectedTags.some(t => t.tag === tagName)) {
tagEl.classList.add('extra-tag');
}
tagEl.dataset.tag = tagName; tagEl.dataset.tag = tagName;
// Show count only if it's > 0 (known count) // Show count only if it's > 0 (known count)
@@ -234,26 +317,28 @@ export class FilterManager {
tagsContainer.appendChild(tagEl); tagsContainer.appendChild(tagEl);
}); });
// Add "No tags" as a special filter at the end // Add "No tags" as a special filter at the end (skip during search)
const noTagsEl = document.createElement('div'); if (!this.tagSearchQuery) {
noTagsEl.className = 'filter-tag tag-filter special-tag'; const noTagsEl = document.createElement('div');
const noTagsLabel = translate('header.filter.noTags', {}, 'No tags'); noTagsEl.className = 'filter-tag tag-filter special-tag';
const noTagsKey = '__no_tags__'; const noTagsLabel = translate('header.filter.noTags', {}, 'No tags');
noTagsEl.dataset.tag = noTagsKey; const noTagsKey = '__no_tags__';
noTagsEl.innerHTML = noTagsLabel; noTagsEl.dataset.tag = noTagsKey;
noTagsEl.innerHTML = noTagsLabel;
noTagsEl.addEventListener('click', async () => { noTagsEl.addEventListener('click', async () => {
const currentState = (this.filters.tags && this.filters.tags[noTagsKey]) || 'none'; const currentState = (this.filters.tags && this.filters.tags[noTagsKey]) || 'none';
const newState = this.getNextTriStateState(currentState); const newState = this.getNextTriStateState(currentState);
this.setTagFilterState(noTagsKey, newState); this.setTagFilterState(noTagsKey, newState);
this.applyTagElementState(noTagsEl, newState); this.applyTagElementState(noTagsEl, newState);
this.updateActiveFiltersCount(); this.updateActiveFiltersCount();
await this.applyFilters(false); await this.applyFilters(false);
}); });
tagsContainer.appendChild(noTagsEl); tagsContainer.appendChild(noTagsEl);
}
this.updateTagSelections(); this.updateTagSelections();
} }
@@ -341,7 +426,7 @@ export class FilterManager {
if (!baseModelTagsContainer) return; if (!baseModelTagsContainer) return;
// Set the API endpoint based on current page // Set the API endpoint based on current page
const apiEndpoint = `/api/lm/${this.currentPage}/base-models?limit=0`; const apiEndpoint = `${getApiEndpoints(this.currentPage).baseModels}?limit=0`;
// Fetch base models // Fetch base models
fetch(apiEndpoint) fetch(apiEndpoint)
@@ -644,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
@@ -721,6 +808,16 @@ export class FilterManager {
tagLogic: 'any' tagLogic: 'any'
}); });
// Clear tag search input and reset search state
if (this.modelTagsSearchInput) {
this.modelTagsSearchInput.value = '';
}
this.tagSearchQuery = '';
if (this.tagSearchAbortController) {
this.tagSearchAbortController.abort();
this.tagSearchAbortController = null;
}
// Update tag logic toggle UI // Update tag logic toggle UI
this.updateTagLogicToggleUI(); this.updateTagLogicToggleUI();
@@ -731,6 +828,10 @@ export class FilterManager {
// Update UI // Update UI
this.updateTagSelections(); this.updateTagSelections();
this.updateActiveFiltersCount(); this.updateActiveFiltersCount();
// Reload tag area to drop any non-top-20 tags from the deactivated preset
if (this.tagsLoaded) {
await this.loadTopTags();
}
this.presetManager.renderPresets(); // Re-render to remove active state this.presetManager.renderPresets(); // Re-render to remove active state
// Remove from local Storage // Remove from local Storage
@@ -885,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 || [])],
@@ -892,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 ?? ''
}; };
} }
+13 -19
View File
@@ -478,11 +478,9 @@ export class FilterPresetManager {
const pageState = getCurrentPageState(); const pageState = getCurrentPageState();
pageState.filters = this.filterManager.cloneFilters(); pageState.filters = this.filterManager.cloneFilters();
// If tags haven't been loaded yet, load them first // Refresh tag display so preset's non-top-20 tags appear inline
if (!this.filterManager.tagsLoaded) { await this.filterManager.loadTopTags();
await this.filterManager.loadTopTags(); this.filterManager.tagsLoaded = true;
this.filterManager.tagsLoaded = true;
}
// Check again after async operation // Check again after async operation
if (requestId !== this.applyPresetRequestId) return; if (requestId !== this.applyPresetRequestId) return;
@@ -745,8 +743,16 @@ export class FilterPresetManager {
presetEl.classList.add('active'); presetEl.classList.add('active');
} }
presetEl.addEventListener('click', (e) => { // Apply preset on click (toggle if already active)
e.stopPropagation(); // Bind to the whole .filter-preset div so clicking anywhere inside triggers apply
presetEl.addEventListener('click', async () => {
this.cancelPendingDelete();
if (this.activePreset === preset.name) {
await this.filterManager.clearFilters();
} else {
await this.applyPreset(preset.name);
}
}); });
const presetName = document.createElement('span'); const presetName = document.createElement('span');
@@ -759,18 +765,6 @@ export class FilterPresetManager {
deleteBtn.innerHTML = '<i class="fas fa-times"></i>'; deleteBtn.innerHTML = '<i class="fas fa-times"></i>';
deleteBtn.title = translate('header.filter.presetDeleteTooltip', {}, 'Delete preset'); deleteBtn.title = translate('header.filter.presetDeleteTooltip', {}, 'Delete preset');
// Apply preset on name click (toggle if already active)
presetName.addEventListener('click', async (e) => {
e.stopPropagation();
this.cancelPendingDelete();
if (this.activePreset === preset.name) {
await this.filterManager.clearFilters();
} else {
await this.applyPreset(preset.name);
}
});
// Two-step delete on delete button click // Two-step delete on delete button click
deleteBtn.addEventListener('click', (e) => { deleteBtn.addEventListener('click', (e) => {
e.stopPropagation(); e.stopPropagation();
+10
View File
@@ -2346,6 +2346,16 @@ export class SettingsManager {
enableMetadataArchiveCheckbox.checked = state.global.settings.enable_metadata_archive_db || false; enableMetadataArchiveCheckbox.checked = state.global.settings.enable_metadata_archive_db || false;
} }
const enableCivarchiveApiCheckbox = document.getElementById('enableCivarchiveApi');
if (enableCivarchiveApiCheckbox) {
enableCivarchiveApiCheckbox.checked = state.global.settings.enable_civarchive_api ?? true;
}
const metadataProviderOrderSelect = document.getElementById('metadataProviderOrder');
if (metadataProviderOrderSelect) {
metadataProviderOrderSelect.value = state.global.settings.metadata_provider_order || 'civitai_archive_sqlite';
}
// Load status // Load status
await this.updateMetadataArchiveStatus(); await this.updateMetadataArchiveStatus();
} catch (error) { } catch (error) {
+211 -41
View File
@@ -24,7 +24,9 @@ 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.progressKeepVisible = false;
this.currentVersionInfo = null; this.currentVersionInfo = null;
this.versionMismatch = false; this.versionMismatch = false;
this.activeNotificationTab = 'updates'; this.activeNotificationTab = 'updates';
@@ -49,43 +51,161 @@ 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 if (this.channelMode === null) {
this.updateBadgeVisibility(); this.channelMode = this.hasGit ? 'nightly' : 'release';
}
this.checkForUpdates().then(() => {
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;
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');
}
}
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() {
@@ -373,7 +493,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 === '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 +502,19 @@ 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 (this.channelMode === null) {
this.updateAvailable = this.isNewerVersion(this.latestVersion, this.currentVersion); this.channelMode = this.hasGit ? 'nightly' : 'release';
}
// Update last check time
this.updateAvailable = data.update_available;
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,
@@ -482,8 +605,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 +745,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 +773,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'
}) })
}); });
@@ -698,7 +848,26 @@ export class UpdateService {
progressText.textContent = text; progressText.textContent = text;
} }
} }
_showSwitchCompleteMessage(version) {
this.showUpdateProgress(true);
this.updateProgress(100, '');
const progressText = document.getElementById('updateProgressText');
if (progressText) {
progressText.innerHTML = `
<div style="text-align: center; color: var(--lora-success);">
<i class="fas fa-check-circle" style="margin-right: 8px;"></i>
${translate('update.completion.successMessage', { version })}
<br><br>
<div style="opacity: 0.95; color: var(--lora-error); font-size: 1em;">
${translate('update.completion.restartMessage')}<br>
${translate('update.completion.reloadMessage')}
</div>
</div>
`;
}
}
showUpdateCompleteMessage(newVersion) { showUpdateCompleteMessage(newVersion) {
const modal = document.getElementById('updateModal'); const modal = document.getElementById('updateModal');
if (!modal) return; if (!modal) return;
@@ -771,6 +940,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 +971,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);
+2
View File
@@ -13,6 +13,8 @@ const DEFAULT_SETTINGS_BASE = Object.freeze({
language: 'en', language: 'en',
show_only_sfw: false, show_only_sfw: false,
enable_metadata_archive_db: false, enable_metadata_archive_db: false,
enable_civarchive_api: true,
metadata_provider_order: 'civitai_archive_sqlite',
proxy_enabled: false, proxy_enabled: false,
proxy_type: 'http', proxy_type: 'http',
proxy_host: '', proxy_host: '',
+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
+40 -5
View File
@@ -141,6 +141,20 @@ const PARAM_TO_WIDGET_CANDIDATES = {
scheduler: ['scheduler'], scheduler: ['scheduler'],
}; };
// ---------------------------------------------------------------------------
// Node-type-specific widget name overrides.
// Keys are ComfyUI node class names (e.g. "GlobalSeed //Inspire").
// Values are partial PARAM_TO_WIDGET_CANDIDATES maps; the per-node candidates
// are tried *before* the global ones. Only the params listed here are
// overridden — every other param still uses the global candidates.
// ---------------------------------------------------------------------------
const NODE_TYPE_WIDGET_OVERRIDES = {
// Inspire Pack — Global Seed node stores the seed in a widget named "value"
'GlobalSeed //Inspire': {
seed: ['value'],
},
};
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
// Parse a combined sampler+scheduler value (space-separated or underscore) // Parse a combined sampler+scheduler value (space-separated or underscore)
// e.g., "Euler a Karras", "DPM++ 2M beta", "er_sde_beta" // e.g., "Euler a Karras", "DPM++ 2M beta", "er_sde_beta"
@@ -235,7 +249,7 @@ function resolveSamplerScheduler(rawValue) {
// Find which gen params can be sent to a given node, matching by widget names // Find which gen params can be sent to a given node, matching by widget names
// Returns array of { widgetName, value } objects // Returns array of { widgetName, value } objects
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
function findMatchingWidgets(nodeWidgetNames, resolvedParams) { function findMatchingWidgets(nodeWidgetNames, resolvedParams, nodeType) {
if (!nodeWidgetNames || !Array.isArray(nodeWidgetNames) || nodeWidgetNames.length === 0) { if (!nodeWidgetNames || !Array.isArray(nodeWidgetNames) || nodeWidgetNames.length === 0) {
return []; return [];
} }
@@ -243,6 +257,26 @@ function findMatchingWidgets(nodeWidgetNames, resolvedParams) {
const widgetSet = new Set(nodeWidgetNames.map(w => String(w).toLowerCase())); const widgetSet = new Set(nodeWidgetNames.map(w => String(w).toLowerCase()));
const updates = []; const updates = [];
// Resolve node-type-specific overrides (if any)
const typeOverrides =
nodeType && typeof nodeType === 'string'
? (NODE_TYPE_WIDGET_OVERRIDES[nodeType] || {})
: {};
/**
* Build the effective candidate list for a parameter:
* type-specific overrides (if any) come first, then the global candidates.
*/
function getCandidates(key) {
const global = PARAM_TO_WIDGET_CANDIDATES[key] || [key];
const extra = typeOverrides[key];
if (extra && Array.isArray(extra) && extra.length > 0) {
// Prepend type-specific candidates; keep global as fallback
return [...extra, ...global];
}
return global;
}
// Simple numeric/string params: seed, steps, cfg // Simple numeric/string params: seed, steps, cfg
const simpleParams = [ const simpleParams = [
{ key: 'seed', value: resolvedParams.seed }, { key: 'seed', value: resolvedParams.seed },
@@ -251,10 +285,10 @@ function findMatchingWidgets(nodeWidgetNames, resolvedParams) {
]; ];
for (const { key, value } of simpleParams) { for (const { key, value } of simpleParams) {
if (value === undefined || value === null || value === '') continue; if (value === undefined || value === null || value === '') continue;
const candidates = PARAM_TO_WIDGET_CANDIDATES[key] || [key]; const candidates = getCandidates(key);
for (const candidate of candidates) { for (const candidate of candidates) {
if (widgetSet.has(candidate.toLowerCase())) { if (widgetSet.has(candidate.toLowerCase())) {
updates.push({ widgetName: candidate, value: String(value) }); updates.push({ widgetName: candidate, value });
break; break;
} }
} }
@@ -262,7 +296,7 @@ function findMatchingWidgets(nodeWidgetNames, resolvedParams) {
// Sampler // Sampler
if (resolvedParams.sampler) { if (resolvedParams.sampler) {
const candidates = PARAM_TO_WIDGET_CANDIDATES.sampler; const candidates = getCandidates('sampler');
for (const candidate of candidates) { for (const candidate of candidates) {
if (widgetSet.has(candidate.toLowerCase())) { if (widgetSet.has(candidate.toLowerCase())) {
updates.push({ widgetName: candidate, value: resolvedParams.sampler }); updates.push({ widgetName: candidate, value: resolvedParams.sampler });
@@ -273,7 +307,7 @@ function findMatchingWidgets(nodeWidgetNames, resolvedParams) {
// Scheduler // Scheduler
if (resolvedParams.scheduler) { if (resolvedParams.scheduler) {
const candidates = PARAM_TO_WIDGET_CANDIDATES.scheduler; const candidates = getCandidates('scheduler');
for (const candidate of candidates) { for (const candidate of candidates) {
if (widgetSet.has(candidate.toLowerCase())) { if (widgetSet.has(candidate.toLowerCase())) {
updates.push({ widgetName: candidate, value: resolvedParams.scheduler }); updates.push({ widgetName: candidate, value: resolvedParams.scheduler });
@@ -290,6 +324,7 @@ export {
SCHEDULER_SUFFIXES, SCHEDULER_SUFFIXES,
SCHEDULER_ONLY_VALUES, SCHEDULER_ONLY_VALUES,
PARAM_TO_WIDGET_CANDIDATES, PARAM_TO_WIDGET_CANDIDATES,
NODE_TYPE_WIDGET_OVERRIDES,
parseCombinedSamplerName, parseCombinedSamplerName,
resolveSamplerScheduler, resolveSamplerScheduler,
findMatchingWidgets, findMatchingWidgets,
+22 -11
View File
@@ -134,7 +134,10 @@ export async function copyToClipboard(text, successMessage = null) {
} }
export function showToast(key, params = {}, type = 'info', fallback = null) { export function showToast(key, params = {}, type = 'info', fallback = null) {
const message = translate(key, params, fallback); // Plain messages (contain spaces) are not i18n dot-notation keys — use verbatim
// to avoid spurious "Translation key not found" warnings from i18next
const isPlainMessage = typeof key === 'string' && /\s/.test(key);
const message = isPlainMessage ? key : translate(key, params, fallback);
const toast = document.createElement('div'); const toast = document.createElement('div');
toast.className = `toast toast-${type}`; toast.className = `toast toast-${type}`;
toast.textContent = message; toast.textContent = message;
@@ -605,7 +608,7 @@ function isNodeEnabled(node) {
if (!node) { if (!node) {
return false; return false;
} }
// ComfyUI node mode: 0 = Normal/Enabled, others = Always/Never/OnEvent // ComfyUI node mode (LGraphEventMode): 0 = Always, 2 = Never, 4 = Bypass
return node.mode === undefined || node.mode === 0; return node.mode === undefined || node.mode === 0;
} }
@@ -656,7 +659,7 @@ async function ensureRelativeModelPath(modelPath, collectionType) {
* @param {string} syntaxType - The type of syntax ('lora' or 'recipe') * @param {string} syntaxType - The type of syntax ('lora' or 'recipe')
* @returns {Promise<boolean>} - Whether the operation was successful * @returns {Promise<boolean>} - Whether the operation was successful
*/ */
export async function sendLoraToWorkflow(loraSyntax, replaceMode = false, syntaxType = 'lora') { export async function sendLoraToWorkflow(loraSyntax, replaceMode = false, syntaxType = 'lora', onComplete = null) {
const registry = await fetchWorkflowRegistry(); const registry = await fetchWorkflowRegistry();
if (!registry) { if (!registry) {
return false; return false;
@@ -681,7 +684,9 @@ export async function sendLoraToWorkflow(loraSyntax, replaceMode = false, syntax
} }
if (nodeKeys.length === 1) { if (nodeKeys.length === 1) {
return await sendLoraToNodes([nodeKeys[0]], loraNodes, loraSyntax, replaceMode, syntaxType); const result = await sendLoraToNodes([nodeKeys[0]], loraNodes, loraSyntax, replaceMode, syntaxType);
if (result && typeof onComplete === 'function') onComplete();
return result;
} }
const actionType = const actionType =
@@ -695,8 +700,11 @@ export async function sendLoraToWorkflow(loraSyntax, replaceMode = false, syntax
showNodeSelector(loraNodes, { showNodeSelector(loraNodes, {
actionType, actionType,
actionMode, actionMode,
onSend: (selectedNodeIds) => onSend: async (selectedNodeIds) => {
sendLoraToNodes(selectedNodeIds, loraNodes, loraSyntax, replaceMode, syntaxType), const result = await sendLoraToNodes(selectedNodeIds, loraNodes, loraSyntax, replaceMode, syntaxType);
if (result && typeof onComplete === 'function') onComplete();
return result;
},
}); });
return true; return true;
} }
@@ -967,7 +975,7 @@ async function sendTextToNodes(nodeIds, nodesMap, text, mode, messages = {}) {
} }
} }
export async function sendEmbeddingToWorkflow(embeddingCode) { export async function sendEmbeddingToWorkflow(embeddingCode, onComplete = null) {
const registry = await fetchWorkflowRegistry(); const registry = await fetchWorkflowRegistry();
if (!registry) { if (!registry) {
return false; return false;
@@ -995,8 +1003,11 @@ export async function sendEmbeddingToWorkflow(embeddingCode) {
missingTargetMessage: translate('uiHelpers.workflow.noTargetNodeSelected', {}, 'No target node selected'), missingTargetMessage: translate('uiHelpers.workflow.noTargetNodeSelected', {}, 'No target node selected'),
}; };
const handleSend = (selectedNodeIds) => const handleSend = async (selectedNodeIds) => {
sendTextToNodes(selectedNodeIds, textNodes, embeddingCode, 'append', messages); const result = await sendTextToNodes(selectedNodeIds, textNodes, embeddingCode, 'append', messages);
if (result && typeof onComplete === 'function') onComplete();
return result;
};
if (nodeKeys.length === 1) { if (nodeKeys.length === 1) {
return await handleSend([nodeKeys[0]]); return await handleSend([nodeKeys[0]]);
@@ -1136,8 +1147,8 @@ export async function sendGenParamsToWorkflow(genParams) {
const node = targetNodes[nodeKey]; const node = targetNodes[nodeKey];
if (!node) continue; if (!node) continue;
const widgetNames = node.widget_names || []; const widgetNames = getWidgetNames(node);
const updates = findMatchingWidgets(widgetNames, raw); const updates = findMatchingWidgets(widgetNames, raw, node.type_name);
if (updates.length === 0) { if (updates.length === 0) {
showToast(`Node "${node.title || node.type}" has no matching widgets for these parameters`, {}, 'warning'); showToast(`Node "${node.title || node.type}" has no matching widgets for these parameters`, {}, 'warning');
+5
View File
@@ -251,10 +251,15 @@
<button class="tag-logic-option" data-value="all" title="{{ t('header.filter.tagLogicAll') }}">{{ t('header.filter.all') }}</button> <button class="tag-logic-option" data-value="all" title="{{ t('header.filter.tagLogicAll') }}">{{ t('header.filter.all') }}</button>
</div> </div>
</div> </div>
<input type="text" id="modelTagsSearchInput" class="filter-search-input"
placeholder="{{ t('header.filter.tagSearchPlaceholder') }}" autocomplete="off">
<div class="filter-tags" id="modelTagsFilter"> <div class="filter-tags" id="modelTagsFilter">
<!-- Top tags will be dynamically inserted here --> <!-- Top tags will be dynamically inserted here -->
<div class="tags-loading">{{ t('common.status.loading') }}</div> <div class="tags-loading">{{ t('common.status.loading') }}</div>
</div> </div>
<div id="modelTagsEmptyState" class="filter-empty-state" hidden>
{{ t('header.filter.noTagMatches') }}
</div>
</div> </div>
{% if current_page == 'loras' or current_page == 'checkpoints' %} {% if current_page == 'loras' or current_page == 'checkpoints' %}
<div class="filter-section"> <div class="filter-section">
+80 -43
View File
@@ -144,6 +144,46 @@
</div> </div>
</div> </div>
<div class="settings-subsection">
<div class="settings-subsection-header">
<h4>{{ t('settings.sections.downloads') }}</h4>
</div>
<div class="setting-item">
<div class="setting-row">
<div class="setting-info">
<label for="downloadBackend">{{ t('settings.downloadBackend.label') }}</label>
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.downloadBackend.help') }}"></i>
<a class="settings-action-link" href="https://github.com/willmiao/ComfyUI-Lora-Manager/wiki/Aria2-Download-Backend-(Experimental)" target="_blank" rel="noopener" aria-label="{{ t('settings.aria2HelpLink') }}" title="{{ t('settings.aria2HelpLink') }}">
<i class="fas fa-question-circle" aria-hidden="true"></i>
</a>
</div>
<div class="setting-control select-control">
<select id="downloadBackend" onchange="settingsManager.saveSelectSetting('downloadBackend', 'download_backend')">
<option value="python">{{ t('settings.downloadBackend.options.python') }}</option>
<option value="aria2">{{ t('settings.downloadBackend.options.aria2') }}</option>
</select>
</div>
</div>
</div>
<div class="setting-item" id="aria2PathSetting" style="display: none;">
<div class="setting-row">
<div class="setting-info">
<label for="aria2cPath">{{ t('settings.aria2cPath.label') }}</label>
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.aria2cPath.help') }}"></i>
</div>
<div class="setting-control">
<div class="text-input-wrapper">
<input type="text"
id="aria2cPath"
placeholder="{{ t('settings.aria2cPath.placeholder') }}"
onblur="settingsManager.saveInputSetting('aria2cPath', 'aria2c_path')"
onkeydown="if(event.key === 'Enter') { this.blur(); }" />
</div>
</div>
</div>
</div>
</div>
<!-- AI Provider Configuration (BYOK) --> <!-- AI Provider Configuration (BYOK) -->
<div class="settings-subsection"> <div class="settings-subsection">
<div class="settings-subsection-header"> <div class="settings-subsection-header">
@@ -250,46 +290,6 @@
{{ provider_models_json | safe }} {{ provider_models_json | safe }}
</script> </script>
<div class="settings-subsection">
<div class="settings-subsection-header">
<h4>{{ t('settings.sections.downloads') }}</h4>
</div>
<div class="setting-item">
<div class="setting-row">
<div class="setting-info">
<label for="downloadBackend">{{ t('settings.downloadBackend.label') }}</label>
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.downloadBackend.help') }}"></i>
<a class="settings-action-link" href="https://github.com/willmiao/ComfyUI-Lora-Manager/wiki/Aria2-Download-Backend-(Experimental)" target="_blank" rel="noopener" aria-label="{{ t('settings.aria2HelpLink') }}" title="{{ t('settings.aria2HelpLink') }}">
<i class="fas fa-question-circle" aria-hidden="true"></i>
</a>
</div>
<div class="setting-control select-control">
<select id="downloadBackend" onchange="settingsManager.saveSelectSetting('downloadBackend', 'download_backend')">
<option value="python">{{ t('settings.downloadBackend.options.python') }}</option>
<option value="aria2">{{ t('settings.downloadBackend.options.aria2') }}</option>
</select>
</div>
</div>
</div>
<div class="setting-item" id="aria2PathSetting" style="display: none;">
<div class="setting-row">
<div class="setting-info">
<label for="aria2cPath">{{ t('settings.aria2cPath.label') }}</label>
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.aria2cPath.help') }}"></i>
</div>
<div class="setting-control">
<div class="text-input-wrapper">
<input type="text"
id="aria2cPath"
placeholder="{{ t('settings.aria2cPath.placeholder') }}"
onblur="settingsManager.saveInputSetting('aria2cPath', 'aria2c_path')"
onkeydown="if(event.key === 'Enter') { this.blur(); }" />
</div>
</div>
</div>
</div>
</div>
<!-- Backup --> <!-- Backup -->
<div class="settings-subsection"> <div class="settings-subsection">
<div class="settings-subsection-header"> <div class="settings-subsection-header">
@@ -1401,7 +1401,26 @@
<div class="settings-input-error-message" id="metadataRefreshSkipPathsError"></div> <div class="settings-input-error-message" id="metadataRefreshSkipPathsError"></div>
</div> </div>
<!-- Metadata Archive --> <!-- CivArchive API provider toggle -->
<div class="setting-item">
<div class="setting-row">
<div class="setting-info">
<label for="enableCivarchiveApi">
{{ t('settings.metadataArchive.enableCivarchiveApi') }}
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.metadataArchive.enableCivarchiveApiHelp') }}"></i>
</label>
</div>
<div class="setting-control">
<label class="toggle-switch">
<input type="checkbox" id="enableCivarchiveApi"
onchange="settingsManager.saveToggleSetting('enableCivarchiveApi', 'enable_civarchive_api')">
<span class="toggle-slider"></span>
</label>
</div>
</div>
</div>
<!-- Metadata Archive DB -->
<div class="setting-item"> <div class="setting-item">
<div class="setting-row"> <div class="setting-row">
<div class="setting-info"> <div class="setting-info">
@@ -1419,13 +1438,13 @@
</div> </div>
</div> </div>
</div> </div>
<div class="setting-item"> <div class="setting-item">
<div class="metadata-archive-status" id="metadataArchiveStatus"> <div class="metadata-archive-status" id="metadataArchiveStatus">
<!-- Status will be populated by JavaScript --> <!-- Status will be populated by JavaScript -->
</div> </div>
</div> </div>
<div class="setting-item"> <div class="setting-item">
<div class="setting-row"> <div class="setting-row">
<div class="setting-info"> <div class="setting-info">
@@ -1444,6 +1463,24 @@
</div> </div>
</div> </div>
</div> </div>
<!-- Metadata provider fallback order -->
<div class="setting-item">
<div class="setting-row">
<div class="setting-info">
<label for="metadataProviderOrder">
{{ t('settings.metadataArchive.providerOrder') }}
<i class="fas fa-info-circle info-icon" data-tooltip="{{ t('settings.metadataArchive.providerOrderHelp') }}"></i>
</label>
</div>
<div class="setting-control select-control">
<select id="metadataProviderOrder" onchange="settingsManager.saveSelectSetting('metadataProviderOrder', 'metadata_provider_order')">
<option value="civitai_archive_sqlite">{{ t('settings.metadataArchive.providerOrderCivitaiArchiveSqlite') }}</option>
<option value="civitai_sqlite_archive">{{ t('settings.metadataArchive.providerOrderCivitaiSqliteArchive') }}</option>
</select>
</div>
</div>
</div>
</div> </div>
</div> </div>
</div> </div>
@@ -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">
+70
View File
@@ -823,3 +823,73 @@ def test_apply_library_settings_ignores_extra_lora_path_overlapping_primary_root
"same lora folder" in record.message.lower() "same lora folder" in record.message.lower()
for record in caplog.records for record in caplog.records
) )
def test_save_paths_removes_stale_empty_default_when_comfyui_exists(
monkeypatch: pytest.MonkeyPatch, tmp_path,
):
"""When an empty-shell 'default' library coexists with 'comfyui', the
stale 'default' entry should be removed and 'comfyui' activated."""
folder_paths = _setup_config_environment(monkeypatch, tmp_path)
class FakeSettingsService:
def __init__(self):
# Replicate the user's settings.json: empty default + populated comfyui
self.libraries = {
"default": {
"folder_paths": {},
"extra_folder_paths": {},
"default_lora_root": "",
"default_checkpoint_root": "",
"default_unet_root": "",
"default_embedding_root": "",
"recipes_path": "",
},
"comfyui": {
"folder_paths": {
key: list(value) for key, value in folder_paths.items()
},
"default_lora_root": folder_paths["loras"][0],
"default_checkpoint_root": folder_paths["checkpoints"][0],
"default_embedding_root": folder_paths["embeddings"][0],
},
}
# No active_library key — get_active_library_name() falls back to
# dict order, returning "default".
self.active_library = "default"
self.delete_calls: list[str] = []
self.upsert_calls: list[tuple[str, dict]] = []
def get_libraries(self):
return dict(self.libraries)
def delete_library(self, name: str):
self.delete_calls.append(name)
self.libraries.pop(name, None)
def rename_library(self, *_):
raise AssertionError("rename_library should not be invoked")
def get_active_library_name(self):
return self.active_library
def upsert_library(self, name: str, **payload):
self.upsert_calls.append((name, payload))
self.libraries[name] = {**payload}
if payload.get("activate"):
self.active_library = name
fake_settings = FakeSettingsService()
monkeypatch.setattr(settings_manager_module, "settings", fake_settings)
config_module.Config()
assert fake_settings.delete_calls == ["default"]
assert "default" not in fake_settings.libraries
assert set(fake_settings.libraries.keys()) == {"comfyui"}
assert len(fake_settings.upsert_calls) == 1
name, payload = fake_settings.upsert_calls[0]
assert name == "comfyui"
assert payload["activate"] is True
assert fake_settings.active_library == "comfyui"
+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()
@@ -113,6 +113,8 @@ function renderControlsDom(pageKey) {
<div id="baseModelEmptyState" hidden></div> <div id="baseModelEmptyState" hidden></div>
<div id="filterPresets" class="filter-presets"></div> <div id="filterPresets" class="filter-presets"></div>
<div id="modelTagsFilter" class="filter-tags"></div> <div id="modelTagsFilter" class="filter-tags"></div>
<input id="modelTagsSearchInput" />
<div id="modelTagsEmptyState" hidden></div>
<button class="clear-filter"></button> <button class="clear-filter"></button>
</div> </div>
<div class="controls"> <div class="controls">
@@ -961,4 +963,198 @@ describe('PageControls favorites, sorting, and duplicates scenarios', () => {
expect(stateModule.state.bulkMode).toBe(true); expect(stateModule.state.bulkMode).toBe(true);
expect(pageState.duplicatesMode).toBe(true); expect(pageState.duplicatesMode).toBe(true);
}); });
describe('tag search', () => {
it('fetches /search-tags when typing in the tag search input (debounced)', async () => {
vi.useFakeTimers();
const searchTagsUrls = [];
global.fetch = vi.fn((url) => {
if (url.includes('/search-tags')) {
searchTagsUrls.push(url);
return Promise.resolve({
ok: true,
json: async () => ({ success: true, tags: [{ tag: 'anime', count: 3 }] }),
});
}
if (url.includes('/top-tags')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, tags: [] }) });
}
if (url.includes('/base-models')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, base_models: [] }) });
}
if (url.includes('/model-types')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, model_types: [] }) });
}
return Promise.resolve({ ok: true, json: async () => ({ success: true }) });
});
renderControlsDom('loras');
const stateModule = await import('../../../static/js/state/index.js');
stateModule.initPageState('loras');
const { FilterManager } = await import('../../../static/js/managers/FilterManager.js');
const manager = new FilterManager({ page: 'loras' });
// Open the panel so tags load
manager.toggleFilterPanel();
await vi.runAllTimersAsync();
const input = document.getElementById('modelTagsSearchInput');
input.value = 'ani';
input.dispatchEvent(new Event('input', { bubbles: true }));
// Before debounce fires, no search-tags call yet
expect(searchTagsUrls.length).toBe(0);
// Advance past the 150ms debounce
vi.advanceTimersByTime(160);
await vi.runAllTimersAsync();
expect(searchTagsUrls.length).toBe(1);
expect(searchTagsUrls[0]).toContain('/search-tags');
expect(searchTagsUrls[0]).toContain('q=ani');
vi.useRealTimers();
});
it('renders selected-but-missing tags in a dedicated group at the top', async () => {
vi.useFakeTimers();
global.fetch = vi.fn((url) => {
if (url.includes('/search-tags')) {
return Promise.resolve({
ok: true,
json: async () => ({ success: true, tags: [{ tag: 'anime', count: 3 }] }),
});
}
if (url.includes('/top-tags')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, tags: [] }) });
}
if (url.includes('/base-models')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, base_models: [] }) });
}
if (url.includes('/model-types')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, model_types: [] }) });
}
return Promise.resolve({ ok: true, json: async () => ({ success: true }) });
});
renderControlsDom('loras');
const stateModule = await import('../../../static/js/state/index.js');
stateModule.initPageState('loras');
const { FilterManager } = await import('../../../static/js/managers/FilterManager.js');
const manager = new FilterManager({ page: 'loras' });
// Pre-seed an active tag filter that won't appear in search results
manager.filters.tags = { 'my-custom-tag': 'include' };
// Open panel and let top-tags load (empty)
manager.toggleFilterPanel();
await vi.runAllTimersAsync();
// Type a search query
const input = document.getElementById('modelTagsSearchInput');
input.value = 'ani';
input.dispatchEvent(new Event('input', { bubbles: true }));
vi.advanceTimersByTime(160);
await vi.runAllTimersAsync();
const container = document.getElementById('modelTagsFilter');
const extraTag = container.querySelector('.filter-tag.extra-tag');
expect(extraTag).not.toBeNull();
expect(extraTag.dataset.tag).toBe('my-custom-tag');
// The search result tag should also be present
const resultTag = container.querySelector('.filter-tag.tag-filter[data-tag="anime"]');
expect(resultTag).not.toBeNull();
vi.useRealTimers();
});
it('shows empty state when search returns no matches and no selected tags', async () => {
vi.useFakeTimers();
global.fetch = vi.fn((url) => {
if (url.includes('/search-tags')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, tags: [] }) });
}
if (url.includes('/top-tags')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, tags: [] }) });
}
if (url.includes('/base-models')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, base_models: [] }) });
}
if (url.includes('/model-types')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, model_types: [] }) });
}
return Promise.resolve({ ok: true, json: async () => ({ success: true }) });
});
renderControlsDom('loras');
const stateModule = await import('../../../static/js/state/index.js');
stateModule.initPageState('loras');
const { FilterManager } = await import('../../../static/js/managers/FilterManager.js');
const manager = new FilterManager({ page: 'loras' });
manager.toggleFilterPanel();
await vi.runAllTimersAsync();
const input = document.getElementById('modelTagsSearchInput');
input.value = 'zzz';
input.dispatchEvent(new Event('input', { bubbles: true }));
vi.advanceTimersByTime(160);
await vi.runAllTimersAsync();
const emptyState = document.getElementById('modelTagsEmptyState');
expect(emptyState.hidden).toBe(false);
vi.useRealTimers();
});
it('reloads top tags when search input is cleared', async () => {
vi.useFakeTimers();
let topTagsCallCount = 0;
global.fetch = vi.fn((url) => {
if (url.includes('/search-tags')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, tags: [{ tag: 'anime', count: 3 }] }) });
}
if (url.includes('/top-tags')) {
topTagsCallCount++;
return Promise.resolve({ ok: true, json: async () => ({ success: true, tags: [] }) });
}
if (url.includes('/base-models')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, base_models: [] }) });
}
if (url.includes('/model-types')) {
return Promise.resolve({ ok: true, json: async () => ({ success: true, model_types: [] }) });
}
return Promise.resolve({ ok: true, json: async () => ({ success: true }) });
});
renderControlsDom('loras');
const stateModule = await import('../../../static/js/state/index.js');
stateModule.initPageState('loras');
const { FilterManager } = await import('../../../static/js/managers/FilterManager.js');
const manager = new FilterManager({ page: 'loras' });
manager.toggleFilterPanel();
await vi.runAllTimersAsync();
const callsAfterOpen = topTagsCallCount;
expect(callsAfterOpen).toBeGreaterThanOrEqual(1);
// Type, then clear
const input = document.getElementById('modelTagsSearchInput');
input.value = 'ani';
input.dispatchEvent(new Event('input', { bubbles: true }));
vi.advanceTimersByTime(160);
await vi.runAllTimersAsync();
input.value = '';
input.dispatchEvent(new Event('input', { bubbles: true }));
vi.advanceTimersByTime(160);
await vi.runAllTimersAsync();
// An additional top-tags call should have happened after clearing
expect(topTagsCallCount).toBeGreaterThan(callsAfterOpen);
vi.useRealTimers();
});
});
}); });
+54 -4
View File
@@ -9,6 +9,7 @@ import {
parseCombinedSamplerName, parseCombinedSamplerName,
resolveSamplerScheduler, resolveSamplerScheduler,
findMatchingWidgets, findMatchingWidgets,
NODE_TYPE_WIDGET_OVERRIDES,
} from '../../../static/js/utils/genParamsMapper.js'; } from '../../../static/js/utils/genParamsMapper.js';
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
@@ -204,9 +205,9 @@ describe('findMatchingWidgets', () => {
it('matches seed to seed widget', () => { it('matches seed to seed widget', () => {
const updates = findMatchingWidgets(['seed', 'steps', 'cfg', 'sampler_name', 'scheduler'], resolved); const updates = findMatchingWidgets(['seed', 'steps', 'cfg', 'sampler_name', 'scheduler'], resolved);
expect(updates).toContainEqual({ widgetName: 'seed', value: '42' }); expect(updates).toContainEqual({ widgetName: 'seed', value: 42 });
expect(updates).toContainEqual({ widgetName: 'steps', value: '30' }); expect(updates).toContainEqual({ widgetName: 'steps', value: 30 });
expect(updates).toContainEqual({ widgetName: 'cfg', value: '7' }); expect(updates).toContainEqual({ widgetName: 'cfg', value: 7 });
expect(updates).toContainEqual({ widgetName: 'sampler_name', value: 'euler_ancestral' }); expect(updates).toContainEqual({ widgetName: 'sampler_name', value: 'euler_ancestral' });
expect(updates).toContainEqual({ widgetName: 'scheduler', value: 'karras' }); expect(updates).toContainEqual({ widgetName: 'scheduler', value: 'karras' });
}); });
@@ -221,7 +222,7 @@ describe('findMatchingWidgets', () => {
const updates = findMatchingWidgets(['noise_seed', 'steps', 'cfg', 'sampler_name', 'scheduler'], resolved); const updates = findMatchingWidgets(['noise_seed', 'steps', 'cfg', 'sampler_name', 'scheduler'], resolved);
const seedUpdate = updates.find(u => u.widgetName === 'noise_seed'); const seedUpdate = updates.find(u => u.widgetName === 'noise_seed');
expect(seedUpdate).toBeDefined(); expect(seedUpdate).toBeDefined();
expect(seedUpdate.value).toBe('42'); expect(seedUpdate.value).toBe(42);
}); });
it('matches rgthree-style sampler widget name', () => { it('matches rgthree-style sampler widget name', () => {
@@ -243,4 +244,53 @@ describe('findMatchingWidgets', () => {
const updates = findMatchingWidgets(['seed', 'steps', 'cfg', 'sampler_name', 'scheduler'], resolved); const updates = findMatchingWidgets(['seed', 'steps', 'cfg', 'sampler_name', 'scheduler'], resolved);
expect(updates.map(u => u.widgetName)).toEqual(['seed', 'steps', 'cfg', 'sampler_name', 'scheduler']); expect(updates.map(u => u.widgetName)).toEqual(['seed', 'steps', 'cfg', 'sampler_name', 'scheduler']);
}); });
// --- node-type-specific overrides ---
it('matches GlobalSeed //Inspire value widget for seed param', () => {
const updates = findMatchingWidgets(
['value', 'mode', 'action', 'last_seed'],
{ seed: 42 },
'GlobalSeed //Inspire'
);
expect(updates).toHaveLength(1);
expect(updates[0]).toEqual({ widgetName: 'value', value: 42 });
});
it('ignores nodeType when it does not match any override entry', () => {
const updates = findMatchingWidgets(
['value', 'mode', 'action', 'last_seed'],
{ seed: 42 },
'SomeOtherNode'
);
expect(updates).toEqual([]);
});
it('still falls back to global candidates when override candidates do not match', () => {
// GlobalSeed override does not include steps — should use global candidate "steps"
const updates = findMatchingWidgets(
['steps', 'cfg', 'sampler_name'],
{ steps: 20 },
'GlobalSeed //Inspire'
);
expect(updates).toHaveLength(1);
expect(updates[0]).toEqual({ widgetName: 'steps', value: 20 });
});
it('prefers overrides when both override and global candidates match', () => {
// If a hypothetical node has both "value" and "seed" widgets AND a
// GlobalSeed override, the override candidate "value" should take precedence
const updates = findMatchingWidgets(
['seed', 'noise_seed', 'value', 'mode'],
{ seed: 99 },
'GlobalSeed //Inspire'
);
expect(updates).toHaveLength(1);
expect(updates[0].widgetName).toBe('value');
});
it('omits nodeType argument and still matches via global candidates', () => {
const updates = findMatchingWidgets(['seed', 'steps', 'cfg'], { seed: 7 });
expect(updates).toHaveLength(1);
expect(updates[0]).toEqual({ widgetName: 'seed', value: 7 });
});
}); });
@@ -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()
+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
+50
View File
@@ -36,3 +36,53 @@ async def test_model_query_handler_rejects_negative_limit_for_base_models():
await handler.get_base_models(SimpleNamespace(query={"limit": "-1"})) await handler.get_base_models(SimpleNamespace(query={"limit": "-1"}))
assert service.received_limit == 20 assert service.received_limit == 20
class DummySearchTagsService:
"""Minimal service stub recording search_tags arguments."""
def __init__(self, result=None):
self.received_query = None
self.received_limit = None
self._result = result or []
async def search_tags(self, query, limit):
self.received_query = query
self.received_limit = limit
return self._result
@pytest.mark.asyncio
async def test_model_query_handler_search_tags_passes_query_and_limit():
service = DummySearchTagsService(result=[{"tag": "anime", "count": 3}])
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
response = await handler.search_tags(
SimpleNamespace(query={"q": "ani", "limit": "50"})
)
payload = json.loads(response.text)
assert payload["success"] is True
assert payload["tags"] == [{"tag": "anime", "count": 3}]
assert service.received_query == "ani"
assert service.received_limit == 50
@pytest.mark.asyncio
async def test_model_query_handler_search_tags_defaults_limit_to_20():
service = DummySearchTagsService()
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
await handler.search_tags(SimpleNamespace(query={}))
assert service.received_limit == 20
@pytest.mark.asyncio
async def test_model_query_handler_search_tags_clamps_negative_limit():
service = DummySearchTagsService()
handler = ModelQueryHandler(service=service, logger=logging.getLogger(__name__))
await handler.search_tags(SimpleNamespace(query={"limit": "-5"}))
assert service.received_limit == 20
+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_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"))
(tmp_path / ".git").mkdir()
async def _fake_zip(*args, **kwargs):
return True, "v9.9.9"
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"] == "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())
+66
View File
@@ -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"
@@ -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()
@@ -0,0 +1,193 @@
"""Unit tests for DownloadQueueService history operations.
Covers the new ``download_id``-based code paths in
``delete_history_item`` and ``retry_from_history``, plus backward
compatibility with ``id``.
"""
from pathlib import Path
import pytest
from py.services.download_queue_service import DownloadQueueService
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_service(tmp_path: Path) -> DownloadQueueService:
"""Create a DownloadQueueService backed by a temporary database."""
return DownloadQueueService(db_path=str(tmp_path / "queue.sqlite"))
async def _seed(
svc: DownloadQueueService,
download_id: str,
status: str = "failed",
) -> tuple[int, str]:
"""Insert a history row and return (autoincrement id, download_id)."""
row_id = await svc.add_to_history(
download_id=download_id,
model_id=1,
model_version_id=100,
model_name="TestModel",
version_name="v1",
status=status,
)
return row_id, download_id
# ---------------------------------------------------------------------------
# delete_history_item
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_delete_by_download_id(tmp_path: Path) -> None:
"""delete_history_item(download_id=...) removes the correct row."""
svc = _make_service(tmp_path)
rid, did = await _seed(svc, "dl-aaa")
deleted = await svc.delete_history_item(download_id=did)
assert deleted is True
# Verify gone from history
history = await svc.get_history()
assert len(history["items"]) == 0
@pytest.mark.asyncio
async def test_delete_by_id_legacy(tmp_path: Path) -> None:
"""delete_history_item(id=...) still works (backward compat)."""
svc = _make_service(tmp_path)
rid, _did = await _seed(svc, "dl-bbb")
deleted = await svc.delete_history_item(id=rid)
assert deleted is True
history = await svc.get_history()
assert len(history["items"]) == 0
@pytest.mark.asyncio
async def test_delete_no_params_returns_false(tmp_path: Path) -> None:
"""Calling delete_history_item with no params returns False."""
svc = _make_service(tmp_path)
await _seed(svc, "dl-ccc")
deleted = await svc.delete_history_item()
assert deleted is False
# Row is still there
history = await svc.get_history()
assert len(history["items"]) == 1
@pytest.mark.asyncio
async def test_delete_download_id_precedence(tmp_path: Path) -> None:
"""When both id and download_id are given, download_id is used."""
svc = _make_service(tmp_path)
# Insert two rows
rid_a, did_a = await _seed(svc, "dl-aaa")
rid_b, did_b = await _seed(svc, "dl-bbb")
# Delete by download_id while also passing the *wrong* id
deleted = await svc.delete_history_item(id=rid_b, download_id=did_a)
assert deleted is True
history = await svc.get_history()
ids_left = [it["id"] for it in history["items"]]
assert rid_a not in ids_left # dl-aaa was deleted
assert rid_b in ids_left # dl-bbb (wrong id) was ignored
# ---------------------------------------------------------------------------
# retry_from_history
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_retry_by_download_id(tmp_path: Path) -> None:
"""retry_from_history(download_id=...) re-queues and deletes history."""
svc = _make_service(tmp_path)
rid, did = await _seed(svc, "dl-fail", status="failed")
item = await svc.retry_from_history(download_id=did)
assert item is not None
assert item["status"] == "queued"
# History row must be deleted (the bug fix)
history = await svc.get_history()
ids_in_history = [it["id"] for it in history["items"]]
assert rid not in ids_in_history
# Queue must contain the new item
queue = await svc.get_queue()
assert len(queue) == 1
@pytest.mark.asyncio
async def test_retry_by_download_id_canceled(tmp_path: Path) -> None:
"""retry_from_history works for 'canceled' status too."""
svc = _make_service(tmp_path)
rid, did = await _seed(svc, "dl-cancel", status="canceled")
item = await svc.retry_from_history(download_id=did)
assert item is not None
assert item["status"] == "queued"
history = await svc.get_history()
assert len(history["items"]) == 0
@pytest.mark.asyncio
async def test_retry_by_id_legacy(tmp_path: Path) -> None:
"""retry_from_history(item_id=...) still works (backward compat)."""
svc = _make_service(tmp_path)
rid, _did = await _seed(svc, "dl-legacy", status="failed")
item = await svc.retry_from_history(item_id=rid)
assert item is not None
assert item["status"] == "queued"
history = await svc.get_history()
assert len(history["items"]) == 0
@pytest.mark.asyncio
async def test_retry_no_params_returns_none(tmp_path: Path) -> None:
"""Calling retry_from_history with no params returns None."""
svc = _make_service(tmp_path)
await _seed(svc, "dl-none", status="failed")
item = await svc.retry_from_history()
assert item is None
# History untouched
history = await svc.get_history()
assert len(history["items"]) == 1
@pytest.mark.asyncio
async def test_retry_non_retryable_status(tmp_path: Path) -> None:
"""retry_from_history returns None for 'completed' status."""
svc = _make_service(tmp_path)
_rid, did = await _seed(svc, "dl-ok", status="completed")
item = await svc.retry_from_history(download_id=did)
assert item is None
# History untouched
history = await svc.get_history()
assert len(history["items"]) == 1
@pytest.mark.asyncio
async def test_retry_unknown_download_id(tmp_path: Path) -> None:
"""retry_from_history returns None for a non-existent download_id."""
svc = _make_service(tmp_path)
await _seed(svc, "dl-real", status="failed")
item = await svc.retry_from_history(download_id="dl-nope")
assert item is None
+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)
+111
View File
@@ -60,3 +60,114 @@ async def test_get_metadata_provider_returns_fallback_as_is(monkeypatch):
provider = await metadata_service.get_metadata_provider() provider = await metadata_service.get_metadata_provider()
assert provider is fallback assert provider is fallback
# ---------------------------------------------------------------------------
# initialize_metadata_providers — provider gating + fallback ordering
# ---------------------------------------------------------------------------
def _stub_settings(**overrides):
"""Minimal settings stub returning configured values."""
base = {
"enable_metadata_archive_db": False,
"enable_civarchive_api": True,
"metadata_provider_order": "civitai_archive_sqlite",
}
base.update(overrides)
return SimpleNamespace(get=lambda key, default=None: base.get(key, default))
async def _run_initialize(monkeypatch, settings):
# Fresh provider manager for each test
monkeypatch.setattr(
metadata_service.ModelMetadataProviderManager,
"get_instance",
AsyncMock(return_value=metadata_service.ModelMetadataProviderManager()),
)
monkeypatch.setattr(
metadata_service, "get_settings_manager", lambda: settings
)
monkeypatch.setattr(
metadata_service.ServiceRegistry,
"get_civitai_client",
AsyncMock(return_value=object()),
)
monkeypatch.setattr(
metadata_service.ServiceRegistry,
"get_civarchive_client",
AsyncMock(return_value=object()),
)
# Make MetadataArchiveManager report a usable db path when enabled
fake_archive = SimpleNamespace(get_database_path=lambda: "/tmp/fake.db")
monkeypatch.setattr(
metadata_service, "MetadataArchiveManager", lambda _base: fake_archive
)
# Pretend the db file exists
monkeypatch.setattr(metadata_service.os.path, "exists", lambda _p: True)
manager = await metadata_service.initialize_metadata_providers()
return manager
def _fallback_provider_order(manager):
"""Return the ordered list of provider labels inside the fallback provider."""
fallback = manager.providers.get("fallback")
assert isinstance(fallback, FallbackMetadataProvider), "expected a fallback provider"
return list(fallback._provider_labels)
@pytest.mark.asyncio
async def test_initialize_providers_default_order(monkeypatch):
settings = _stub_settings(enable_metadata_archive_db=True)
manager = await _run_initialize(monkeypatch, settings)
assert _fallback_provider_order(manager) == ["civitai_api", "civarchive_api", "sqlite"]
@pytest.mark.asyncio
async def test_initialize_providers_prefer_sqlite_order(monkeypatch):
settings = _stub_settings(
enable_metadata_archive_db=True,
metadata_provider_order="civitai_sqlite_archive",
)
manager = await _run_initialize(monkeypatch, settings)
assert _fallback_provider_order(manager) == ["civitai_api", "sqlite", "civarchive_api"]
@pytest.mark.asyncio
async def test_initialize_providers_disables_civarchive(monkeypatch):
settings = _stub_settings(
enable_metadata_archive_db=True,
enable_civarchive_api=False,
)
manager = await _run_initialize(monkeypatch, settings)
# civarchive_api must not be registered at all
assert "civarchive_api" not in manager.providers
assert _fallback_provider_order(manager) == ["civitai_api", "sqlite"]
@pytest.mark.asyncio
async def test_initialize_providers_skips_unavailable_sqlite_in_preset(monkeypatch):
# Preset wants sqlite before civarchive, but archive db is disabled ->
# sqlite is unavailable and must be skipped, civarchive stays.
settings = _stub_settings(
enable_metadata_archive_db=False,
metadata_provider_order="civitai_sqlite_archive",
)
manager = await _run_initialize(monkeypatch, settings)
assert _fallback_provider_order(manager) == ["civitai_api", "civarchive_api"]
@pytest.mark.asyncio
async def test_initialize_providers_single_provider_when_only_civitai(monkeypatch):
# Both archive db and civarchive disabled -> only civitai_api remains,
# which takes the single-provider path (registered as default, no fallback).
settings = _stub_settings(
enable_metadata_archive_db=False,
enable_civarchive_api=False,
)
manager = await _run_initialize(monkeypatch, settings)
assert "fallback" not in manager.providers
assert manager.default_provider == "civitai_api"
+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
+307
View File
@@ -667,3 +667,310 @@ async def test_log_duplicate_filename_summary_silent_when_no_duplicates(tmp_path
# No warning should be logged when there are no duplicates # No warning should be logged when there are no duplicates
for record in caplog.records: for record in caplog.records:
assert "Duplicate filename conflict detected" not in record.message assert "Duplicate filename conflict detected" not in record.message
# ── _cache_entries_differ ────────────────────────────────────────────
@pytest.mark.parametrize(
"a_tags, b_tags, expect_differ",
[
(["alpha", "beta"], ["beta", "alpha"], False), # order-insensitive
(["alpha"], ["alpha", "beta"], True), # count differs
([], ["alpha"], True),
(None, [], False), # None ≈ []
(["alpha"], None, True),
],
)
def test_cache_entries_differ_tags(a_tags, b_tags, expect_differ):
base = {"file_path": "/m/a.safetensors", "model_name": "A", "size": 1}
entry_a = {**base, "tags": a_tags}
entry_b = {**base, "tags": b_tags}
assert ModelScanner._cache_entries_differ(entry_a, entry_b) == expect_differ
def test_cache_entries_differ_identical():
entry = {
"file_path": "/m/a.safetensors", "model_name": "A", "size": 1,
"tags": ["x"], "civitai": {"id": 1}, "notes": "hi",
}
assert ModelScanner._cache_entries_differ(entry, dict(entry)) is False
def test_cache_entries_differ_field_changed():
a = {"file_path": "/m/a.safetensors", "model_name": "A", "size": 1}
b = {**a, "model_name": "B"}
assert ModelScanner._cache_entries_differ(a, b) is True
def test_cache_entries_differ_extra_key():
a = {"file_path": "/m/a.safetensors", "model_name": "A"}
b = {**a, "extra_field": "value"}
assert ModelScanner._cache_entries_differ(a, b) is True
# ── sync_cache_from_metadata ─────────────────────────────────────────
def _make_cache_entry(**overrides) -> dict:
entry = {
"file_path": "/m/a.safetensors",
"model_name": "TestModel",
"file_name": "a",
"folder": "",
"size": 100,
"modified": 10.0,
"sha256": "abc123",
"base_model": "SD1.5",
"preview_url": "",
"preview_nsfw_level": 0,
"from_civitai": True,
"favorite": False,
"notes": "old note",
"usage_tips": "{}",
"metadata_source": None,
"exclude": False,
"db_checked": False,
"last_checked_at": 0.0,
"tags": ["alpha"],
"civitai": {"id": 111, "modelId": 222, "name": "v1"},
"civitai_deleted": False,
"skip_metadata_refresh": False,
"hf_url": "",
"license_flags": 113,
"hash_status": "completed",
}
entry.update(overrides)
return entry
@pytest.mark.asyncio
async def test_sync_cache_no_change(tmp_path: Path):
"""When metadata matches the cache entry, return False and mutate nothing."""
scanner = DummyScanner(tmp_path)
entry = _make_cache_entry()
scanner._cache = ModelCache(
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
)
await scanner._cache.resort()
scanner._tags_count = {"alpha": 1}
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
# metadata_dict that would produce the identical cache entry
metadata_dict = {
"file_path": "/m/a.safetensors",
"model_name": "TestModel",
"file_name": "a",
"folder": "",
"size": 100,
"modified": 10.0,
"sha256": "abc123",
"base_model": "SD1.5",
"preview_url": "",
"preview_nsfw_level": 0,
"from_civitai": True,
"favorite": False,
"notes": "old note",
"usage_tips": "{}",
"tags": ["alpha"],
"civitai": {"id": 111, "modelId": 222, "name": "v1"},
"hf_url": "",
}
changed = await scanner.sync_cache_from_metadata(
"/m/a.safetensors", metadata_dict
)
assert changed is False
# Verify cache was NOT mutated
cached = await scanner.get_cached_data()
assert cached.raw_data[0]["notes"] == "old note"
@pytest.mark.asyncio
async def test_sync_cache_in_place_update(tmp_path: Path):
"""When metadata differs, update the cache entry in-place."""
scanner = DummyScanner(tmp_path)
entry = _make_cache_entry(notes="old note", tags=["alpha"], model_name="OldName")
scanner._cache = ModelCache(
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
)
await scanner._cache.resort()
scanner._tags_count = {"alpha": 1}
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
# Capture the exact dict object in raw_data before sync
original_entry_ref = scanner._cache.raw_data[0]
metadata_dict = {
"file_path": "/m/a.safetensors",
"model_name": "NewName",
"file_name": "a",
"folder": "",
"size": 100,
"modified": 10.0,
"sha256": "abc123",
"base_model": "SD1.5",
"preview_url": "",
"preview_nsfw_level": 0,
"from_civitai": True,
"favorite": False,
"notes": "new note",
"usage_tips": "{}",
"tags": ["beta", "gamma"],
"civitai": {"id": 111, "modelId": 222, "name": "v1"},
"hf_url": "",
}
changed = await scanner.sync_cache_from_metadata(
"/m/a.safetensors", metadata_dict
)
assert changed is True
cached = await scanner.get_cached_data()
updated = cached.raw_data[0]
# In-place: the same dict object persisted in raw_data
assert updated is original_entry_ref
assert updated["notes"] == "new note"
assert updated["model_name"] == "NewName"
assert sorted(updated["tags"]) == ["beta", "gamma"]
# Tag counts updated incrementally
assert scanner._tags_count.get("alpha", 0) == 0
assert scanner._tags_count.get("beta", 0) == 1
assert scanner._tags_count.get("gamma", 0) == 1
@pytest.mark.asyncio
async def test_sync_cache_not_in_cache_delegates(tmp_path: Path):
"""When the file_path is not in the cache at all, fall back to full update."""
scanner = DummyScanner(tmp_path)
scanner._cache = ModelCache(raw_data=[], folders=[], name_display_mode="model_name")
await scanner._cache.resort()
metadata_dict = {
"file_path": "/m/b.safetensors",
"model_name": "BrandNew",
"file_name": "b",
"folder": "",
"size": 200,
"modified": 20.0,
"sha256": "def456",
"base_model": "SDXL",
"preview_url": "",
"preview_nsfw_level": 0,
"from_civitai": True,
"favorite": False,
"notes": "",
"usage_tips": "{}",
"tags": [],
"civitai": {},
"hf_url": "",
}
changed = await scanner.sync_cache_from_metadata(
"/m/b.safetensors", metadata_dict
)
assert changed is True
cached = await scanner.get_cached_data()
assert len(cached.raw_data) == 1
assert cached.raw_data[0]["model_name"] == "BrandNew"
@pytest.mark.asyncio
async def test_sync_cache_conditional_resort_skipped(tmp_path: Path, monkeypatch):
"""When only non-sort-key fields change, resort() is NOT called."""
scanner = DummyScanner(tmp_path)
entry = _make_cache_entry(notes="old note", model_name="SameName")
scanner._cache = ModelCache(
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
)
await scanner._cache.resort()
scanner._cache._last_sort = ("name", "asc") # name sort is active
scanner._tags_count = {"alpha": 1}
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
# Track resort calls
resort_called = False
original_resort = scanner._cache.resort
async def tracking_resort():
nonlocal resort_called
resort_called = True
await original_resort()
monkeypatch.setattr(scanner._cache, "resort", tracking_resort)
metadata_dict = {
"file_path": "/m/a.safetensors",
"model_name": "SameName", # unchanged — no resort needed
"file_name": "a",
"folder": "",
"size": 100,
"modified": 10.0,
"sha256": "abc123",
"base_model": "SD1.5",
"preview_url": "",
"preview_nsfw_level": 0,
"from_civitai": True,
"favorite": False,
"notes": "updated note", # changed, but not sort-relevant
"usage_tips": "{}",
"tags": ["alpha"],
"civitai": {"id": 111, "modelId": 222, "name": "v1"},
"hf_url": "",
}
changed = await scanner.sync_cache_from_metadata(
"/m/a.safetensors", metadata_dict
)
assert changed is True
assert resort_called is False
@pytest.mark.asyncio
async def test_sync_cache_conditional_resort_triggered(tmp_path: Path, monkeypatch):
"""When the sort-key field changes, resort() IS called."""
scanner = DummyScanner(tmp_path)
entry = _make_cache_entry(model_name="OldName")
scanner._cache = ModelCache(
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
)
await scanner._cache.resort()
scanner._cache._last_sort = ("name", "asc")
scanner._tags_count = {"alpha": 1}
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
resort_calls = 0
original_resort = scanner._cache.resort
async def tracking_resort():
nonlocal resort_calls
resort_calls += 1
await original_resort()
monkeypatch.setattr(scanner._cache, "resort", tracking_resort)
metadata_dict = {
"file_path": "/m/a.safetensors",
"model_name": "NewName", # changed — should trigger resort
"file_name": "a",
"folder": "",
"size": 100,
"modified": 10.0,
"sha256": "abc123",
"base_model": "SD1.5",
"preview_url": "",
"preview_nsfw_level": 0,
"from_civitai": True,
"favorite": False,
"notes": "old note",
"usage_tips": "{}",
"tags": ["alpha"],
"civitai": {"id": 111, "modelId": 222, "name": "v1"},
"hf_url": "",
}
changed = await scanner.sync_cache_from_metadata(
"/m/a.safetensors", metadata_dict
)
assert changed is True
assert resort_calls == 1
@@ -225,3 +225,119 @@ def test_incremental_updates_only_touch_changed_rows(tmp_path: Path, monkeypatch
assert second['metadata_source'] == 'archive_db' assert second['metadata_source'] == 'archive_db'
assert second['civitai_deleted'] is True assert second['civitai_deleted'] is True
assert second['civitai']['creator']['username'] == 'builder_v2' assert second['civitai']['creator']['username'] == 'builder_v2'
# ── update_single_model ───────────────────────────────────────────────
def test_update_single_model_insert(tmp_path: Path, monkeypatch):
"""Insert a brand-new model row via update_single_model."""
monkeypatch.setenv('LORA_MANAGER_DISABLE_PERSISTENT_CACHE', '0')
db_path = tmp_path / 'cache.sqlite'
store = PersistentModelCache(db_path=str(db_path))
file_path = (tmp_path / 'x.safetensors').as_posix()
new_item = {
'file_path': file_path,
'file_name': 'x',
'model_name': 'Model X',
'folder': '',
'size': 42,
'modified': 1.0,
'sha256': 'sha-x',
'base_model': 'SDXL',
'preview_url': '',
'preview_nsfw_level': 0,
'from_civitai': True,
'favorite': True,
'notes': 'test note',
'usage_tips': '{}',
'metadata_source': None,
'exclude': False,
'db_checked': False,
'last_checked_at': 0.0,
'tags': ['test', 'new'],
'civitai': None,
'civitai_deleted': False,
'skip_metadata_refresh': False,
'license_flags': DEFAULT_LICENSE_FLAGS,
'hash_status': 'completed',
'hf_url': '',
}
store.update_single_model('dummy', new_item)
persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data}
assert file_path in items
assert items[file_path]['model_name'] == 'Model X'
assert items[file_path]['favorite'] is True
assert sorted(items[file_path]['tags']) == ['new', 'test']
def test_update_single_model_update_tags(tmp_path: Path, monkeypatch):
"""Tags are updated incrementally: old tags removed, new tags added."""
monkeypatch.setenv('LORA_MANAGER_DISABLE_PERSISTENT_CACHE', '0')
db_path = tmp_path / 'cache.sqlite'
store = PersistentModelCache(db_path=str(db_path))
file_path = (tmp_path / 'y.safetensors').as_posix()
base = {
'file_path': file_path, 'file_name': 'y', 'model_name': 'Y',
'folder': '', 'size': 1, 'modified': 1.0, 'sha256': 'sha-y',
'base_model': '', 'preview_url': '', 'preview_nsfw_level': 0,
'from_civitai': True, 'favorite': False, 'notes': '', 'usage_tips': '{}',
'metadata_source': None, 'exclude': False, 'db_checked': False,
'last_checked_at': 0.0, 'civitai': None, 'civitai_deleted': False,
'skip_metadata_refresh': False, 'license_flags': DEFAULT_LICENSE_FLAGS,
'hash_status': 'completed', 'hf_url': '',
}
# First insert with tags [alpha, beta]
store.update_single_model('dummy', {**base, 'tags': ['alpha', 'beta']})
# Now update: replace with [beta, gamma]
old_item = {'file_path': file_path, 'tags': ['alpha', 'beta'], 'sha256': 'sha-y'}
new_item = {**base, 'tags': ['beta', 'gamma']}
store.update_single_model('dummy', new_item, old_item=old_item)
persisted = store.load_cache('dummy')
assert persisted is not None
items = {item['file_path']: item for item in persisted.raw_data}
assert sorted(items[file_path]['tags']) == ['beta', 'gamma']
def test_update_single_model_update_hash(tmp_path: Path, monkeypatch):
"""When sha256 changes, the hash_index is updated incrementally."""
monkeypatch.setenv('LORA_MANAGER_DISABLE_PERSISTENT_CACHE', '0')
db_path = tmp_path / 'cache.sqlite'
store = PersistentModelCache(db_path=str(db_path))
file_path = (tmp_path / 'z.safetensors').as_posix()
base = {
'file_path': file_path, 'file_name': 'z', 'model_name': 'Z',
'folder': '', 'size': 1, 'modified': 1.0, 'base_model': '',
'preview_url': '', 'preview_nsfw_level': 0, 'from_civitai': True,
'favorite': False, 'notes': '', 'usage_tips': '{}',
'metadata_source': None, 'exclude': False, 'db_checked': False,
'last_checked_at': 0.0, 'tags': [], 'civitai': None,
'civitai_deleted': False, 'skip_metadata_refresh': False,
'license_flags': DEFAULT_LICENSE_FLAGS, 'hash_status': 'completed', 'hf_url': '',
}
store.update_single_model('dummy', {**base, 'sha256': 'old-hash'})
old_item = {'file_path': file_path, 'tags': [], 'sha256': 'old-hash'}
new_item = {**base, 'sha256': 'new-hash'}
store.update_single_model('dummy', new_item, old_item=old_item)
persisted = store.load_cache('dummy')
assert persisted is not None
# old hash should be gone from hash_index
old_hash_pairs = [p for p in persisted.hash_rows if p[0] == 'old-hash']
assert len(old_hash_pairs) == 0
# new hash should be present
new_hash_pairs = [p for p in persisted.hash_rows if p[0] == 'new-hash']
assert len(new_hash_pairs) == 1
assert new_hash_pairs[0][1] == file_path
+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()

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