Compare commits

...

55 Commits

Author SHA1 Message Date
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
Will Miao bf6a614e0d chore(release): bump version to v1.1.7 2026-07-13 22:18:16 +08:00
Will Miao feab01cd9c fix(preview): hide license icons for models without CivitAI metadata 2026-07-13 19:49:10 +08:00
Will Miao 966024e534 fix(registry): force re-registration on WS refresh to prevent timeout, demote empty-registry log to debug
- workflow_registry.js: add force param to refreshRegistry(), bypass fingerprint
  dedup when responding to lora_registry_refresh WS message. Without this, the
  backend's wait_for_all() times out after 0.5s because the frontend skips the
  register-nodes POST when the workflow fingerprint hasn't changed (common after
  ComfyUI restart with an empty or unchanged workflow).
- misc_handlers.py: demote 'No nodes registered after refresh' from WARNING to
  DEBUG — empty workflows are a normal operational state, not a warning-worthy
  condition.
2026-07-13 19:10:48 +08:00
Will Miao 2018722cc8 fix(registry): handle compound subgraph node IDs, add proactive node push from graph hooks
- Handle compound node IDs (e.g. "252:0") from expanded group subgraphs
  to fix 400 Bad Request on workflows with group nodes
- Frontend proactively pushes node data via afterConfigureGraph and
  LiteGraph hooks (onNodeAdded/onNodeRemoved/graphChanged), eliminating
  WebSocket round-trip latency for most "Send to Workflow" operations
- Add content-fingerprint dedup to skip duplicate register-nodes POSTs
- Fast-path cache returns immediately when tabs are registered (including
  0-node registrations), avoiding unnecessary WS refresh cycles
- Distinguish "Empty Registry" from other errors in standalone UI toast
- Reduce WS refresh timeout 2s→0.5s, add cooldown and lock to prevent
  concurrent refresh storms
- All [LM:Registry] logs at DEBUG level
2026-07-13 18:02:26 +08:00
Will Miao 9d85c2a44a fix(ui): prevent tags widget from auto-resizing in Vue mode when tags change 2026-07-13 14:55:40 +08:00
Will Miao 03dd047e62 fix(download): return 200 instead of 500 when user cancels download 2026-07-13 11:47:48 +08:00
Will Miao 86b547c1e0 fix(locales): add missing downloadStopped key to toast.downloads section 2026-07-13 11:35:48 +08:00
Will Miao bab9752c8b fix(download): close modal before progress overlay and fix downloadId ReferenceError on cancel 2026-07-13 11:29:47 +08:00
Will Miao 774cc1be86 fix(download): use file ID for exact match, add debug logging for multi-file selection (#1023)
- Frontend: send file.id in file_params, use null instead of hardcoded defaults
- Backend: priority matching (ID exact → primary → lenient metadata)
- Lenient metadata: only compare fields present on both sides (fixes GGUF size mismatch)
- Add debug logs at key points: entry, file_params received, match result, anomaly signals
2026-07-13 11:15:03 +08:00
Will Miao 234b73c8a2 feat(ui): add cancel button to download progress modal 2026-07-13 09:40:53 +08:00
Will Miao abd06c48f4 fix(settings): reject checkpoints↔unet path overlap in extra folder paths with inline error UI
Changes:
- Backend: _validate_folder_paths() now checks checkpoints↔unet overlap
  within the same library using os.path.realpath() for symlink resolution
- Backend: set() calls _validate_folder_paths() for both folder_paths and
  extra_folder_paths before writing
- Backend: extracted _normalize_path_set() helper to eliminate duplicated
  normalization logic
- Frontend: inline error display with red border + error message below the
  conflicting input, no save triggered
- Frontend: path normalization (strip trailing slash, lowercase) in pre-check
  to reduce false negatives vs backend realpath
- Frontend: asymmetric error UX — message only on the user-edited side,
  red border on the pre-existing conflict side
- CSS: has-error styles with hardcoded rgba fallback for older browsers
- i18n: checkpointUnetOverlap + checkpointUnetOverlapInline keys added to
  all 10 locale files
2026-07-13 08:22:40 +08:00
Will Miao 6ca411e4e4 fix(ui): make loras widget fixed-size with user-controlled node resize
Remove dynamic height calculation that auto-resized the node when
LoRAs are added or removed. The widget now stays at the size the user
sets via the node resize handle, scrolling when content overflows.

- Drop updateWidgetHeight() and hardcoded entry-count height math
- Set --comfy-widget-min-height once (200px) instead of recalculating
- In Vue mode: add contain:layout+size to break the ResizeObserver
  feedback loop that forced node growth with content (CSS via
  .lm-loras-container.lm-vue-node scoped to vueNodesMode only)
- Remove unused "Node 2.0: Maximum visible LoRA entries" setting
2026-07-12 22:35:58 +08:00
Will Miao 6470021e77 feat(settings): persist LORA_MANAGER_PORTABLE to settings.json on first use (#1018) 2026-07-12 09:32:30 +08:00
Will Miao 71658ab37b feat(settings): add LORA_MANAGER_PORTABLE env var for per-instance settings isolation (#1018) 2026-07-12 07:44:31 +08:00
Will Miao 4f016a8024 feat(fetch): skip CivArchive API for HuggingFace-sourced models
- Bulk refresh filter now excludes models with hf_url
- Individual refresh for HF models only checks CivitAI API
- CivArchive client validates model IDs before querying
2026-07-11 20:29:54 +08:00
Will Miao f362ed585b fix(preview): gracefully handle deleted preview files - image fallback, cache cleanup, quieter logs
- Add onerror handler on <img> previews to fallback to no-preview.png
- Fire async cache cleanup when preview file returns 404
- Add ModelCache.clear_preview_by_path() for safe stale-url removal
- Downgrade /api/lm/previews 404 log from warning to debug
2026-07-10 21:25:07 +08:00
112 changed files with 27118 additions and 20915 deletions
+1
View File
@@ -102,6 +102,7 @@ npm run test:coverage # Generate coverage report
- ComfyUI: `app.registerExtension()`, `node.addDOMWidget(name, type, element, options)`
- Event handlers via `addEventListener` or widget callbacks
- Shared utilities: `web/comfyui/utils.js`
- Dual-mode rendering patterns (canvas vs Vue): see `docs/comfyui-dual-mode-widgets.md`
### Vue Composables Pattern
+2 -2
View File
File diff suppressed because one or more lines are too long
+13
View File
@@ -15,6 +15,9 @@ try: # pragma: no cover - import fallback for pytest collection
from .py.nodes.lora_pool import LoraPoolLM
from .py.nodes.lora_randomizer import LoraRandomizerLM
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.metadata_collector import init as init_metadata_collector
except (
ImportError
@@ -56,6 +59,13 @@ except (
"py.nodes.lora_randomizer"
).LoraRandomizerLM
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
init_metadata_collector = importlib.import_module("py.metadata_collector").init
NODE_CLASS_MAPPINGS = {
@@ -75,6 +85,9 @@ NODE_CLASS_MAPPINGS = {
LoraPoolLM.NAME: LoraPoolLM,
LoraRandomizerLM.NAME: LoraRandomizerLM,
LoraCyclerLM.NAME: LoraCyclerLM,
LoraInfoLM.NAME: LoraInfoLM,
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
CreateHookLoraLM.NAME: CreateHookLoraLM,
}
WEB_DIRECTORY = "./web/comfyui"
+391 -362
View File
File diff suppressed because it is too large Load Diff
+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
+2202 -2189
View File
File diff suppressed because it is too large Load Diff
+18 -5
View File
@@ -233,7 +233,7 @@
"presetNamePlaceholder": "Preset name...",
"baseModel": "Base Model",
"baseModelSearchPlaceholder": "Search base models...",
"modelTags": "Tags (Top 20)",
"modelTags": "Tags",
"modelTypes": "Model Types",
"license": "License",
"noCreditRequired": "No Credit Required",
@@ -241,6 +241,8 @@
"allowSellingGeneratedContentTooltip": "Allow selling generated images",
"noCreditRequiredTooltip": "Use the model without crediting the creator",
"noTags": "No tags",
"tagSearchPlaceholder": "Search tags...",
"noTagMatches": "No tags match the current search.",
"autoTags": "Auto Tags",
"noBaseModelMatches": "No base models match the current search.",
"clearAll": "Clear All Filters",
@@ -505,7 +507,9 @@
"saveSuccess": "Extra folder paths updated. Restart required to apply changes.",
"saveError": "Failed to update extra folder paths: {message}",
"validation": {
"duplicatePath": "This path is already configured"
"duplicatePath": "This path is already configured",
"checkpointUnetOverlap": "Cannot use the same path for both checkpoints and diffusion models: {paths}",
"checkpointUnetOverlapInline": "This path is also used for a different model type. Use separate folders for checkpoints and diffusion models."
}
},
"priorityTags": {
@@ -638,7 +642,13 @@
"preparing": "Preparing download...",
"connecting": "Connecting to download server...",
"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": {
"enableProxy": "Enable App-level Proxy",
@@ -1205,7 +1215,9 @@
"preparing": "Preparing download...",
"downloadedPreview": "Downloaded preview image",
"downloadingFile": "Downloading {type} file",
"finalizing": "Finalizing download..."
"finalizing": "Finalizing download...",
"cancelling": "Cancelling download...",
"cancelled": "Download cancelled"
},
"progress": {
"currentFile": "Current file:",
@@ -2013,7 +2025,8 @@
"imagesCompleted": "Example images {action} completed",
"imagesFailed": "Example images {action} failed",
"loadError": "Error loading downloads: {message}",
"downloadError": "Download error: {message}"
"downloadError": "Download error: {message}",
"downloadStopped": "Download cancelled"
},
"import": {
"folderTreeFailed": "Failed to load folder tree",
+2202 -2189
View File
File diff suppressed because it is too large Load Diff
+2202 -2189
View File
File diff suppressed because it is too large Load Diff
+2202 -2189
View File
File diff suppressed because it is too large Load Diff
+2202 -2189
View File
File diff suppressed because it is too large Load Diff
+2202 -2189
View File
File diff suppressed because it is too large Load Diff
+2202 -2189
View File
File diff suppressed because it is too large Load Diff
+2202 -2189
View File
File diff suppressed because it is too large Load Diff
+2202 -2189
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):
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")
if not isinstance(extra_folder_paths, dict):
return
@@ -233,10 +239,6 @@ class Config:
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:
logger.info(
"Found extra LoRA roots:"
@@ -357,6 +359,47 @@ class Config:
"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(
comfy_library.get("default_lora_root", ""),
list(self.loras_roots or []),
+6 -1
View File
@@ -41,7 +41,12 @@ async def api_json_error(
if exc.status < 400:
raise
logger.warning(
# Preview 404 is routine (file deleted from disk) — not worth a warning.
logger_method = logger.warning
if request.path.startswith("/api/lm/previews") and exc.status == 404:
logger_method = logger.debug
logger_method(
"API %s %s returned HTTP %d: %s",
request.method,
request.path,
+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 logging
import re
import comfy.sd # type: ignore
import comfy.utils # type: ignore
@@ -14,6 +13,7 @@ from .utils import (
extract_lora_name,
get_loras_list,
nunchaku_load_lora,
parse_lora_syntax,
)
logger = logging.getLogger(__name__)
@@ -189,25 +189,10 @@ class LoraTextLoaderLM:
RETURN_NAMES = ("MODEL", "CLIP", "trigger_words", "loaded_loras")
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):
"""Load LoRAs based on text syntax input."""
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_entries.append({
"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),)
+20
View File
@@ -36,6 +36,7 @@ any_type = AnyType("*")
# Common methods extracted from lora_loader.py and lora_stacker.py
import os
import re
import logging
import copy
import sys
@@ -69,6 +70,25 @@ def extract_lora_name(lora_path):
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):
"""Helper to extract loras list from either old or new kwargs format"""
if "loras" not in kwargs:
+361 -34
View File
@@ -573,12 +573,18 @@ class NodeRegistry:
tab_nodes[nd["unique_id"]] = nd
async with self._lock:
prev_count = len(self._tab_nodes.get(sid, {}))
self._tab_nodes[sid] = tab_nodes
self._waiting_clients.discard(sid)
if not self._waiting_clients:
self._ready.set()
total_tabs = len(self._tab_nodes)
logger.debug("Registered %s nodes from client %s", len(nodes), sid)
if len(nodes) != prev_count or len(nodes) > 0:
logger.debug(
"[LM:Registry] stored %s nodes (was %s) for client %s (total tabs: %s)",
len(nodes), prev_count, sid, total_tabs,
)
def prepare_for_refresh(self, active_sids: list[str]) -> None:
"""Set the list of client IDs we expect to hear from during the next refresh cycle."""
@@ -601,10 +607,17 @@ class NodeRegistry:
longer connected."""
async with self._lock:
# Garbage-collect stale entries (disconnected tabs)
stale_sids = []
if active_sids is not None:
for sid in list(self._tab_nodes):
if sid not in active_sids:
stale_sids.append(sid)
del self._tab_nodes[sid]
if stale_sids:
logger.debug(
"[LM:Registry] GC pruned %s disconnected tabs: %s",
len(stale_sids), stale_sids,
)
merged: dict[str, dict] = {}
tab_info: dict[str, dict] = {}
@@ -1557,7 +1570,11 @@ class SettingsHandler:
else:
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()
if key in self._PROXY_KEYS:
@@ -1771,6 +1788,124 @@ class LoraCodeHandler:
logger.error("Failed to update lora code: %s", exc, exc_info=True)
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:
async def get_trained_words(self, request: web.Request) -> web.Response:
@@ -3116,6 +3251,8 @@ class NodeRegistryHandler:
self._node_registry = node_registry
self._prompt_server = prompt_server
self._standalone_mode = standalone_mode
self._refresh_lock = asyncio.Lock()
self._last_slow_path_ts: float = 0.0
async def register_nodes(self, request: web.Request) -> web.Response:
try:
@@ -3162,7 +3299,12 @@ class NodeRegistryHandler:
)
graph_name = node.get("graph_name")
try:
node["node_id"] = int(node_id)
# Handle compound node IDs from expanded group subgraphs,
# e.g. "252:0" → 0 (parent scope is already in graph_id)
if isinstance(node_id, str) and ":" in node_id:
node["node_id"] = int(node_id.rsplit(":", 1)[-1])
else:
node["node_id"] = int(node_id)
except (TypeError, ValueError):
return web.json_response(
{
@@ -3203,42 +3345,101 @@ class NodeRegistryHandler:
status=503,
)
# Snapshot of currently-connected ComfyUI tabs
active_sids = list(self._prompt_server.instance.sockets.keys())
self._node_registry.prepare_for_refresh(active_sids)
try:
self._prompt_server.instance.send_sync("lora_registry_refresh", {})
logger.debug(
"Sent registry refresh request (expecting %s clients)", len(active_sids)
)
except Exception as exc:
logger.error("Failed to send registry refresh message: %s", exc)
return web.json_response(
{
"success": False,
"error": "Communication Error",
"message": f"Failed to communicate with ComfyUI frontend: {exc}",
},
status=500,
)
if not await self._node_registry.wait_for_all(timeout=2.0):
logger.warning(
"Registry refresh timeout after 2s (%s/%s clients responded)",
len(active_sids) - self._node_registry.pending_client_count,
len(active_sids),
)
# Re-read current sockets after the wait: a tab may have connected
# while we were waiting, and we don't want to garbage-collect it.
current_sids = set(self._prompt_server.instance.sockets.keys())
# Fast path: if the frontend has already pushed node data (via
# afterConfigureGraph / graphChanged hooks), return it immediately
# without triggering a WebSocket round-trip.
registry_info = await self._node_registry.get_merged_registry(
active_sids=current_sids
)
if registry_info["tab_count"] > 0:
logger.debug(
"[LM:Registry] fast path: %s nodes across %s tabs %s",
registry_info["node_count"],
registry_info["tab_count"],
dict(registry_info.get("tabs", {})),
)
return web.json_response({"success": True, "data": registry_info})
# Slow path: registry is empty — trigger refresh via WebSocket.
# Serialize with an async lock so concurrent callers don't all
# trigger separate WS refresh cycles. The second caller will
# re-check the fast path and (usually) find populated data.
async with self._refresh_lock:
# Re-check after acquiring the lock — another concurrent call
# may have populated the cache while we were waiting.
registry_info = await self._node_registry.get_merged_registry(
active_sids=current_sids
)
if registry_info["tab_count"] > 0:
logger.debug(
"[LM:Registry] fast path after lock wait: %s nodes across %s tabs",
registry_info["node_count"],
registry_info["tab_count"],
)
return web.json_response({"success": True, "data": registry_info})
# Cooldown: if the slow path ran recently (< 2 s) and
# returned empty, skip another WS round-trip.
elapsed = time.monotonic() - self._last_slow_path_ts
if elapsed < 2.0:
logger.debug(
"[LM:Registry] slow path cooldown (%.1fs since last refresh), returning empty",
elapsed,
)
return web.json_response(
{
"success": False,
"error": "Empty Registry",
"message": "No workflow nodes found — ensure ComfyUI is open and the extension is loaded.",
},
status=408,
)
logger.debug(
"[LM:Registry] slow path: cache empty, triggering WS refresh (%s connected tabs: %s)",
len(current_sids), list(current_sids)[:5],
)
active_sids = list(current_sids)
self._node_registry.prepare_for_refresh(active_sids)
try:
self._prompt_server.instance.send_sync("lora_registry_refresh", {})
logger.debug(
"Sent registry refresh request (expecting %s clients)", len(active_sids)
)
except Exception as exc:
logger.error("Failed to send registry refresh message: %s", exc)
return web.json_response(
{
"success": False,
"error": "Communication Error",
"message": f"Failed to communicate with ComfyUI frontend: {exc}",
},
status=500,
)
if not await self._node_registry.wait_for_all(timeout=0.5):
logger.warning(
"Registry refresh timeout after 0.5s (%s/%s clients responded)",
len(active_sids) - self._node_registry.pending_client_count,
len(active_sids),
)
# Re-read current sockets after the wait: a tab may have connected
# while we were waiting, and we don't want to garbage-collect it.
current_sids = set(self._prompt_server.instance.sockets.keys())
registry_info = await self._node_registry.get_merged_registry(
active_sids=current_sids
)
self._last_slow_path_ts = time.monotonic()
if registry_info["node_count"] == 0:
logger.warning("No nodes registered after refresh")
logger.debug(
"[LM:Registry] refresh OK — %s connected tab(s) but 0 compatible nodes found",
registry_info["tab_count"],
)
return web.json_response(
{
"success": False,
@@ -3274,7 +3475,7 @@ class NodeRegistryHandler:
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(
{"success": False, "error": "Missing value parameter"}, status=400
)
@@ -3352,6 +3553,130 @@ class NodeRegistryHandler:
logger.error("Failed to update node widget: %s", exc, exc_info=True)
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:
"""Aggregate handlers into a lookup compatible with the registrar."""
@@ -3418,10 +3743,12 @@ class MiscHandlerSet:
"update_usage_stats": self.usage_stats.update_usage_stats,
"get_usage_stats": self.usage_stats.get_usage_stats,
"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_model_example_files": self.model_examples.get_model_example_files,
"register_nodes": self.node_registry.register_nodes,
"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,
"check_model_exists": self.model_library.check_model_exists,
"check_models_exist": self.model_library.check_models_exist,
+60 -14
View File
@@ -973,6 +973,8 @@ class ModelQueryHandler:
limit = int(request.query.get("limit", "20"))
if limit < 0:
limit = 20
elif limit > 200:
limit = 20
top_tags = await self._service.get_top_tags(limit)
return web.json_response({"success": True, "tags": top_tags})
except Exception as exc:
@@ -981,6 +983,22 @@ class ModelQueryHandler:
{"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:
try:
limit = int(request.query.get("limit", "20"))
@@ -1275,9 +1293,13 @@ class ModelQueryHandler:
text=f"{self._service.model_type.capitalize()} file name is required",
status=400,
)
notes = await self._service.get_model_notes(model_name)
if notes is not None:
return web.json_response({"success": True, "notes": notes})
result = await self._service.get_model_notes(model_name)
if result is not None:
return web.json_response({
"success": True,
"notes": result["notes"],
"file_path": result["file_path"],
})
return web.json_response(
{
"success": False,
@@ -1313,9 +1335,20 @@ class ModelQueryHandler:
}
if include_license_flags:
model_data = await self._service.get_model_info_by_name(model_name)
license_flags = (model_data or {}).get("license_flags")
if license_flags is not None:
response_payload["license_flags"] = int(license_flags)
# Only return license_flags when real CivitAI model license
# data exists. This mirrors ModelModal's guard
# (modelData?.civitai?.model) so the preview tooltip never
# shows misleading license icons for HF or other models
# without actual license metadata.
civitai_data = (model_data or {}).get("civitai") or {}
has_license_data = (
isinstance(civitai_data, dict)
and isinstance(civitai_data.get("model"), dict)
)
if has_license_data:
license_flags = (model_data or {}).get("license_flags")
if license_flags is not None:
response_payload["license_flags"] = int(license_flags)
# Include the user's license icon style preference so the
# ComfyUI tooltip can pick the right set without a separate
# API call.
@@ -1772,14 +1805,20 @@ class ModelDownloadHandler:
async def delete_download_history_item(self, request: web.Request) -> web.Response:
try:
item_id = int(request.query.get("id", "0"))
if not item_id:
download_id = request.query.get("download_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(
{"success": False, "error": "id is required"}, status=400
{"success": False, "error": "id or download_id is required"},
status=400,
)
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})
except Exception as exc:
self._logger.error(
@@ -1789,14 +1828,20 @@ class ModelDownloadHandler:
async def retry_download_from_history(self, request: web.Request) -> web.Response:
try:
item_id = int(request.query.get("id", "0"))
if not item_id:
download_id = request.query.get("download_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(
{"success": False, "error": "id is required"}, status=400
{"success": False, "error": "id or download_id is required"},
status=400,
)
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:
return web.json_response(
{"success": False, "error": "History item not found or not retryable"},
@@ -2920,6 +2965,7 @@ class ModelHandlerSet:
"bulk_delete_models": self.management.bulk_delete_models,
"verify_duplicates": self.management.verify_duplicates,
"get_top_tags": self.query.get_top_tags,
"search_tags": self.query.search_tags,
"get_base_models": self.query.get_base_models,
"get_model_types": self.query.get_model_types,
"scan_models": self.query.scan_models,
+31
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import asyncio
import logging
import mimetypes
import urllib.parse
@@ -53,6 +54,7 @@ class PreviewHandler:
if not resolved.is_file():
logger.debug("Preview file not found at %s", str(resolved))
asyncio.create_task(self._cleanup_stale_preview_url(normalized))
raise web.HTTPNotFound(text="Preview file not found")
# aiohttp's FileResponse handles range requests, content headers, and
@@ -69,6 +71,35 @@ class PreviewHandler:
resp.headers["Cache-Control"] = "public, max-age=86400"
return resp
async def _cleanup_stale_preview_url(self, normalized_preview_path: str) -> None:
"""Fire-and-forget: clear stale preview_url from all model caches.
When a preview file is no longer on disk, remove its reference from
every cached entry so subsequent list API responses return an empty
``preview_url``, letting the frontend show the no-preview placeholder.
"""
try:
from ...services.service_registry import ServiceRegistry
for service_name in ("lora_scanner", "checkpoint_scanner", "embedding_scanner"):
scanner = ServiceRegistry.get_service_sync(service_name)
if scanner is None or not hasattr(scanner, "_cache"):
continue
cache = getattr(scanner, "_cache", None)
if cache is None or not hasattr(cache, "clear_preview_by_path"):
continue
cleared = await cache.clear_preview_by_path(normalized_preview_path)
if cleared and hasattr(scanner, "_persist_current_cache"):
await scanner._persist_current_cache()
logger.info(
"Cleared stale preview_url for %d %s entries (%s)",
cleared,
service_name,
normalized_preview_path,
)
except Exception as exc:
logger.debug("Failed to clean up stale preview_url: %s", exc)
async def _stream_file(
self, request: web.Request, path: Path
) -> web.StreamResponse:
+55 -6
View File
@@ -72,6 +72,7 @@ class RecipeHandlerSet:
"save_recipe": self.management.save_recipe,
"delete_recipe": self.management.delete_recipe,
"get_top_tags": self.query.get_top_tags,
"search_tags": self.query.search_tags,
"get_base_models": self.query.get_base_models,
"get_roots": self.query.get_roots,
"get_folders": self.query.get_folders,
@@ -317,12 +318,11 @@ class RecipeQueryHandler:
raise RuntimeError("Recipe scanner unavailable")
limit = int(request.query.get("limit", "20"))
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
if limit < 0:
limit = 20
elif limit > 200:
limit = 20
tag_counts = await self._get_recipe_tag_counts(recipe_scanner)
sorted_tags = [
{"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)
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:
try:
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("GET", "/api/lm/get-usage-stats", "get_usage_stats"),
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/model-example-files", "get_model_example_files"),
RouteDefinition("POST", "/api/lm/register-nodes", "register_nodes"),
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/check-model-exists", "check_model_exists"),
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"
),
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}/model-types", "get_model_types"),
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("DELETE", "/api/lm/recipe/{recipe_id}", "delete_recipe"),
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/roots", "get_roots"),
RouteDefinition("GET", "/api/lm/recipes/folders", "get_folders"),
+36 -4
View File
@@ -804,6 +804,12 @@ class BaseModelService(ABC):
"""Get top tags sorted by frequency"""
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]:
"""Get base models sorted by frequency"""
return await self.scanner.get_base_models(limit)
@@ -955,13 +961,21 @@ class BaseModelService(ABC):
return unified_tree
async def get_model_notes(self, model_name: str) -> Optional[str]:
"""Get notes for a specific model file"""
async def get_model_notes(self, model_name: str) -> Optional[dict]:
"""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()
for model in cache.raw_data:
if model["file_name"] == model_name:
return model.get("notes", "")
file_name = model.get("file_name", "")
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
@@ -1084,6 +1098,11 @@ class BaseModelService(ABC):
Listing/search endpoints return lightweight cache entries; this method performs
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(
file_path, self.metadata_class
@@ -1101,6 +1120,19 @@ class BaseModelService(ABC):
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", {}))
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.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
async with self._hash_calculation_lock:
@@ -125,6 +132,7 @@ class CheckpointScanner(ModelScanner):
and metadata.hash_status == "completed"
and metadata.sha256
):
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
return metadata.sha256
task = self._hash_calculation_tasks.get(real_path)
@@ -175,6 +183,9 @@ class CheckpointScanner(ModelScanner):
# Check if hash is already calculated
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
# Update status to calculating
@@ -193,6 +204,20 @@ class CheckpointScanner(ModelScanner):
# Update hash index
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}")
return sha256
+14
View File
@@ -304,6 +304,20 @@ class CivArchiveClient:
version_id = file_data.get("model_version_id") or file_data.get("modelVersionId")
if model_id is None or version_id is None:
continue
# CivitAI / CivArchive model IDs are small integers (typically ≤ 7
# digits). Reject suspiciously large values that indicate the API
# returned a malformed payload (e.g. a hash reinterpreted as an ID)
# to avoid pointless HTTP 500 errors from CivArchive.
_MAX_VALID_CIVITAI_ID = 100_000_000
try:
if int(model_id) >= _MAX_VALID_CIVITAI_ID or int(version_id) >= _MAX_VALID_CIVITAI_ID:
logger.debug(
"Skipping implausible CivArchive model_id=%s / version_id=%s",
model_id, version_id,
)
continue
except (TypeError, ValueError):
continue
resolved = await self.get_model_version(model_id, version_id)
if resolved:
return resolved
+89 -43
View File
@@ -230,6 +230,12 @@ class DownloadManager:
Returns:
Dict with download result
"""
logger.debug(
"[download] download_from_civitai called: model_id=%s, model_version_id=%s, "
"source=%s, file_params=%s",
model_id, model_version_id, source, file_params,
)
# Validate that at least one identifier is provided
if not model_id and not model_version_id:
return {
@@ -250,6 +256,7 @@ class DownloadManager:
"source": source,
"file_params": copy.deepcopy(file_params) if file_params is not None else None,
"progress": 0,
"status": "queued",
"transfer_backend": self._get_model_download_backend(),
"bytes_downloaded": 0,
@@ -289,8 +296,8 @@ class DownloadManager:
return result
except asyncio.CancelledError:
return {
"success": False,
"error": "Download was cancelled",
"success": True,
"cancelled": True,
"download_id": task_id,
}
finally:
@@ -675,7 +682,10 @@ class DownloadManager:
u for u in download_urls if not u.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
]
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")
if download_url:
download_urls.append(normalize_civitai_download_url(download_url))
@@ -1379,7 +1389,17 @@ class DownloadManager:
# Update save directory with relative path if provided
if relative_path:
base_save_dir = save_dir
save_dir = os.path.join(save_dir, relative_path)
# Security: validate path containment after joining
resolved_dir = os.path.realpath(os.path.normpath(save_dir))
base_dir = os.path.realpath(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
os.makedirs(save_dir, exist_ok=True)
@@ -1421,14 +1441,35 @@ class DownloadManager:
# If file_params is provided, try to find matching file
if file_params and model_version_id:
target_file_id = file_params.get("id")
target_type = file_params.get("type", "Model")
target_format = file_params.get("format", "SafeTensor")
target_size = file_params.get("size", "full")
target_format = file_params.get("format")
target_size = file_params.get("size")
target_fp = file_params.get("fp")
is_primary = file_params.get("isPrimary", False)
if is_primary:
# Find primary file
logger.debug(
"[download] file_params received: id=%s, type=%s, format=%s, size=%s, fp=%s, isPrimary=%s, "
"model_version_id=%s, total_files=%d",
target_file_id, target_type, target_format, target_size, target_fp, is_primary,
model_version_id, len(files),
)
if target_file_id:
target_id_str = str(target_file_id)
for f in files:
f_id = f.get("id")
if str(f_id) == target_id_str:
file_info = f
logger.debug(
"[download] MATCH by ID: id=%s name='%s'",
f_id, f.get("name"),
)
break
if not file_info:
logger.debug("[download] No file found with id=%s", target_file_id)
elif is_primary:
file_info = next(
(
f
@@ -1439,28 +1480,41 @@ class DownloadManager:
None,
)
else:
# Match by metadata
# Lenient metadata match: only compare fields present on both sides
for f in files:
f_type = f.get("type", "")
f_meta = f.get("metadata", {})
# Check type match
if f_type != target_type:
continue
# Check metadata match
if f_meta.get("format") != target_format:
f_meta = f.get("metadata", {})
f_format = f_meta.get("format") or f.get("format")
f_size = f_meta.get("size") or f.get("size")
f_fp = f_meta.get("fp") or f.get("fp")
if target_format and f_format != target_format:
continue
if f_meta.get("size") != target_size:
if target_size and f_size and f_size != target_size:
continue
if target_fp and f_meta.get("fp") != target_fp:
if target_fp and f_fp and f_fp != target_fp:
continue
file_info = f
break
if not file_info:
logger.debug(
"[download] No match found via file_params — falling back to primary file lookup",
)
elif not file_params:
logger.debug(
"[download] No file_params provided (null/None) — will use primary file lookup. "
"model_version_id=%s, total_files=%d",
model_version_id, len(files),
)
# Fallback to primary file if no match found
if not file_info:
logger.debug("[download] Looking for primary file as fallback")
file_info = next(
(
f
@@ -1469,38 +1523,18 @@ class DownloadManager:
),
None,
)
if file_info:
logger.debug(
"[download] Fallback primary file selected: id=%s, name=%s",
file_info.get("id"), file_info.get("name"),
)
else:
logger.debug("[download] No primary file found in fallback lookup")
if not file_info:
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
# 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)
)
download_urls = self._build_download_urls_from_file_info(file_info, source=source)
if not download_urls:
return {"success": False, "error": "No mirror URL found"}
@@ -1803,6 +1837,9 @@ class DownloadManager:
model_tags, model_type
)
if not first_tag:
first_tag = "no tags" # Default if no tags available
# Format the template with available data
formatted_path = path_template
formatted_path = formatted_path.replace("{base_model}", mapped_base_model)
@@ -1818,6 +1855,15 @@ class DownloadManager:
if model_type == "embedding":
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
async def _execute_download(
+54 -19
View File
@@ -74,6 +74,8 @@ class DownloadQueueService:
);
CREATE INDEX IF NOT EXISTS idx_dh_completed ON download_history(completed_at DESC);
CREATE INDEX IF NOT EXISTS idx_dh_status ON download_history(status);
CREATE UNIQUE INDEX IF NOT EXISTS idx_dh_download_id
ON download_history(download_id) WHERE download_id IS NOT NULL;
"""
@classmethod
@@ -154,13 +156,23 @@ class DownloadQueueService:
"""Insert a new download into the queue.
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()
file_params_json = json.dumps(file_params) if file_params is not None else None
async with self._lock:
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(
"""
INSERT OR IGNORE INTO download_queue (
@@ -380,7 +392,7 @@ class DownloadQueueService:
)
conn.execute(
"""
INSERT INTO download_history (
INSERT OR IGNORE INTO download_history (
download_id, model_id, model_version_id, model_name,
version_name, thumbnail_url, status, error, file_path,
bytes_downloaded, total_bytes, completed_at
@@ -537,17 +549,27 @@ class DownloadQueueService:
"offset": offset,
}
async def delete_history_item(self, id: int) -> bool:
"""Delete a single history entry by its *id*.
async def delete_history_item(
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.
"""
async with self._lock:
conn = self._get_conn()
cursor = conn.execute(
"DELETE FROM download_history WHERE id = ?",
(id,),
)
if download_id:
cursor = conn.execute(
"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()
return cursor.rowcount > 0
@@ -604,21 +626,34 @@ class DownloadQueueService:
# 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.
Looks up the history record by its primary key. If the status is
``failed`` or ``canceled`` a new queue entry is created with the
same model metadata and a fresh download id, and the original
history entry is **deleted** to prevent exponential growth when
the retried item is later canceled or fails again and re-retried.
Looks up the history record by *download_id* (preferred) or
*item_id*. If the status is ``failed`` or ``canceled`` a new
queue entry is created with the same model metadata and a fresh
download id, and the original history entry is **deleted** to
prevent exponential growth when the retried item is later
canceled or fails again and re-retried.
"""
async with self._lock:
conn = self._get_conn()
row = conn.execute(
"SELECT * FROM download_history WHERE id = ?",
(item_id,),
).fetchone()
if download_id:
row = conn.execute(
"SELECT * FROM download_history WHERE download_id = ?",
(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:
return None
status = str(row["status"])
@@ -650,7 +685,7 @@ class DownloadQueueService:
)
conn.execute(
"DELETE FROM download_history WHERE id = ?",
(item_id,),
(row["id"],),
)
conn.commit()
queued = conn.execute(
+19 -10
View File
@@ -270,14 +270,14 @@ class Downloader:
Note: This is private and caller MUST hold self._session_lock.
"""
# Close existing session if any
if self._session is not None:
try:
await self._session.close()
except Exception as e: # pragma: no cover
logger.warning(f"Error closing previous session: {e}")
finally:
self._session = None
# Snapshot and clear old session reference before creating the new
# one. This ensures self._session is always valid (or None, which
# triggers a fresh creation) and avoids a race where concurrent
# requests hold a reference to a session whose connector has been
# torn down by a premature close() call — the root cause of the
# intermittent "NoneType has no attribute connect" crash.
old_session = self._session
self._session = None
# Check for app-level proxy settings
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._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(
"Created new HTTP session with proxy settings. App-level proxy: %s, System-level proxy (trust_env): %s",
bool(proxy_url),
@@ -753,7 +760,8 @@ class Downloader:
else:
resume_offset = 0
total_size = 0
await self._create_session()
async with self._session_lock:
await self._create_session()
continue
return False, integrity_error
@@ -843,7 +851,8 @@ class Downloader:
logger.info(f"Will resume from byte {resume_offset}")
# Refresh session to get new connection
await self._create_session()
async with self._session_lock:
await self._create_session()
continue
else:
logger.error(f"Max retries exceeded for download: {e}")
+7 -3
View File
@@ -271,12 +271,16 @@ class LoraService(BaseModelService):
return letters
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()
for lora in cache.raw_data:
if lora["file_name"] == lora_name:
civitai_data = lora.get("civitai", {})
file_name = lora.get("file_name", "")
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 []
+69 -16
View File
@@ -15,6 +15,17 @@ from .service_registry import ServiceRegistry
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():
"""Initialize and configure all metadata providers based on settings"""
provider_manager = await ModelMetadataProviderManager.get_instance()
@@ -26,7 +37,9 @@ async def initialize_metadata_providers():
# Get settings
settings_manager = get_settings_manager()
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 = []
# Initialize archive database provider if enabled
@@ -59,27 +72,48 @@ async def initialize_metadata_providers():
except Exception as e:
logger.error(f"Failed to initialize Civitai API metadata provider: {e}")
# Register CivArchive provider, and all add to fallback providers
try:
civarchive_client = await ServiceRegistry.get_civarchive_client()
civarchive_provider = CivArchiveModelMetadataProvider(civarchive_client)
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}")
# Register CivArchive provider when enabled. Civitai API is always
# preferred (better metadata); CivArchive mainly recovers metadata for
# models deleted from Civitai, so it can be turned off to avoid its long
# rate-limit windows entirely.
if enable_civarchive_api:
try:
civarchive_client = await ServiceRegistry.get_civarchive_client()
civarchive_provider = CivArchiveModelMetadataProvider(civarchive_client)
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
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.extend([p for p in providers if p[0] == 'civitai_api'])
ordered_providers.extend([p for p in providers if p[0] == 'civarchive_api'])
ordered_providers.extend([p for p in providers if p[0] == 'sqlite'])
for name in desired_order:
ordered_providers.extend([p for p in providers if p[0] == name])
# 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:
fallback_provider = FallbackMetadataProvider(ordered_providers)
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:
# Only one provider available, set it as default
provider_name, provider = providers[0]
@@ -96,11 +130,30 @@ async def update_metadata_providers():
# Get current settings
settings_manager = get_settings_manager()
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
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
except Exception as e:
logger.error(f"Failed to update metadata providers: {e}")
+15 -1
View File
@@ -209,7 +209,21 @@ class MetadataSyncService:
error_msg = "CivitAI model is deleted and no archive provider is available"
return False, error_msg
else:
provider_attempts.append((None, await self._get_default_provider()))
is_hf_source = bool(model_data.get("hf_url"))
if is_hf_source:
# HF-sourced model: only check CivitAI API directly.
# CivArchive is almost guaranteed to have no record, and
# hitting it wastes rate-limit budget.
# Use a distinct provider name ("civitai_api" not None) so
# downstream code does NOT interpret a "Model not found"
# response as civitai_api_not_found — which would mark the
# model civitai_deleted=True when it was never on CivitAI.
try:
provider_attempts.append(("civitai_api", await self._get_provider("civitai_api")))
except Exception as exc: # pragma: no cover - provider resolution fault
logger.debug("Unable to resolve civitai_api provider: %s", exc)
if not provider_attempts:
provider_attempts.append((None, await self._get_default_provider()))
civitai_metadata: Optional[Dict[str, Any]] = None
metadata_provider: Optional[MetadataProviderProtocol] = None
+22 -1
View File
@@ -337,4 +337,25 @@ class ModelCache:
else:
return False # Model not found
return True
return True
async def clear_preview_by_path(self, preview_file_path: str) -> int:
"""Clear ``preview_url`` for every cached entry referencing a file path.
When a preview file has been deleted from disk, this removes its
reference from all matching cache entries so the next list-API
response returns an empty ``preview_url`` instead of a stale URL
that produces 404s.
Returns the number of entries that were updated.
"""
normalized = preview_file_path.replace("\\", "/")
cleared = 0
async with self._lock:
for item in self.raw_data:
cached_url = item.get("preview_url", "")
if cached_url.replace("\\", "/") == normalized:
item["preview_url"] = ""
item["preview_nsfw_level"] = 0
cleared += 1
return cleared
+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.constants import AUTO_ORGANIZE_BATCH_SIZE
from ..services.settings_manager import get_settings_manager
from ..services.model_lifecycle_service import _require_path_in_library_roots
logger = logging.getLogger(__name__)
@@ -493,6 +494,9 @@ class ModelMoveService:
Dictionary with move result
"""
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:
# Find the model in cache to get metadata
cache = await self.scanner.get_cached_data()
+40
View File
@@ -48,6 +48,35 @@ async def delete_model_artifacts(
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.realpath()`` to resolve symlinks before comparing,
so symlink-based escapes are also caught. 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.realpath(os.path.normpath(file_path))
for root in roots:
root_resolved = os.path.realpath(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:
"""Co-ordinate destructive and mutating model operations."""
@@ -74,6 +103,8 @@ class ModelLifecycleService:
if not file_path:
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()
cached_entry = None
@@ -182,6 +213,8 @@ class ModelLifecycleService:
if not file_path:
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 = await self._metadata_loader(metadata_path)
metadata["exclude"] = True
@@ -229,6 +262,8 @@ class ModelLifecycleService:
if not file_path:
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):
raise ValueError("Model file does not exist")
@@ -270,6 +305,9 @@ class ModelLifecycleService:
if not file_paths:
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)
async def rename_model(
@@ -280,6 +318,8 @@ class ModelLifecycleService:
if not file_path or not new_file_name:
raise ValueError("File path and new file name are required")
_require_path_in_library_roots(file_path, self._scanner, label="File path")
invalid_chars = {"/", "\\", ":", "*", "?", '"', "<", ">", "|"}
if any(char in new_file_name for char in invalid_chars):
raise ValueError("Invalid characters in file name")
+249 -2
View File
@@ -14,7 +14,7 @@ from ..utils.metadata_manager import MetadataManager
from ..utils.civitai_utils import resolve_license_info
from .model_cache import ModelCache
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 .websocket_manager import ws_manager
from .persistent_model_cache import get_persistent_cache
@@ -227,6 +227,11 @@ class ModelScanner:
entry: Dict[str, Any] = {
'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 '',
'model_name': get_value('model_name', '') or '',
'folder': normalized_folder,
@@ -1389,6 +1394,9 @@ class ModelScanner:
base_name = os.path.splitext(os.path.basename(source_path))[0]
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)
@@ -1561,6 +1569,218 @@ class ModelScanner:
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:
"""Check if a model with given hash exists"""
return self._hash_index.has_hash(sha256.lower())
@@ -1613,7 +1833,32 @@ class ModelScanner:
if limit == 0:
return sorted_tags
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]]:
"""Get base models sorted by count. If limit is 0, return all."""
cache = await self.get_cached_data()
@@ -1729,6 +1974,8 @@ class ModelScanner:
break
try:
_require_path_in_library_roots(file_path, self, label="File path")
target_dir = os.path.dirname(file_path)
base_name = os.path.basename(file_path)
file_name, main_extension = os.path.splitext(base_name)
+89
View File
@@ -587,6 +587,95 @@ class PersistentModelCache:
placeholders = ", ".join(["?"] * len(self._MODEL_COLUMNS))
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]]:
tag_rows = conn.execute(
"SELECT file_path, tag FROM model_tags WHERE model_type = ?",
+6 -2
View File
@@ -1,7 +1,6 @@
import asyncio
from typing import Iterable, List, Dict, Optional
from dataclasses import dataclass, field
from operator import itemgetter
from natsort import natsorted
@@ -149,5 +148,10 @@ class RecipeCache:
)
if not name_only:
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,
)
+2 -1
View File
@@ -216,11 +216,12 @@ class RecipePersistenceService:
"preview_nsfw_level",
"favorite",
"gen_params",
"base_model",
)
if not any(key in updates for key in allowed_fields):
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):
+56 -1
View File
@@ -65,6 +65,8 @@ DEFAULT_SETTINGS: Dict[str, Any] = {
"onboarding_completed": False,
"dismissed_banners": [],
"enable_metadata_archive_db": False,
"enable_civarchive_api": True,
"metadata_provider_order": "civitai_archive_sqlite",
"proxy_enabled": False,
"proxy_host": "",
"proxy_port": "",
@@ -152,6 +154,11 @@ class SettingsManager:
self._check_environment_variables()
self._collect_configuration_warnings()
if os.environ.get("LORA_MANAGER_PORTABLE", "0") == "1":
if not self.settings.get("use_portable_settings"):
self.settings["use_portable_settings"] = True
self._save_settings()
if self._needs_initial_save:
self._save_settings()
self._needs_initial_save = False
@@ -625,12 +632,37 @@ class SettingsManager:
return False
@staticmethod
def _normalize_path_set(paths: Iterable[str]) -> set[str]:
"""Normalize an iterable of paths for set-based overlap comparison.
Resolves symlinks via ``os.path.realpath`` when the path exists on disk,
then applies ``os.path.normcase`` + ``os.path.normpath`` for consistent
cross-platform comparison. Non-string / empty entries are skipped.
"""
result: set[str] = set()
for p in paths:
if not isinstance(p, str):
continue
stripped = p.strip()
if not stripped:
continue
if os.path.exists(stripped):
stripped = os.path.normpath(os.path.realpath(stripped))
result.add(os.path.normcase(stripped))
return result
def _validate_folder_paths(
self,
library_name: str,
folder_paths: Mapping[str, Iterable[str]],
) -> None:
"""Ensure folder paths do not overlap with other libraries."""
"""Ensure folder paths do not overlap with other libraries.
Also detects checkpoints unet path overlap within the same library
(including via symlink resolution), which is a configuration error since
these model types must use separate physical folders.
"""
libraries = self.settings.get("libraries", {})
normalized_new: Dict[str, Dict[str, str]] = {}
for key, values in folder_paths.items():
@@ -668,6 +700,22 @@ class SettingsManager:
f"Folder path(s) {collisions} already assigned to library '{other_name}'"
)
# Checkpoints ↔ unet overlap within the same library
ckpt_paths = folder_paths.get("checkpoints", []) or []
unet_paths = folder_paths.get("unet", []) or []
if ckpt_paths and unet_paths:
ckpt_real = self._normalize_path_set(ckpt_paths)
unet_real = self._normalize_path_set(unet_paths)
overlap = ckpt_real & unet_real
if overlap:
collisions = ", ".join(sorted(overlap))
raise ValueError(
f"Path(s) {collisions} are configured for both "
f"'checkpoints' and 'unet' (diffusion models). "
f"These model types must use separate physical folders. "
f"Please remove one of the conflicting entries."
)
def _update_active_library_entry(
self,
*,
@@ -1542,8 +1590,12 @@ class SettingsManager:
portable_switch_pending = True
self._prepare_portable_switch(value)
if key == "folder_paths" and isinstance(value, Mapping):
active_name = self.get_active_library_name()
self._validate_folder_paths(active_name, value)
self._update_active_library_entry(folder_paths=value) # type: ignore[arg-type]
elif key == "extra_folder_paths" and isinstance(value, Mapping):
active_name = self.get_active_library_name()
self._validate_folder_paths(active_name, value)
self._update_active_library_entry(extra_folder_paths=value) # type: ignore[arg-type]
elif key == "default_lora_root":
self._update_active_library_entry(default_lora_root=str(value))
@@ -1797,6 +1849,9 @@ class SettingsManager:
if key in self.settings:
minimal[key] = copy.deepcopy(self.settings[key])
if self.settings.get("use_portable_settings"):
minimal["use_portable_settings"] = True
if self._seed_template:
for key, value in self._seed_template.items():
minimal.setdefault(key, copy.deepcopy(value))
@@ -51,6 +51,10 @@ class BulkMetadataRefreshUseCase:
if not model.get("skip_metadata_refresh", False)
and not self._is_in_skip_path(model.get("folder", ""), skip_paths)
and (not model.get("civitai") or not model["civitai"].get("id"))
# Skip models downloaded from Hugging Face — they are not on
# CivitAI / CivArchive. Users can still refresh them individually
# via the right-click context menu.
and not model.get("hf_url", "")
and not (
# Skip models confirmed not on CivitAI when no need to retry
model.get("from_civitai") is False
+1
View File
@@ -12,6 +12,7 @@ NODE_TYPES = {
"Lora Loader (LoraManager)": 1,
"Lora Stacker (LoraManager)": 2,
"WanVideo Lora Select (LoraManager)": 3,
"Create Hook LoRA (LoraManager)": 4,
}
# 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,
)
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
+6 -1
View File
@@ -12,6 +12,7 @@ from platformdirs import user_config_dir
APP_NAME = "ComfyUI-LoRA-Manager"
_LM_PORTABLE_ENV = "LORA_MANAGER_PORTABLE"
_LOGGER = logging.getLogger(__name__)
@@ -100,7 +101,11 @@ def ensure_settings_file(logger: Optional[logging.Logger] = None) -> str:
def _should_use_portable_settings(path: str, logger: logging.Logger) -> bool:
"""Return ``True`` when the repository settings file enables portable mode."""
"""Return ``True`` when the env var forces it or the settings file enables it."""
if os.environ.get(_LM_PORTABLE_ENV, "0") == "1":
logger.debug("Portable mode enabled via %s", _LM_PORTABLE_ENV)
return True
if not os.path.exists(path):
return False
+6
View File
@@ -488,6 +488,12 @@ def calculate_relative_path_for_model(
if model_type == "embedding":
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
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-lora-manager"
description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!"
version = "1.1.6"
version = "1.1.9"
license = {file = "LICENSE"}
dependencies = [
"aiohttp",
+4
View File
@@ -1,6 +1,10 @@
import os
import sys
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.error_middleware import api_json_error
from py.utils.settings_paths import ensure_settings_file
@@ -1562,6 +1562,29 @@ input:checked + .toggle-slider:before {
box-shadow: 0 0 0 2px rgba(var(--lora-accent-rgb, 79, 70, 229), 0.1);
}
.extra-folder-path-row .path-controls .extra-folder-path-input.has-error {
border-color: var(--lora-error);
background-color: rgba(220, 53, 69, 0.08);
background-color: rgba(from var(--lora-error) r g b / 0.08);
}
.extra-folder-path-row .path-controls .extra-folder-path-input.has-error:focus {
box-shadow: 0 0 0 2px rgba(220, 53, 69, 0.15);
box-shadow: 0 0 0 2px rgba(from var(--lora-error) r g b / 0.15);
}
.extra-folder-path-error {
color: var(--lora-error);
font-size: 0.8em;
margin-top: 4px;
line-height: 1.4;
display: none;
}
.extra-folder-path-error.visible {
display: block;
}
.extra-folder-path-row .path-controls .remove-path-btn {
width: 32px;
height: 32px;
+5
View File
@@ -274,6 +274,11 @@
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 */
.filter-tag.special-tag.active,
.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
*/
export function getApiEndpoints(modelType) {
if (!Object.values(MODEL_TYPES).includes(modelType)) {
throw new Error(`Invalid model type: ${modelType}`);
}
return {
// Base CRUD operations
list: `/api/lm/${modelType}/list`,
@@ -93,6 +89,7 @@ export function getApiEndpoints(modelType) {
// Query operations
scan: `/api/lm/${modelType}/scan`,
topTags: `/api/lm/${modelType}/top-tags`,
searchTags: `/api/lm/${modelType}/search-tags`,
baseModels: `/api/lm/${modelType}/base-models`,
roots: `/api/lm/${modelType}/roots`,
folders: `/api/lm/${modelType}/folders`,
+12
View File
@@ -112,6 +112,18 @@ export class BaseModelApiClient {
}
}
async cancelDownload(downloadId) {
try {
const response = await fetch(
`${DOWNLOAD_ENDPOINTS.cancelGet}?download_id=${encodeURIComponent(downloadId)}`
);
return await response.json();
} catch (error) {
console.error('Error cancelling download:', error);
return { success: false, error: error.message };
}
}
async loadMoreWithVirtualScroll(resetPage = false, updateFolders = false) {
const pageState = this.getPageState();
@@ -152,7 +152,9 @@ export class LoraContextMenu extends BaseContextMenu {
sendLoraToWorkflow(replaceMode) {
const card = this.currentCard;
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');
}
+1 -1
View File
@@ -358,7 +358,7 @@ class RecipeCard {
<div class="delete-preview">
${isVideo ?
`<video src="${previewUrl}" controls muted loop playsinline style="max-width: 100%;"></video>` :
`<img src="${previewUrl}" alt="${this.recipe.title}">`
`<img src="${previewUrl}" alt="${this.recipe.title}" onerror="this.onerror=null; this.src='/loras_static/images/no-preview.png'">`
}
</div>
<div class="delete-info">
+2 -2
View File
@@ -757,7 +757,7 @@ class RecipeModal {
`<video class="thumbnail-video" autoplay loop muted playsinline>
<source src="${lora.preview_url}" type="video/mp4">
</video>` :
`<img src="${lora.preview_url || '/loras_static/images/no-preview.png'}" alt="LoRA preview">`;
`<img src="${lora.preview_url || '/loras_static/images/no-preview.png'}" alt="LoRA preview" onerror="this.onerror=null; this.src='/loras_static/images/no-preview.png'">`;
let loraItemClass = 'recipe-lora-item';
if (existsLocally) {
@@ -1606,7 +1606,7 @@ class RecipeModal {
<video class="thumbnail-video" autoplay loop muted playsinline>
<source src="${previewUrl}" type="video/mp4">
</video>
` : `<img src="${previewUrl}" alt="Checkpoint preview">`;
` : `<img src="${previewUrl}" alt="Checkpoint preview" onerror="this.onerror=null; this.src='/loras_static/images/no-preview.png'">`;
const badge = existsLocally ? `
<div class="local-badge">
+1 -1
View File
@@ -643,7 +643,7 @@ export function createModelCard(model, modelType) {
<div class="card-preview ${shouldBlur ? 'blurred' : ''}">
${isVideo ?
`<video ${videoAttrs.join(' ')} style="pointer-events: none;"></video>` :
`<img src="${versionedPreviewUrl}" alt="${model.model_name}">`
`<img src="${versionedPreviewUrl}" alt="${model.model_name}" onerror="this.onerror=null; this.src='/loras_static/images/no-preview.png'">`
}
<div class="card-header">
${shouldBlur ?
@@ -432,7 +432,7 @@ function renderMediaMarkup(version) {
return `
<div class="version-media">
<img src="${escapeHtml(version.previewUrl)}" alt="${escapeHtml(version.name || 'preview')}">
<img src="${escapeHtml(version.previewUrl)}" alt="${escapeHtml(version.name || 'preview')}" onerror="this.onerror=null; this.src='/loras_static/images/no-preview.png'">
</div>
`;
}
+17 -5
View File
@@ -397,6 +397,7 @@ export class BulkManager {
const updated = {
...existing,
fileName: card.dataset.file_name ?? existing.fileName,
folder: card.dataset.folder ?? existing.folder,
usageTips: card.dataset.usage_tips ?? existing.usageTips,
modelName: card.dataset.name ?? existing.modelName,
};
@@ -494,7 +495,8 @@ export class BulkManager {
if (metadata) {
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 {
missingLoras.push(filepath);
}
@@ -537,7 +539,8 @@ export class BulkManager {
if (metadata) {
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 {
missingLoras.push(filepath);
}
@@ -553,7 +556,8 @@ export class BulkManager {
return;
}
await sendLoraToWorkflow(loraSyntaxes.join(', '), replaceMode, 'lora');
const exitBulkMode = () => { if (state.bulkMode) this.toggleBulkMode(); };
await sendLoraToWorkflow(loraSyntaxes.join(', '), replaceMode, 'lora', exitBulkMode);
}
async _sendAllEmbeddingsToWorkflow() {
@@ -575,7 +579,8 @@ export class BulkManager {
}
const joinedCode = embeddingCodes.join(', ');
await sendEmbeddingToWorkflow(joinedCode);
const exitBulkMode = () => { if (state.bulkMode) this.toggleBulkMode(); };
await sendEmbeddingToWorkflow(joinedCode, exitBulkMode);
}
showBulkDeleteModal() {
@@ -674,6 +679,7 @@ export class BulkManager {
const modelId = this.parseModelId(item?.civitai?.modelId);
metadataCache.set(item.file_path, {
fileName: item.file_name,
folder: item.folder || '',
usageTips: item.usage_tips || '{}',
modelName: item.name || item.file_name,
...(modelId !== null ? { modelId } : {})
@@ -1659,13 +1665,19 @@ export class BulkManager {
cancelled = true;
});
const isRecipesPage = state.currentPageType === 'recipes';
for (const filepath of state.selectedModels) {
if (cancelled) {
showToast('toast.api.operationCancelled', {}, 'info');
break;
}
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++;
} catch (error) {
errorCount++;
@@ -196,6 +196,17 @@ export class BulkMissingLoraDownloadManager {
let completedDownloads = 0;
let failedDownloads = 0;
let currentLoraProgress = 0;
let cancelled = false;
loadingManager.showCancelButton(async () => {
if (cancelled) return;
cancelled = true;
try {
await this.loraApiClient.cancelDownload(batchDownloadId);
} catch (e) {
console.error('Cancel request failed:', e);
}
});
// Set up WebSocket message handler
ws.onmessage = (event) => {
@@ -207,6 +218,11 @@ export class BulkMissingLoraDownloadManager {
return;
}
if (data.status === 'cancelled') {
cancelled = true;
return;
}
// Process progress updates
if (data.status === 'progress' && data.download_id && data.download_id.startsWith(batchDownloadId)) {
currentLoraProgress = data.progress;
@@ -249,6 +265,8 @@ export class BulkMissingLoraDownloadManager {
// Download each LoRA sequentially
for (let i = 0; i < lorasToDownload.length; i++) {
if (cancelled) break;
const lora = lorasToDownload[i];
currentLoraProgress = 0;
@@ -275,11 +293,13 @@ export class BulkMissingLoraDownloadManager {
modelId,
versionId,
loraRoot,
'', // Empty relative path, use default paths
'',
useDefaultPaths,
batchDownloadId
);
if (cancelled) break;
if (!response.success) {
console.error(`Failed to download LoRA ${lora.name || lora.file_name}: ${response.error}`);
failedDownloads++;
@@ -288,8 +308,10 @@ export class BulkMissingLoraDownloadManager {
updateProgress(100, completedDownloads, '');
}
} catch (error) {
console.error(`Error downloading LoRA ${lora.name || lora.file_name}:`, error);
failedDownloads++;
if (!cancelled) {
console.error(`Error downloading LoRA ${lora.name || lora.file_name}:`, error);
failedDownloads++;
}
}
}
@@ -300,7 +322,10 @@ export class BulkMissingLoraDownloadManager {
loadingManager.hide();
// Show completion message
if (failedDownloads === 0) {
if (cancelled) {
showToast('toast.downloads.downloadStopped', {}, 'info',
`Download cancelled. ${completedDownloads} item(s) completed.`);
} else if (failedDownloads === 0) {
showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success');
} else {
showToast('toast.loras.downloadPartialSuccess', {
+123 -23
View File
@@ -728,14 +728,23 @@ export class DownloadManager {
confirmFileSelection() {
const selectedRadio = document.querySelector('#fileSelectionList input[type="radio"]:checked');
if (!selectedRadio) return;
if (!selectedRadio) {
console.warn('[download] confirmFileSelection: no radio button checked');
return;
}
const version = this.currentVersion;
if (!version) return;
if (!version) {
console.warn('[download] confirmFileSelection: no currentVersion set');
return;
}
const modelFiles = (version.files || []).filter(f => f.type === 'Model' || f.type === 'UNet' || f.type === 'Diffusion Model');
this.selectedFile = modelFiles.find(f => f.id.toString() === selectedRadio.value);
console.log('[download] confirmFileSelection: selected file id=%s, name="%s", type="%s", metadata=%o',
this.selectedFile?.id, this.selectedFile?.name, this.selectedFile?.type, this.selectedFile?.metadata);
document.getElementById('fileSelectionStep').style.display = 'none';
document.getElementById('locationStep').style.display = 'block';
this.proceedToLocationContent();
@@ -872,16 +881,26 @@ export class DownloadManager {
const displayName = versionName || `#${versionId}`;
let ws = null;
let updateProgress = () => { };
let cancelled = false;
const downloadId = Date.now().toString();
try {
this.loadingManager.restoreProgressBar();
updateProgress = this.loadingManager.showDownloadProgress(1);
updateProgress(0, 0, displayName);
const downloadId = Date.now().toString();
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${downloadId}`);
this.loadingManager.showCancelButton(async () => {
if (cancelled) return;
cancelled = true;
try {
await this.apiClient.cancelDownload(downloadId);
} catch (e) {
console.error('Cancel request failed:', e);
}
});
ws.onmessage = event => {
const data = JSON.parse(event.data);
@@ -890,6 +909,12 @@ export class DownloadManager {
return;
}
if (data.status === 'cancelled') {
cancelled = true;
this.loadingManager.setStatus(translate('modals.download.status.cancelled', {}, 'Download cancelled'));
return;
}
if (data.status === 'progress' && data.download_id === downloadId) {
const metrics = {
bytesDownloaded: data.bytes_downloaded,
@@ -928,6 +953,10 @@ export class DownloadManager {
fileParams
);
if (cancelled) {
return false;
}
if (response?.skipped) {
this.loadingManager.setStatus(translate('modals.download.status.finalizing'));
updateProgress(100, 0, displayName);
@@ -968,8 +997,12 @@ export class DownloadManager {
return true;
} catch (error) {
console.error('Failed to download model version:', error);
showToast('toast.downloads.downloadError', { message: error?.message }, 'error');
if (cancelled) {
console.log('Download cancelled by user:', downloadId);
} else {
console.error('Failed to download model version:', error);
showToast('toast.downloads.downloadError', { message: error?.message }, 'error');
}
return false;
} finally {
try {
@@ -989,16 +1022,33 @@ export class DownloadManager {
const totalFiles = this.hfSelectedFiles.length;
const updateProgress = this.loadingManager.showDownloadProgress(totalFiles);
let cancelled = false;
let currentDownloadId = null;
this.loadingManager.showCancelButton(async () => {
if (cancelled) return;
cancelled = true;
if (currentDownloadId) {
try {
await this.apiClient.cancelDownload(currentDownloadId);
} catch (e) {
console.error('Cancel request failed:', e);
}
}
});
try {
let completedDownloads = 0;
for (let i = 0; i < totalFiles; i++) {
if (cancelled) break;
const filename = this.hfSelectedFiles[i];
updateProgress(0, completedDownloads, filename);
this.loadingManager.setStatus(`Downloading ${filename}...`);
const downloadId = Date.now().toString() + '_' + i;
currentDownloadId = Date.now().toString() + '_' + i;
const wsProtocol = window.location.protocol === 'https:' ? 'wss://' : 'ws://';
const ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${downloadId}`);
const ws = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${currentDownloadId}`);
try {
await new Promise((resolve, reject) => {
@@ -1006,12 +1056,13 @@ export class DownloadManager {
ws.onerror = reject;
});
// Capture completed count at WS creation time so progress
// updates arriving after completedDownloads increments still
// show the correct "N / total" position.
const snapshotCompleted = completedDownloads;
ws.onmessage = (event) => {
const data = JSON.parse(event.data);
if (data.status === 'cancelled') {
cancelled = true;
return;
}
if (data.status === 'progress') {
const metrics = {
bytesDownloaded: data.bytes_downloaded,
@@ -1029,9 +1080,11 @@ export class DownloadManager {
modelRoot,
relativePath: targetFolder,
useDefaultPaths,
download_id: downloadId,
download_id: currentDownloadId,
});
if (cancelled) break;
if (response?.success) {
completedDownloads++;
updateProgress(100, completedDownloads, filename);
@@ -1041,13 +1094,19 @@ export class DownloadManager {
}
}
showToast('toast.loras.downloadCompleted', {}, 'success');
// Reload page data — model is already in scanner cache via backend
if (cancelled) {
showToast('toast.downloads.downloadStopped', {}, 'info',
`Download cancelled. ${completedDownloads} item(s) completed.`);
} else {
showToast('toast.loras.downloadCompleted', {}, 'success');
}
await resetAndReload(true);
return true;
} catch (error) {
console.error('Failed to download HF model:', error);
showToast('toast.downloads.downloadError', { message: error?.message }, 'error');
if (!cancelled) {
console.error('Failed to download HF model:', error);
showToast('toast.downloads.downloadError', { message: error?.message }, 'error');
}
return false;
} finally {
this.loadingManager.hide();
@@ -1426,12 +1485,23 @@ export class DownloadManager {
}
const fileParams = this.selectedFile ? {
id: this.selectedFile.id,
type: this.selectedFile.type || 'Model',
format: this.selectedFile.metadata?.format || 'SafeTensor',
size: this.selectedFile.metadata?.size || 'full',
fp: this.selectedFile.metadata?.fp,
format: this.selectedFile.metadata?.format || null,
size: this.selectedFile.metadata?.size || null,
fp: this.selectedFile.metadata?.fp || null,
} : null;
if (fileParams) {
console.log('[download] startDownload (single): fileParams built from selectedFile — id=%s, type=%s, format=%s, size=%s, fp=%s',
fileParams.id, fileParams.type, fileParams.format, fileParams.size, fileParams.fp);
} else {
console.log('[download] startDownload (single): this.selectedFile is null — no file selection, will download primary/default file. version=%s has %d files',
this.currentVersion?.id, (this.currentVersion?.files || []).length);
}
modalManager.closeModal('downloadModal');
return this.executeDownloadWithProgress({
modelId: this.modelId,
versionId: this.currentVersion.id,
@@ -1470,11 +1540,27 @@ export class DownloadManager {
let completedDownloads = 0;
let failedDownloads = 0;
let cancelled = false;
loadingManager.showCancelButton(async () => {
if (cancelled) return;
cancelled = true;
try {
await this.apiClient.cancelDownload(batchDownloadId);
} catch (e) {
console.error('Cancel request failed:', e);
}
});
ws.onmessage = (event) => {
const data = JSON.parse(event.data);
if (data.type === 'download_id') return;
if (data.status === 'cancelled') {
cancelled = true;
return;
}
if (data.status === 'progress' && data.download_id?.startsWith(batchDownloadId)) {
const current = downloadItems[completedDownloads + failedDownloads];
const name = current?.selectedVersion?.name || current?.displayName || current?.filename || `#${completedDownloads + failedDownloads + 1}`;
@@ -1493,6 +1579,8 @@ export class DownloadManager {
});
for (let i = 0; i < downloadItems.length; i++) {
if (cancelled) break;
const item = downloadItems[i];
const name = item.displayName || item.filename || (item.selectedVersion?.name || `Model #${item.modelId}`);
const isHf = item.source === 'huggingface';
@@ -1503,7 +1591,6 @@ export class DownloadManager {
try {
let response;
if (isHf) {
// Per-file WebSocket for real-time progress
const downloadId = Date.now().toString() + '_hf_' + i;
const wsHf = new WebSocket(`${wsProtocol}${window.location.host}/ws/download-progress?id=${downloadId}`);
try {
@@ -1537,6 +1624,8 @@ export class DownloadManager {
wsHf.close();
}
} else {
console.log('[download] batch download: fileParams NOT passed for modelId=%s, versionId=%s — backend will use primary file',
item.modelId, item.selectedVersion?.id);
response = await this.apiClient.downloadModel(
item.modelId,
item.selectedVersion.id,
@@ -1548,6 +1637,8 @@ export class DownloadManager {
);
}
if (cancelled) break;
if (!response.success) {
failedDownloads++;
} else {
@@ -1555,15 +1646,20 @@ export class DownloadManager {
updateProgress(100, completedDownloads, '');
}
} catch (err) {
console.error(`Failed to download ${name}:`, err);
failedDownloads++;
if (!cancelled) {
console.error(`Failed to download ${name}:`, err);
failedDownloads++;
}
}
}
ws.close();
loadingManager.hide();
if (failedDownloads === 0) {
if (cancelled) {
showToast('toast.downloads.downloadStopped', {}, 'info',
`Download cancelled. ${completedDownloads} item(s) completed.`);
} else if (failedDownloads === 0) {
showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success');
} else {
showToast('toast.loras.downloadPartialSuccess', {
@@ -1581,6 +1677,10 @@ export class DownloadManager {
modelRoot = '',
targetFolder = ''
} = {}) {
console.warn('[download] downloadVersionWithDefaults: NO fileParams will be sent — backend will always use primary file. '
+ 'modelType=%s, modelId=%s, versionId=%s, versionName="%s"',
modelType, modelId, versionId, versionName);
try {
this.apiClient = getModelApiClient(modelType);
} catch (error) {
+123 -24
View File
@@ -1,6 +1,7 @@
import { getCurrentPageState } from '../state/index.js';
import { showToast, updatePanelPositions } from '../utils/uiHelpers.js';
import { getModelApiClient } from '../api/modelApiFactory.js';
import { getApiEndpoints } from '../api/apiConfig.js';
import { removeStorageItem, setStorageItem, getStorageItem } from '../utils/storageHelpers.js';
import { MODEL_TYPE_DISPLAY_NAMES } from '../utils/constants.js';
import { translate } from '../utils/i18nHelpers.js';
@@ -24,6 +25,12 @@ export class FilterManager {
this.baseModelOptions = [];
this.tagsLoaded = false;
// Tag search state
this.modelTagsSearchInput = document.getElementById('modelTagsSearchInput');
this.tagSearchDebounceTimer = null;
this.tagSearchAbortController = null;
this.tagSearchQuery = '';
// Initialize preset manager
this.presetManager = new FilterPresetManager({
page: this.currentPage,
@@ -123,6 +130,60 @@ export class FilterManager {
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) {
@@ -146,15 +207,24 @@ export class FilterManager {
}
async loadTopTags() {
// Abort any in-flight tag search request
if (this.tagSearchAbortController) {
this.tagSearchAbortController.abort();
this.tagSearchAbortController = null;
}
this.tagSearchQuery = '';
try {
// Show loading state
const tagsContainer = document.getElementById('modelTagsFilter');
const emptyState = document.getElementById('modelTagsEmptyState');
if (!tagsContainer) return;
if (emptyState) emptyState.hidden = true;
tagsContainer.innerHTML = '<div class="tags-loading">Loading tags...</div>';
// 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);
if (!response.ok) throw new Error('Failed to fetch tags');
@@ -179,29 +249,38 @@ export class FilterManager {
createTagFilterElements(tags) {
const tagsContainer = document.getElementById('modelTagsFilter');
const emptyState = document.getElementById('modelTagsEmptyState');
if (!tagsContainer) return;
tagsContainer.innerHTML = '';
if (emptyState) emptyState.hidden = true;
// Collect existing tag names from the API response
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) {
Object.keys(this.filters.tags).forEach(tagName => {
// Skip special tags like __no_tags__
if (tagName.startsWith('__')) return;
if (!existingTagNames.has(tagName)) {
// Add this tag to the list with count 0 (unknown)
tags.push({ tag: tagName, count: 0 });
missingSelectedTags.push({ tag: tagName, count: 0 });
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) {
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;
}
@@ -209,6 +288,10 @@ export class FilterManager {
const tagEl = document.createElement('div');
tagEl.className = 'filter-tag tag-filter';
const tagName = tag.tag;
if (missingSelectedTags.some(t => t.tag === tagName)) {
tagEl.classList.add('extra-tag');
}
tagEl.dataset.tag = tagName;
// Show count only if it's > 0 (known count)
@@ -234,26 +317,28 @@ export class FilterManager {
tagsContainer.appendChild(tagEl);
});
// Add "No tags" as a special filter at the end
const noTagsEl = document.createElement('div');
noTagsEl.className = 'filter-tag tag-filter special-tag';
const noTagsLabel = translate('header.filter.noTags', {}, 'No tags');
const noTagsKey = '__no_tags__';
noTagsEl.dataset.tag = noTagsKey;
noTagsEl.innerHTML = noTagsLabel;
// Add "No tags" as a special filter at the end (skip during search)
if (!this.tagSearchQuery) {
const noTagsEl = document.createElement('div');
noTagsEl.className = 'filter-tag tag-filter special-tag';
const noTagsLabel = translate('header.filter.noTags', {}, 'No tags');
const noTagsKey = '__no_tags__';
noTagsEl.dataset.tag = noTagsKey;
noTagsEl.innerHTML = noTagsLabel;
noTagsEl.addEventListener('click', async () => {
const currentState = (this.filters.tags && this.filters.tags[noTagsKey]) || 'none';
const newState = this.getNextTriStateState(currentState);
this.setTagFilterState(noTagsKey, newState);
this.applyTagElementState(noTagsEl, newState);
noTagsEl.addEventListener('click', async () => {
const currentState = (this.filters.tags && this.filters.tags[noTagsKey]) || 'none';
const newState = this.getNextTriStateState(currentState);
this.setTagFilterState(noTagsKey, 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();
}
@@ -341,7 +426,7 @@ export class FilterManager {
if (!baseModelTagsContainer) return;
// 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(apiEndpoint)
@@ -721,6 +806,16 @@ export class FilterManager {
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
this.updateTagLogicToggleUI();
@@ -731,6 +826,10 @@ export class FilterManager {
// Update UI
this.updateTagSelections();
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
// Remove from local Storage
+13 -19
View File
@@ -478,11 +478,9 @@ export class FilterPresetManager {
const pageState = getCurrentPageState();
pageState.filters = this.filterManager.cloneFilters();
// If tags haven't been loaded yet, load them first
if (!this.filterManager.tagsLoaded) {
await this.filterManager.loadTopTags();
this.filterManager.tagsLoaded = true;
}
// Refresh tag display so preset's non-top-20 tags appear inline
await this.filterManager.loadTopTags();
this.filterManager.tagsLoaded = true;
// Check again after async operation
if (requestId !== this.applyPresetRequestId) return;
@@ -745,8 +743,16 @@ export class FilterPresetManager {
presetEl.classList.add('active');
}
presetEl.addEventListener('click', (e) => {
e.stopPropagation();
// Apply preset on click (toggle if already active)
// 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');
@@ -759,18 +765,6 @@ export class FilterPresetManager {
deleteBtn.innerHTML = '<i class="fas fa-times"></i>';
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
deleteBtn.addEventListener('click', (e) => {
e.stopPropagation();
+4
View File
@@ -281,6 +281,10 @@ export class LoadingManager {
// Initialize transfer stats with empty data
updateTransferStats();
if (this.cancelButton) {
this.loadingContent.appendChild(this.cancelButton);
}
// Return update function
return (currentProgress, currentIndex = 0, currentName = '', metrics = {}) => {
// Update current item progress
+95 -1
View File
@@ -1693,13 +1693,15 @@ export class SettingsManager {
<input type="text" class="extra-folder-path-input"
placeholder="${translate('settings.extraFolderPaths.pathPlaceholder', {}, '/path/to/models')}" value="${path}"
onblur="settingsManager.updateExtraFolderPaths('${modelType}')"
onfocus="settingsManager.clearExtraFolderPathError(this)"
onkeydown="if(event.key === 'Enter') { this.blur(); }" />
<button type="button" class="remove-path-btn"
onclick="this.parentElement.parentElement.remove(); settingsManager.updateExtraFolderPaths('${modelType}')"
onclick="settingsManager.removeExtraFolderPathRow(this, '${modelType}')"
title="${translate('common.actions.delete', {}, 'Delete')}">
<i class="fas fa-times"></i>
</button>
</div>
<div class="extra-folder-path-error"></div>
`;
container.appendChild(row);
@@ -1713,7 +1715,63 @@ export class SettingsManager {
}
}
clearExtraFolderPathError(input) {
input.classList.remove('has-error');
const row = input.closest('.extra-folder-path-row');
if (row) {
const errEl = row.querySelector('.extra-folder-path-error');
if (errEl) {
errEl.classList.remove('visible');
errEl.textContent = '';
}
}
}
_clearAllExtraFolderPathErrors() {
document.querySelectorAll('.extra-folder-path-input.has-error').forEach((input) => {
input.classList.remove('has-error');
});
document.querySelectorAll('.extra-folder-path-error.visible').forEach((el) => {
el.classList.remove('visible');
el.textContent = '';
});
}
_markExtraFolderPathsError(modelType, overlappingPaths, showMessage = false) {
const container = document.getElementById(`extraFolderPaths-${modelType}`);
if (!container) return;
const inputs = container.querySelectorAll('.extra-folder-path-input');
inputs.forEach((input) => {
const val = input.value.trim();
if (val && overlappingPaths.includes(val)) {
input.classList.add('has-error');
if (showMessage) {
const row = input.closest('.extra-folder-path-row');
if (row) {
const errEl = row.querySelector('.extra-folder-path-error');
if (errEl) {
errEl.textContent = translate('settings.extraFolderPaths.validation.checkpointUnetOverlapInline', {}, 'This path is also used for a different model type. Use separate folders for checkpoints and diffusion models.');
errEl.classList.add('visible');
}
}
}
}
});
}
removeExtraFolderPathRow(btn, modelType) {
const row = btn.closest('.extra-folder-path-row');
if (row) {
row.remove();
this.updateExtraFolderPaths(modelType);
}
}
async updateExtraFolderPaths(changedModelType) {
// Clear previous errors
this._clearAllExtraFolderPathErrors();
const extraFolderPaths = {};
// Collect paths for all model types
@@ -1734,6 +1792,32 @@ export class SettingsManager {
extraFolderPaths[modelType] = paths;
});
// Client-side pre-check: checkpoints and unet must not share the same path.
// Normalise paths to reduce false negatives vs the backend's realpath + normcase.
const normalise = (p) => p.replace(/[/\\]+$/, '').toLowerCase();
const ckptSet = new Set((extraFolderPaths.checkpoints || []).map(normalise));
const unetSet = new Set((extraFolderPaths.unet || []).map(normalise));
const ckptOverlap = (extraFolderPaths.checkpoints || []).filter(p => p && unetSet.has(normalise(p)));
const unetOverlap = (extraFolderPaths.unet || []).filter(p => p && ckptSet.has(normalise(p)));
const hasOverlap = ckptOverlap.length > 0 || unetOverlap.length > 0;
if (hasOverlap) {
// Error message only on the side the user just edited.
// The other side gets red border only (passive conflict indicator).
if (changedModelType === 'checkpoints') {
this._markExtraFolderPathsError('checkpoints', ckptOverlap, true);
this._markExtraFolderPathsError('unet', unetOverlap, false);
} else if (changedModelType === 'unet') {
this._markExtraFolderPathsError('unet', unetOverlap, true);
this._markExtraFolderPathsError('checkpoints', ckptOverlap, false);
} else {
// Pre-existing conflict from direct config edit — mark both without messages
this._markExtraFolderPathsError('checkpoints', ckptOverlap, false);
this._markExtraFolderPathsError('unet', unetOverlap, false);
}
return;
}
// Check if paths have actually changed
const currentPaths = state.global.settings.extra_folder_paths || {};
const pathsChanged = JSON.stringify(currentPaths) !== JSON.stringify(extraFolderPaths);
@@ -2262,6 +2346,16 @@ export class SettingsManager {
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
await this.updateMetadataArchiveStatus();
} catch (error) {
+29 -8
View File
@@ -168,6 +168,18 @@ export class DownloadManager {
let failedDownloads = 0;
let accessFailures = 0;
let currentLoraProgress = 0;
let cancelled = false;
this.importManager.loadingManager.showCancelButton(async () => {
if (cancelled) return;
cancelled = true;
try {
const loraClient = getModelApiClient(MODEL_TYPES.LORA);
await loraClient.cancelDownload(batchDownloadId);
} catch (e) {
console.error('Cancel request failed:', e);
}
});
// Set up progress tracking for current download
ws.onmessage = (event) => {
@@ -179,6 +191,11 @@ export class DownloadManager {
return;
}
if (data.status === 'cancelled') {
cancelled = true;
return;
}
// Process progress updates for our current active download
if (data.status === 'progress' && data.download_id && data.download_id.startsWith(batchDownloadId)) {
// Update current LoRA progress
@@ -221,6 +238,8 @@ export class DownloadManager {
const useDefaultPaths = getStorageItem('use_default_path_loras', false);
for (let i = 0; i < this.importManager.downloadableLoRAs.length; i++) {
if (cancelled) break;
const lora = this.importManager.downloadableLoRAs[i];
// Reset current LoRA progress for new download
@@ -241,15 +260,13 @@ export class DownloadManager {
batchDownloadId
);
if (cancelled) break;
if (!response.success) {
console.error(`Failed to download LoRA ${lora.name}: ${response.error}`);
failedDownloads++;
// Continue with next download
} else {
completedDownloads++;
// Update progress to show completion of current LoRA
updateProgress(100, completedDownloads, '');
if (completedDownloads + failedDownloads < this.importManager.downloadableLoRAs.length) {
@@ -259,9 +276,10 @@ export class DownloadManager {
}
}
} catch (downloadError) {
console.error(`Error downloading LoRA ${lora.name}:`, downloadError);
failedDownloads++;
// Continue with next download
if (!cancelled) {
console.error(`Error downloading LoRA ${lora.name}:`, downloadError);
failedDownloads++;
}
}
}
@@ -269,7 +287,10 @@ export class DownloadManager {
ws.close();
// Show appropriate completion message based on results
if (failedDownloads === 0) {
if (cancelled) {
showToast('toast.downloads.downloadStopped', {}, 'info',
`Download cancelled. ${completedDownloads} item(s) completed.`);
} else if (failedDownloads === 0) {
showToast('toast.loras.allDownloadSuccessful', { count: completedDownloads }, 'success');
} else {
if (accessFailures > 0) {
+2
View File
@@ -13,6 +13,8 @@ const DEFAULT_SETTINGS_BASE = Object.freeze({
language: 'en',
show_only_sfw: false,
enable_metadata_archive_db: false,
enable_civarchive_api: true,
metadata_provider_order: 'civitai_archive_sqlite',
proxy_enabled: false,
proxy_type: 'http',
proxy_host: '',
+6 -3
View File
@@ -369,21 +369,24 @@ export function getMatureBlurThreshold(settings = {}) {
export const NODE_TYPES = {
LORA_LOADER: 1,
LORA_STACKER: 2,
WAN_VIDEO_LORA_SELECT: 3
WAN_VIDEO_LORA_SELECT: 3,
HOOK_LORA: 4
};
// Node type names to IDs mapping
export const NODE_TYPE_NAMES = {
"Lora Loader (LoraManager)": NODE_TYPES.LORA_LOADER,
"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
export const NODE_TYPE_ICONS = {
[NODE_TYPES.LORA_LOADER]: "fas fa-l",
[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
+40 -5
View File
@@ -141,6 +141,20 @@ const PARAM_TO_WIDGET_CANDIDATES = {
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)
// 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
// Returns array of { widgetName, value } objects
// ---------------------------------------------------------------------------
function findMatchingWidgets(nodeWidgetNames, resolvedParams) {
function findMatchingWidgets(nodeWidgetNames, resolvedParams, nodeType) {
if (!nodeWidgetNames || !Array.isArray(nodeWidgetNames) || nodeWidgetNames.length === 0) {
return [];
}
@@ -243,6 +257,26 @@ function findMatchingWidgets(nodeWidgetNames, resolvedParams) {
const widgetSet = new Set(nodeWidgetNames.map(w => String(w).toLowerCase()));
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
const simpleParams = [
{ key: 'seed', value: resolvedParams.seed },
@@ -251,10 +285,10 @@ function findMatchingWidgets(nodeWidgetNames, resolvedParams) {
];
for (const { key, value } of simpleParams) {
if (value === undefined || value === null || value === '') continue;
const candidates = PARAM_TO_WIDGET_CANDIDATES[key] || [key];
const candidates = getCandidates(key);
for (const candidate of candidates) {
if (widgetSet.has(candidate.toLowerCase())) {
updates.push({ widgetName: candidate, value: String(value) });
updates.push({ widgetName: candidate, value });
break;
}
}
@@ -262,7 +296,7 @@ function findMatchingWidgets(nodeWidgetNames, resolvedParams) {
// Sampler
if (resolvedParams.sampler) {
const candidates = PARAM_TO_WIDGET_CANDIDATES.sampler;
const candidates = getCandidates('sampler');
for (const candidate of candidates) {
if (widgetSet.has(candidate.toLowerCase())) {
updates.push({ widgetName: candidate, value: resolvedParams.sampler });
@@ -273,7 +307,7 @@ function findMatchingWidgets(nodeWidgetNames, resolvedParams) {
// Scheduler
if (resolvedParams.scheduler) {
const candidates = PARAM_TO_WIDGET_CANDIDATES.scheduler;
const candidates = getCandidates('scheduler');
for (const candidate of candidates) {
if (widgetSet.has(candidate.toLowerCase())) {
updates.push({ widgetName: candidate, value: resolvedParams.scheduler });
@@ -290,6 +324,7 @@ export {
SCHEDULER_SUFFIXES,
SCHEDULER_ONLY_VALUES,
PARAM_TO_WIDGET_CANDIDATES,
NODE_TYPE_WIDGET_OVERRIDES,
parseCombinedSamplerName,
resolveSamplerScheduler,
findMatchingWidgets,
+24 -11
View File
@@ -134,7 +134,10 @@ export async function copyToClipboard(text, successMessage = 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');
toast.className = `toast toast-${type}`;
toast.textContent = message;
@@ -552,6 +555,8 @@ async function fetchWorkflowRegistry() {
if (!registryData.success) {
if (registryData.error === 'Standalone Mode Active') {
showToast('toast.general.cannotInteractStandalone', {}, 'warning');
} else if (registryData.error === 'Empty Registry') {
showToast('uiHelpers.workflow.noSupportedNodes', {}, 'warning');
} else {
showToast('toast.general.failedWorkflowInfo', {}, 'error');
}
@@ -603,7 +608,7 @@ function isNodeEnabled(node) {
if (!node) {
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;
}
@@ -654,7 +659,7 @@ async function ensureRelativeModelPath(modelPath, collectionType) {
* @param {string} syntaxType - The type of syntax ('lora' or 'recipe')
* @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();
if (!registry) {
return false;
@@ -679,7 +684,9 @@ export async function sendLoraToWorkflow(loraSyntax, replaceMode = false, syntax
}
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 =
@@ -693,8 +700,11 @@ export async function sendLoraToWorkflow(loraSyntax, replaceMode = false, syntax
showNodeSelector(loraNodes, {
actionType,
actionMode,
onSend: (selectedNodeIds) =>
sendLoraToNodes(selectedNodeIds, loraNodes, loraSyntax, replaceMode, syntaxType),
onSend: async (selectedNodeIds) => {
const result = await sendLoraToNodes(selectedNodeIds, loraNodes, loraSyntax, replaceMode, syntaxType);
if (result && typeof onComplete === 'function') onComplete();
return result;
},
});
return true;
}
@@ -965,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();
if (!registry) {
return false;
@@ -993,8 +1003,11 @@ export async function sendEmbeddingToWorkflow(embeddingCode) {
missingTargetMessage: translate('uiHelpers.workflow.noTargetNodeSelected', {}, 'No target node selected'),
};
const handleSend = (selectedNodeIds) =>
sendTextToNodes(selectedNodeIds, textNodes, embeddingCode, 'append', messages);
const handleSend = async (selectedNodeIds) => {
const result = await sendTextToNodes(selectedNodeIds, textNodes, embeddingCode, 'append', messages);
if (result && typeof onComplete === 'function') onComplete();
return result;
};
if (nodeKeys.length === 1) {
return await handleSend([nodeKeys[0]]);
@@ -1134,8 +1147,8 @@ export async function sendGenParamsToWorkflow(genParams) {
const node = targetNodes[nodeKey];
if (!node) continue;
const widgetNames = node.widget_names || [];
const updates = findMatchingWidgets(widgetNames, raw);
const widgetNames = getWidgetNames(node);
const updates = findMatchingWidgets(widgetNames, raw, node.type_name);
if (updates.length === 0) {
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>
</div>
</div>
<input type="text" id="modelTagsSearchInput" class="filter-search-input"
placeholder="{{ t('header.filter.tagSearchPlaceholder') }}" autocomplete="off">
<div class="filter-tags" id="modelTagsFilter">
<!-- Top tags will be dynamically inserted here -->
<div class="tags-loading">{{ t('common.status.loading') }}</div>
</div>
<div id="modelTagsEmptyState" class="filter-empty-state" hidden>
{{ t('header.filter.noTagMatches') }}
</div>
</div>
{% if current_page == 'loras' or current_page == 'checkpoints' %}
<div class="filter-section">
@@ -112,6 +112,10 @@
<a href="https://github.com/willmiao/ComfyUI-Lora-Manager/wiki/Priority-Tags-Configuration-Guide" target="_blank">
Priority Tags Configuration Guide
<span class="new-content-badge inline">{{ t('help.documentation.newBadge') }}</span>
<li>
<a href="https://github.com/willmiao/ComfyUI-Lora-Manager/wiki/AI-Provider-Setup" target="_blank">
AI Provider Setup
<span class="new-content-badge inline">{{ t('help.documentation.newBadge') }}</span>
</a>
</li>
</ul>
+80 -43
View File
@@ -144,6 +144,46 @@
</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) -->
<div class="settings-subsection">
<div class="settings-subsection-header">
@@ -250,46 +290,6 @@
{{ provider_models_json | safe }}
</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 -->
<div class="settings-subsection">
<div class="settings-subsection-header">
@@ -1401,7 +1401,26 @@
<div class="settings-input-error-message" id="metadataRefreshSkipPathsError"></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-row">
<div class="setting-info">
@@ -1419,13 +1438,13 @@
</div>
</div>
</div>
<div class="setting-item">
<div class="metadata-archive-status" id="metadataArchiveStatus">
<!-- Status will be populated by JavaScript -->
</div>
</div>
<div class="setting-item">
<div class="setting-row">
<div class="setting-info">
@@ -1444,6 +1463,24 @@
</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>
+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()
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.model_management'] = comfy_mock.model_management
sys.modules['comfy.comfy_types'] = comfy_mock.comfy_types
sys.modules['comfy.hooks'] = MockModule("comfy.hooks")
execution_mock = MockModule("execution")
execution_mock.PromptExecutor = mock.MagicMock()
@@ -113,6 +113,8 @@ function renderControlsDom(pageKey) {
<div id="baseModelEmptyState" hidden></div>
<div id="filterPresets" class="filter-presets"></div>
<div id="modelTagsFilter" class="filter-tags"></div>
<input id="modelTagsSearchInput" />
<div id="modelTagsEmptyState" hidden></div>
<button class="clear-filter"></button>
</div>
<div class="controls">
@@ -961,4 +963,198 @@ describe('PageControls favorites, sorting, and duplicates scenarios', () => {
expect(stateModule.state.bulkMode).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,
resolveSamplerScheduler,
findMatchingWidgets,
NODE_TYPE_WIDGET_OVERRIDES,
} from '../../../static/js/utils/genParamsMapper.js';
// ---------------------------------------------------------------------------
@@ -204,9 +205,9 @@ describe('findMatchingWidgets', () => {
it('matches seed to seed widget', () => {
const updates = findMatchingWidgets(['seed', 'steps', 'cfg', 'sampler_name', 'scheduler'], resolved);
expect(updates).toContainEqual({ widgetName: 'seed', value: '42' });
expect(updates).toContainEqual({ widgetName: 'steps', value: '30' });
expect(updates).toContainEqual({ widgetName: 'cfg', value: '7' });
expect(updates).toContainEqual({ widgetName: 'seed', value: 42 });
expect(updates).toContainEqual({ widgetName: 'steps', value: 30 });
expect(updates).toContainEqual({ widgetName: 'cfg', value: 7 });
expect(updates).toContainEqual({ widgetName: 'sampler_name', value: 'euler_ancestral' });
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 seedUpdate = updates.find(u => u.widgetName === 'noise_seed');
expect(seedUpdate).toBeDefined();
expect(seedUpdate.value).toBe('42');
expect(seedUpdate.value).toBe(42);
});
it('matches rgthree-style sampler widget name', () => {
@@ -243,4 +244,53 @@ describe('findMatchingWidgets', () => {
const updates = findMatchingWidgets(['seed', 'steps', 'cfg', 'sampler_name', 'scheduler'], resolved);
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 });
});
});
+48
View File
@@ -728,6 +728,54 @@ async def test_register_nodes_includes_capabilities():
assert stored_node["widget_names"] == ["ckpt_name"]
@pytest.mark.asyncio
async def test_register_nodes_accepts_compound_node_ids():
"""Subgraph nodes from expanded group nodes have compound IDs like '252:0'."""
node_registry = NodeRegistry()
handler = NodeRegistryHandler(
node_registry=node_registry,
prompt_server=FakePromptServer,
standalone_mode=False,
)
request = FakeRequest(
json_data={
"nodes": [
{
"node_id": "252:0",
"graph_id": "252",
"type": "CheckpointLoaderSimple",
"title": "Checkpoint Loader (subgraph)",
},
{
"node_id": "252:1",
"graph_id": "252",
"type": "CLIPLoader",
"title": "CLIP Loader (subgraph)",
},
],
"client_id": "test-client-1",
}
)
response = await handler.register_nodes(request)
payload = json.loads(response.text)
assert response.status == 200
assert payload["success"] is True
assert "2 nodes registered" in payload["message"]
registry = await node_registry.get_merged_registry()
assert registry["node_count"] == 2
nodes_map = registry["nodes"]
assert "252:0" in nodes_map
assert "252:1" in nodes_map
assert nodes_map["252:0"]["id"] == 0
assert nodes_map["252:0"]["graph_id"] == "252"
assert nodes_map["252:1"]["id"] == 1
@pytest.mark.asyncio
async def test_update_node_widget_sends_payload():
send_calls: list[tuple[str, dict]] = []
+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"}))
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
@@ -1189,6 +1189,65 @@ def test_relative_path_sanitizes_model_and_version_placeholders():
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_distribute_preview_to_entries_moves_and_copies(tmp_path):
"""Test that preview distribution moves file to first entry and copies to others."""
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
+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()
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"
+154 -1
View File
@@ -3,11 +3,164 @@ from pathlib import Path
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.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_rejects_symlink_escape(self, tmp_path):
root = tmp_path / "loras"
root.mkdir()
model = root / "model.safetensors"
model.write_text("")
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)])
with pytest.raises(ValueError, match="outside configured library"):
_require_path_in_library_roots(str(symlink), 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:
def __init__(self, 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
for record in caplog.records:
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['civitai_deleted'] is True
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
+56
View File
@@ -0,0 +1,56 @@
"""Tests for settings path resolution."""
import json
import logging
import os
import pytest
from py.utils.settings_paths import _should_use_portable_settings
class TestShouldUsePortableSettings:
"""Tests for _should_use_portable_settings()."""
@pytest.mark.parametrize(
"env_value, settings_flag, expected",
[
("1", False, True), # env = 1 overrides settings.json false
("1", True, True), # env = 1 matches settings.json true
("0", False, False), # env = 0 → rely on settings.json
("0", True, True), # env = 0 → rely on settings.json
("", False, False), # unset → rely on settings.json
("", True, True), # unset → rely on settings.json
],
)
def test_env_var_overrides_settings(self, tmp_path, env_value, settings_flag, expected):
"""The LORA_MANAGER_PORTABLE env var takes precedence over settings.json."""
settings_file = tmp_path / "settings.json"
settings_file.write_text(
json.dumps({"use_portable_settings": settings_flag})
)
with pytest.MonkeyPatch.context() as mp:
if env_value:
mp.setenv("LORA_MANAGER_PORTABLE", env_value)
else:
mp.delenv("LORA_MANAGER_PORTABLE", raising=False)
result = _should_use_portable_settings(str(settings_file), logging.getLogger())
assert result == expected
def test_missing_file_without_env(self, tmp_path):
"""Without env var, missing settings file returns False."""
missing = tmp_path / "nonexistent.json"
result = _should_use_portable_settings(str(missing), logging.getLogger())
assert result is False
def test_missing_file_with_env(self, tmp_path):
"""With env var, even a missing settings file returns True."""
missing = tmp_path / "nonexistent.json"
with pytest.MonkeyPatch.context() as mp:
mp.setenv("LORA_MANAGER_PORTABLE", "1")
result = _should_use_portable_settings(str(missing), logging.getLogger())
assert result is True
+32
View File
@@ -114,6 +114,38 @@ def test_calculate_relative_path_sanitizes_model_and_version_names(isolated_sett
assert relative_path == "Fancy_Model/Version_One"
def test_calculate_relative_path_sanitizes_leading_slash(isolated_settings):
"""Test that empty base_model does NOT produce a leading slash in the path."""
isolated_settings["download_path_templates"]["lora"] = "{base_model}/{first_tag}"
model_data = {
"base_model": "",
"tags": [],
"civitai": {"id": 1, "creator": {"username": "Author"}},
}
relative_path = calculate_relative_path_for_model(model_data, "lora")
assert not relative_path.startswith("/")
assert relative_path == "no tags"
def test_calculate_relative_path_sanitizes_double_slashes(isolated_settings):
"""Test that empty substitutions don't produce double slashes."""
isolated_settings["download_path_templates"]["lora"] = "{base_model}/{first_tag}/{author}"
model_data = {
"base_model": "",
"tags": [],
"civitai": {"id": 1, "creator": {"username": "Author"}},
}
relative_path = calculate_relative_path_for_model(model_data, "lora")
assert "//" not in relative_path
assert relative_path == "no tags/Author"
def test_calculate_recipe_fingerprint_filters_and_sorts():
loras = [
{"hash": "ABC", "strength": 0.1234},
@@ -0,0 +1,638 @@
<template>
<div class="lora-info-widget" :class="{ 'lm-vue-node': isVueMode }" @wheel="onWheel">
<template v-if="loraName">
<!-- Tab bar -->
<div class="lora-info-tabs">
<label
class="lora-info-tab"
:class="{ active: activeTab === 'notes' }"
>
<input
type="radio"
v-model="activeTab"
value="notes"
class="lora-info-tab-input"
/>
<span class="lora-info-tab-label">Notes</span>
</label>
<label
class="lora-info-tab"
:class="{ active: activeTab === 'description' }"
>
<input
type="radio"
v-model="activeTab"
value="description"
class="lora-info-tab-input"
@change="onDescriptionTabActivated"
/>
<span class="lora-info-tab-label">Description</span>
</label>
</div>
<!-- Notes tab content -->
<div v-show="activeTab === 'notes'" class="tab-content notes-tab">
<div class="info-field">
<label class="info-label">Filename</label>
<div class="lora-filename">{{ loraName }}</div>
</div>
<div class="info-field notes-field">
<label class="info-label">Notes</label>
<textarea
v-model="notes"
class="lora-notes lm-wheel-scrollable"
placeholder="Add notes about this LoRA..."
:disabled="saving"
></textarea>
</div>
<button
class="save-btn"
:disabled="notes === originalNotes || saving"
@click="saveNotes"
>
{{ saving ? 'Saving...' : 'Save' }}
</button>
</div>
<!-- Description tab content -->
<div v-show="activeTab === 'description'" class="tab-content description-tab lm-wheel-scrollable">
<!-- Loading state -->
<div v-if="descriptionLoading" class="description-state">
<i class="fas fa-spinner fa-spin"></i>
<span>Loading description...</span>
</div>
<!-- Error state -->
<div v-else-if="descriptionError" class="description-state error">
<span>Failed to load description</span>
</div>
<!-- Empty state (loaded but no content) -->
<div v-else-if="!hasDescription" class="description-state placeholder">
<span>No description available</span>
</div>
<!-- Description content -->
<div v-else class="description-content">
<div v-if="versionDescription" class="description-section">
<label class="info-label">About this version</label>
<div class="description-text" v-html="versionDescription"></div>
</div>
<div v-if="modelDescription" class="description-section">
<label class="info-label">Model Description</label>
<div class="description-text" v-html="modelDescription"></div>
</div>
</div>
</div>
</template>
<div v-else class="placeholder">No LoRA selected</div>
</div>
</template>
<script setup lang="ts">
import { onMounted, ref, computed, watch } from 'vue'
interface LoraInfoWidget {
serializeValue?: () => Promise<unknown>
value?: unknown
onSetValue?: (v: unknown) => void
callback?: unknown
options?: {
getValue?: () => LoraInfoWidgetValue
setValue?: (v: unknown) => void
}
node?: { widgets?: Array<{ id?: string }>; widgets_values?: Array<unknown> }
id?: string
_setLoraInfo?: (data: { name: string; notes: string; filePath: string; activeTab?: string } | null) => void
__pendingLoraInfo?: { name: string; notes: string; filePath: string; activeTab?: string } | null
}
interface LoraInfoWidgetValue {
name?: string
notes?: string
filePath?: string
activeTab?: string
}
const props = defineProps<{
widget: LoraInfoWidget
node: { id: number }
api: { fetchApi: (url: string, options?: RequestInit) => Promise<Response> }
app: { extensionManager: { toast: { add: (opts: Record<string, unknown>) => void } } }
isVueMode?: boolean
}>()
const loraName = ref<string>('')
const notes = ref<string>('')
const originalNotes = ref<string>('')
const filePath = ref<string>('')
const saving = ref<boolean>(false)
const activeTab = ref<string>('notes')
// Description tab state
const versionDescription = ref<string>('')
const modelDescription = ref<string>('')
const descriptionLoading = ref<boolean>(false)
const descriptionError = ref<boolean>(false)
const descriptionLoaded = ref<boolean>(false)
const hasDescription = computed(() =>
!!(versionDescription.value || modelDescription.value)
)
// Reset and auto-fetch description state when the LoRA selection changes
watch(filePath, (newPath) => {
descriptionLoaded.value = false
descriptionError.value = false
versionDescription.value = ''
modelDescription.value = ''
if (newPath && activeTab.value === 'description') {
fetchDescription()
}
})
function onDescriptionTabActivated() {
if (!descriptionLoaded.value && filePath.value) {
fetchDescription()
}
}
async function fetchDescription() {
if (descriptionLoading.value || !filePath.value) return
descriptionLoading.value = true
descriptionError.value = false
try {
const response = await props.api.fetchApi(
`/lm/loras/metadata?file_path=${encodeURIComponent(filePath.value)}`,
{ method: 'GET' }
)
if (!response.ok) {
throw new Error(`Failed to fetch metadata: ${response.statusText}`)
}
const data = await response.json()
if (data.success && data.metadata) {
versionDescription.value = data.metadata.description || ''
modelDescription.value = data.metadata.model?.description || ''
descriptionLoaded.value = true
} else {
// Successful response but no metadata treat as empty, not error
descriptionLoaded.value = true
}
} catch (e) {
console.error('[LoraInfoWidget] Failed to fetch description:', e)
descriptionError.value = true
// Don't set descriptionLoaded allow retry on next tab switch
} finally {
descriptionLoading.value = false
}
}
async function saveNotes() {
if (notes.value === originalNotes.value || saving.value) return
if (!filePath.value) return
saving.value = true
try {
const response = await props.api.fetchApi('/lm/loras/save-metadata', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ file_path: filePath.value, notes: notes.value })
})
const result = await response.json()
if (result.success) {
props.app.extensionManager.toast.add({
severity: 'success',
summary: 'Saved',
detail: 'Notes updated successfully',
life: 2000
})
originalNotes.value = notes.value
} else {
props.app.extensionManager.toast.add({
severity: 'error',
summary: 'Error',
detail: result.message || result.error || 'Failed to save notes',
life: 3000
})
}
} catch (e) {
console.error('[LoraInfoWidget] Failed to save notes:', e)
props.app.extensionManager.toast.add({
severity: 'error',
summary: 'Error',
detail: (e as Error).message || 'Failed to save notes',
life: 3000
})
} finally {
saving.value = false
}
}
function onWheel(event: WheelEvent) {
const target = event.target as HTMLElement | null
if (!target) return
const comfyApp = (window as unknown as { app?: { canvas?: { processMouseWheel?: (e: WheelEvent) => void } } }).app
if (!comfyApp?.canvas?.processMouseWheel) return
// Always pass pinch-to-zoom to canvas
if (event.ctrlKey) {
event.preventDefault()
event.stopPropagation()
comfyApp.canvas.processMouseWheel(event)
return
}
// Horizontal scroll: pass to canvas
if (Math.abs(event.deltaX) > Math.abs(event.deltaY)) {
event.preventDefault()
event.stopPropagation()
comfyApp.canvas.processMouseWheel(event)
return
}
// Check if the target is inside a scrollable area (notes textarea or description tab)
const scrollableEl = target.closest('.lora-notes, .description-tab') as HTMLElement | null
if (scrollableEl) {
const canScrollY = scrollableEl.scrollHeight > scrollableEl.clientHeight
if (canScrollY) {
// Let native scroll handle it, but stop propagation to prevent canvas zoom
event.stopPropagation()
return
}
}
// Forward to canvas for zoom
event.preventDefault()
event.stopPropagation()
comfyApp.canvas.processMouseWheel(event)
}
onMounted(() => {
// Build current state snapshot for serialization
const buildValue = (): LoraInfoWidgetValue => ({
name: loraName.value,
notes: notes.value,
filePath: filePath.value,
activeTab: activeTab.value,
})
// Set value from external source (workflow load, paste, etc.)
const applyValue = (v: unknown) => {
if (v && typeof v === 'object') {
const data = v as LoraInfoWidgetValue
// Set activeTab before filePath so the filePath watcher sees the correct tab
// and triggers fetchDescription() when restoring description tab
if (data.activeTab !== undefined) activeTab.value = data.activeTab
if (data.name !== undefined) loraName.value = data.name
if (data.notes !== undefined) {
notes.value = data.notes
originalNotes.value = data.notes
}
if (data.filePath !== undefined) filePath.value = data.filePath
}
}
// ComponentWidgetImpl.value getter/setter delegates to options.getValue/options.setValue.
// These must be set for workflow JSON persistence (LGraphNode.serialize/configure) to work.
props.widget.options.getValue = buildValue
props.widget.options.setValue = applyValue
// Also set serializeValue for prompt/API serialization path (executionUtil.ts)
props.widget.serializeValue = async () => buildValue()
// Handle external value updates (e.g., loading workflow, paste)
props.widget.onSetValue = applyValue
// Restore from saved value. Because configure() may call widget.value = data
// before onMounted fires (and before options.setValue is assigned), we check
// widgets_values directly in case the value was already pushed.
const widgetIndex = props.widget.node?.widgets?.findIndex(
(w: { id?: string }) => w.id === props.widget.id
)
let restored = false
if (widgetIndex !== undefined && widgetIndex >= 0) {
const savedValue = props.widget.node?.widgets_values?.[widgetIndex]
if (savedValue && typeof savedValue === 'object') {
applyValue(savedValue)
restored = true
}
}
// Fallback: if configure() ran after onMounted, widget.value (via options.getValue)
// already has the saved data. Only use this path if the widgets_values lookup didn't restore.
if (!restored && props.widget.value && typeof props.widget.value === 'object') {
applyValue(props.widget.value)
}
// Expose setLoraInfo on the widget object for external callers (e.g., lora_info.js).
// Accepts null to clear the display (when selection is deselected).
props.widget._setLoraInfo = (data: { name: string; notes: string; filePath: string; activeTab?: string } | null) => {
if (data) {
loraName.value = data.name
notes.value = data.notes
originalNotes.value = data.notes
filePath.value = data.filePath
// Preserve existing activeTab unless explicitly provided
if (data.activeTab !== undefined) {
activeTab.value = data.activeTab
}
} else {
loraName.value = ''
notes.value = ''
originalNotes.value = ''
filePath.value = ''
// Do NOT reset activeTab on deselection user's tab preference persists
}
}
// Consume any data pushed before the Vue component mounted (race condition fix)
if (props.widget.__pendingLoraInfo) {
props.widget._setLoraInfo(props.widget.__pendingLoraInfo)
delete props.widget.__pendingLoraInfo
}
})
</script>
<style scoped>
.lora-info-widget {
padding: 12px;
background: rgba(40, 44, 52, 0.6);
border-radius: 4px;
height: 100%;
display: flex;
flex-direction: column;
box-sizing: border-box;
overflow: hidden;
}
/* Vue node mode: prevent content from pushing node size via ResizeObserver.
contain:layout size tells the browser the element's intrinsic size is
determined solely by CSS not by descendant content. This breaks the
feedback loop where content grows ResizeObserver resizes content
reflows repeat. Same technique used by tags_widget.js + lm_styles.css. */
.lora-info-widget.lm-vue-node {
contain: layout size;
}
/* ── Tab bar ── */
.lora-info-tabs {
display: flex;
gap: 0;
margin-bottom: 10px;
border-bottom: 1px solid var(--border-color, #444);
flex-shrink: 0;
}
.lora-info-tab {
flex: 1;
text-align: center;
cursor: pointer;
padding: 6px 0;
position: relative;
}
.lora-info-tab-input {
position: absolute;
opacity: 0;
width: 0;
height: 0;
}
.lora-info-tab-label {
font-size: 12px;
font-weight: 500;
color: var(--fg-color, #fff);
opacity: 0.5;
transition: opacity 0.15s;
}
.lora-info-tab:hover .lora-info-tab-label {
opacity: 0.75;
}
.lora-info-tab.active .lora-info-tab-label {
opacity: 1;
}
.lora-info-tab.active::after {
content: '';
position: absolute;
bottom: -1px;
left: 25%;
right: 25%;
height: 2px;
background: rgba(66, 153, 225, 0.8);
border-radius: 1px;
}
/* ── Tab content ── */
.tab-content {
flex: 1;
min-height: 0;
overflow: hidden;
}
.notes-tab {
display: flex;
flex-direction: column;
}
.description-tab {
display: flex;
flex-direction: column;
overflow-y: auto;
min-height: 0;
}
/* ── Info fields (shared) ── */
.info-field {
display: flex;
flex-direction: column;
gap: 4px;
}
.info-label {
font-size: 10px;
font-weight: 600;
text-transform: uppercase;
letter-spacing: 0.05em;
color: var(--fg-color, #fff);
opacity: 0.6;
}
.lora-filename {
font-size: 13px;
font-weight: 500;
color: var(--fg-color, #fff);
word-break: break-all;
margin-bottom: 8px;
/* Override node-level grab cursor and user-select:none from .lg-node.cursor-grab */
cursor: auto;
user-select: text;
-webkit-user-select: text;
}
.notes-field {
flex: 1;
min-height: 0;
}
.lora-notes {
width: 100%;
flex: 1;
min-height: 60px;
padding: 8px;
border-radius: 4px;
border: 1px solid var(--border-color, #444);
background: var(--comfy-input-bg, #333);
color: var(--fg-color, #fff);
font-size: 12px;
resize: none;
box-sizing: border-box;
font-family: inherit;
outline: none;
}
.lora-notes:focus {
border-color: var(--comfy-input-border, #444);
}
.lora-notes:disabled {
opacity: 0.6;
cursor: not-allowed;
}
.save-btn {
width: 100%;
margin-top: 8px;
padding: 6px 12px;
border-radius: 4px;
border: 1px solid rgba(66, 153, 225, 0.4);
background: rgba(66, 153, 225, 0.15);
color: var(--fg-color, #fff);
font-size: 12px;
cursor: pointer;
transition: all 0.2s;
box-sizing: border-box;
flex-shrink: 0;
}
.save-btn:hover:not(:disabled) {
background: rgba(66, 153, 225, 0.25);
border-color: rgba(66, 153, 225, 0.6);
}
.save-btn:disabled {
opacity: 0.4;
cursor: not-allowed;
background: rgba(66, 153, 225, 0.05);
border-color: rgba(226, 232, 240, 0.1);
}
/* ── Description states ── */
.description-state {
display: flex;
align-items: center;
justify-content: center;
gap: 8px;
padding: 24px 16px;
color: var(--fg-color, #fff);
opacity: 0.5;
font-size: 12px;
min-height: 0;
flex-shrink: 0;
}
.description-state.error {
opacity: 0.7;
color: #f87171;
}
/* ── Description content ── */
.description-content {
min-height: 0;
}
.description-section {
margin-bottom: 14px;
}
.description-section:last-child {
margin-bottom: 0;
}
.description-text {
padding: 8px 0;
font-size: 12px;
line-height: 1.5;
color: var(--fg-color, #fff);
opacity: 0.85;
word-break: break-word;
/* Override node-level grab cursor and user-select:none from .lg-node.cursor-grab */
cursor: auto;
user-select: text;
-webkit-user-select: text;
}
.description-text :deep(p) {
margin: 0 0 8px 0;
}
.description-text :deep(p:last-child) {
margin-bottom: 0;
}
.description-text :deep(a) {
color: rgba(66, 153, 225, 0.9);
}
.description-text :deep(ul),
.description-text :deep(ol) {
padding-left: 20px;
margin: 4px 0;
}
.description-text :deep(h1),
.description-text :deep(h2),
.description-text :deep(h3) {
font-size: 13px;
margin: 10px 0 4px 0;
font-weight: 600;
opacity: 0.95;
}
.description-text :deep(code) {
background: rgba(255, 255, 255, 0.08);
padding: 1px 4px;
border-radius: 3px;
font-size: 11px;
}
.description-text :deep(img) {
max-width: 100%;
border-radius: 4px;
}
/* ── Placeholder (shared) ── */
.placeholder {
font-style: italic;
color: rgba(226, 232, 240, 0.5);
text-align: center;
padding: 16px 0;
font-size: 12px;
}
/* ── Spinner (Font Awesome) ── */
.fa-spinner {
animation: fa-spin 1s linear infinite;
}
@keyframes fa-spin {
0% { transform: rotate(0deg); }
100% { transform: rotate(360deg); }
}
</style>
+185 -23
View File
@@ -5,6 +5,7 @@ import LoraRandomizerWidget from '@/components/LoraRandomizerWidget.vue'
import LoraCyclerWidget from '@/components/LoraCyclerWidget.vue'
import JsonDisplayWidget from '@/components/JsonDisplayWidget.vue'
import AutocompleteTextWidget from '@/components/AutocompleteTextWidget.vue'
import LoraInfoWidget from '@/components/LoraInfoWidget.vue'
import { createVueWidgetCleanup } from './vue-widget-cleanup'
import type { LoraPoolConfig, RandomizerConfig, CyclerConfig } from './composables/types'
import {
@@ -23,6 +24,8 @@ const LORA_CYCLER_WIDGET_MIN_HEIGHT = 408
const LORA_CYCLER_WIDGET_MAX_HEIGHT = LORA_CYCLER_WIDGET_MIN_HEIGHT
const JSON_DISPLAY_WIDGET_MIN_WIDTH = 300
const JSON_DISPLAY_WIDGET_MIN_HEIGHT = 200
const LORA_INFO_WIDGET_MIN_WIDTH = 300
const LORA_INFO_WIDGET_MIN_HEIGHT = 200
const AUTOCOMPLETE_TEXT_WIDGET_MIN_HEIGHT = 60
const AUTOCOMPLETE_TEXT_WIDGET_MAX_HEIGHT = 100
// Per-modelType min size hints for node initial sizing.
@@ -71,7 +74,7 @@ function forwardMiddleMouseToCanvas(container: HTMLElement) {
})
}
const vueApps = new Map<number, VueApp>()
const vueApps = new Map<number | string, VueApp>()
let autocompleteTextWidgetInstanceId = 0
export function createAutocompleteTextWidgetInstanceId() {
@@ -402,7 +405,6 @@ function createJsonDisplayWidget(node) {
return { widget }
}
// Store nodeData options per widget type for autocomplete widgets
const widgetInputOptions: Map<string, { placeholder?: string }> = new Map()
function getSerializableWidgetNames(node: any): string[] {
@@ -642,6 +644,75 @@ if (app.ui?.settings) {
}, 100)
}
// @ts-ignore
function createLoraInfoWidget(node: any) {
const container = document.createElement('div')
container.id = `lora-info-widget-${node.id}`
container.style.width = '100%'
container.style.height = '100%'
container.style.display = 'flex'
container.style.flexDirection = 'column'
container.style.overflow = 'hidden'
forwardMiddleMouseToCanvas(container)
let internalValue: { name?: string; notes?: string; filePath?: string; activeTab?: string } | undefined
const widget = node.addDOMWidget(
'lora_info_display',
'LORA_INFO_DISPLAY',
container,
{
getValue() {
return internalValue
},
setValue(v: { name?: string; notes?: string; filePath?: string; activeTab?: string }) {
internalValue = v
if (typeof widget.onSetValue === 'function') {
widget.onSetValue(v)
}
},
serialize: true,
getMinHeight() {
return LORA_INFO_WIDGET_MIN_HEIGHT
}
}
)
const vueApp = createApp(LoraInfoWidget, {
widget,
node,
api,
app,
isVueMode: typeof LiteGraph !== 'undefined' && LiteGraph.vueNodesMode,
})
vueApp.use(PrimeVue, {
unstyled: true,
ripple: false
})
vueApp.mount(container)
vueApps.set(node.id + 40000, vueApp) // Offset to avoid collision
widget.computeLayoutSize = () => {
const minWidth = LORA_INFO_WIDGET_MIN_WIDTH
const minHeight = LORA_INFO_WIDGET_MIN_HEIGHT
return { minHeight, minWidth }
}
widget.onRemove = () => {
const vueApp = vueApps.get(node.id + 40000)
if (vueApp) {
vueApp.unmount()
vueApps.delete(node.id + 40000)
}
}
return { widget }
}
// Factory function for creating autocomplete text widgets
// @ts-ignore
function createAutocompleteTextWidgetFactory(
@@ -651,16 +722,30 @@ function createAutocompleteTextWidgetFactory(
inputOptions: { placeholder?: string } = {}
) {
const metadataWidgetName = `__lm_autocomplete_meta_${widgetName}`
const instanceId = createAutocompleteTextWidgetInstanceId()
const container = document.createElement('div')
container.id = `autocomplete-text-widget-${instanceId}`
container.style.width = '100%'
container.style.height = '100%'
container.style.display = 'flex'
container.style.flexDirection = 'column'
container.style.overflow = 'hidden'
forwardMiddleMouseToCanvas(container)
let container: HTMLElement | null = null
const existingContainers = document.querySelectorAll<HTMLElement>(
'[id^="autocomplete-text-widget-"]'
)
for (const el of existingContainers) {
if (el.children.length === 0) {
container = el
break
}
}
if (!container) {
const instanceId = String(createAutocompleteTextWidgetInstanceId())
container = document.createElement('div')
container.id = `autocomplete-text-widget-${instanceId}`
container.style.width = '100%'
container.style.height = '100%'
container.style.display = 'flex'
container.style.flexDirection = 'column'
container.style.overflow = 'hidden'
forwardMiddleMouseToCanvas(container)
}
// Store textarea reference on the container element so cloned widgets can access it
// This is necessary because when widgets are promoted to subgraph nodes,
@@ -739,15 +824,10 @@ function createAutocompleteTextWidgetFactory(
})
vueApp.mount(container)
const appKey = instanceId
const appKey = container.id
vueApps.set(appKey, vueApp)
if (maxHeight) {
// Set only minHeight as a true minimum — remove maxHeight so the
// textarea can grow when the user resizes it in app mode (where
// [&_textarea]:resize-y applies). Graph mode (canvas & Vue render)
// is unaffected because LiteGraph's layout system still governs
// the widget area size.
container.style.minHeight = `${AUTOCOMPLETE_TEXT_WIDGET_MIN_HEIGHT}px`
}
@@ -759,10 +839,14 @@ function createAutocompleteTextWidgetFactory(
)
}
widget.onRemove = createVueWidgetCleanup(vueApp, () => {
const vueCleanup = createVueWidgetCleanup(vueApp, () => {
vueApps.delete(appKey)
})
widget.onRemove = () => {
vueCleanup()
}
// Return minWidth/minHeight hints so ComfyUI's _initialMinSize mechanism
// sets a sensible initial node width (and height for prompt/embeddings).
// loras modelType retains its existing height constraints (getMaxHeight: 100).
@@ -804,7 +888,75 @@ app.registerExtension({
updateDownstreamLoaders(node)
} : null
return addLorasWidgetCache(node, 'loras', { isRandomizerNode }, callback)
const opts: { isRandomizerNode?: boolean; onSelectionChange?: (selection: any) => void } = {
isRandomizerNode,
}
if (isRandomizerNode) {
opts.onSelectionChange = async (selection: any) => {
if (!selection?.name || !selection?.active) return
// Walk outputs to find directly connected Lora Info nodes
const infoNodes: any[] = []
if (node.outputs) {
for (const output of node.outputs) {
if (!output?.links?.length) continue
for (const linkId of output.links) {
const links = node.graph?.links
if (!links) continue
const link = Array.isArray(links) ? links[linkId] : links.get?.(linkId)
if (!link) continue
const targetNode = node.graph?.getNodeById?.(link.target_id)
if (targetNode?.comfyClass === 'Lora Info (LoraManager)') {
infoNodes.push(targetNode)
}
}
}
}
if (infoNodes.length === 0) return
// Bump request token to guard against stale async responses
for (const infoNode of infoNodes) {
infoNode.__loraInfoReqId = (infoNode.__loraInfoReqId || 0) + 1
}
const reqIdSnapshot = new Map<any, number>()
for (const infoNode of infoNodes) {
reqIdSnapshot.set(infoNode, infoNode.__loraInfoReqId)
}
// Fetch notes via the real ComfyUI api
let infoData: any
try {
const response = await api.fetchApi(
`/lm/loras/get-notes?name=${encodeURIComponent(selection.name)}`,
{ method: 'GET' }
)
if (response?.ok) {
const data = await response.json()
infoData = {
name: selection.name,
notes: data?.notes || '',
filePath: data?.file_path || '',
}
} else {
infoData = { name: selection.name, notes: '[Error loading notes]', filePath: '' }
}
} catch {
infoData = { name: selection.name, notes: '[Error loading notes]', filePath: '' }
}
for (const infoNode of infoNodes) {
if (infoNode.__loraInfoReqId !== reqIdSnapshot.get(infoNode)) {
continue
}
if (typeof infoNode._setLoraInfo === 'function') {
infoNode._setLoraInfo(infoData)
}
}
}
}
return addLorasWidgetCache(node, 'loras', opts, callback)
},
// Autocomplete text widget for LoRAs (used by Lora Loader, Lora Stacker, WanVideo Lora Select)
// @ts-ignore
@@ -823,7 +975,7 @@ app.registerExtension({
AUTOCOMPLETE_TEXT_PROMPT(node) {
const options = widgetInputOptions.get(`${node.comfyClass}:text`) || {}
return createAutocompleteTextWidgetFactory(node, 'text', 'prompt', options)
}
},
}
},
@@ -868,9 +1020,7 @@ app.registerExtension({
info.widgets_values = [...(info.widgets_values ?? []), null]
}
const result = originalConfigure?.apply(this, arguments)
return result
return originalConfigure?.apply(this, arguments)
}
}
@@ -903,5 +1053,17 @@ app.registerExtension({
createJsonDisplayWidget(this)
}
}
// Add the Lora Info display widget
if (nodeData.name === 'Lora Info (LoraManager)') {
const onNodeCreated = nodeType.prototype.onNodeCreated
nodeType.prototype.onNodeCreated = function () {
onNodeCreated?.apply(this, [])
// Create the lora info display widget
createLoraInfoWidget(this)
}
}
}
})
+1
View File
@@ -16,6 +16,7 @@ export const LORA_PROVIDER_NODE_TYPES = [
"Lora Stacker (LoraManager)",
"Lora Randomizer (LoraManager)",
"Lora Cycler (LoraManager)",
"Create Hook LoRA (LoraManager)",
] as const;
/**
@@ -0,0 +1,418 @@
/**
* Tests for LoraInfoWidget tab switching, lazy description loading,
* state serialization roundtrip, and activeTab persistence.
*/
import { nextTick } from 'vue'
import { shallowMount } from '@vue/test-utils'
import { describe, expect, it, vi, beforeEach, afterEach } from 'vitest'
import LoraInfoWidget from '@/components/LoraInfoWidget.vue'
import { setupFetchMock, resetFetchMock } from '../setup'
// ── Helpers ──
function createMockFetchApi(overrides: {
response?: unknown
ok?: boolean
error?: string
} = {}) {
const { response = { success: true, metadata: {} }, ok = true } = overrides
return vi.fn().mockResolvedValue({
ok,
json: () => Promise.resolve(response),
})
}
function createMockToast() {
return { add: vi.fn() }
}
function createMockWidget(value?: unknown) {
type PendingInfo = { name: string; notes: string; filePath: string; activeTab?: string } | null
const widget = {
options: {} as { getValue?: () => unknown; setValue?: (v: unknown) => void },
serializeValue: (async () => null) as () => Promise<unknown>,
value: (value ?? undefined) as unknown,
onSetValue: undefined as unknown as ((v: unknown) => void),
_setLoraInfo: undefined as unknown as (data: Record<string, unknown> | null) => void,
__pendingLoraInfo: undefined as unknown as PendingInfo | undefined,
}
return widget
}
interface MountOptions {
initialValue?: Record<string, unknown>
}
type TestWidget = ReturnType<typeof createMockWidget>
function mountWidget(options: MountOptions = {}) {
const fetchApi = createMockFetchApi()
const widget = createMockWidget(options.initialValue)
const node = { id: 1 }
const app = { extensionManager: { toast: createMockToast() } }
const wrapper = shallowMount(LoraInfoWidget, {
props: { widget, node, api: { fetchApi }, app },
})
return { wrapper, widget: widget as TestWidget, fetchApi, app }
}
// ── Tests ──
describe('LoraInfoWidget', () => {
beforeEach(() => {
setupFetchMock()
})
afterEach(() => {
resetFetchMock()
})
describe('initial state', () => {
it('shows placeholder when no LoRA is selected', () => {
const { wrapper } = mountWidget()
expect(wrapper.text()).toContain('No LoRA selected')
})
it('shows Notes tab by default when LoRA is set', async () => {
const { wrapper, widget } = mountWidget()
widget._setLoraInfo!({ name: 'test.safetensors', notes: '', filePath: '/path/test.safetensors' })
await nextTick()
expect(wrapper.text()).toContain('test.safetensors')
expect(wrapper.find('.notes-tab').isVisible()).toBe(true)
expect(wrapper.find('.description-tab').isVisible()).toBe(false)
})
})
describe('tab switching', () => {
it('switches to Description tab and back to Notes', async () => {
const { wrapper, widget } = mountWidget()
widget._setLoraInfo!({ name: 'test.safetensors', notes: '', filePath: '/path/test.safetensors' })
await nextTick()
const tabs = wrapper.findAll('.lora-info-tab')
// Click Description tab
const descriptionTab = wrapper.findAll('.lora-info-tab-input')[1]
await descriptionTab.setValue('description')
await nextTick()
expect(tabs[1].classes()).toContain('active')
expect(wrapper.text()).toContain('No description available')
// Switch back to Notes
const notesTab = wrapper.findAll('.lora-info-tab-input')[0]
await notesTab.setValue('notes')
await nextTick()
expect(tabs[0].classes()).toContain('active')
expect(wrapper.text()).toContain('test.safetensors')
})
})
describe('description lazy loading', () => {
it('fetches metadata when Description tab is activated', async () => {
const fetchApi = createMockFetchApi({
response: {
success: true,
metadata: {
description: '<p>Version desc</p>',
model: { description: '<p>Model desc</p>' },
},
},
})
const widget = createMockWidget()
const wrapper = shallowMount(LoraInfoWidget, {
props: {
widget,
node: { id: 1 },
api: { fetchApi },
app: { extensionManager: { toast: createMockToast() } },
},
})
widget._setLoraInfo!({ name: 'test.safetensors', notes: '', filePath: '/path/test.safetensors' })
await nextTick()
// Switch to Description tab
const descriptionTab = wrapper.findAll('.lora-info-tab-input')[1]
await descriptionTab.setValue('description')
await nextTick()
await nextTick() // flush async fetch
expect(fetchApi).toHaveBeenCalledWith(
expect.stringContaining('/lm/loras/metadata'),
expect.objectContaining({ method: 'GET' })
)
expect(wrapper.html()).toContain('Version desc')
expect(wrapper.html()).toContain('Model desc')
})
it('shows loading state while fetching', async () => {
// Use a never-resolving promise to simulate loading
const fetchApi = vi.fn().mockReturnValue(new Promise(() => {}))
const widget = createMockWidget()
const wrapper = shallowMount(LoraInfoWidget, {
props: {
widget,
node: { id: 1 },
api: { fetchApi },
app: { extensionManager: { toast: createMockToast() } },
},
})
widget._setLoraInfo!({ name: 'test.safetensors', notes: '', filePath: '/path/test.safetensors' })
await nextTick()
const descriptionTab = wrapper.findAll('.lora-info-tab-input')[1]
await descriptionTab.setValue('description')
await nextTick()
expect(wrapper.text()).toContain('Loading description')
})
it('shows error state when fetch fails', async () => {
const fetchApi = vi.fn().mockRejectedValue(new Error('Network error'))
const widget = createMockWidget()
const wrapper = shallowMount(LoraInfoWidget, {
props: {
widget,
node: { id: 1 },
api: { fetchApi },
app: { extensionManager: { toast: createMockToast() } },
},
})
widget._setLoraInfo!({ name: 'test.safetensors', notes: '', filePath: '/path/test.safetensors' })
await nextTick()
const descriptionTab = wrapper.findAll('.lora-info-tab-input')[1]
await descriptionTab.setValue('description')
await nextTick()
await nextTick()
expect(wrapper.text()).toContain('Failed to load description')
})
it('shows empty state when metadata has no descriptions', async () => {
const fetchApi = createMockFetchApi({
response: {
success: true,
metadata: {
description: '',
model: {},
},
},
})
const widget = createMockWidget()
const wrapper = shallowMount(LoraInfoWidget, {
props: {
widget,
node: { id: 1 },
api: { fetchApi },
app: { extensionManager: { toast: createMockToast() } },
},
})
widget._setLoraInfo!({ name: 'test.safetensors', notes: '', filePath: '/path/test.safetensors' })
await nextTick()
const descriptionTab = wrapper.findAll('.lora-info-tab-input')[1]
await descriptionTab.setValue('description')
await nextTick()
await nextTick()
expect(wrapper.text()).toContain('No description available')
})
it('caches description and does not re-fetch on second activation', async () => {
const fetchApi = createMockFetchApi({
response: {
success: true,
metadata: {
description: '<p>Version desc</p>',
model: { description: '<p>Model desc</p>' },
},
},
})
const widget = createMockWidget()
const wrapper = shallowMount(LoraInfoWidget, {
props: {
widget,
node: { id: 1 },
api: { fetchApi },
app: { extensionManager: { toast: createMockToast() } },
},
})
widget._setLoraInfo!({ name: 'test.safetensors', notes: '', filePath: '/path/test.safetensors' })
await nextTick()
// First activation
const descriptionTab = wrapper.findAll('.lora-info-tab-input')[1]
await descriptionTab.setValue('description')
await nextTick()
await nextTick()
expect(fetchApi).toHaveBeenCalledTimes(1)
// Switch away and back
const notesTab = wrapper.findAll('.lora-info-tab-input')[0]
await notesTab.setValue('notes')
await nextTick()
await descriptionTab.setValue('description')
await nextTick()
// Should NOT have called fetch again
expect(fetchApi).toHaveBeenCalledTimes(1)
})
it('re-fetches when LoRA selection changes', async () => {
const fetchApi = createMockFetchApi({
response: {
success: true,
metadata: {
description: '<p>Version desc</p>',
model: { description: '<p>Model desc</p>' },
},
},
})
const widget = createMockWidget()
const wrapper = shallowMount(LoraInfoWidget, {
props: {
widget,
node: { id: 1 },
api: { fetchApi },
app: { extensionManager: { toast: createMockToast() } },
},
})
widget._setLoraInfo!({ name: 'first.safetensors', notes: '', filePath: '/path/first.safetensors' })
await nextTick()
const descriptionTab = wrapper.findAll('.lora-info-tab-input')[1]
await descriptionTab.setValue('description')
await nextTick()
await nextTick()
expect(fetchApi).toHaveBeenCalledTimes(1)
// Select a different LoRA — resets description state
widget._setLoraInfo!({ name: 'second.safetensors', notes: '', filePath: '/path/second.safetensors' })
await nextTick()
// Should show loading again (not cached)
await descriptionTab.setValue('description')
await nextTick()
await nextTick()
expect(fetchApi).toHaveBeenCalledTimes(2)
})
})
describe('serialization roundtrip', () => {
it('serializeValue includes activeTab', async () => {
const { wrapper, widget } = mountWidget()
widget._setLoraInfo!({ name: 'test.safetensors', notes: 'my notes', filePath: '/path/test.safetensors' })
await nextTick()
// Switch to Description tab
const descriptionTab = wrapper.findAll('.lora-info-tab-input')[1]
await descriptionTab.setValue('description')
await nextTick()
const serialized = await widget.serializeValue!()
expect(serialized).toMatchObject({
name: 'test.safetensors',
notes: 'my notes',
filePath: '/path/test.safetensors',
activeTab: 'description',
})
})
it('onSetValue restores activeTab from workflow value', async () => {
const { wrapper } = mountWidget({
initialValue: {
name: 'saved.safetensors',
notes: 'saved notes',
filePath: '/path/saved.safetensors',
activeTab: 'description',
},
})
await nextTick()
// Description tab should be visible (activeTab restored to 'description')
expect(wrapper.find('.description-tab').isVisible()).toBe(true)
expect(wrapper.text()).toContain('saved.safetensors')
})
it('defaults to notes tab when activeTab is missing in saved value', async () => {
const { wrapper } = mountWidget({
initialValue: {
name: 'legacy.safetensors',
notes: 'legacy notes',
filePath: '/path/legacy.safetensors',
// No activeTab — legacy workflow
},
})
await nextTick()
expect(wrapper.find('.notes-tab').isVisible()).toBe(true)
})
})
describe('_setLoraInfo race condition guard', () => {
it('consumes __pendingLoraInfo pushed before mount', async () => {
const widget = createMockWidget()
widget.__pendingLoraInfo = {
name: 'pending.safetensors',
notes: 'pending notes',
filePath: '/path/pending.safetensors',
}
const wrapper = shallowMount(LoraInfoWidget, {
props: {
widget,
node: { id: 1 },
api: { fetchApi: createMockFetchApi() },
app: { extensionManager: { toast: createMockToast() } },
},
})
await nextTick()
expect(widget.__pendingLoraInfo).toBeUndefined()
expect(wrapper.text()).toContain('pending.safetensors')
})
it('preserves activeTab when _setLoraInfo called with null (deselection)', async () => {
const { wrapper, widget } = mountWidget()
widget._setLoraInfo!({ name: 'test.safetensors', notes: '', filePath: '/path/test.safetensors' })
await nextTick()
// Switch to Description tab
const descriptionTab = wrapper.findAll('.lora-info-tab-input')[1]
await descriptionTab.setValue('description')
await nextTick()
// Deselect — template shows placeholder (no tab bar rendered)
widget._setLoraInfo!(null)
await nextTick()
// Placeholder shown
expect(wrapper.text()).toContain('No LoRA selected')
// Re-select — activeTab should still be 'description'
widget._setLoraInfo!({ name: 'second.safetensors', notes: '', filePath: '/path/second.safetensors' })
await nextTick()
const tabs = wrapper.findAll('.lora-info-tab')
expect(tabs[1].classes()).toContain('active')
})
})
})
+143
View File
@@ -0,0 +1,143 @@
import { app } from "../../scripts/app.js";
import {
getActiveLorasFromNode,
updateConnectedTriggerWords,
chainCallback,
mergeLoras,
getWidgetByName,
getWidgetSerializedValue,
} from "./utils.js";
import { addLorasWidget } from "./loras_widget.js";
import { applyLoraValuesToText, debounce } from "./lora_syntax_utils.js";
import { applySelectionHighlight } from "./trigger_word_highlight.js";
import { updateConnectedLoraInfoNodes } from "./lora_info.js";
app.registerExtension({
name: "LoraManager.CreateHookLora",
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeType.comfyClass === "Create Hook LoRA (LoraManager)") {
chainCallback(nodeType.prototype, "onNodeCreated", function () {
// Enable widget serialization so loras widget state is persisted
this.serialize_widgets = true;
this.addInput("prev_hooks", "HOOKS", {
shape: 7,
});
// Flags to prevent callback loops between text widget ↔ loras widget
let isUpdating = false;
let isSyncingInput = false;
// Get the text input widget (AUTOCOMPLETE_TEXT_LORAS type, created by Vue widgets)
const inputWidget = getWidgetByName(this, "text");
if (!inputWidget) {
console.warn(
"LoRA Manager: text widget not found for Create Hook LoRA"
);
return;
}
this.inputWidget = inputWidget;
const scheduleInputSync = debounce((lorasValue) => {
if (isSyncingInput) {
return;
}
isSyncingInput = true;
isUpdating = true;
try {
const nextText = applyLoraValuesToText(
inputWidget.value,
lorasValue
);
if (inputWidget.value !== nextText) {
inputWidget.value = nextText;
}
} finally {
isUpdating = false;
isSyncingInput = false;
}
});
// Create the LoRA list widget
const result = addLorasWidget(
this,
"loras",
{
onSelectionChange: (selection) => {
applySelectionHighlight(this, selection);
updateConnectedLoraInfoNodes(this, selection);
},
},
(value) => {
// Prevent recursive calls
if (isUpdating) return;
isUpdating = true;
try {
// Update connected trigger word toggles with active LoRA names
const activeLoraNames = new Set();
value.forEach((lora) => {
if (lora.active) {
activeLoraNames.add(lora.name);
}
});
updateConnectedTriggerWords(this, activeLoraNames);
} finally {
isUpdating = false;
}
scheduleInputSync(value);
}
);
this.lorasWidget = result.widget;
// Set up callback for the text input widget to trigger merge logic
inputWidget.callback = (value) => {
if (isUpdating) return;
isUpdating = true;
try {
const currentLoras = this.lorasWidget?.value || [];
const mergedLoras = mergeLoras(value, currentLoras);
if (this.lorasWidget) {
this.lorasWidget.value = mergedLoras;
}
// Update connected trigger word toggles
const activeLoraNames = getActiveLorasFromNode(this);
updateConnectedTriggerWords(this, activeLoraNames);
} finally {
isUpdating = false;
}
};
});
}
},
async loadedGraphNode(node) {
if (node.comfyClass === "Create Hook LoRA (LoraManager)") {
// Restore saved loras widget values on workflow load
let existingLoras = [];
if (node.widgets_values && node.widgets_values.length > 0) {
const savedValue = getWidgetSerializedValue(node, "loras");
existingLoras = savedValue || [];
}
// Merge the loras data from text widget with saved values
const inputWidget =
node.inputWidget || getWidgetByName(node, "text");
if (!inputWidget) {
console.warn(
"LoRA Manager: text widget not found while restoring Create Hook LoRA"
);
return;
}
const mergedLoras = mergeLoras(inputWidget.value, existingLoras);
node.lorasWidget.value = mergedLoras;
}
},
});
+17
View File
@@ -120,10 +120,27 @@
outline: none;
}
/* Vue node mode: prevent content from pushing node size via ResizeObserver.
contain:size breaks the feedback loop the container's intrinsic size
is determined solely by CSS, not by how many LoRAs are inside. */
.lm-loras-container.lm-vue-node {
height: 100%;
min-height: var(--comfy-widget-min-height, 200px);
contain: layout size;
}
.lm-loras-container:focus {
outline: none;
}
/* Vue node mode: prevent content from pushing node size via ResizeObserver.
Same technique as .lm-loras-container.lm-vue-node above. */
.comfy-tags-container.lm-vue-node {
height: 100%;
min-height: var(--comfy-widget-min-height, 150px);
contain: layout size;
}
.lm-lora-empty-state {
text-align: center;
padding: 20px 0;
+182
View File
@@ -0,0 +1,182 @@
import { app } from "../../scripts/app.js";
import { api } from "../../scripts/api.js";
import {
getLinkFromGraph,
chainCallback,
} from "./utils.js";
const LORA_INFO_CLASS = "Lora Info (LoraManager)";
/**
* Find Lora Info nodes directly connected to the given node's outputs.
* Mirrors the getConnectedTriggerToggleNodes pattern from utils.js.
* @param {object} node - The source node to check outputs from
* @returns {object[]} Array of connected Lora Info node instances
*/
export function getConnectedLoraInfoNodes(node) {
const connectedNodes = [];
if (!node?.outputs) {
return connectedNodes;
}
for (const output of node.outputs) {
if (!output?.links?.length) {
continue;
}
for (const linkId of output.links) {
const link = getLinkFromGraph(node.graph, linkId);
if (!link) {
continue;
}
const targetNode = node.graph?.getNodeById?.(link.target_id);
if (targetNode && targetNode.comfyClass === LORA_INFO_CLASS) {
connectedNodes.push(targetNode);
}
}
}
return connectedNodes;
}
/**
* Fetch notes for the selected lora and push them to all directly connected
* Lora Info nodes (no recursive chain traversal only direct connections).
* @param {object} node - The source LoRA Loader/Stacker node
* @param {object|null} selection - The current lora selection {name, active, entry}
*/
export async function updateConnectedLoraInfoNodes(node, selection) {
if (!node) {
return;
}
const infoNodes = getConnectedLoraInfoNodes(node);
if (infoNodes.length === 0) {
return;
}
// No selection or inactive — clear the display on all connected info nodes
if (!selection?.name || !selection?.active) {
for (const infoNode of infoNodes) {
infoNode.__loraInfoReqId = (infoNode.__loraInfoReqId || 0) + 1;
if (typeof infoNode._setLoraInfo === "function") {
infoNode._setLoraInfo(null);
} else {
infoNode.__pendingLoraInfo = null;
}
}
return;
}
// Bump request token on each info node to guard against stale async responses
for (const infoNode of infoNodes) {
infoNode.__loraInfoReqId = (infoNode.__loraInfoReqId || 0) + 1;
}
const reqIdSnapshot = new Map();
for (const infoNode of infoNodes) {
reqIdSnapshot.set(infoNode, infoNode.__loraInfoReqId);
}
// Fetch notes for the selected lora
try {
const response = await api.fetchApi(
`/lm/loras/get-notes?name=${encodeURIComponent(selection.name)}`,
{ method: "GET" }
);
if (!response?.ok) {
throw new Error(`Failed to fetch notes for ${selection.name}`);
}
const data = await response.json();
const infoData = {
name: selection.name,
notes: data?.notes || "",
filePath: data?.file_path || "",
};
for (const infoNode of infoNodes) {
// Discard if a newer request has been issued for this node
if (infoNode.__loraInfoReqId !== reqIdSnapshot.get(infoNode)) {
continue;
}
if (typeof infoNode._setLoraInfo === "function") {
infoNode._setLoraInfo(infoData);
} else {
infoNode.__pendingLoraInfo = infoData;
}
}
} catch (error) {
console.error("Error fetching notes for lora info:", error);
const errorData = {
name: selection.name,
notes: "[Error loading notes]",
filePath: "",
};
for (const infoNode of infoNodes) {
if (infoNode.__loraInfoReqId !== reqIdSnapshot.get(infoNode)) {
continue;
}
if (typeof infoNode._setLoraInfo === "function") {
infoNode._setLoraInfo(errorData);
} else {
infoNode.__pendingLoraInfo = errorData;
}
}
}
}
app.registerExtension({
name: "LoraManager.LoraInfo",
beforeRegisterNodeDef(nodeType, nodeData) {
if (nodeData.name !== LORA_INFO_CLASS) {
return;
}
chainCallback(nodeType.prototype, "onNodeCreated", function () {
// Add wire-only input for receiving connections from LoRA nodes
this.addInput("lora_source", "*", { shape: 7 });
// Forward lora info data to the Vue widget when available.
this._setLoraInfo = function (data) {
const widget = this.widgets?.find(
(w) => w.type === "LORA_INFO_DISPLAY"
);
if (widget) {
if (typeof widget._setLoraInfo === "function") {
widget._setLoraInfo(data);
} else {
widget.__pendingLoraInfo = data;
}
}
};
});
// When the lora_source wire is disconnected, clear the display.
const origOnConnectionsChange = nodeType.prototype.onConnectionsChange;
nodeType.prototype.onConnectionsChange = function (type, index, connected, link_info) {
if (origOnConnectionsChange) {
origOnConnectionsChange.apply(this, arguments);
}
// type 1 = input connection change; disconnected = !connected
if (type === 1 && !connected) {
const input = this.inputs?.[index];
if (input?.name === "lora_source") {
// Check if any lora_source input still has a connection
const hasLoraSourceConnection = this.inputs?.some(
(inp) => inp.name === "lora_source" && inp.link != null
);
if (!hasLoraSourceConnection) {
this._setLoraInfo?.(null);
}
}
}
};
},
});
+5 -2
View File
@@ -13,6 +13,7 @@ import {
import { addLorasWidget } from "./loras_widget.js";
import { applyLoraValuesToText, debounce } from "./lora_syntax_utils.js";
import { applySelectionHighlight } from "./trigger_word_highlight.js";
import { updateConnectedLoraInfoNodes } from "./lora_info.js";
app.registerExtension({
name: "LoraManager.LoraLoader",
@@ -185,8 +186,10 @@ app.registerExtension({
this,
"loras",
{
onSelectionChange: (selection) =>
applySelectionHighlight(this, selection),
onSelectionChange: (selection) => {
applySelectionHighlight(this, selection);
updateConnectedLoraInfoNodes(this, selection);
},
},
(value) => {
// Prevent recursive calls

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