mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-14 09:43:22 -03:00
Compare commits
127 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 8e45c22d7a | |||
| 191c4e03cd | |||
| ab4154c57d | |||
| 28e93d12ff | |||
| 75e63c758b | |||
| 823f71f269 | |||
| 042dd4088d | |||
| eaa791a9eb | |||
| 2228627ff4 | |||
| 4c647ad9c8 | |||
| 8ca3e6c33f | |||
| dd6bdbf297 | |||
| b47dde87e4 | |||
| 99e65cccd8 | |||
| 3bdacb8f46 | |||
| b4f9c224d3 | |||
| 5ec0399c81 | |||
| b464fdc333 | |||
| 53825500db | |||
| f2ac790752 | |||
| 0d8805cdee | |||
| 656e24ac9b | |||
| 6718b37403 | |||
| c9e5e784fc | |||
| f92f958682 | |||
| f63fab0676 | |||
| cfc4903c0c | |||
| a527a847fe | |||
| 91b0bf8933 | |||
| 66d1c96783 | |||
| 986128076e | |||
| 1de0a53241 | |||
| 0ec7eaf606 | |||
| d9fcb0e92b | |||
| f49b4ba4db | |||
| 84e708328b | |||
| 125bed3f09 | |||
| 077e70169d | |||
| e6dc169a05 | |||
| f34c02756d | |||
| 1e4c315481 | |||
| a8283a0d00 | |||
| 55896669fc | |||
| e341e0b9d2 | |||
| e6538c83bb | |||
| 92e1285ea5 | |||
| 2aabd1d90e | |||
| 7b8b778f83 | |||
| 7c8dc57d55 | |||
| fe95fae5f2 | |||
| ce8a95abf7 | |||
| c8e7e543d6 | |||
| a9dbb15ffa | |||
| cf64043f7d | |||
| ccaff92c18 | |||
| 585b5c922a | |||
| ea80c2224c | |||
| 8b0f56c1a6 | |||
| 8022d12f03 | |||
| 3939f7f91b | |||
| aebf2e37dd | |||
| f53f859a71 | |||
| d916375abe | |||
| 57983df4bd | |||
| c68d7559a0 | |||
| 9a8f5bf2d6 | |||
| a8d742b031 | |||
| c27e4d1bfc | |||
| d15a8aa9a2 | |||
| 74a7d12ca4 | |||
| 2f94a9773e | |||
| 37bdfa21ea | |||
| f0bf2728c9 | |||
| dc715aa273 | |||
| 7ee2361e87 | |||
| e04c22f83f | |||
| 681cc13e90 | |||
| 090e0297d4 | |||
| 6f71335be4 | |||
| 7f51812c1e | |||
| a9dc4d7b9d | |||
| 5d50ddb5d4 | |||
| f86198d234 | |||
| ffe65d983c | |||
| b0b5be913c | |||
| 01efcbc584 | |||
| 02c249917a | |||
| 419bbc90b2 | |||
| b0c4510fdb | |||
| bf6a614e0d | |||
| feab01cd9c | |||
| 966024e534 | |||
| 2018722cc8 | |||
| 9d85c2a44a | |||
| 03dd047e62 | |||
| 86b547c1e0 | |||
| bab9752c8b | |||
| 774cc1be86 | |||
| 234b73c8a2 | |||
| abd06c48f4 | |||
| 6ca411e4e4 | |||
| 6470021e77 | |||
| 71658ab37b | |||
| 4f016a8024 | |||
| f362ed585b | |||
| 196172624f | |||
| 316702b7ab | |||
| a7625b009f | |||
| 5d4a33c90d | |||
| 041a6b8525 | |||
| 2638109ad6 | |||
| b019326747 | |||
| 54b44131b6 | |||
| a1d948025c | |||
| a90b2514ba | |||
| cb4ad27813 | |||
| 637831248b | |||
| 00228deaaa | |||
| 2373edf73c | |||
| e0e1b804a7 | |||
| fecbe8241f | |||
| 5983eaa1ce | |||
| 8bee8f4069 | |||
| 817fe21b3e | |||
| 3494037d20 | |||
| 3c83e78d9f | |||
| d7291f73c9 |
@@ -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
|
||||
|
||||
@@ -136,7 +137,13 @@ npm run test:coverage # Generate coverage report
|
||||
- Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json)
|
||||
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
||||
- Run `python scripts/sync_translation_keys.py` after adding UI strings to `locales/en.json`
|
||||
- Symlinks require normalized paths
|
||||
- Symlinks require normalized paths.
|
||||
**Business paths vs real paths**: All stored paths and operation routing use the
|
||||
original paths as they appear under configured model roots — symlinks are NOT
|
||||
resolved. `os.path.realpath` is only for scanner dedup and the symlink cache.
|
||||
Any path passed to `os.remove`/`os.rename`/`shutil.move` or validated by a
|
||||
containment check MUST use the business path (i.e. `os.path.abspath`, not
|
||||
`realpath`).
|
||||
|
||||
## Git / Commit Messages
|
||||
|
||||
|
||||
+18
@@ -15,6 +15,10 @@ 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.nodes.metadata_overwrite import MetadataOverwriteLM
|
||||
from .py.metadata_collector import init as init_metadata_collector
|
||||
except (
|
||||
ImportError
|
||||
@@ -56,6 +60,16 @@ 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
|
||||
MetadataOverwriteLM = importlib.import_module(
|
||||
"py.nodes.metadata_overwrite"
|
||||
).MetadataOverwriteLM
|
||||
init_metadata_collector = importlib.import_module("py.metadata_collector").init
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -75,6 +89,10 @@ NODE_CLASS_MAPPINGS = {
|
||||
LoraPoolLM.NAME: LoraPoolLM,
|
||||
LoraRandomizerLM.NAME: LoraRandomizerLM,
|
||||
LoraCyclerLM.NAME: LoraCyclerLM,
|
||||
LoraInfoLM.NAME: LoraInfoLM,
|
||||
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
|
||||
CreateHookLoraLM.NAME: CreateHookLoraLM,
|
||||
MetadataOverwriteLM.NAME: MetadataOverwriteLM,
|
||||
}
|
||||
|
||||
WEB_DIRECTORY = "./web/comfyui"
|
||||
|
||||
+346
-295
File diff suppressed because it is too large
Load Diff
@@ -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
+2221
-2178
File diff suppressed because it is too large
Load Diff
+50
-7
@@ -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",
|
||||
@@ -704,7 +714,9 @@
|
||||
"versionsCount": "Local Versions",
|
||||
"versionsCountDesc": "Most versions first",
|
||||
"versionsCountAsc": "Fewest versions first",
|
||||
"versionIdDesc": "Newest version first"
|
||||
"versionIdDesc": "Newest version first",
|
||||
"random": "Random",
|
||||
"randomAction": "Randomize (shuffle)"
|
||||
},
|
||||
"refresh": {
|
||||
"title": "Refresh model list",
|
||||
@@ -786,7 +798,9 @@
|
||||
"contextMenu": {
|
||||
"refreshMetadata": "Refresh Civitai Data",
|
||||
"checkUpdates": "Check Updates",
|
||||
"relinkCivitai": "Re-link to Civitai",
|
||||
"linkModel": "Link Model",
|
||||
"linkCivitai": "Link to Civitai",
|
||||
"linkHuggingFace": "Link to HuggingFace",
|
||||
"copySyntax": "Copy LoRA Syntax",
|
||||
"copyFilename": "Copy Model Filename",
|
||||
"copyRecipeSyntax": "Copy Recipe Syntax",
|
||||
@@ -1203,7 +1217,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:",
|
||||
@@ -1319,6 +1335,14 @@
|
||||
"pathPlaceholder": "Type folder path or select from tree below...",
|
||||
"root": "Root"
|
||||
},
|
||||
"linkHuggingFace": {
|
||||
"title": "Link to HuggingFace",
|
||||
"infoText": "Paste the HuggingFace repository URL to associate this model with its source. This enables AI-powered metadata enrichment.",
|
||||
"urlLabel": "HuggingFace Repository URL:",
|
||||
"urlPlaceholder": "https://huggingface.co/user/repo",
|
||||
"helpText": "Enter the full URL of the HuggingFace repository.",
|
||||
"confirmAction": "Save & Link"
|
||||
},
|
||||
"relinkCivitai": {
|
||||
"title": "Re-link to Civitai",
|
||||
"warning": "Warning:",
|
||||
@@ -1526,6 +1550,7 @@
|
||||
"empty": "No version history available for this model yet.",
|
||||
"error": "Failed to load versions.",
|
||||
"missingModelId": "This model is missing a Civitai model id.",
|
||||
"hfGroupInfo": "This is a HuggingFace model group. Open the library to see all versions in the grid.",
|
||||
"confirm": {
|
||||
"delete": "Delete this version from your library?"
|
||||
},
|
||||
@@ -1729,6 +1754,12 @@
|
||||
"checkingMessage": "Please wait while we check for the latest version.",
|
||||
"showNotifications": "Show update notifications",
|
||||
"latestBadge": "Latest",
|
||||
"latestMain": "Latest main",
|
||||
"channel": "Update Channel",
|
||||
"channels": {
|
||||
"release": "Release",
|
||||
"nightly": "Nightly"
|
||||
},
|
||||
"updateProgress": {
|
||||
"preparing": "Preparing update...",
|
||||
"installing": "Installing update...",
|
||||
@@ -1749,6 +1780,15 @@
|
||||
"warning": "Warning: Nightly builds may contain experimental features and could be unstable.",
|
||||
"enable": "Enable Nightly Updates"
|
||||
},
|
||||
"channelSwitch": {
|
||||
"nightlyTitle": "Switch to Nightly Channel",
|
||||
"nightlyMessage": "Switching to Nightly will initialize a Git repository and track the latest main branch commits. Updates will be more frequent but may be unstable. You can switch back to Release at any time.",
|
||||
"releaseTitle": "Switch to Release Channel",
|
||||
"releaseMessage": "Switching to Release will checkout the latest stable release tag. You can switch back to Nightly at any time.",
|
||||
"switching": "Switching to {channel} channel...",
|
||||
"completed": "Successfully switched to {channel} channel",
|
||||
"failed": "Failed to switch channel"
|
||||
},
|
||||
"banners": {
|
||||
"recent": "Recent messages",
|
||||
"empty": "No recent banners yet.",
|
||||
@@ -2003,7 +2043,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",
|
||||
@@ -2048,6 +2089,8 @@
|
||||
"contentRatingFailed": "Failed to set content rating: {message}",
|
||||
"relinkSuccess": "Model successfully re-linked to Civitai",
|
||||
"relinkFailed": "Error: {message}",
|
||||
"linkHfSuccess": "Model successfully linked to HuggingFace",
|
||||
"linkHfFailed": "Error: {message}",
|
||||
"fetchMetadataFirst": "Please fetch metadata from CivitAI first",
|
||||
"noCivitaiInfo": "No CivitAI information available",
|
||||
"missingHash": "Model hash not available"
|
||||
|
||||
+2221
-2178
File diff suppressed because it is too large
Load Diff
+2221
-2178
File diff suppressed because it is too large
Load Diff
+2221
-2178
File diff suppressed because it is too large
Load Diff
+2221
-2178
File diff suppressed because it is too large
Load Diff
+2221
-2178
File diff suppressed because it is too large
Load Diff
+2221
-2178
File diff suppressed because it is too large
Load Diff
+2221
-2178
File diff suppressed because it is too large
Load Diff
+2221
-2178
File diff suppressed because it is too large
Load Diff
+47
-4
@@ -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 []),
|
||||
|
||||
@@ -1,5 +1,11 @@
|
||||
"""Constants used by the metadata collector"""
|
||||
|
||||
# Sentinel value for clip_skip to distinguish "unconnected / widget default"
|
||||
# from "user wired value 0". Both ComfyUI CLIPSetLastLayer (-24..-1) and
|
||||
# A1111 conventions treat 0 as meaningless for clip skipping, but users may
|
||||
# explicitly wire 0 to the overwrite node to express "no clip skip / default".
|
||||
CLIP_SKIP_SENTINEL = -25
|
||||
|
||||
# Metadata categories
|
||||
MODELS = "models"
|
||||
PROMPTS = "prompts"
|
||||
@@ -9,6 +15,14 @@ EMBEDDINGS = "embeddings"
|
||||
SIZE = "size"
|
||||
IMAGES = "images"
|
||||
IS_SAMPLER = "is_sampler" # New constant to mark sampler nodes
|
||||
OVERWRITE = "overwrite" # Manual metadata overwrite from MetadataOverwriteLM node
|
||||
|
||||
# Field names that the MetadataOverwriteLM node and its extractor share
|
||||
METADATA_OVERWRITE_FIELDS = (
|
||||
"prompt", "negative_prompt", "seed", "steps", "cfg_scale",
|
||||
"sampler", "scheduler", "model", "loras", "size",
|
||||
"clip_skip", "additional_data",
|
||||
)
|
||||
|
||||
# Complete list of categories to track
|
||||
METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES]
|
||||
METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES, OVERWRITE]
|
||||
|
||||
@@ -83,7 +83,8 @@ class MetadataHook:
|
||||
|
||||
# Record inputs before execution
|
||||
if node_id is not None:
|
||||
registry.record_node_execution(node_id, class_type, input_data_all, None)
|
||||
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
|
||||
except Exception as e:
|
||||
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
||||
|
||||
@@ -114,7 +115,8 @@ class MetadataHook:
|
||||
|
||||
# Record outputs after execution
|
||||
if node_id is not None:
|
||||
registry.update_node_execution(node_id, class_type, results)
|
||||
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
|
||||
except Exception as e:
|
||||
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
||||
|
||||
@@ -135,10 +137,13 @@ class MetadataHook:
|
||||
# Store the dynprompt reference for node lookups
|
||||
if hasattr(prompt, 'original_prompt'):
|
||||
registry.set_current_prompt(prompt)
|
||||
|
||||
|
||||
# Store extra_data for accessing full workflow node properties
|
||||
registry.set_extra_data(extra_data)
|
||||
|
||||
# Execute the original function
|
||||
return original_execute(*args, **kwargs)
|
||||
|
||||
|
||||
# Replace the functions
|
||||
execution._map_node_over_list = map_node_over_list_with_metadata
|
||||
execution.execute = execute_with_prompt_tracking
|
||||
@@ -163,7 +168,8 @@ class MetadataHook:
|
||||
class_type = obj.__class__.__name__
|
||||
node_id = unique_id
|
||||
if node_id is not None:
|
||||
registry.record_node_execution(node_id, class_type, input_data_all, None)
|
||||
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
|
||||
except Exception as e:
|
||||
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
||||
|
||||
@@ -180,7 +186,8 @@ class MetadataHook:
|
||||
class_type = obj.__class__.__name__
|
||||
node_id = unique_id
|
||||
if node_id is not None:
|
||||
registry.update_node_execution(node_id, class_type, results)
|
||||
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
|
||||
except Exception as e:
|
||||
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
||||
|
||||
@@ -202,6 +209,9 @@ class MetadataHook:
|
||||
if hasattr(prompt, 'original_prompt'):
|
||||
registry.set_current_prompt(prompt)
|
||||
|
||||
# Store extra_data for accessing full workflow node properties
|
||||
registry.set_extra_data(extra_data)
|
||||
|
||||
# Execute the original function
|
||||
return await original_execute(*args, **kwargs)
|
||||
|
||||
|
||||
@@ -1,15 +1,68 @@
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from .constants import IMAGES
|
||||
|
||||
# Check if running in standalone mode
|
||||
standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1" or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
|
||||
|
||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER
|
||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER, OVERWRITE
|
||||
from .node_extractors import NODE_EXTRACTORS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Keys that identify metadata hint marks stored in node.properties.lm_marker_role
|
||||
_META_MARK_PREFIX = "meta_"
|
||||
_MARK_PRIMARY_MODEL = "primary_model"
|
||||
_MARK_PRIMARY_SAMPLER = "primary_sampler"
|
||||
_MARK_POSITIVE_PROMPT = "positive_prompt"
|
||||
_MARK_NEGATIVE_PROMPT = "negative_prompt"
|
||||
|
||||
class MetadataProcessor:
|
||||
"""Process and format collected metadata"""
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _get_user_marks(metadata):
|
||||
"""Scan workflow nodes (from extra_data.extra_pnginfo.workflow) for user-assigned
|
||||
metadata hint marks stored in node.properties.lm_marker_role.
|
||||
|
||||
Returns a dict mapping mark type keys to node IDs.
|
||||
Example: {'primary_model': '42', 'primary_sampler': '17'}
|
||||
"""
|
||||
marks: dict[str, str] = {}
|
||||
|
||||
# Primary source: extra_data.extra_pnginfo.workflow.nodes (has full properties)
|
||||
extra_data = metadata.get("extra_data")
|
||||
if extra_data and isinstance(extra_data, dict):
|
||||
extra_pnginfo = extra_data.get("extra_pnginfo", {})
|
||||
if isinstance(extra_pnginfo, dict):
|
||||
workflow = extra_pnginfo.get("workflow", {})
|
||||
nodes = workflow.get("nodes", [])
|
||||
for node in nodes:
|
||||
node_id = str(node.get("id", ""))
|
||||
role = node.get("properties", {}).get("lm_marker_role", "")
|
||||
if role.startswith(_META_MARK_PREFIX):
|
||||
mark_type = role[len(_META_MARK_PREFIX):]
|
||||
if mark_type in marks:
|
||||
logger.warning(
|
||||
"Duplicate meta hint '%s': node %s (previous: %s), "
|
||||
"last match wins",
|
||||
mark_type, node_id, marks[mark_type],
|
||||
)
|
||||
marks[mark_type] = node_id
|
||||
|
||||
# Fallback: try prompt.original_prompt (API-only submissions may not have workflow)
|
||||
if not marks:
|
||||
prompt = metadata.get("current_prompt")
|
||||
if prompt and getattr(prompt, "original_prompt", None):
|
||||
for node_id, node_data in prompt.original_prompt.items():
|
||||
role = node_data.get("properties", {}).get("lm_marker_role", "")
|
||||
if role.startswith(_META_MARK_PREFIX):
|
||||
mark_type = role[len(_META_MARK_PREFIX):]
|
||||
marks[mark_type] = node_id
|
||||
|
||||
return marks
|
||||
|
||||
@staticmethod
|
||||
def find_primary_sampler(metadata, downstream_id=None):
|
||||
"""
|
||||
@@ -471,20 +524,57 @@ class MetadataProcessor:
|
||||
"checkpoint": None,
|
||||
"loras": "",
|
||||
"size": None,
|
||||
"clip_skip": None
|
||||
"clip_skip": None,
|
||||
"additional_data": "",
|
||||
}
|
||||
|
||||
# Get the prompt object for node relationship tracing
|
||||
prompt = metadata.get("current_prompt")
|
||||
|
||||
# Find the primary KSampler node
|
||||
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id)
|
||||
|
||||
# Directly get checkpoint from metadata instead of tracing
|
||||
# Pass primary_sampler_id to avoid redundant calculation
|
||||
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id)
|
||||
if checkpoint:
|
||||
params["checkpoint"] = checkpoint
|
||||
|
||||
# ---- User marks: override heuristic inference with user-assigned hints ----
|
||||
user_marks = MetadataProcessor._get_user_marks(metadata)
|
||||
|
||||
# Find the primary KSampler node (user mark takes priority)
|
||||
primary_sampler_id = None
|
||||
primary_sampler = None
|
||||
if _MARK_PRIMARY_SAMPLER in user_marks:
|
||||
marked_id = user_marks[_MARK_PRIMARY_SAMPLER]
|
||||
sampler_data = metadata.get(SAMPLING, {}).get(marked_id)
|
||||
if sampler_data and sampler_data.get(IS_SAMPLER):
|
||||
primary_sampler_id = marked_id
|
||||
primary_sampler = sampler_data
|
||||
else:
|
||||
logger.warning(
|
||||
"User-marked primary sampler %s has no runtime metadata, "
|
||||
"falling back to heuristic",
|
||||
marked_id,
|
||||
)
|
||||
if primary_sampler is None:
|
||||
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id)
|
||||
|
||||
# Resolve checkpoint / model (user mark takes priority)
|
||||
if _MARK_PRIMARY_MODEL in user_marks:
|
||||
marked_id = user_marks[_MARK_PRIMARY_MODEL]
|
||||
if marked_id in metadata.get(MODELS, {}):
|
||||
params["checkpoint"] = metadata[MODELS][marked_id].get("name")
|
||||
else:
|
||||
extra_data = metadata.get("extra_data")
|
||||
extra_pnginfo = extra_data.get("extra_pnginfo", {}) if extra_data and isinstance(extra_data, dict) else {}
|
||||
workflow = extra_pnginfo.get("workflow", {}) if isinstance(extra_pnginfo, dict) else {}
|
||||
node_type = "unknown"
|
||||
for n in workflow.get("nodes", []):
|
||||
if str(n.get("id", "")) == marked_id:
|
||||
node_type = n.get("type", "unknown")
|
||||
break
|
||||
logger.warning(
|
||||
"User-marked primary model %s (type=%s, registered=%s) has no runtime metadata, "
|
||||
"falling back to heuristic",
|
||||
marked_id, node_type, node_type in NODE_EXTRACTORS,
|
||||
)
|
||||
if params["checkpoint"] is None:
|
||||
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id)
|
||||
if checkpoint:
|
||||
params["checkpoint"] = checkpoint
|
||||
|
||||
# Check if guidance parameter exists in any sampling node
|
||||
for node_id, sampler_info in metadata.get(SAMPLING, {}).items():
|
||||
@@ -539,7 +629,22 @@ class MetadataProcessor:
|
||||
|
||||
# For SamplerCustom, handle any additional parameters
|
||||
MetadataProcessor.handle_custom_advanced_sampler(metadata, prompt, primary_sampler_id, params)
|
||||
|
||||
|
||||
# ---- User marks: override prompts with explicitly tagged nodes ----
|
||||
prompts_data = metadata.get(PROMPTS, {})
|
||||
if _MARK_POSITIVE_PROMPT in user_marks:
|
||||
pos_id = user_marks[_MARK_POSITIVE_PROMPT]
|
||||
if pos_id in prompts_data:
|
||||
prompt_text = prompts_data[pos_id].get("text") or prompts_data[pos_id].get("positive_text")
|
||||
if prompt_text:
|
||||
params["prompt"] = prompt_text
|
||||
if _MARK_NEGATIVE_PROMPT in user_marks:
|
||||
neg_id = user_marks[_MARK_NEGATIVE_PROMPT]
|
||||
if neg_id in prompts_data:
|
||||
prompt_text = prompts_data[neg_id].get("text") or prompts_data[neg_id].get("negative_text")
|
||||
if prompt_text:
|
||||
params["negative_prompt"] = prompt_text
|
||||
|
||||
# Size extraction is same for all sampler types
|
||||
# Check if the sampler itself has size information (from latent_image)
|
||||
if primary_sampler_id in metadata.get(SIZE, {}):
|
||||
@@ -568,7 +673,26 @@ class MetadataProcessor:
|
||||
break
|
||||
if params["clip_skip"] is None:
|
||||
params["clip_skip"] = "1"
|
||||
|
||||
|
||||
# ---- Apply manual metadata overwrites ----
|
||||
for overwrite_info in metadata.get(OVERWRITE, {}).values():
|
||||
overwrite_params = overwrite_info.get("parameters", {})
|
||||
for key, value in overwrite_params.items():
|
||||
if key == "clip_skip":
|
||||
# Accept any value from overwrite node (sentinel -25 already
|
||||
# filtered upstream). Needed because falsy check treats 0
|
||||
# as "not set" even though 0 is a valid wired input here.
|
||||
params[key] = value
|
||||
elif value: # truthy check — only overwrite when user provided a real value
|
||||
params[key] = value
|
||||
|
||||
# Bridge: the overwrite node exposes the field as "model" (more accurate),
|
||||
# but the internal pipeline key remains "checkpoint" for backward compatibility
|
||||
# with A1111 metadata format and downstream consumers.
|
||||
if params.get("model"):
|
||||
params["checkpoint"] = params["model"]
|
||||
del params["model"]
|
||||
|
||||
return params
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import time
|
||||
from nodes import NODE_CLASS_MAPPINGS # type: ignore
|
||||
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
|
||||
from .constants import METADATA_CATEGORIES, IMAGES
|
||||
from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE
|
||||
|
||||
|
||||
class MetadataRegistry:
|
||||
@@ -61,6 +61,7 @@ class MetadataRegistry:
|
||||
{
|
||||
"execution_order": [],
|
||||
"current_prompt": None, # Will store the prompt object
|
||||
"extra_data": None, # Will store the API extra_data for workflow metadata
|
||||
"timestamp": time.time(),
|
||||
}
|
||||
)
|
||||
@@ -75,6 +76,11 @@ class MetadataRegistry:
|
||||
# Store the prompt in the metadata for later relationship tracing
|
||||
self.prompt_metadata[self.current_prompt_id]["current_prompt"] = prompt
|
||||
|
||||
def set_extra_data(self, extra_data):
|
||||
"""Store the API extra_data (contains extra_pnginfo.workflow with node properties)"""
|
||||
if self.current_prompt_id and self.current_prompt_id in self.prompt_metadata:
|
||||
self.prompt_metadata[self.current_prompt_id]["extra_data"] = extra_data
|
||||
|
||||
def get_metadata(self, prompt_id=None):
|
||||
"""Get collected metadata for a prompt"""
|
||||
key = prompt_id if prompt_id is not None else self.current_prompt_id
|
||||
@@ -122,20 +128,28 @@ class MetadataRegistry:
|
||||
cache_key = f"{node_id}:{class_type}"
|
||||
|
||||
# Check if this node type is relevant for metadata collection
|
||||
if class_type in NODE_EXTRACTORS:
|
||||
if class_type in NODE_EXTRACTORS or cache_key in self.node_cache:
|
||||
# Check if we have cached metadata for this node
|
||||
if cache_key in self.node_cache:
|
||||
cached_data = self.node_cache[cache_key]
|
||||
|
||||
# Detect bypass (mode=4) / mute (mode=2) — these nodes
|
||||
# were intentionally disabled and should not contribute
|
||||
# overwrite values from a previous execution's cache.
|
||||
node_mode = node_data.get("mode", 0)
|
||||
node_is_disabled = node_mode in (2, 4)
|
||||
|
||||
# Apply cached metadata to the current metadata
|
||||
for category in self.metadata_categories:
|
||||
if category == OVERWRITE and node_is_disabled:
|
||||
continue
|
||||
if category in cached_data and node_id in cached_data[category]:
|
||||
if node_id not in metadata[category]:
|
||||
metadata[category][node_id] = cached_data[category][
|
||||
node_id
|
||||
]
|
||||
|
||||
def record_node_execution(self, node_id, class_type, inputs, outputs):
|
||||
def record_node_execution(self, node_id, class_type, inputs, outputs, return_types=None):
|
||||
"""Record information about a node's execution"""
|
||||
if not self.current_prompt_id:
|
||||
return
|
||||
@@ -158,17 +172,18 @@ class MetadataRegistry:
|
||||
|
||||
# Extract node-specific metadata
|
||||
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
||||
extractor.extract(
|
||||
node_id,
|
||||
processed_inputs,
|
||||
outputs,
|
||||
self.prompt_metadata[self.current_prompt_id],
|
||||
)
|
||||
if extractor is GenericNodeExtractor:
|
||||
extractor.extract(node_id, processed_inputs, outputs,
|
||||
self.prompt_metadata[self.current_prompt_id],
|
||||
return_types=return_types)
|
||||
else:
|
||||
extractor.extract(node_id, processed_inputs, outputs,
|
||||
self.prompt_metadata[self.current_prompt_id])
|
||||
|
||||
# Cache this node's metadata
|
||||
self._cache_node_metadata(node_id, class_type)
|
||||
|
||||
def update_node_execution(self, node_id, class_type, outputs):
|
||||
def update_node_execution(self, node_id, class_type, outputs, return_types=None):
|
||||
"""Update node metadata with output information"""
|
||||
if not self.current_prompt_id:
|
||||
return
|
||||
@@ -179,9 +194,17 @@ class MetadataRegistry:
|
||||
# Use the same extractor to update with outputs
|
||||
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
||||
if hasattr(extractor, "update"):
|
||||
extractor.update(
|
||||
node_id, processed_outputs, self.prompt_metadata[self.current_prompt_id]
|
||||
)
|
||||
if extractor is GenericNodeExtractor:
|
||||
extractor.update(
|
||||
node_id, processed_outputs,
|
||||
self.prompt_metadata[self.current_prompt_id],
|
||||
return_types=return_types,
|
||||
)
|
||||
else:
|
||||
extractor.update(
|
||||
node_id, processed_outputs,
|
||||
self.prompt_metadata[self.current_prompt_id],
|
||||
)
|
||||
|
||||
# Update the cached metadata for this node
|
||||
self._cache_node_metadata(node_id, class_type)
|
||||
|
||||
@@ -2,7 +2,8 @@ import json
|
||||
import os
|
||||
import re
|
||||
|
||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER
|
||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE
|
||||
from .overwrite_utils import collect_overwrite_params
|
||||
|
||||
|
||||
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
||||
@@ -31,11 +32,78 @@ class NodeMetadataExtractor:
|
||||
pass
|
||||
|
||||
class GenericNodeExtractor(NodeMetadataExtractor):
|
||||
"""Default extractor for nodes without specific handling"""
|
||||
"""Fallback extractor with type-signature-based detection.
|
||||
|
||||
When a node is not in the NODE_EXTRACTORS registry, the hook layer
|
||||
passes ``return_types`` from ``obj.RETURN_TYPES``:
|
||||
|
||||
* ``MODEL`` output: common input fields (ckpt_name, unet_name, etc.)
|
||||
are checked for a model file name and stored as checkpoint metadata.
|
||||
* ``CONDITIONING`` output: common text input fields are checked for
|
||||
prompt text and stored as prompt metadata.
|
||||
"""
|
||||
|
||||
# Input field names that carry a model path in loader-style nodes.
|
||||
_MODEL_NAME_FIELDS = (
|
||||
"ckpt_name", "unet_name", "model_path", "model_name", "gguf_name",
|
||||
)
|
||||
|
||||
# Extensions used by checkpoint_scanner.py — only record values that look
|
||||
# like real model filenames to avoid capturing unrelated string fields.
|
||||
_MODEL_EXTENSIONS = {
|
||||
".ckpt", ".pt", ".pt2", ".bin", ".pth", ".safetensors", ".pkl", ".sft", ".gguf",
|
||||
}
|
||||
|
||||
# Input field names that may carry prompt text in encoder-style nodes.
|
||||
_TEXT_FIELDS = ("text", "clip_l", "t5xxl", "prompt", "positive", "negative")
|
||||
|
||||
@staticmethod
|
||||
def extract(node_id, inputs, outputs, metadata):
|
||||
pass
|
||||
|
||||
def extract(node_id, inputs, outputs, metadata, return_types=None):
|
||||
if return_types is None:
|
||||
return
|
||||
|
||||
# — MODEL loader detection (checkpoint / UNET / GGUF) —
|
||||
if "MODEL" in return_types or any("MODEL" in str(t) for t in return_types):
|
||||
for field in GenericNodeExtractor._MODEL_NAME_FIELDS:
|
||||
val = inputs.get(field)
|
||||
if val and isinstance(val, str) and val.strip():
|
||||
name = val.strip()
|
||||
if not any(name.lower().endswith(ext) for ext in GenericNodeExtractor._MODEL_EXTENSIONS):
|
||||
continue
|
||||
_store_checkpoint_metadata(metadata, node_id, name)
|
||||
return
|
||||
|
||||
# — CONDITIONING encoder detection (CLIPTextEncode, Flux, custom) —
|
||||
if "CONDITIONING" in return_types or any("CONDITIONING" in str(t) for t in return_types):
|
||||
text = None
|
||||
for field in GenericNodeExtractor._TEXT_FIELDS:
|
||||
val = inputs.get(field)
|
||||
if val and isinstance(val, str) and val.strip():
|
||||
text = val.strip()
|
||||
break
|
||||
if text:
|
||||
prompt_data = metadata.setdefault(PROMPTS, {})
|
||||
prompt_data[node_id] = {
|
||||
"text": text,
|
||||
"node_id": node_id,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def update(node_id, outputs, metadata, return_types=None):
|
||||
if return_types is None:
|
||||
return
|
||||
if "CONDITIONING" not in return_types and not any(
|
||||
"CONDITIONING" in str(t) for t in return_types
|
||||
):
|
||||
return
|
||||
if node_id not in metadata.get(PROMPTS, {}):
|
||||
return
|
||||
if outputs and isinstance(outputs, list) and len(outputs) > 0:
|
||||
if isinstance(outputs[0], tuple) and len(outputs[0]) > 0:
|
||||
cond = outputs[0][0]
|
||||
if cond is not None:
|
||||
metadata[PROMPTS][node_id]["conditioning"] = cond
|
||||
|
||||
class CheckpointLoaderExtractor(NodeMetadataExtractor):
|
||||
@staticmethod
|
||||
def extract(node_id, inputs, outputs, metadata):
|
||||
@@ -1154,6 +1222,28 @@ class CR_ApplyControlNetStackExtractor(NodeMetadataExtractor):
|
||||
metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive
|
||||
metadata[PROMPTS][node_id]["negative_encoded"] = transformed_negative
|
||||
|
||||
class MetadataOverwriteExtractor(NodeMetadataExtractor):
|
||||
"""Extract manually specified metadata from MetadataOverwriteLM node.
|
||||
|
||||
Stores truthy input values under the OVERWRITE category so that
|
||||
extract_generation_params can merge them over the inferred params.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def extract(node_id, inputs, outputs, metadata):
|
||||
if not inputs:
|
||||
return
|
||||
|
||||
overwrite_params = collect_overwrite_params(inputs)
|
||||
|
||||
if overwrite_params:
|
||||
metadata.setdefault(OVERWRITE, {})
|
||||
metadata[OVERWRITE][node_id] = {
|
||||
"parameters": overwrite_params,
|
||||
"node_id": node_id,
|
||||
}
|
||||
|
||||
|
||||
# Registry of node-specific extractors
|
||||
# Keys are node class names
|
||||
NODE_EXTRACTORS = {
|
||||
@@ -1221,5 +1311,7 @@ NODE_EXTRACTORS = {
|
||||
"CFGGuider": CFGGuiderExtractor, # Add CFGGuider
|
||||
# Image
|
||||
"VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor
|
||||
# Metadata overwrite
|
||||
"MetadataOverwriteLM": MetadataOverwriteExtractor,
|
||||
# Add other nodes as needed
|
||||
}
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
"""Shared helpers for Metadata Overwrite node metadata collection.
|
||||
|
||||
Used by both the MetadataOverwriteLM node (execution time) and the
|
||||
MetadataOverwriteExtractor (hook time) so the conversion/filtering logic
|
||||
cannot drift between the two paths.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Any, Dict
|
||||
|
||||
from ..utils.utils import model_patcher_to_name
|
||||
from .constants import CLIP_SKIP_SENTINEL, METADATA_OVERWRITE_FIELDS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def collect_overwrite_params(values: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Convert node input values into non-default overwrite parameters.
|
||||
|
||||
For most fields, a falsy value (empty string, 0) means "not set" and is
|
||||
skipped. clip_skip uses a dedicated sentinel (-25) so that a wired value
|
||||
of 0 is preserved. The ``model`` field accepts either a manual string or
|
||||
a wired MODEL (ModelPatcher) connection; in the latter case the source
|
||||
model name is extracted from the patcher's ``cached_patcher_init`` and
|
||||
stored as a ComfyUI-style relative path.
|
||||
"""
|
||||
result: Dict[str, Any] = {}
|
||||
for key in METADATA_OVERWRITE_FIELDS:
|
||||
value = values.get(key)
|
||||
if key == "model" and not isinstance(value, str):
|
||||
value = model_patcher_to_name(value)
|
||||
if value is None:
|
||||
logger.warning(
|
||||
"Could not extract model name from wired MODEL input "
|
||||
"(no cached_patcher_init); model metadata overwrite skipped"
|
||||
)
|
||||
if key == "clip_skip":
|
||||
if value != CLIP_SKIP_SENTINEL:
|
||||
result[key] = value
|
||||
elif value:
|
||||
result[key] = value
|
||||
return result
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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
@@ -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"],
|
||||
|
||||
@@ -1,26 +1,102 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
_STACK_INPUT_PATTERN = re.compile(r"^lora_stack(?:_([ab])|(\d+))$")
|
||||
|
||||
|
||||
def _is_stack_input(name: str) -> bool:
|
||||
return bool(_STACK_INPUT_PATTERN.match(name))
|
||||
|
||||
|
||||
def _stack_slot_number(name: str) -> int:
|
||||
"""Numeric slot used to order stack inputs; legacy a/b map to 1/2."""
|
||||
match = _STACK_INPUT_PATTERN.match(name)
|
||||
if not match:
|
||||
return -1
|
||||
letter, digits = match.group(1), match.group(2)
|
||||
if digits is not None:
|
||||
return int(digits)
|
||||
return 1 if letter == "a" else 2
|
||||
|
||||
|
||||
class _LoraStackOptionalInputs:
|
||||
"""Lookup that preserves explicit optional inputs and dynamic lora_stack slots."""
|
||||
|
||||
def __init__(self, explicit_inputs: dict[str, tuple[str, dict[str, Any]]]) -> None:
|
||||
self._explicit_inputs = explicit_inputs
|
||||
|
||||
def __contains__(self, item: object) -> bool:
|
||||
if not isinstance(item, str):
|
||||
return False
|
||||
return item in self._explicit_inputs or _is_stack_input(item)
|
||||
|
||||
def __getitem__(self, key: str) -> tuple[str, dict[str, Any]]:
|
||||
if key in self._explicit_inputs:
|
||||
return self._explicit_inputs[key]
|
||||
if _is_stack_input(key):
|
||||
return (
|
||||
"LORA_STACK",
|
||||
{
|
||||
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||
},
|
||||
)
|
||||
raise KeyError(key)
|
||||
|
||||
|
||||
class LoraStackCombinerLM:
|
||||
NAME = "Lora Stack Combiner (LoraManager)"
|
||||
CATEGORY = "Lora Manager/stackers"
|
||||
DESCRIPTION = (
|
||||
"Combines multiple LoRA stacks into a single stack. "
|
||||
"Supports dynamic inputs: connect a stack to add more inputs."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
optional_inputs: dict[str, tuple[str, dict[str, Any]]] = {
|
||||
"lora_stack1": (
|
||||
"LORA_STACK",
|
||||
{
|
||||
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||
},
|
||||
),
|
||||
"lora_stack2": (
|
||||
"LORA_STACK",
|
||||
{
|
||||
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||
},
|
||||
),
|
||||
}
|
||||
|
||||
stack = inspect.stack()
|
||||
if len(stack) > 2 and stack[2].function == "get_input_info":
|
||||
optional_inputs = _LoraStackOptionalInputs(optional_inputs) # type: ignore[assignment]
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"lora_stack_a": ("LORA_STACK",),
|
||||
"lora_stack_b": ("LORA_STACK",),
|
||||
},
|
||||
"required": {},
|
||||
"optional": optional_inputs,
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LORA_STACK",)
|
||||
RETURN_NAMES = ("LORA_STACK",)
|
||||
FUNCTION = "combine_stacks"
|
||||
|
||||
def combine_stacks(self, lora_stack_a, lora_stack_b):
|
||||
combined_stack = []
|
||||
def combine_stacks(self, lora_stack1=None, lora_stack2=None, **kwargs):
|
||||
stacks = {
|
||||
"lora_stack1": lora_stack1,
|
||||
"lora_stack2": lora_stack2,
|
||||
}
|
||||
for key, value in kwargs.items():
|
||||
if _is_stack_input(key) and value is not None:
|
||||
stacks[key] = value
|
||||
|
||||
if lora_stack_a:
|
||||
combined_stack.extend(lora_stack_a)
|
||||
if lora_stack_b:
|
||||
combined_stack.extend(lora_stack_b)
|
||||
combined_stack = []
|
||||
for key in sorted(stacks, key=_stack_slot_number):
|
||||
stack = stacks[key]
|
||||
if stack:
|
||||
combined_stack.extend(stack)
|
||||
|
||||
return (combined_stack,)
|
||||
|
||||
@@ -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),)
|
||||
@@ -0,0 +1,169 @@
|
||||
"""Metadata Overwrite node — allows users to manually specify generation parameters
|
||||
that override the automatically collected/inferred metadata.
|
||||
|
||||
Most inputs have falsy defaults (empty string / 0) which are skipped.
|
||||
clip_skip uses a sentinel default (-25) so that a wired value of 0 is
|
||||
preserved — both ComfyUI and A1111 conventions have no meaningful 0 value,
|
||||
but users may wire 0 to express "no clip skip / default".
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..metadata_collector.constants import CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL
|
||||
from ..metadata_collector.overwrite_utils import collect_overwrite_params
|
||||
|
||||
|
||||
class MetadataOverwriteLM:
|
||||
NAME = "Metadata Overwrite (LoraManager)"
|
||||
CATEGORY = "Lora Manager/utils"
|
||||
DESCRIPTION = (
|
||||
"Manually specify generation parameters to override automatically collected "
|
||||
"metadata. Only filled/connected inputs will take effect — empty defaults "
|
||||
"are ignored."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"optional": {
|
||||
"prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Positive prompt. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"negative_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Negative prompt. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"seed": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 0xFFFFFFFFFFFFFFFF,
|
||||
"control_after_generate": False,
|
||||
"tooltip": "Seed value. Only overwrites when > 0.",
|
||||
},
|
||||
),
|
||||
"steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 10000,
|
||||
"tooltip": "Number of steps. Only overwrites when > 0.",
|
||||
},
|
||||
),
|
||||
"cfg_scale": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0,
|
||||
"tooltip": "CFG scale. Only overwrites when > 0.",
|
||||
},
|
||||
),
|
||||
"sampler": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Sampler name. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"scheduler": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Scheduler name. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"model": (
|
||||
"STRING,MODEL",
|
||||
{
|
||||
"default": "",
|
||||
"widgetType": "STRING",
|
||||
"tooltip": (
|
||||
"The checkpoint or diffusion model (UNet) used "
|
||||
"for generation. Fill in the name manually or "
|
||||
"connect a MODEL output — the model name is then "
|
||||
"extracted automatically. Only overwrites when "
|
||||
"non-empty."
|
||||
),
|
||||
},
|
||||
),
|
||||
"loras": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": (
|
||||
"LoRA syntax, e.g. <lora:name:strength> "
|
||||
"or <lora:name:model_strength:clip_strength>, "
|
||||
"separated by spaces. Only overwrites when non-empty."
|
||||
),
|
||||
},
|
||||
),
|
||||
"size": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"Image size in WIDTHxHEIGHT format (e.g. 512x768). "
|
||||
"Only overwrites when non-empty."
|
||||
),
|
||||
},
|
||||
),
|
||||
"clip_skip": (
|
||||
"INT",
|
||||
{
|
||||
"default": _CLIP_SKIP_SENTINEL,
|
||||
"min": -25,
|
||||
"max": 24,
|
||||
"tooltip": (
|
||||
"Clip skip (ComfyUI: -24..-1, A1111: 1+). "
|
||||
"Default -25 means not set — any other value "
|
||||
"overwrites."
|
||||
),
|
||||
},
|
||||
),
|
||||
"additional_data": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": (
|
||||
"Additional data to embed in the image metadata. "
|
||||
"Inserted between Clip skip and Model hash in the "
|
||||
"A1111-compatible parameters string. "
|
||||
'Example: "Copyright": "Some license info"'
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("METADATA",)
|
||||
RETURN_NAMES = ("metadata",)
|
||||
FUNCTION = "collect_metadata"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def collect_metadata(self, **kwargs: Any) -> tuple[dict[str, Any]]:
|
||||
"""Collect non-default input values into a metadata dict.
|
||||
|
||||
For most fields, a falsy value (empty string, 0) means "not set"
|
||||
and is skipped. clip_skip uses a dedicated sentinel (-25) so that
|
||||
a wired value of 0 is preserved and reaches the metadata pipeline.
|
||||
|
||||
The ``model`` field accepts either a manual string or a wired MODEL
|
||||
(ModelPatcher) connection; in the latter case the underlying model
|
||||
name is extracted from the patcher's ``cached_patcher_init`` and
|
||||
stored as a ComfyUI-style relative path.
|
||||
"""
|
||||
return (collect_overwrite_params(kwargs),)
|
||||
+346
-127
@@ -16,6 +16,156 @@ from PIL import Image, PngImagePlugin
|
||||
import piexif
|
||||
import logging
|
||||
|
||||
# Civitai-compatible sampler name mapping: ComfyUI internal → A1111 display name
|
||||
CIVITAI_SAMPLER_MAP = {
|
||||
"euler": "Euler",
|
||||
"euler_ancestral": "Euler a",
|
||||
"lms": "LMS",
|
||||
"heun": "Heun",
|
||||
"dpm_2": "DPM2",
|
||||
"dpm_2_ancestral": "DPM2 a",
|
||||
"dpmpp_2s_ancestral": "DPM++ 2S a",
|
||||
"dpmpp_2m": "DPM++ 2M",
|
||||
"dpmpp_sde": "DPM++ SDE",
|
||||
"dpmpp_sde_gpu": "DPM++ SDE",
|
||||
"dpmpp_2m_sde": "DPM++ 2M SDE",
|
||||
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
|
||||
"dpmpp_3m_sde": "DPM++ 3M SDE",
|
||||
"dpm_fast": "DPM fast",
|
||||
"dpm_adaptive": "DPM adaptive",
|
||||
"ddim": "DDIM",
|
||||
"plms": "PLMS",
|
||||
"uni_pc_bh2": "UniPC",
|
||||
"uni_pc": "UniPC",
|
||||
"lcm": "LCM",
|
||||
}
|
||||
|
||||
# Base model display name → AIR URN slug
|
||||
# Sourced from civitai source: src/shared/constants/basemodel.constants.ts
|
||||
BASE_MODEL_AIR_SLUG = {
|
||||
# Stable Diffusion family
|
||||
"SD 1.4": "sd1",
|
||||
"SD 1.5": "sd1",
|
||||
"SD 1.5 LCM": "sd1",
|
||||
"SD 1.5 Hyper": "sd1",
|
||||
"SD 2.0": "sd2",
|
||||
"SD 2.0 768": "sd2",
|
||||
"SD 2.1": "sd2",
|
||||
"SD 2.1 768": "sd2",
|
||||
"SD 2.1 Unclip": "sd2",
|
||||
"SD 3.0": "sd3",
|
||||
"SD 3.5": "sd35",
|
||||
"SD 3.5 Large": "sd35",
|
||||
"SD 3.5 Large Turbo": "sd35",
|
||||
"SD 3.5 Medium": "sd35",
|
||||
"SDXL 0.9": "sdxl",
|
||||
"SDXL 1.0": "sdxl",
|
||||
"SDXL 1.0 LCM": "sdxl",
|
||||
"SDXL Lightning": "sdxl",
|
||||
"SDXL Hyper": "sdxl",
|
||||
"SDXL Turbo": "sdxl",
|
||||
"SDXL Distilled": "sdxldistilled",
|
||||
"Stable Cascade": "scascade",
|
||||
"Stable Video Diffusion": "svd",
|
||||
"SVD": "svd",
|
||||
"SVD XT": "svdxt",
|
||||
|
||||
# SDXL community fine-tunes
|
||||
"Pony": "pony",
|
||||
"Pony Diffusion": "pony",
|
||||
"Illustrious": "illustrious",
|
||||
"NoobAI": "noobai",
|
||||
"Animagine": "illustrious",
|
||||
|
||||
# Flux family
|
||||
"Flux.1": "flux1",
|
||||
"Flux.1 D": "flux1",
|
||||
"Flux.1 S": "flux1",
|
||||
"Flux.1 Krea": "fluxkrea",
|
||||
"Flux.1 Kontext": "flux1kontext",
|
||||
"Flux.2": "flux2",
|
||||
"Flux.2 D": "flux2",
|
||||
"Flux.2 Klein 9B": "flux2klein_9b",
|
||||
"Flux.2 Klein 9B Base": "flux2klein_9b_base",
|
||||
"Flux.2 Klein 4B": "flux2klein_4b",
|
||||
"Flux.2 Klein 4B Base": "flux2klein_4b_base",
|
||||
|
||||
# Other image models (sorted alphabetically)
|
||||
"AuraFlow": "auraflow",
|
||||
"Chroma": "chroma",
|
||||
"HiDream": "hidream",
|
||||
"HiDream-O1": "hidream-o1",
|
||||
"Hunyuan DiT": "hydit1",
|
||||
"Hunyuan Video": "hyv1",
|
||||
"Kolors": "kolors",
|
||||
"Lumina": "lumina",
|
||||
"Mochi": "mochi",
|
||||
"ODOR": "odor",
|
||||
"PixArt Alpha": "pixarta",
|
||||
"PixArt Sigma": "pixarte",
|
||||
"Playground v2": "playgroundv2",
|
||||
"Playground v2.5": "playgroundv2",
|
||||
"Pony Diffusion V7": "ponyv7",
|
||||
|
||||
# Video models
|
||||
"CogVideoX": "cogvideox",
|
||||
"LTX Video": "ltxv",
|
||||
"LTX Video 2": "ltxv2",
|
||||
"LTX Video 2.3": "ltxv23",
|
||||
"Wan Video": "wanvideo",
|
||||
"Wan Video 1.3B T2V": "wanvideo_13b_t2v",
|
||||
"Wan Video 14B T2V": "wanvideo_14b_t2v",
|
||||
"Wan Video 14B I2V 480p": "wanvideo_14b_i2v_480p",
|
||||
"Wan Video 14B I2V 720p": "wanvideo_14b_i2v_720p",
|
||||
|
||||
# Third-party / proprietary image models
|
||||
"Boogu": "boogu",
|
||||
"Ernie": "ernie",
|
||||
"Grok": "grok",
|
||||
"HappyHorse": "happyhorse",
|
||||
"Ideogram": "ideogram",
|
||||
"Ideogram 4.0": "ideogram",
|
||||
"Imagen": "imagen4",
|
||||
"Imagen 4": "imagen4",
|
||||
"Krea": "krea2",
|
||||
"Krea 2": "krea2",
|
||||
"Lens": "lens",
|
||||
"MAI": "mai",
|
||||
"Nano Banana": "nanobanana",
|
||||
"OpenAI": "openai",
|
||||
"Reve": "reve",
|
||||
"Reve 2": "reve",
|
||||
"Reve 2.1": "reve",
|
||||
"Seedream": "seedream",
|
||||
"Sora": "sora2",
|
||||
"Sora 2": "sora2",
|
||||
"Veo": "veo3",
|
||||
"Veo 2": "veo3",
|
||||
"Veo 3": "veo3",
|
||||
"ZImageTurbo": "zimageturbo",
|
||||
"ZImageBase": "zimagebase",
|
||||
"ZImage": "zimagebase",
|
||||
|
||||
# Third-party video models
|
||||
"Hailuo by MiniMax": "minimax",
|
||||
"Haiper": "haiper",
|
||||
"Kling": "kling",
|
||||
"Lightricks": "lightricks",
|
||||
"Seedance": "seedance",
|
||||
"Vidu": "vidu",
|
||||
|
||||
# Qwen family
|
||||
"Qwen": "qwen",
|
||||
"Qwen 2": "qwen2",
|
||||
|
||||
# Anima
|
||||
"Anima": "anima",
|
||||
|
||||
# Special
|
||||
"Upscaler": "upscaler",
|
||||
"Other": "other",
|
||||
}
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -70,11 +220,29 @@ class SaveImageLM:
|
||||
"tooltip": "Compression quality for JPEG and lossy WebP formats (1-100). Higher values mean better quality but larger files.",
|
||||
},
|
||||
),
|
||||
"webp_method": (
|
||||
"INT",
|
||||
{
|
||||
"default": 6,
|
||||
"min": 0,
|
||||
"max": 6,
|
||||
"tooltip": "WebP compression method (0-6). 0=fastest/largest, 6=slowest/smallest. Only applies when file_format is 'webp'.",
|
||||
},
|
||||
),
|
||||
"jpeg_subsampling": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 2,
|
||||
"tooltip": "JPEG chroma subsampling level. 0=4:4:4 (best quality), 1=4:2:2, 2=4:2:0 (smallest files). Only applies when file_format is 'jpeg'.",
|
||||
},
|
||||
),
|
||||
"embed_workflow": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": False,
|
||||
"tooltip": "Embeds the complete workflow data into the image metadata. Only works with PNG and WebP formats.",
|
||||
"tooltip": "When enabled, saved images store the complete workflow. Drag the image back into ComfyUI to restore the original node graph. PNG and WebP only.",
|
||||
},
|
||||
),
|
||||
"save_with_metadata": (
|
||||
@@ -142,148 +310,194 @@ class SaveImageLM:
|
||||
|
||||
return None
|
||||
|
||||
def format_metadata(self, metadata_dict):
|
||||
"""Format metadata in the requested format similar to userComment example"""
|
||||
if not metadata_dict:
|
||||
return ""
|
||||
def _resolve_model_cache_entry(self, scanner_type: str, name: str):
|
||||
"""Resolve model hash, civitai metadata, and base_model from scanner cache.
|
||||
Returns (hash_str, civitai_dict, base_model_str). All values are empty defaults when not found."""
|
||||
scanner = ServiceRegistry.get_service_sync(scanner_type)
|
||||
if scanner is None or not name:
|
||||
return "", {}, ""
|
||||
|
||||
# Helper function to only add parameter if value is not None
|
||||
def add_param_if_not_none(param_list, label, value):
|
||||
if value is not None:
|
||||
param_list.append(f"{label}: {value}")
|
||||
entry = self._get_cached_model_by_name(scanner, name)
|
||||
if entry is None:
|
||||
basename = os.path.splitext(os.path.basename(name))[0]
|
||||
hash_val = scanner.get_hash_by_filename(basename)
|
||||
return (hash_val or "").lower(), {}, ""
|
||||
|
||||
hash_val = (entry.get("sha256") or "").lower()
|
||||
civitai = entry.get("civitai") or {}
|
||||
base_model = entry.get("base_model") or ""
|
||||
return hash_val, civitai, base_model
|
||||
|
||||
@staticmethod
|
||||
def _get_civitai_sampler_name(sampler_name: str, scheduler: str) -> str:
|
||||
if sampler_name in CIVITAI_SAMPLER_MAP:
|
||||
civitai_name = CIVITAI_SAMPLER_MAP[sampler_name]
|
||||
if scheduler == "karras":
|
||||
civitai_name += " Karras"
|
||||
elif scheduler == "exponential":
|
||||
civitai_name += " Exponential"
|
||||
return civitai_name
|
||||
else:
|
||||
if scheduler and scheduler != "normal":
|
||||
return f"{sampler_name}_{scheduler}"
|
||||
return sampler_name
|
||||
|
||||
@staticmethod
|
||||
def _build_air_string(base_model: str, model_type: str, model_id: int, version_id: int) -> str:
|
||||
slug = BASE_MODEL_AIR_SLUG.get(base_model, "other")
|
||||
type_lower = model_type.lower() if model_type else "other"
|
||||
return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}"
|
||||
|
||||
def format_metadata(self, metadata_dict: dict) -> str:
|
||||
"""Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources."""
|
||||
if not metadata_dict: return ""
|
||||
|
||||
# Extract the prompt and negative prompt
|
||||
prompt = metadata_dict.get("prompt", "")
|
||||
negative_prompt = metadata_dict.get("negative_prompt", "")
|
||||
|
||||
# Extract loras from the prompt if present
|
||||
steps = metadata_dict.get("steps")
|
||||
cfg = metadata_dict.get("guidance")
|
||||
if cfg is None:
|
||||
cfg = metadata_dict.get("cfg_scale")
|
||||
if cfg is None:
|
||||
cfg = metadata_dict.get("cfg")
|
||||
seed = metadata_dict.get("seed")
|
||||
size = metadata_dict.get("size")
|
||||
sampler = metadata_dict.get("sampler") or ""
|
||||
scheduler = metadata_dict.get("scheduler") or "normal"
|
||||
checkpoint = metadata_dict.get("checkpoint") or ""
|
||||
loras_text = metadata_dict.get("loras", "")
|
||||
lora_hashes = {}
|
||||
clip_skip = metadata_dict.get("clip_skip")
|
||||
|
||||
# If loras are found, add them on a new line after the prompt
|
||||
# Parse LoRA entries from <lora:name:strength> format
|
||||
lora_entries: list[tuple[str, float]] = []
|
||||
if loras_text:
|
||||
prompt_with_loras = f"{prompt}\n{loras_text}"
|
||||
for match in re.findall(r"<lora:([^:]+):([^>]+)>", loras_text):
|
||||
lora_name, strength_str = match
|
||||
try:
|
||||
strength = float(strength_str)
|
||||
except (ValueError, TypeError):
|
||||
strength = 1.0
|
||||
lora_entries.append((lora_name, strength))
|
||||
|
||||
# Extract lora names from the format <lora:name:strength>
|
||||
lora_matches = re.findall(r"<lora:([^:]+):([^>]+)>", loras_text)
|
||||
# Resolve checkpoint hash and Civitai data from local cache
|
||||
ckpt_hash, ckpt_civitai, ckpt_base_model = "", {}, ""
|
||||
ckpt_display_name = ""
|
||||
if checkpoint:
|
||||
ckpt_hash, ckpt_civitai, ckpt_base_model = self._resolve_model_cache_entry(
|
||||
"checkpoint_scanner", checkpoint
|
||||
)
|
||||
ckpt_display_name = os.path.splitext(os.path.basename(checkpoint))[0]
|
||||
|
||||
# Get hash for each lora
|
||||
for lora_name, strength in lora_matches:
|
||||
hash_value = self.get_lora_hash(lora_name)
|
||||
if hash_value:
|
||||
lora_hashes[lora_name] = hash_value
|
||||
else:
|
||||
prompt_with_loras = prompt
|
||||
# Resolve LoRA hash and Civitai data from local cache
|
||||
loras_data: list[dict] = []
|
||||
for lora_name, strength in lora_entries:
|
||||
lora_hash, lora_civitai, lora_base_model = self._resolve_model_cache_entry(
|
||||
"lora_scanner", lora_name
|
||||
)
|
||||
loras_data.append({
|
||||
"name": lora_name,
|
||||
"strength": strength,
|
||||
"hash": lora_hash,
|
||||
"civitai": lora_civitai,
|
||||
"base_model": lora_base_model,
|
||||
})
|
||||
|
||||
# Format the first part (prompt and loras)
|
||||
metadata_parts = [prompt_with_loras]
|
||||
# Build Hashes JSON (A1111 / Civitai standard format)
|
||||
hashes: dict[str, str] = {}
|
||||
if ckpt_hash:
|
||||
hashes["model"] = ckpt_hash[:10].upper()
|
||||
for lora in loras_data:
|
||||
if lora["hash"]:
|
||||
hashes[f"LORA:{lora['name']}"] = lora["hash"][:10].upper()
|
||||
|
||||
# Add negative prompt
|
||||
# Build Civitai resources JSON array
|
||||
civitai_resources: list[dict] = []
|
||||
if ckpt_civitai.get("id", 0) > 0:
|
||||
ckpt_resource: dict = {}
|
||||
ckpt_type = (ckpt_civitai.get("model") or {}).get("type", "Checkpoint")
|
||||
model_id = ckpt_civitai.get("modelId", 0)
|
||||
version_id = ckpt_civitai.get("id", 0)
|
||||
if model_id and version_id:
|
||||
ckpt_resource["air"] = self._build_air_string(
|
||||
ckpt_base_model, ckpt_type, int(model_id), int(version_id)
|
||||
)
|
||||
elif version_id:
|
||||
ckpt_resource["modelVersionId"] = int(version_id)
|
||||
if ckpt_civitai.get("name"):
|
||||
ckpt_resource["versionName"] = ckpt_civitai["name"]
|
||||
if ckpt_resource:
|
||||
civitai_resources.append(ckpt_resource)
|
||||
|
||||
for lora in loras_data:
|
||||
lora_civitai = lora["civitai"]
|
||||
if not lora_civitai or lora_civitai.get("id", 0) <= 0:
|
||||
continue
|
||||
lora_resource: dict = {"weight": lora["strength"]}
|
||||
lora_type = (lora_civitai.get("model") or {}).get("type", "LORA")
|
||||
model_id = lora_civitai.get("modelId", 0)
|
||||
version_id = lora_civitai.get("id", 0)
|
||||
if model_id and version_id:
|
||||
lora_resource["air"] = self._build_air_string(
|
||||
lora["base_model"], lora_type, int(model_id), int(version_id)
|
||||
)
|
||||
elif version_id:
|
||||
lora_resource["modelVersionId"] = int(version_id)
|
||||
if lora_civitai.get("name"):
|
||||
lora_resource["versionName"] = lora_civitai["name"]
|
||||
civitai_resources.append(lora_resource)
|
||||
|
||||
sampler_name = CIVITAI_SAMPLER_MAP.get(sampler, sampler) if sampler else None
|
||||
|
||||
scheduler_mapping = {
|
||||
"normal": "Normal",
|
||||
"karras": "Karras",
|
||||
"exponential": "Exponential",
|
||||
"sgm_uniform": "SGM Uniform",
|
||||
"sgm_quadratic": "SGM Quadratic",
|
||||
}
|
||||
scheduler_name = scheduler_mapping.get(scheduler, scheduler) if scheduler else None
|
||||
|
||||
# Build output lines
|
||||
lines = [prompt] if prompt else [""]
|
||||
if negative_prompt:
|
||||
metadata_parts.append(f"Negative prompt: {negative_prompt}")
|
||||
lines.append(f"Negative prompt: {negative_prompt}")
|
||||
|
||||
# Format the second part (generation parameters)
|
||||
params = []
|
||||
|
||||
# Add standard parameters in the correct order
|
||||
if "steps" in metadata_dict:
|
||||
add_param_if_not_none(params, "Steps", metadata_dict.get("steps"))
|
||||
|
||||
# Combine sampler and scheduler information
|
||||
sampler_name = None
|
||||
scheduler_name = None
|
||||
|
||||
if "sampler" in metadata_dict:
|
||||
sampler = metadata_dict.get("sampler")
|
||||
# Convert ComfyUI sampler names to user-friendly names
|
||||
sampler_mapping = {
|
||||
"euler": "Euler",
|
||||
"euler_ancestral": "Euler a",
|
||||
"dpm_2": "DPM2",
|
||||
"dpm_2_ancestral": "DPM2 a",
|
||||
"heun": "Heun",
|
||||
"dpm_fast": "DPM fast",
|
||||
"dpm_adaptive": "DPM adaptive",
|
||||
"lms": "LMS",
|
||||
"dpmpp_2s_ancestral": "DPM++ 2S a",
|
||||
"dpmpp_sde": "DPM++ SDE",
|
||||
"dpmpp_sde_gpu": "DPM++ SDE",
|
||||
"dpmpp_2m": "DPM++ 2M",
|
||||
"dpmpp_2m_sde": "DPM++ 2M SDE",
|
||||
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
|
||||
"ddim": "DDIM",
|
||||
}
|
||||
sampler_name = sampler_mapping.get(sampler, sampler)
|
||||
|
||||
if "scheduler" in metadata_dict:
|
||||
scheduler = metadata_dict.get("scheduler")
|
||||
scheduler_mapping = {
|
||||
"normal": "Simple",
|
||||
"karras": "Karras",
|
||||
"exponential": "Exponential",
|
||||
"sgm_uniform": "SGM Uniform",
|
||||
"sgm_quadratic": "SGM Quadratic",
|
||||
}
|
||||
scheduler_name = scheduler_mapping.get(scheduler, scheduler)
|
||||
|
||||
# Add combined sampler and scheduler information
|
||||
params: list[str] = []
|
||||
if steps is not None:
|
||||
params.append(f"Steps: {steps}")
|
||||
if sampler_name:
|
||||
if scheduler_name:
|
||||
params.append(f"Sampler: {sampler_name} {scheduler_name}")
|
||||
else:
|
||||
params.append(f"Sampler: {sampler_name}")
|
||||
if cfg is not None:
|
||||
params.append(f"CFG scale: {cfg}")
|
||||
if seed is not None:
|
||||
params.append(f"Seed: {seed}")
|
||||
if size:
|
||||
params.append(f"Size: {size}")
|
||||
if clip_skip is not None:
|
||||
try:
|
||||
params.append(f"Clip skip: {abs(int(clip_skip))}")
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
additional_data = metadata_dict.get("additional_data", "")
|
||||
if additional_data:
|
||||
params.append(additional_data)
|
||||
if ckpt_hash:
|
||||
params.append(f"Model hash: {ckpt_hash[:10].upper()}")
|
||||
if ckpt_display_name:
|
||||
params.append(f"Model: {ckpt_display_name}")
|
||||
if hashes:
|
||||
params.append(f"Hashes: {json.dumps(hashes, separators=(',', ':'))}")
|
||||
params.append("Version: ComfyUI")
|
||||
if civitai_resources:
|
||||
params.append(
|
||||
f"Civitai resources: {json.dumps(civitai_resources, separators=(',', ':'))}"
|
||||
)
|
||||
|
||||
# CFG scale (Use guidance if available, otherwise fall back to cfg_scale or cfg)
|
||||
if "guidance" in metadata_dict:
|
||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("guidance"))
|
||||
elif "cfg_scale" in metadata_dict:
|
||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("cfg_scale"))
|
||||
elif "cfg" in metadata_dict:
|
||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("cfg"))
|
||||
|
||||
# Seed
|
||||
if "seed" in metadata_dict:
|
||||
add_param_if_not_none(params, "Seed", metadata_dict.get("seed"))
|
||||
|
||||
# Size
|
||||
if "size" in metadata_dict:
|
||||
add_param_if_not_none(params, "Size", metadata_dict.get("size"))
|
||||
|
||||
# Model info
|
||||
if "checkpoint" in metadata_dict:
|
||||
# Ensure checkpoint is a string before processing
|
||||
checkpoint = metadata_dict.get("checkpoint")
|
||||
if checkpoint is not None:
|
||||
# Get model hash
|
||||
model_hash = self.get_checkpoint_hash(checkpoint)
|
||||
|
||||
# Extract basename without path
|
||||
checkpoint_name = os.path.basename(checkpoint)
|
||||
# Remove extension if present
|
||||
checkpoint_name = os.path.splitext(checkpoint_name)[0]
|
||||
|
||||
# Add model hash if available
|
||||
if model_hash:
|
||||
params.append(
|
||||
f"Model hash: {model_hash[:10]}, Model: {checkpoint_name}"
|
||||
)
|
||||
else:
|
||||
params.append(f"Model: {checkpoint_name}")
|
||||
|
||||
# Add LoRA hashes if available
|
||||
if lora_hashes:
|
||||
lora_hash_parts = []
|
||||
for lora_name, hash_value in lora_hashes.items():
|
||||
lora_hash_parts.append(f"{lora_name}: {hash_value[:10]}")
|
||||
|
||||
if lora_hash_parts:
|
||||
params.append(f'Lora hashes: "{", ".join(lora_hash_parts)}"')
|
||||
|
||||
# Combine all parameters with commas
|
||||
metadata_parts.append(", ".join(params))
|
||||
|
||||
# Join all parts with a new line
|
||||
return "\n".join(metadata_parts)
|
||||
lines.append(", ".join(params))
|
||||
return "\n".join(lines)
|
||||
|
||||
# credit to nkchocoai
|
||||
# Add format_filename method to handle pattern substitution
|
||||
@@ -573,6 +787,8 @@ class SaveImageLM:
|
||||
extra_pnginfo=None,
|
||||
lossless_webp=True,
|
||||
quality=100,
|
||||
webp_method=6,
|
||||
jpeg_subsampling=0,
|
||||
embed_workflow=False,
|
||||
save_with_metadata=True,
|
||||
add_counter_to_filename=True,
|
||||
@@ -627,15 +843,14 @@ class SaveImageLM:
|
||||
elif file_format == "jpeg":
|
||||
file = base_filename + ".jpg"
|
||||
file_extension = ".jpg"
|
||||
save_kwargs = {"quality": quality, "optimize": True}
|
||||
save_kwargs = {"quality": quality, "optimize": True, "subsampling": jpeg_subsampling}
|
||||
elif file_format == "webp":
|
||||
file = base_filename + ".webp"
|
||||
file_extension = ".webp"
|
||||
# Add optimization param to control performance
|
||||
save_kwargs = {
|
||||
"quality": quality,
|
||||
"lossless": lossless_webp,
|
||||
"method": 0,
|
||||
"method": webp_method,
|
||||
}
|
||||
else:
|
||||
raise ValueError(f"Unsupported file format: {file_format}")
|
||||
@@ -722,6 +937,8 @@ class SaveImageLM:
|
||||
extra_pnginfo=None,
|
||||
lossless_webp=True,
|
||||
quality=100,
|
||||
webp_method=6,
|
||||
jpeg_subsampling=0,
|
||||
embed_workflow=False,
|
||||
save_with_metadata=True,
|
||||
add_counter_to_filename=True,
|
||||
@@ -751,6 +968,8 @@ class SaveImageLM:
|
||||
extra_pnginfo,
|
||||
lossless_webp,
|
||||
quality,
|
||||
webp_method,
|
||||
jpeg_subsampling,
|
||||
embed_workflow,
|
||||
save_with_metadata,
|
||||
add_counter_to_filename,
|
||||
|
||||
@@ -7,6 +7,21 @@ from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_c
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _reload_gguf_unet(
|
||||
unet_path: str, weight_dtype: str, disable_dynamic: bool = False
|
||||
) -> object:
|
||||
"""Reload a GGUF diffusion model from disk (cached_patcher_init factory).
|
||||
|
||||
Mirrors the GGUF branch of UNETLoaderLM.load_unet so ModelPatcher
|
||||
deepclone/dynamic machinery can rebuild GGUF models with the correct
|
||||
GGMLOps. ``disable_dynamic`` is accepted for signature compatibility
|
||||
with core ComfyUI loaders.
|
||||
"""
|
||||
loader = UNETLoaderLM()
|
||||
model, = loader._load_gguf_unet(unet_path, unet_path, weight_dtype)
|
||||
return model
|
||||
|
||||
|
||||
class UNETLoaderLM:
|
||||
"""UNET Loader with support for extra folder paths
|
||||
|
||||
@@ -196,6 +211,12 @@ class UNETLoaderLM:
|
||||
# Wrap with GGUFModelPatcher
|
||||
model = GGUFModelPatcher.clone(model)
|
||||
|
||||
# Register a reload factory so the MODEL carries its source path
|
||||
# (cached_patcher_init) like core ComfyUI loaders do — required
|
||||
# for model-name extraction downstream and for ModelPatcher
|
||||
# deepclone/dynamic machinery.
|
||||
model.cached_patcher_init = (_reload_gguf_unet, (unet_path, weight_dtype))
|
||||
|
||||
return (model,)
|
||||
|
||||
except Exception as e:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -60,21 +60,21 @@ class AgentHandler:
|
||||
skill_name = request.match_info.get("skill_name", "")
|
||||
if not skill_name:
|
||||
return web.json_response(
|
||||
{"error": "Skill name is required"}, status_code=400
|
||||
{"error": "Skill name is required"}, status=400
|
||||
)
|
||||
|
||||
try:
|
||||
body = await request.json()
|
||||
except Exception:
|
||||
return web.json_response(
|
||||
{"error": "Invalid JSON body"}, status_code=400
|
||||
{"error": "Invalid JSON body"}, status=400
|
||||
)
|
||||
|
||||
model_paths = body.get("model_paths", [])
|
||||
if not model_paths or not isinstance(model_paths, list):
|
||||
return web.json_response(
|
||||
{"error": "model_paths must be a non-empty array"},
|
||||
status_code=400,
|
||||
status=400,
|
||||
)
|
||||
|
||||
service = await self._ensure_service()
|
||||
@@ -161,5 +161,5 @@ class AgentHandler:
|
||||
# TODO: implement cooperative cancellation in AgentService
|
||||
return web.json_response(
|
||||
{"status": "acknowledged", "note": "Cancellation not yet implemented"},
|
||||
status_code=200,
|
||||
status=200,
|
||||
)
|
||||
|
||||
@@ -122,8 +122,12 @@ async def _save_hf_metadata(dest_path: str, repo: str, model_root: str) -> None:
|
||||
metadata._unknown_fields["hf_url"] = hf_url
|
||||
metadata.from_civitai = False # HF models are not from CivitAI
|
||||
|
||||
metadata_dict = metadata.to_dict()
|
||||
if "trainedWords" in metadata_dict and not metadata_dict["trainedWords"]:
|
||||
del metadata_dict["trainedWords"]
|
||||
|
||||
# 3. Save metadata atomically
|
||||
await MetadataManager.save_metadata(dest_path, metadata)
|
||||
await MetadataManager.save_metadata(dest_path, metadata_dict)
|
||||
logger.info("Saved HF metadata (with hf_url) for %s", dest_path)
|
||||
|
||||
# 4. Determine relative folder path for cache
|
||||
@@ -147,9 +151,117 @@ async def _save_hf_metadata(dest_path: str, repo: str, model_root: str) -> None:
|
||||
logger.warning("Failed to save HF metadata for %s: %s", dest_path, exc)
|
||||
|
||||
|
||||
def _find_matching_root(dest_dir: str) -> str | None:
|
||||
"""Walk up *dest_dir* to find which configured scanner root it belongs to."""
|
||||
norm = os.path.normpath(dest_dir).replace(os.sep, "/")
|
||||
all_roots = []
|
||||
for root_list in (
|
||||
config.loras_roots or [],
|
||||
config.extra_loras_roots or [],
|
||||
config.checkpoints_roots or [],
|
||||
config.extra_checkpoints_roots or [],
|
||||
config.unet_roots or [],
|
||||
config.extra_unet_roots or [],
|
||||
config.embeddings_roots or [],
|
||||
config.extra_embeddings_roots or [],
|
||||
):
|
||||
all_roots.extend([os.path.normpath(p).replace(os.sep, "/") for p in root_list])
|
||||
# Find the longest matching prefix
|
||||
match: str | None = None
|
||||
for root in all_roots:
|
||||
if norm.startswith(root):
|
||||
if match is None or len(root) > len(match):
|
||||
match = root
|
||||
return match
|
||||
|
||||
|
||||
async def _add_to_scanner_cache(dest_path: str, metadata: dict[str, Any]) -> None:
|
||||
model_dir = os.path.dirname(dest_path)
|
||||
model_root = _find_matching_root(model_dir)
|
||||
if not model_root:
|
||||
raise ValueError(f"File path {dest_path} is not within any configured scanner root")
|
||||
scanner_getter_name = _infer_model_type(model_root)[1]
|
||||
scanner_getter = getattr(ServiceRegistry, scanner_getter_name, None)
|
||||
if scanner_getter is None:
|
||||
raise RuntimeError(f"Scanner getter '{scanner_getter_name}' not found in ServiceRegistry")
|
||||
scanner = await scanner_getter()
|
||||
if scanner is None:
|
||||
raise RuntimeError(f"Scanner '{scanner_getter_name}' returned None")
|
||||
await scanner.update_single_model_cache(dest_path, dest_path, metadata)
|
||||
|
||||
|
||||
class HfHandler:
|
||||
"""Handle Hugging Face model browsing and download."""
|
||||
|
||||
async def set_hf_url(self, request: web.Request) -> web.Response:
|
||||
try:
|
||||
payload: dict[str, Any] = await request.json()
|
||||
except json.JSONDecodeError:
|
||||
return web.json_response({"success": False, "error": "Invalid JSON"}, status=400)
|
||||
|
||||
file_path = (payload.get("file_path") or "").strip()
|
||||
hf_url = (payload.get("hf_url") or "").strip()
|
||||
|
||||
if not file_path or not hf_url:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Missing required fields: 'file_path' and 'hf_url'"},
|
||||
status=400,
|
||||
)
|
||||
|
||||
m = re.match(r"^https?://huggingface\.co/([^/]+/[^/]+)/?$", hf_url)
|
||||
if not m:
|
||||
return web.json_response(
|
||||
{
|
||||
"success": False,
|
||||
"error": "Invalid HuggingFace URL. Expected format: https://huggingface.co/user/repo",
|
||||
},
|
||||
status=400,
|
||||
)
|
||||
|
||||
if not os.path.isfile(file_path):
|
||||
return web.json_response(
|
||||
{"success": False, "error": f"File not found: {file_path}"},
|
||||
status=404,
|
||||
)
|
||||
|
||||
model_root = _find_matching_root(os.path.dirname(file_path))
|
||||
if not model_root:
|
||||
return web.json_response(
|
||||
{
|
||||
"success": False,
|
||||
"error": "File is not within any configured model directory. Cannot link to HuggingFace.",
|
||||
},
|
||||
status=400,
|
||||
)
|
||||
|
||||
try:
|
||||
existing = await MetadataManager.load_metadata_payload(file_path)
|
||||
if existing.get("hf_url") == hf_url:
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"message": "hf_url already set",
|
||||
"hf_url": hf_url,
|
||||
})
|
||||
|
||||
existing["hf_url"] = hf_url
|
||||
existing["from_civitai"] = False
|
||||
await MetadataManager.save_metadata(file_path, existing)
|
||||
|
||||
await _add_to_scanner_cache(file_path, existing)
|
||||
|
||||
logger.info("Set hf_url=%s for %s", hf_url, file_path)
|
||||
return web.json_response({
|
||||
"success": True,
|
||||
"message": f"hf_url set to {hf_url}",
|
||||
"hf_url": hf_url,
|
||||
})
|
||||
except Exception as exc:
|
||||
logger.error("Failed to set hf_url for %s: %s", file_path, exc)
|
||||
return web.json_response(
|
||||
{"success": False, "error": str(exc)},
|
||||
status=500,
|
||||
)
|
||||
|
||||
async def get_hf_repo_files(self, request: web.Request) -> web.Response:
|
||||
"""List model-weight files from a HF repo with real file sizes.
|
||||
|
||||
@@ -251,8 +363,8 @@ class HfHandler:
|
||||
if ".." in (author, repo_name) or "." in (author, repo_name):
|
||||
return web.json_response({"error": f"Invalid repo format: {repo}"}, status=400)
|
||||
|
||||
# Validate filename — must not contain path separators or ..
|
||||
if "/" in filename or "\\" in filename or ".." in filename:
|
||||
# Validate filename — must not contain path traversal
|
||||
if ".." in filename:
|
||||
return web.json_response({"error": "Invalid filename"}, status=400)
|
||||
|
||||
# Validate relative_path — must not be absolute or escape base directory
|
||||
@@ -262,35 +374,17 @@ class HfHandler:
|
||||
if ".." in relative_path.split("/") or "\\" in relative_path:
|
||||
return web.json_response({"error": "Invalid relative_path"}, status=400)
|
||||
|
||||
# Validate model_root — must not contain path traversal
|
||||
if not os.path.isabs(model_root):
|
||||
# For relative model_root, check it doesn't escape
|
||||
resolved_model_root = os.path.realpath(
|
||||
os.path.join(os.getcwd(), "models", model_root)
|
||||
)
|
||||
# Use model_root directly as the base directory — same approach as
|
||||
# CivitAI's download path (download_manager.py). No realpath, no
|
||||
# allowed-roots validation, no path-traversal check; those are
|
||||
# unnecessary when the frontend sends the path from its own dropdown
|
||||
# (populated from scanner roots). Using the "business path" directly
|
||||
# keeps dest_path consistent with scanner roots so that later folder
|
||||
# derivation (in _save_hf_metadata) works correctly.
|
||||
if os.path.isabs(model_root):
|
||||
base_dir = os.path.normpath(model_root)
|
||||
else:
|
||||
resolved_model_root = os.path.realpath(model_root)
|
||||
|
||||
# Verify model_root is within a configured scanner root
|
||||
allowed_roots = set()
|
||||
for root_list in (
|
||||
config.loras_roots or [],
|
||||
config.extra_loras_roots or [],
|
||||
config.checkpoints_roots or [],
|
||||
config.extra_checkpoints_roots or [],
|
||||
config.unet_roots or [],
|
||||
config.extra_unet_roots or [],
|
||||
config.embeddings_roots or [],
|
||||
config.extra_embeddings_roots or [],
|
||||
):
|
||||
for r in root_list:
|
||||
allowed_roots.add(os.path.realpath(r))
|
||||
|
||||
if not any(resolved_model_root == root or resolved_model_root.startswith(root + os.sep) for root in allowed_roots):
|
||||
logger.warning("Invalid model_root rejected: %s", model_root)
|
||||
return web.json_response({"error": f"Invalid model_root: {model_root}"}, status=400)
|
||||
|
||||
base_dir = resolved_model_root
|
||||
base_dir = os.path.normpath(os.path.join(os.getcwd(), "models", model_root))
|
||||
|
||||
if use_default_paths:
|
||||
target_dir = os.path.join(base_dir, "huggingface", author, repo_name)
|
||||
@@ -299,15 +393,12 @@ class HfHandler:
|
||||
else:
|
||||
target_dir = base_dir
|
||||
|
||||
os.makedirs(target_dir, exist_ok=True)
|
||||
dest_path = os.path.join(target_dir, filename)
|
||||
# Strip HF repo subdirectory — "diffusion_models/xxx.safetensors"
|
||||
# is an HF repo convention, not meaningful for local storage.
|
||||
file_base = os.path.basename(filename)
|
||||
|
||||
# Resolve symlinks and check for path traversal escape
|
||||
real_dest = os.path.realpath(dest_path)
|
||||
real_base = os.path.realpath(target_dir)
|
||||
if not real_dest.startswith(real_base + os.sep):
|
||||
logger.warning("Path traversal blocked: %s -> %s", dest_path, real_dest)
|
||||
return web.json_response({"error": "Path traversal detected"}, status=400)
|
||||
os.makedirs(target_dir, exist_ok=True)
|
||||
dest_path = os.path.join(target_dir, file_base)
|
||||
|
||||
# Check if already exists (simple skip)
|
||||
if os.path.exists(dest_path) and os.path.getsize(dest_path) > 0:
|
||||
|
||||
@@ -38,7 +38,12 @@ from ...services.settings_manager import get_settings_manager
|
||||
from ...services.websocket_manager import ws_manager
|
||||
from ...services.downloader import get_downloader
|
||||
from ...services.errors import ResourceNotFoundError
|
||||
from ...services.llm_service import get_provider_model_ids, fetch_ollama_models
|
||||
from ...services.llm_service import (
|
||||
PROVIDER_PRESETS,
|
||||
fetch_ollama_models,
|
||||
get_all_provider_models,
|
||||
get_provider_model_ids,
|
||||
)
|
||||
from ...services.cache_health_monitor import CacheHealthMonitor, CacheHealthStatus
|
||||
from ...utils.models import BaseModelMetadata
|
||||
from ...utils.constants import (
|
||||
@@ -568,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."""
|
||||
@@ -596,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] = {}
|
||||
@@ -1544,6 +1562,11 @@ class SettingsHandler:
|
||||
{"success": False, "error": validation_error}
|
||||
)
|
||||
|
||||
if key == "update_channel" and value not in ("release", "nightly"):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "update_channel must be 'release' or 'nightly'"}
|
||||
)
|
||||
|
||||
if value == "__DELETE__" and key in (
|
||||
"proxy_username",
|
||||
"proxy_password",
|
||||
@@ -1552,7 +1575,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:
|
||||
@@ -1625,6 +1652,20 @@ class SettingsHandler:
|
||||
def _is_dedicated_example_images_folder(self, folder_path: str) -> bool:
|
||||
return is_valid_example_images_root(folder_path)
|
||||
|
||||
async def get_provider_models(self, request: web.Request) -> web.Response:
|
||||
"""Return the model catalog for all preset providers.
|
||||
|
||||
This endpoint is called asynchronously by the settings UI so that
|
||||
page rendering never blocks on the remote model catalog fetch.
|
||||
"""
|
||||
catalog_provider_ids = [p for p in PROVIDER_PRESETS if p != "custom"]
|
||||
try:
|
||||
provider_models = await get_all_provider_models(catalog_provider_ids)
|
||||
return web.json_response({"success": True, "models": provider_models})
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to fetch provider models: %s", exc)
|
||||
return web.json_response({"success": False, "models": {}, "error": str(exc)})
|
||||
|
||||
|
||||
class UsageStatsHandler:
|
||||
def __init__(self, usage_stats_factory: UsageStatsFactory = UsageStats) -> None:
|
||||
@@ -1752,6 +1793,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:
|
||||
@@ -2431,6 +2590,8 @@ class ModelLibraryHandler:
|
||||
status=400,
|
||||
)
|
||||
|
||||
cursor = request.query.get("cursor")
|
||||
|
||||
metadata_provider = await self._metadata_provider_factory()
|
||||
if not metadata_provider:
|
||||
return web.json_response(
|
||||
@@ -2439,7 +2600,7 @@ class ModelLibraryHandler:
|
||||
)
|
||||
|
||||
try:
|
||||
models = await metadata_provider.get_user_models(username)
|
||||
result = await metadata_provider.get_user_models(username, cursor)
|
||||
except NotImplementedError:
|
||||
return web.json_response(
|
||||
{
|
||||
@@ -2449,14 +2610,35 @@ class ModelLibraryHandler:
|
||||
status=501,
|
||||
)
|
||||
|
||||
if models is None:
|
||||
if result is None:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Failed to fetch user models"},
|
||||
status=502,
|
||||
)
|
||||
|
||||
if isinstance(result, dict):
|
||||
models = result.get("items")
|
||||
next_cursor = result.get("nextCursor")
|
||||
else:
|
||||
# Defensive: tolerate providers that still return a raw list
|
||||
models = result
|
||||
next_cursor = None
|
||||
|
||||
if not isinstance(models, list):
|
||||
models = []
|
||||
if next_cursor is not None and not isinstance(next_cursor, str):
|
||||
next_cursor = str(next_cursor)
|
||||
|
||||
estimated_total = None
|
||||
if cursor is None:
|
||||
get_count = getattr(metadata_provider, "get_creator_model_count", None)
|
||||
if get_count is not None:
|
||||
try:
|
||||
estimated_total = await get_count(username)
|
||||
except Exception: # best-effort only
|
||||
estimated_total = None
|
||||
if not isinstance(estimated_total, int):
|
||||
estimated_total = None
|
||||
|
||||
lora_scanner = await self._service_registry.get_lora_scanner()
|
||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||
@@ -2476,6 +2658,7 @@ class ModelLibraryHandler:
|
||||
versions: list[dict] = []
|
||||
history_service = await self._get_download_history_service()
|
||||
model_ids: list[int] = []
|
||||
model_count = 0
|
||||
for model in models:
|
||||
try:
|
||||
model_ids.append(int(model.get("id")))
|
||||
@@ -2509,6 +2692,8 @@ class ModelLibraryHandler:
|
||||
if model_type not in normalized_allowed_types:
|
||||
continue
|
||||
|
||||
model_count += 1
|
||||
|
||||
scanner = type_scanner_map.get(model_type)
|
||||
if scanner is None:
|
||||
return web.json_response(
|
||||
@@ -2574,7 +2759,15 @@ class ModelLibraryHandler:
|
||||
)
|
||||
|
||||
return web.json_response(
|
||||
{"success": True, "username": username, "versions": versions}
|
||||
{
|
||||
"success": True,
|
||||
"username": username,
|
||||
"versions": versions,
|
||||
"modelCount": model_count,
|
||||
"nextCursor": next_cursor,
|
||||
"hasMore": next_cursor is not None,
|
||||
"estimatedTotal": estimated_total,
|
||||
}
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
|
||||
@@ -3097,6 +3290,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:
|
||||
@@ -3143,7 +3338,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(
|
||||
{
|
||||
@@ -3184,42 +3384,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,
|
||||
@@ -3255,7 +3514,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
|
||||
)
|
||||
@@ -3333,6 +3592,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."""
|
||||
@@ -3395,13 +3778,16 @@ class MiscHandlerSet:
|
||||
"get_settings_libraries": self.settings.get_libraries,
|
||||
"activate_library": self.settings.activate_library,
|
||||
"get_llm_models": self.settings.get_llm_models,
|
||||
"get_provider_models": self.settings.get_provider_models,
|
||||
"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,
|
||||
@@ -3428,6 +3814,7 @@ class MiscHandlerSet:
|
||||
# Hugging Face handlers
|
||||
"get_hf_repo_files": self.hf_handler.get_hf_repo_files,
|
||||
"download_hf_model": self.hf_handler.download_hf_model,
|
||||
"set_hf_url": self.hf_handler.set_hf_url,
|
||||
# Agent skill handlers
|
||||
"get_agent_skills": self.agent_handler.get_agent_skills,
|
||||
"execute_agent_skill": self.agent_handler.execute_agent_skill,
|
||||
|
||||
@@ -154,10 +154,13 @@ class ModelPageView:
|
||||
)
|
||||
self._template_env._i18n_filter_added = True # type: ignore[attr-defined]
|
||||
|
||||
from ...services.llm_service import PROVIDER_PRESETS, get_all_provider_models
|
||||
from ...services.llm_service import PROVIDER_PRESETS
|
||||
|
||||
catalog_provider_ids = [p for p in PROVIDER_PRESETS if p != "custom"]
|
||||
provider_models = await get_all_provider_models(catalog_provider_ids)
|
||||
# Provider presets are embedded directly (local, no await needed).
|
||||
# Provider model catalogs are fetched asynchronously by the
|
||||
# frontend via GET /api/lm/llm/provider-models so page rendering
|
||||
# never blocks on the remote model catalog (which can take up to
|
||||
# 30s on cold cache).
|
||||
|
||||
template_context = {
|
||||
"is_initializing": is_initializing,
|
||||
@@ -167,7 +170,7 @@ class ModelPageView:
|
||||
"t": self._server_i18n.get_translation,
|
||||
"version": self._get_app_version(),
|
||||
"provider_presets_json": json.dumps(PROVIDER_PRESETS),
|
||||
"provider_models_json": json.dumps(provider_models),
|
||||
"provider_models_json": "{}",
|
||||
}
|
||||
|
||||
if not is_initializing:
|
||||
@@ -391,12 +394,14 @@ class ModelListingHandler:
|
||||
)
|
||||
|
||||
# View-local-versions filter: show all local versions of a specific model
|
||||
# Accepts either a CivitAI modelId (int) or a HF group key like "hf:user/repo"
|
||||
civitai_model_id = request.query.get("civitai_model_id")
|
||||
if civitai_model_id is not None:
|
||||
try:
|
||||
civitai_model_id = int(civitai_model_id)
|
||||
except (TypeError, ValueError):
|
||||
civitai_model_id = None
|
||||
# Keep as string — could be an HF group key (e.g. "hf:user/repo")
|
||||
pass
|
||||
|
||||
return {
|
||||
"page": page,
|
||||
@@ -534,6 +539,7 @@ class ModelManagementHandler:
|
||||
# Update model_data with new hash
|
||||
model_data["sha256"] = sha256
|
||||
model_data["hash_status"] = "completed"
|
||||
hash_status = "completed"
|
||||
else:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "No SHA256 hash found"}, status=400
|
||||
@@ -541,6 +547,32 @@ class ModelManagementHandler:
|
||||
|
||||
await MetadataManager.hydrate_model_data(model_data)
|
||||
|
||||
# hydrate_model_data replaces model_data with .metadata.json content,
|
||||
# which may lack sha256. Restore from cache and persist the fix.
|
||||
if not model_data.get("sha256"):
|
||||
if sha256:
|
||||
model_data["sha256"] = sha256
|
||||
model_data["hash_status"] = model_data.get("hash_status", hash_status)
|
||||
data_to_save = model_data.copy()
|
||||
data_to_save.pop("folder", None)
|
||||
await MetadataManager.save_metadata(file_path, data_to_save)
|
||||
else:
|
||||
sha256 = await calculate_sha256(file_path)
|
||||
if sha256:
|
||||
model_data["sha256"] = sha256.lower()
|
||||
model_data["hash_status"] = "completed"
|
||||
data_to_save = model_data.copy()
|
||||
data_to_save.pop("folder", None)
|
||||
await MetadataManager.save_metadata(file_path, data_to_save)
|
||||
else:
|
||||
return web.json_response(
|
||||
{
|
||||
"success": False,
|
||||
"error": "Failed to compute SHA256 hash for model",
|
||||
},
|
||||
status=500,
|
||||
)
|
||||
|
||||
success, error = await self._metadata_sync.fetch_and_update_model(
|
||||
sha256=model_data["sha256"],
|
||||
file_path=file_path,
|
||||
@@ -563,7 +595,12 @@ class ModelManagementHandler:
|
||||
{"success": False, "error": OFFLINE_FRIENDLY_MESSAGE},
|
||||
status=503,
|
||||
)
|
||||
self._logger.error("Error fetching from CivitAI: %s", exc, exc_info=True)
|
||||
self._logger.error(
|
||||
"Error fetching from CivitAI for %s: %s",
|
||||
locals().get("file_path", "unknown"),
|
||||
exc,
|
||||
exc_info=True,
|
||||
)
|
||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||
|
||||
async def relink_civitai(self, request: web.Request) -> web.Response:
|
||||
@@ -970,6 +1007,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:
|
||||
@@ -978,6 +1017,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"))
|
||||
@@ -1272,9 +1327,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,
|
||||
@@ -1310,9 +1369,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.
|
||||
@@ -1769,14 +1839,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(
|
||||
@@ -1786,14 +1862,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"},
|
||||
@@ -2917,6 +2999,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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
@@ -2218,6 +2267,31 @@ class RecipeManagementHandler:
|
||||
"Failed to download image for recipe: %s", exc
|
||||
)
|
||||
|
||||
# Fallback: try to locate a custom image on disk using model_hash + image id
|
||||
if image_bytes is None:
|
||||
image_id = image_data.get("id") or ""
|
||||
if image_id and model_hash:
|
||||
from ...utils.example_images_paths import get_model_folder
|
||||
model_folder = get_model_folder(model_hash)
|
||||
if model_folder and os.path.exists(model_folder):
|
||||
for fname in os.listdir(model_folder):
|
||||
if f"custom_{image_id}" in fname:
|
||||
ext = os.path.splitext(fname)[1].lower()
|
||||
if ext not in (".jpg", ".jpeg", ".png", ".webp", ".gif"):
|
||||
continue
|
||||
fpath = os.path.join(model_folder, fname)
|
||||
if os.path.isfile(fpath):
|
||||
try:
|
||||
with open(fpath, "rb") as f:
|
||||
image_bytes = f.read()
|
||||
extension = ext
|
||||
except Exception as exc:
|
||||
self._logger.warning(
|
||||
"Failed to read custom image file %s: %s",
|
||||
fpath, exc,
|
||||
)
|
||||
break
|
||||
|
||||
prompt = (
|
||||
(parsed.get("gen_params") or {}).get("prompt") or ""
|
||||
)
|
||||
|
||||
@@ -23,6 +23,7 @@ MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||
RouteDefinition("GET", "/api/lm/settings", "get_settings"),
|
||||
RouteDefinition("POST", "/api/lm/settings", "update_settings"),
|
||||
RouteDefinition("GET", "/api/lm/llm/models", "get_llm_models"),
|
||||
RouteDefinition("GET", "/api/lm/llm/provider-models", "get_provider_models"),
|
||||
RouteDefinition("GET", "/api/lm/doctor/diagnostics", "get_doctor_diagnostics"),
|
||||
RouteDefinition("POST", "/api/lm/doctor/repair-cache", "repair_doctor_cache"),
|
||||
RouteDefinition("POST", "/api/lm/doctor/resolve-filename-conflicts", "resolve_doctor_filename_conflicts"),
|
||||
@@ -38,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"),
|
||||
@@ -102,6 +105,9 @@ MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
||||
RouteDefinition(
|
||||
"POST", "/api/lm/download-hf-model", "download_hf_model"
|
||||
),
|
||||
RouteDefinition(
|
||||
"POST", "/api/lm/set-hf-url", "set_hf_url"
|
||||
),
|
||||
# Agent skill endpoints
|
||||
RouteDefinition(
|
||||
"GET", "/api/lm/agent/skills", "get_agent_skills"
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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"),
|
||||
|
||||
+313
-45
@@ -38,6 +38,84 @@ def _clean_excludes() -> List[str]:
|
||||
return excludes
|
||||
|
||||
|
||||
def _stage_preserved_items(plugin_root: str) -> tuple[str, list[str]]:
|
||||
"""Move preserved user-data items to a temp directory outside *plugin_root*.
|
||||
|
||||
This ensures that ``git reset --hard``, ``git clean -fd``, and ZIP-based
|
||||
replacement cannot touch these files even when ``-e`` exclusion patterns
|
||||
are mishandled (e.g. on Windows where forward-slash patterns may not
|
||||
match backslash-prefixed paths in some Git builds, or where file locks
|
||||
prevent deletion/recreation).
|
||||
|
||||
Returns:
|
||||
``(backup_root, staged_names)``: the temp directory path and the
|
||||
list of item names that were successfully moved.
|
||||
"""
|
||||
backup_root = tempfile.mkdtemp(prefix='lora_manager_update_')
|
||||
staged: list[str] = []
|
||||
for name in _PRESERVE_DIRS:
|
||||
src = os.path.join(plugin_root, name)
|
||||
if not os.path.lexists(src):
|
||||
continue
|
||||
dst = os.path.join(backup_root, name)
|
||||
try:
|
||||
shutil.move(src, dst)
|
||||
staged.append(name)
|
||||
logger.debug("Staged '%s' for update safety", name)
|
||||
except OSError:
|
||||
# ``shutil.move`` may fail on Windows if a file handle inside
|
||||
# the directory is still open (e.g. a SQLite WAL file). Fall
|
||||
# back to copy-then-remove.
|
||||
logger.debug("Move failed for '%s', falling back to copy", name)
|
||||
try:
|
||||
if os.path.isdir(src) and not os.path.islink(src):
|
||||
shutil.copytree(src, dst, symlinks=True)
|
||||
shutil.rmtree(src, ignore_errors=True)
|
||||
else:
|
||||
shutil.copy2(src, dst)
|
||||
os.remove(src)
|
||||
staged.append(name)
|
||||
logger.info("Copied (then removed) '%s' for update safety", name)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Could not stage '%s': %s (will rely on git -e / skip lists)", name, exc
|
||||
)
|
||||
return backup_root, staged
|
||||
|
||||
|
||||
def _restore_preserved_items(plugin_root: str, backup_root: str, staged: list[str]) -> None:
|
||||
"""Move staged items back from *backup_root* into *plugin_root*.
|
||||
|
||||
Any leftover placeholder at the destination (created by git checkout or
|
||||
ZIP extraction) is removed before the move.
|
||||
"""
|
||||
for name in staged:
|
||||
src = os.path.join(backup_root, name)
|
||||
dst = os.path.join(plugin_root, name)
|
||||
try:
|
||||
if os.path.lexists(dst):
|
||||
if os.path.isdir(dst) and not os.path.islink(dst):
|
||||
shutil.rmtree(dst, ignore_errors=True)
|
||||
else:
|
||||
os.remove(dst)
|
||||
shutil.move(src, dst)
|
||||
logger.debug("Restored '%s' after update", name)
|
||||
except OSError:
|
||||
logger.debug("Move failed restoring '%s', falling back to copy", name)
|
||||
try:
|
||||
if os.path.isdir(src) and not os.path.islink(src):
|
||||
shutil.copytree(src, dst, symlinks=True, dirs_exist_ok=True)
|
||||
shutil.rmtree(src, ignore_errors=True)
|
||||
else:
|
||||
shutil.copy2(src, dst)
|
||||
os.remove(src)
|
||||
logger.info("Copied '%s' back after update", name)
|
||||
except Exception as exc:
|
||||
logger.error("Failed to restore '%s': %s", name, exc)
|
||||
shutil.rmtree(backup_root, ignore_errors=True)
|
||||
|
||||
|
||||
|
||||
class UpdateRoutes:
|
||||
"""Routes for handling plugin update checks"""
|
||||
|
||||
@@ -47,6 +125,7 @@ class UpdateRoutes:
|
||||
app.router.add_get('/api/lm/check-updates', UpdateRoutes.check_updates)
|
||||
app.router.add_get('/api/lm/version-info', UpdateRoutes.get_version_info)
|
||||
app.router.add_post('/api/lm/perform-update', UpdateRoutes.perform_update)
|
||||
app.router.add_post('/api/lm/switch-channel', UpdateRoutes.switch_channel)
|
||||
|
||||
@staticmethod
|
||||
async def check_updates(request):
|
||||
@@ -65,10 +144,17 @@ class UpdateRoutes:
|
||||
|
||||
# Fetch remote version from GitHub
|
||||
if nightly:
|
||||
remote_version, changelog = await UpdateRoutes._get_nightly_version()
|
||||
releases = None
|
||||
local_hash = git_info.get('short_hash', '')
|
||||
nightly_version, releases_result = await asyncio.gather(
|
||||
UpdateRoutes._get_nightly_version(local_hash),
|
||||
UpdateRoutes._get_remote_version()
|
||||
)
|
||||
remote_version, _, behind_by, commit_date = nightly_version
|
||||
_, changelog, releases = releases_result
|
||||
else:
|
||||
remote_version, changelog, releases = await UpdateRoutes._get_remote_version()
|
||||
behind_by = 0
|
||||
commit_date = ''
|
||||
|
||||
# Compare versions
|
||||
if nightly:
|
||||
@@ -81,6 +167,10 @@ class UpdateRoutes:
|
||||
remote_version.replace('v', '')
|
||||
)
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||
has_git = os.path.exists(os.path.join(plugin_root, '.git'))
|
||||
|
||||
response_data = {
|
||||
'success': True,
|
||||
'current_version': local_version,
|
||||
@@ -88,13 +178,13 @@ class UpdateRoutes:
|
||||
'update_available': update_available,
|
||||
'changelog': changelog,
|
||||
'git_info': git_info,
|
||||
'nightly': nightly
|
||||
'nightly': nightly,
|
||||
'has_git': has_git,
|
||||
'releases': releases,
|
||||
'behind_by': behind_by,
|
||||
'commit_date': commit_date
|
||||
}
|
||||
|
||||
# Include releases list for stable mode
|
||||
if releases is not None:
|
||||
response_data['releases'] = releases
|
||||
|
||||
return web.json_response(response_data)
|
||||
|
||||
except NETWORK_EXCEPTIONS as e:
|
||||
@@ -126,9 +216,14 @@ class UpdateRoutes:
|
||||
# Format: version-short_hash
|
||||
version_string = f"{local_version}-{short_hash}"
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||
has_git = os.path.exists(os.path.join(plugin_root, '.git'))
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'version': version_string
|
||||
'version': version_string,
|
||||
'has_git': has_git
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
@@ -156,20 +251,22 @@ class UpdateRoutes:
|
||||
if os.path.exists(settings_path):
|
||||
with open(settings_path, 'r', encoding='utf-8') as f:
|
||||
settings_backup = f.read()
|
||||
logger.info("Backed up settings.json")
|
||||
logger.debug("Backed up settings.json (%d bytes)", len(settings_backup))
|
||||
|
||||
git_folder = os.path.join(plugin_root, '.git')
|
||||
if os.path.exists(git_folder):
|
||||
# Git update
|
||||
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
|
||||
else:
|
||||
# Fallback: Download ZIP and replace files
|
||||
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
|
||||
try:
|
||||
git_folder = os.path.join(plugin_root, '.git')
|
||||
if os.path.exists(git_folder):
|
||||
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
|
||||
else:
|
||||
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||
finally:
|
||||
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
|
||||
|
||||
if settings_backup and success:
|
||||
with open(settings_path, 'w', encoding='utf-8') as f:
|
||||
f.write(settings_backup)
|
||||
logger.info("Restored settings.json")
|
||||
logger.debug("Restored settings.json content (%d bytes)", len(settings_backup))
|
||||
|
||||
if success:
|
||||
return web.json_response({
|
||||
@@ -190,6 +287,164 @@ class UpdateRoutes:
|
||||
'error': str(e)
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
async def switch_channel(request):
|
||||
"""
|
||||
Switch between release and nightly update channels.
|
||||
|
||||
ZIP/CNR install → Nightly: git init + checkout main (one-way upgrade)
|
||||
Git install → Release: git checkout latest tag (.git preserved)
|
||||
ZIP/CNR install → Release: ZIP download (no .git, stays in ZIP mode)
|
||||
Git install → Nightly: git checkout main + pull
|
||||
"""
|
||||
try:
|
||||
body = await request.json() if request.has_body else {}
|
||||
channel = body.get('channel', '')
|
||||
|
||||
if channel not in ('release', 'nightly'):
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Invalid channel: {channel}. Must be "release" or "nightly".'
|
||||
})
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||
|
||||
settings_path = ensure_settings_file(logger)
|
||||
settings_backup = None
|
||||
if os.path.exists(settings_path):
|
||||
with open(settings_path, 'r', encoding='utf-8') as f:
|
||||
settings_backup = f.read()
|
||||
logger.debug("Backed up settings.json before channel switch (%d bytes)", len(settings_backup))
|
||||
|
||||
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
|
||||
try:
|
||||
git_folder = os.path.join(plugin_root, '.git')
|
||||
|
||||
if channel == 'nightly':
|
||||
git_backup = None
|
||||
if os.path.exists(git_folder):
|
||||
git_backup = UpdateRoutes._backup_git(git_folder, 'nightly')
|
||||
|
||||
success = False
|
||||
new_version = ''
|
||||
try:
|
||||
if os.path.exists(git_folder):
|
||||
success, new_version = await UpdateRoutes._perform_git_update(
|
||||
plugin_root, nightly=True
|
||||
)
|
||||
else:
|
||||
success, new_version = UpdateRoutes._init_git_repo(plugin_root)
|
||||
finally:
|
||||
UpdateRoutes._restore_git(git_backup, git_folder, success, 'nightly')
|
||||
else:
|
||||
success = False
|
||||
new_version = ''
|
||||
if os.path.exists(git_folder):
|
||||
success, new_version = await UpdateRoutes._perform_git_update(
|
||||
plugin_root, nightly=False
|
||||
)
|
||||
else:
|
||||
tracking_file = os.path.join(plugin_root, '.tracking')
|
||||
if os.path.exists(tracking_file):
|
||||
os.remove(tracking_file)
|
||||
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||
finally:
|
||||
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
|
||||
|
||||
if settings_backup and success:
|
||||
with open(settings_path, 'w', encoding='utf-8') as f:
|
||||
f.write(settings_backup)
|
||||
logger.debug("Restored settings.json content after channel switch (%d bytes)", len(settings_backup))
|
||||
|
||||
if success:
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
'channel': channel,
|
||||
'new_version': new_version,
|
||||
'message': f'Switched to {channel} channel'
|
||||
})
|
||||
else:
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': f'Failed to switch to {channel} channel'
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to switch channel: %s", e, exc_info=True)
|
||||
return web.json_response({
|
||||
'success': False,
|
||||
'error': str(e)
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
def _init_git_repo(plugin_root: str) -> tuple[bool, str]:
|
||||
"""
|
||||
Initialize a Git repository in a ZIP-installed plugin folder.
|
||||
Clones the remote history and checks out main branch.
|
||||
"""
|
||||
try:
|
||||
import git
|
||||
except ImportError:
|
||||
logger.error(
|
||||
"GitPython is not available: cannot initialize git repo. "
|
||||
"Install git or set $GIT_PYTHON_GIT_EXECUTABLE to the git binary path."
|
||||
)
|
||||
return False, ""
|
||||
|
||||
clean_excludes = _clean_excludes()
|
||||
|
||||
try:
|
||||
repo = git.Repo.init(plugin_root)
|
||||
origin = repo.create_remote(
|
||||
'origin',
|
||||
'https://github.com/willmiao/ComfyUI-Lora-Manager.git'
|
||||
)
|
||||
origin.fetch()
|
||||
|
||||
repo.create_head('main', origin.refs.main)
|
||||
repo.git.checkout('main', '--force')
|
||||
repo.git.reset('--hard')
|
||||
repo.git.clean('-fd', *clean_excludes)
|
||||
|
||||
tracking_file = os.path.join(plugin_root, '.tracking')
|
||||
if os.path.exists(tracking_file):
|
||||
os.remove(tracking_file)
|
||||
logger.info("Removed .tracking file (now in git mode)")
|
||||
|
||||
new_version = f"main-{repo.head.commit.hexsha[:7]}"
|
||||
logger.info("Initialized git repo on main branch: %s", new_version)
|
||||
return True, new_version
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to initialize git repo: %s", e, exc_info=True)
|
||||
return False, ""
|
||||
|
||||
@staticmethod
|
||||
def _backup_git(git_folder, label):
|
||||
try:
|
||||
backup_dir = tempfile.mkdtemp()
|
||||
backup = os.path.join(backup_dir, '.git')
|
||||
shutil.copytree(git_folder, backup)
|
||||
logger.info("Backed up .git before switching to %s", label)
|
||||
return backup
|
||||
except Exception as e:
|
||||
logger.error("Failed to backup .git before %s switch: %s", label, e)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _restore_git(git_backup, git_folder, success, label):
|
||||
if git_backup and not success:
|
||||
try:
|
||||
if os.path.exists(git_folder):
|
||||
shutil.rmtree(git_folder)
|
||||
shutil.copytree(git_backup, git_folder)
|
||||
logger.info("Restored .git after failed %s switch", label)
|
||||
except Exception as e:
|
||||
logger.error("Failed to restore .git after %s switch: %s", label, e)
|
||||
if git_backup:
|
||||
shutil.rmtree(os.path.dirname(git_backup), ignore_errors=True)
|
||||
|
||||
@staticmethod
|
||||
async def _download_and_replace_zip(plugin_root: str) -> tuple[bool, str]:
|
||||
"""
|
||||
@@ -244,8 +499,7 @@ class UpdateRoutes:
|
||||
except Exception:
|
||||
logger.debug("Could not close downloaded-version history database", exc_info=True)
|
||||
|
||||
# Skip settings.json, civitai, model cache and runtime cache folders
|
||||
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=['settings.json', 'civitai', 'model_cache', 'cache', 'wildcards', 'backups', 'stats'])
|
||||
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=list(_PRESERVE_DIRS))
|
||||
|
||||
# Extract ZIP to temp dir
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
@@ -255,7 +509,7 @@ class UpdateRoutes:
|
||||
extracted_root = next(os.scandir(tmp_dir)).path
|
||||
|
||||
# Copy files, skipping user data that should be preserved
|
||||
skip_items = {'settings.json', 'civitai', 'wildcards', 'backups', 'stats'}
|
||||
skip_items = set(_PRESERVE_DIRS)
|
||||
for item in os.listdir(extracted_root):
|
||||
if item in skip_items:
|
||||
continue
|
||||
@@ -272,7 +526,7 @@ class UpdateRoutes:
|
||||
# for ComfyUI Manager to work properly
|
||||
tracking_info_file = os.path.join(plugin_root, '.tracking')
|
||||
tracking_files = []
|
||||
skip_tracked = {'civitai', 'wildcards', 'backups', 'stats'}
|
||||
skip_tracked = set(_PRESERVE_DIRS) - {'settings.json'}
|
||||
for root, dirs, files in os.walk(extracted_root):
|
||||
# Skip user data directories and their contents
|
||||
rel_root = os.path.relpath(root, extracted_root)
|
||||
@@ -295,7 +549,8 @@ class UpdateRoutes:
|
||||
except Exception as e:
|
||||
logger.error(f"ZIP update failed: {e}", exc_info=True)
|
||||
return False, ""
|
||||
|
||||
|
||||
@staticmethod
|
||||
def _clean_plugin_folder(plugin_root, skip_files=None):
|
||||
skip_files = skip_files or []
|
||||
for item in os.listdir(plugin_root):
|
||||
@@ -308,41 +563,54 @@ class UpdateRoutes:
|
||||
os.remove(path)
|
||||
|
||||
@staticmethod
|
||||
async def _get_nightly_version() -> tuple[str, List[str]]:
|
||||
"""
|
||||
Fetch latest commit from main branch
|
||||
"""
|
||||
async def _get_nightly_version(local_hash: str = "") -> tuple[str, List[str], int, str]:
|
||||
repo_owner = "willmiao"
|
||||
repo_name = "ComfyUI-Lora-Manager"
|
||||
|
||||
# Use GitHub API to fetch the latest commit from main branch
|
||||
|
||||
github_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/commits/main"
|
||||
|
||||
|
||||
try:
|
||||
downloader = await get_downloader()
|
||||
success, data = await downloader.make_request('GET', github_url, custom_headers={'Accept': 'application/vnd.github+json'})
|
||||
|
||||
success, data = await downloader.make_request(
|
||||
'GET', github_url,
|
||||
custom_headers={'Accept': 'application/vnd.github+json'}
|
||||
)
|
||||
|
||||
if not success:
|
||||
logger.warning(f"Failed to fetch GitHub commit: {data}")
|
||||
return "main", []
|
||||
|
||||
commit_sha = data.get('sha', '')[:7] # Short hash
|
||||
logger.warning("Failed to fetch GitHub commit: %s", data)
|
||||
return "main", [], 0, ""
|
||||
|
||||
commit_sha = data.get('sha', '')[:7]
|
||||
commit_message = data.get('commit', {}).get('message', '')
|
||||
|
||||
# Format as "main-{short_hash}"
|
||||
commit_date = data.get('commit', {}).get('committer', {}).get('date', '')[:10]
|
||||
|
||||
version = f"main-{commit_sha}"
|
||||
|
||||
# Use commit message as changelog
|
||||
changelog = [commit_message] if commit_message else []
|
||||
|
||||
return version, changelog
|
||||
|
||||
|
||||
behind_by = 0
|
||||
if local_hash and local_hash not in ('unknown', 'stable'):
|
||||
compare_url = (
|
||||
f"https://api.github.com/repos/{repo_owner}/{repo_name}"
|
||||
f"/compare/{local_hash}...main"
|
||||
)
|
||||
c_ok, c_data = await downloader.make_request(
|
||||
'GET', compare_url,
|
||||
custom_headers={'Accept': 'application/vnd.github+json'}
|
||||
)
|
||||
if c_ok:
|
||||
if c_data.get('status') in ('ahead', 'diverged'):
|
||||
behind_by = c_data.get('ahead_by', 0)
|
||||
else:
|
||||
behind_by = c_data.get('behind_by', 0)
|
||||
|
||||
return version, changelog, behind_by, commit_date
|
||||
|
||||
except NETWORK_EXCEPTIONS as e:
|
||||
logger.warning("Unable to reach GitHub for nightly version: %s", e)
|
||||
return "main", []
|
||||
return "main", [], 0, ""
|
||||
except Exception as e:
|
||||
logger.error(f"Error fetching nightly version: {e}", exc_info=True)
|
||||
return "main", []
|
||||
logger.error("Error fetching nightly version: %s", e, exc_info=True)
|
||||
return "main", [], 0, ""
|
||||
|
||||
@staticmethod
|
||||
def _compare_nightly_versions(local_git_info: Dict[str, str], remote_version: str) -> bool:
|
||||
|
||||
@@ -201,6 +201,13 @@ class Aria2Downloader:
|
||||
"auto-file-renaming": "false",
|
||||
"file-allocation": "none",
|
||||
}
|
||||
|
||||
# Pass proxy to aria2 so the actual file transfer goes through the
|
||||
# same proxy used by the aiohttp-based URL resolution step above.
|
||||
downloader = await get_downloader()
|
||||
if downloader.proxy_url:
|
||||
options["all-proxy"] = downloader.proxy_url
|
||||
|
||||
if request_headers:
|
||||
options["header"] = [
|
||||
f"{key}: {value}" for key, value in request_headers.items()
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
from abc import ABC, abstractmethod
|
||||
import asyncio
|
||||
import re
|
||||
from typing import Any, Dict, List, Optional, Type, TYPE_CHECKING
|
||||
import random
|
||||
from typing import Any, Dict, List, Optional, Type, Union, TYPE_CHECKING
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
@@ -109,12 +110,15 @@ class BaseModelService(ABC):
|
||||
if civitai_model_id is not None:
|
||||
sorted_data = [
|
||||
item for item in sorted_data
|
||||
if self._extract_model_id(item) == civitai_model_id
|
||||
if self._extract_group_key(item) == civitai_model_id
|
||||
]
|
||||
# VLM mode: always sort by version ID descending (newest version first),
|
||||
# regardless of the current sort_by preference.
|
||||
# Fall back to modified timestamp for non-CivitAI sources.
|
||||
sorted_data.sort(
|
||||
key=lambda x: self._extract_version_id(x) or 0,
|
||||
key=lambda x: self._extract_version_id(x)
|
||||
or x.get("modified", 0)
|
||||
or 0,
|
||||
reverse=True,
|
||||
)
|
||||
|
||||
@@ -129,18 +133,21 @@ class BaseModelService(ABC):
|
||||
ufs = self.settings.get("version_grouping", "same_base")
|
||||
group_by_base = ufs == "same_base"
|
||||
|
||||
dedup_map = {} # (modelId [,base_model]) -> (item, version_id)
|
||||
dedup_map = {} # (modelId [,base_model]) -> (item, version_or_modified)
|
||||
version_counter = {} # same-key -> count
|
||||
standalone = []
|
||||
for item in sorted_data:
|
||||
mid = self._extract_model_id(item)
|
||||
mid = self._extract_group_key(item)
|
||||
if mid is None:
|
||||
standalone.append(item)
|
||||
continue
|
||||
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
||||
# Count all versions per key
|
||||
version_counter[key] = version_counter.get(key, 0) + 1
|
||||
vid = self._extract_version_id(item) or 0
|
||||
# Prefer CivitAI version_id; fall back to modified timestamp
|
||||
vid = self._extract_version_id(item)
|
||||
if vid is None:
|
||||
vid = item.get("modified", 0) or 0
|
||||
if key not in dedup_map or vid > dedup_map[key][1]:
|
||||
dedup_map[key] = (item, vid)
|
||||
# Attach version_count to each surviving grouped item (shallow copy
|
||||
@@ -174,16 +181,19 @@ class BaseModelService(ABC):
|
||||
model_groups: Dict[Any, List[Dict]] = {}
|
||||
ungrouped_standalone: List[Dict] = []
|
||||
for item in sorted_data:
|
||||
mid = self._extract_model_id(item)
|
||||
mid = self._extract_group_key(item)
|
||||
if mid is None:
|
||||
ungrouped_standalone.append(item)
|
||||
continue
|
||||
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
||||
model_groups.setdefault(key, []).append(item)
|
||||
# Sort versions within each group by version id descending
|
||||
# Sort versions within each group by version id (descending);
|
||||
# fall back to modified timestamp for non-CivitAI sources.
|
||||
for items in model_groups.values():
|
||||
items.sort(
|
||||
key=lambda x: self._extract_version_id(x) or 0,
|
||||
key=lambda x: self._extract_version_id(x)
|
||||
or x.get("modified", 0)
|
||||
or 0,
|
||||
reverse=True,
|
||||
)
|
||||
# Sort groups by version count
|
||||
@@ -381,6 +391,12 @@ class BaseModelService(ABC):
|
||||
(item.get("model_name") or item.get("file_name") or "").lower(),
|
||||
item.get("file_path", "").lower(),
|
||||
)
|
||||
elif key_name == "random":
|
||||
# Seeded random shuffle: same seed -> same order (stable pagination)
|
||||
rng = random.Random(sort_params.seed or "random")
|
||||
result = list(data)
|
||||
rng.shuffle(result)
|
||||
return result
|
||||
elif key_name == "size":
|
||||
key_fn = lambda item: (
|
||||
int(item.get("size", 0) or 0),
|
||||
@@ -697,6 +713,33 @@ class BaseModelService(ABC):
|
||||
|
||||
return annotated
|
||||
|
||||
@staticmethod
|
||||
def _extract_hf_group_key(item: Dict) -> Optional[str]:
|
||||
"""Extract `hf:{owner}/{repo}` from item's ``hf_url``, or None."""
|
||||
hf_url = item.get("hf_url") if isinstance(item, dict) else None
|
||||
if not hf_url or not isinstance(hf_url, str):
|
||||
return None
|
||||
m = re.match(
|
||||
r"https?://huggingface\.co/([^/]+/[^/]+)", hf_url.strip()
|
||||
)
|
||||
if not m:
|
||||
return None
|
||||
return f"hf:{m.group(1)}"
|
||||
|
||||
@staticmethod
|
||||
def _extract_group_key(item: Dict) -> Union[int, str, None]:
|
||||
"""Return the group identity key: CivitAI modelId (int) or HF repo (str).
|
||||
|
||||
Preference order:
|
||||
1. CivitAI ``modelId`` (int)
|
||||
2. HF repo identity ``hf:{owner}/{repo}`` (str)
|
||||
3. ``None`` (no known grouping source)
|
||||
"""
|
||||
mid = BaseModelService._extract_model_id(item)
|
||||
if mid is not None:
|
||||
return mid
|
||||
return BaseModelService._extract_hf_group_key(item)
|
||||
|
||||
@staticmethod
|
||||
def _extract_model_id(item: Dict) -> Optional[int]:
|
||||
civitai = item.get("civitai") if isinstance(item, dict) else None
|
||||
@@ -804,6 +847,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 +1004,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 +1141,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 +1163,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]:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -2,6 +2,7 @@ import asyncio
|
||||
import copy
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from typing import Any, Optional, Dict, Tuple, List, Sequence
|
||||
from .connectivity_guard import (
|
||||
@@ -19,6 +20,12 @@ from ..utils.civitai_utils import resolve_license_payload
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Best-effort cache for creator model counts, keyed by lowercase username.
|
||||
# Values are (monotonic timestamp, count or None); None results are cached
|
||||
# too so repeated failures don't hammer the API.
|
||||
_CREATOR_COUNT_CACHE_TTL_SECONDS = 600
|
||||
_creator_model_count_cache: Dict[str, Tuple[float, Optional[int]]] = {}
|
||||
|
||||
|
||||
class CivitaiClient:
|
||||
_instance = None
|
||||
@@ -743,17 +750,34 @@ class CivitaiClient:
|
||||
|
||||
return all_versions if all_versions else None
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
"""Fetch all models for a specific Civitai user."""
|
||||
async def get_user_models(
|
||||
self, username: str, cursor: Optional[str] = None
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Fetch one page (up to 100 models) for a specific Civitai user.
|
||||
|
||||
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
|
||||
or None on failure. Pass ``cursor`` (from a previous response's
|
||||
``nextCursor``) to fetch subsequent pages.
|
||||
"""
|
||||
if not username:
|
||||
return None
|
||||
|
||||
params: Dict[str, Any] = {
|
||||
"username": username,
|
||||
"nsfw": "true",
|
||||
"limit": 100,
|
||||
"sort": "Newest",
|
||||
"period": "AllTime",
|
||||
}
|
||||
if cursor:
|
||||
params["cursor"] = cursor
|
||||
|
||||
try:
|
||||
success, result = await self._make_request(
|
||||
"GET",
|
||||
f"{self.base_url}/models",
|
||||
use_auth=True,
|
||||
params={"username": username, "nsfw": "true"},
|
||||
params=params,
|
||||
)
|
||||
|
||||
if not success:
|
||||
@@ -765,7 +789,7 @@ class CivitaiClient:
|
||||
|
||||
items = result.get("items") if isinstance(result, dict) else None
|
||||
if not isinstance(items, list):
|
||||
return []
|
||||
items = []
|
||||
|
||||
for model in items:
|
||||
versions = model.get("modelVersions")
|
||||
@@ -774,9 +798,68 @@ class CivitaiClient:
|
||||
for version in versions:
|
||||
self._remove_comfy_metadata(version)
|
||||
|
||||
return items
|
||||
next_cursor: Optional[str] = None
|
||||
metadata = result.get("metadata") if isinstance(result, dict) else None
|
||||
if isinstance(metadata, dict):
|
||||
raw_cursor = metadata.get("nextCursor")
|
||||
if raw_cursor is not None:
|
||||
next_cursor = str(raw_cursor)
|
||||
|
||||
return {"items": items, "nextCursor": next_cursor}
|
||||
except RateLimitError:
|
||||
raise
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error("Error fetching models for %s: %s", username, exc)
|
||||
return None
|
||||
|
||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||
"""Best-effort lookup of a creator's published model count.
|
||||
|
||||
Uses the ``/creators`` endpoint (a contains-match query), picking the
|
||||
entry whose username matches exactly (case-insensitive). Returns None
|
||||
on any failure; never raises. Results (including None) are cached
|
||||
for ``_CREATOR_COUNT_CACHE_TTL_SECONDS``.
|
||||
"""
|
||||
if not username:
|
||||
return None
|
||||
|
||||
cache_key = username.lower()
|
||||
cached = _creator_model_count_cache.get(cache_key)
|
||||
if cached is not None:
|
||||
cached_at, cached_count = cached
|
||||
if time.monotonic() - cached_at < _CREATOR_COUNT_CACHE_TTL_SECONDS:
|
||||
return cached_count
|
||||
|
||||
count: Optional[int] = None
|
||||
try:
|
||||
success, result = await self._make_request(
|
||||
"GET",
|
||||
f"{self.base_url}/creators",
|
||||
use_auth=True,
|
||||
params={"query": username, "limit": 10},
|
||||
)
|
||||
|
||||
if success and isinstance(result, dict):
|
||||
creators = result.get("items")
|
||||
if isinstance(creators, list):
|
||||
for creator in creators:
|
||||
if not isinstance(creator, dict):
|
||||
continue
|
||||
creator_name = creator.get("username")
|
||||
if not isinstance(creator_name, str):
|
||||
continue
|
||||
if creator_name.lower() != cache_key:
|
||||
continue
|
||||
model_count = creator.get("modelCount")
|
||||
if isinstance(model_count, (int, float)) and not isinstance(
|
||||
model_count, bool
|
||||
):
|
||||
count = int(model_count)
|
||||
break
|
||||
except Exception as exc: # best-effort only, never propagate
|
||||
logger.debug(
|
||||
"Failed to fetch creator model count for %s: %s", username, exc
|
||||
)
|
||||
|
||||
_creator_model_count_cache[cache_key] = (time.monotonic(), count)
|
||||
return count
|
||||
|
||||
@@ -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.abspath(os.path.normpath(save_dir))
|
||||
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
|
||||
if not resolved_dir.startswith(base_dir + os.sep) and resolved_dir != base_dir:
|
||||
logger.warning(
|
||||
"Path traversal detected: %s escapes %s",
|
||||
resolved_dir, base_dir,
|
||||
)
|
||||
return {"success": False, "error": "Download path is outside allowed directory"}
|
||||
# Create directory if it doesn't exist
|
||||
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(
|
||||
|
||||
@@ -31,7 +31,7 @@ class DownloadQueueService:
|
||||
_instance: Optional[DownloadQueueService] = None
|
||||
_class_lock: asyncio.Lock = asyncio.Lock()
|
||||
|
||||
_SCHEMA = """
|
||||
_SCHEMA_TABLES = """
|
||||
CREATE TABLE IF NOT EXISTS download_queue (
|
||||
download_id TEXT PRIMARY KEY,
|
||||
model_id INTEGER,
|
||||
@@ -76,6 +76,11 @@ class DownloadQueueService:
|
||||
CREATE INDEX IF NOT EXISTS idx_dh_status ON download_history(status);
|
||||
"""
|
||||
|
||||
_CREATE_UNIQUE_INDEX = """
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_dh_download_id
|
||||
ON download_history(download_id) WHERE download_id IS NOT NULL;
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
async def get_instance(cls) -> DownloadQueueService:
|
||||
"""Return the singleton instance, creating it if necessary."""
|
||||
@@ -113,10 +118,39 @@ class DownloadQueueService:
|
||||
if self._schema_initialized:
|
||||
return
|
||||
with self._connect() as conn:
|
||||
conn.executescript(self._SCHEMA)
|
||||
conn.executescript(self._SCHEMA_TABLES)
|
||||
|
||||
# Creating the unique index on download_history.download_id can
|
||||
# fail if pre-existing rows have duplicate values (e.g. from a
|
||||
# previous version that lacked the index). Deduplicate first so
|
||||
# that the migration does not crash on startup.
|
||||
if not self._index_exists(conn, "idx_dh_download_id"):
|
||||
self._remove_duplicate_download_ids(conn)
|
||||
conn.executescript(self._CREATE_UNIQUE_INDEX)
|
||||
|
||||
conn.commit()
|
||||
self._schema_initialized = True
|
||||
|
||||
@staticmethod
|
||||
def _index_exists(conn: sqlite3.Connection, name: str) -> bool:
|
||||
return conn.execute(
|
||||
"SELECT 1 FROM sqlite_master WHERE type='index' AND name=?",
|
||||
(name,),
|
||||
).fetchone() is not None
|
||||
|
||||
@staticmethod
|
||||
def _remove_duplicate_download_ids(conn: sqlite3.Connection) -> None:
|
||||
conn.execute("""
|
||||
DELETE FROM download_history
|
||||
WHERE id NOT IN (
|
||||
SELECT MIN(id)
|
||||
FROM download_history
|
||||
WHERE download_id IS NOT NULL
|
||||
GROUP BY download_id
|
||||
)
|
||||
AND download_id IS NOT NULL
|
||||
""")
|
||||
|
||||
def get_database_path(self) -> str:
|
||||
"""Return the resolved database file path."""
|
||||
return self._db_path
|
||||
@@ -154,13 +188,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 +424,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 +581,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 +658,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 +717,7 @@ class DownloadQueueService:
|
||||
)
|
||||
conn.execute(
|
||||
"DELETE FROM download_history WHERE id = ?",
|
||||
(item_id,),
|
||||
(row["id"],),
|
||||
)
|
||||
conn.commit()
|
||||
queued = conn.execute(
|
||||
|
||||
+56
-10
@@ -46,6 +46,30 @@ def is_ssl_cert_verify_error(exc: BaseException) -> bool:
|
||||
return "CERTIFICATE_VERIFY_FAILED" in str(exc)
|
||||
|
||||
|
||||
def _parse_retry_after(value: str) -> int:
|
||||
"""Parse a Retry-After header value into seconds.
|
||||
|
||||
Supports both integer seconds and HTTP-date formats.
|
||||
Returns a default of 60 seconds on invalid/missing input.
|
||||
"""
|
||||
if not value or not value.strip():
|
||||
return 60
|
||||
|
||||
value = value.strip()
|
||||
try:
|
||||
return max(1, int(value))
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
try:
|
||||
parsed = parsedate_to_datetime(value)
|
||||
now = datetime.now().astimezone()
|
||||
delta = (parsed - now).total_seconds()
|
||||
return max(1, int(delta))
|
||||
except (ValueError, OverflowError, OSError):
|
||||
return 60
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DownloadProgress:
|
||||
"""Snapshot of a download transfer at a moment in time."""
|
||||
@@ -246,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
|
||||
@@ -348,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),
|
||||
@@ -729,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
|
||||
@@ -819,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}")
|
||||
@@ -911,6 +944,19 @@ class Downloader:
|
||||
elif response.status == 404:
|
||||
error_msg = "File not found"
|
||||
return False, error_msg, None
|
||||
elif response.status == 429:
|
||||
raw_retry_after = response.headers.get("Retry-After")
|
||||
retry_after = _parse_retry_after(raw_retry_after or "")
|
||||
if raw_retry_after:
|
||||
logger.warning(
|
||||
"Rate limited (429) for %s, Retry-After: %ss", url, retry_after
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Rate limited (429) for %s, no Retry-After header; defaulting to %ss",
|
||||
url, retry_after,
|
||||
)
|
||||
return False, f"Rate limited (429), retry after {retry_after}s", None
|
||||
else:
|
||||
error_msg = f"Download failed with status {response.status}"
|
||||
return False, error_msg, None
|
||||
|
||||
+138
-84
@@ -27,18 +27,27 @@ _MODEL_CATALOG_URL = "https://models.dev/api.json"
|
||||
|
||||
# In-memory cache: maps provider slug -> list of model ID strings.
|
||||
_catalog_cache: Optional[Dict[str, List[str]]] = None
|
||||
|
||||
# Per-model max output token limits parsed from the catalog.
|
||||
# ``{provider_id: {model_id: max_output_tokens}}``.
|
||||
_model_output_limits: Dict[str, Dict[str, int]] = {}
|
||||
|
||||
_CATALOG_TIMEOUT = aiohttp.ClientTimeout(total=30)
|
||||
|
||||
|
||||
async def _load_model_catalog() -> Dict[str, List[str]]:
|
||||
"""Fetch and parse the model catalog, returning ``{provider_id: [model_id, ...]}``.
|
||||
"""Fetch and parse the model catalog.
|
||||
|
||||
Returns ``{provider_id: [model_id, ...]}`` and also populates
|
||||
:data:`_model_output_limits` with per-model ``limit.output`` values
|
||||
for use by :func:`_get_model_max_output`.
|
||||
|
||||
The JSON at ``_MODEL_CATALOG_URL`` is a dict keyed by provider slug; each
|
||||
value has a ``models`` sub-dict keyed by model ID. Only the model IDs are
|
||||
kept. The result is cached in memory after the first successful fetch.
|
||||
value has a ``models`` sub-dict keyed by model ID. The result is cached
|
||||
in memory after the first successful fetch.
|
||||
Subsequent calls return the cached data immediately.
|
||||
"""
|
||||
global _catalog_cache
|
||||
global _catalog_cache, _model_output_limits
|
||||
if _catalog_cache is not None:
|
||||
return _catalog_cache
|
||||
|
||||
@@ -58,25 +67,52 @@ async def _load_model_catalog() -> Dict[str, List[str]]:
|
||||
return _catalog_cache or {}
|
||||
|
||||
result: Dict[str, List[str]] = {}
|
||||
output_limits: Dict[str, Dict[str, int]] = {}
|
||||
for provider_id, provider_info in data.items():
|
||||
if not isinstance(provider_info, dict):
|
||||
continue
|
||||
models_dict = provider_info.get("models")
|
||||
if not isinstance(models_dict, dict):
|
||||
continue
|
||||
model_ids = [str(mid) for mid in models_dict.keys() if isinstance(mid, str)]
|
||||
model_ids: List[str] = []
|
||||
provider_limits: Dict[str, int] = {}
|
||||
for mid, model_info in models_dict.items():
|
||||
if not isinstance(mid, str):
|
||||
continue
|
||||
model_ids.append(mid)
|
||||
if isinstance(model_info, dict):
|
||||
limit = model_info.get("limit")
|
||||
if isinstance(limit, dict):
|
||||
output = limit.get("output")
|
||||
if isinstance(output, (int, float)) and output > 0:
|
||||
provider_limits[mid] = int(output)
|
||||
if model_ids:
|
||||
result[provider_id] = model_ids
|
||||
if provider_limits:
|
||||
output_limits[provider_id] = provider_limits
|
||||
|
||||
_catalog_cache = result
|
||||
_model_output_limits = output_limits
|
||||
logger.debug(
|
||||
"Loaded model catalog: %d providers, %d total models",
|
||||
"Loaded model catalog: %d providers, %d total models "
|
||||
"(%d providers have output limits)",
|
||||
len(result),
|
||||
sum(len(m) for m in result.values()),
|
||||
len(output_limits),
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def _get_model_max_output(provider: str, model: str) -> Optional[int]:
|
||||
"""Return the model's max output token limit from the catalog, or ``None``.
|
||||
|
||||
Returns ``None`` when the provider or model is not found in the catalog
|
||||
(e.g. local Ollama models, custom models, or user-typed model names).
|
||||
Callers should fall back to a safe default.
|
||||
"""
|
||||
return _model_output_limits.get(provider, {}).get(model)
|
||||
|
||||
|
||||
# Short timeout for Ollama's local API
|
||||
_OLLAMA_API_TIMEOUT = aiohttp.ClientTimeout(total=8)
|
||||
|
||||
@@ -246,15 +282,17 @@ class LLMService:
|
||||
def is_configured(self) -> bool:
|
||||
"""Return ``True`` when the LLM provider is minimally configured.
|
||||
|
||||
A provider is considered configured when ``llm_model`` is set and
|
||||
A provider is considered configured when ``llm_model`` is set,
|
||||
an API key is configured for providers that require one (e.g.
|
||||
Ollama does not).
|
||||
Ollama does not), and an API base URL is set for providers that
|
||||
have no preset default (e.g. ``custom``).
|
||||
"""
|
||||
|
||||
cfg = self._get_config()
|
||||
has_model = bool(cfg["model"])
|
||||
has_key = bool(cfg["api_key"]) or not self._provider_requires_key(cfg["provider"])
|
||||
return has_model and has_key
|
||||
has_base = bool(cfg["api_base"]) or bool(_PROVIDER_DEFAULTS.get(cfg["provider"]))
|
||||
return has_model and has_key and has_base
|
||||
|
||||
def _resolve_api_base(self, provider: str, api_base: str) -> str:
|
||||
"""Resolve the API base URL for the given provider.
|
||||
@@ -278,20 +316,26 @@ class LLMService:
|
||||
def _ensure_configured(self) -> Dict[str, Any]:
|
||||
"""Validate configuration and return it, or raise.
|
||||
|
||||
A provider is considered configured when ``llm_model`` is set and
|
||||
(for non-Ollama) an API key is configured.
|
||||
A provider is considered configured when ``llm_model`` is set,
|
||||
an API key is configured for providers that require one, and
|
||||
an API base URL is set for providers without a preset default.
|
||||
"""
|
||||
|
||||
cfg = self._get_config()
|
||||
has_model = bool(cfg["model"])
|
||||
needs_key = self._provider_requires_key(cfg["provider"])
|
||||
has_key = bool(cfg["api_key"]) or not needs_key
|
||||
if not (has_model and has_key):
|
||||
has_base = bool(cfg["api_base"]) or bool(_PROVIDER_DEFAULTS.get(cfg["provider"]))
|
||||
if not (has_model and has_key and has_base):
|
||||
parts = []
|
||||
if not has_model:
|
||||
parts.append("No LLM model specified")
|
||||
if not has_key and needs_key:
|
||||
parts.append("No LLM API key configured")
|
||||
if not has_base:
|
||||
parts.append(
|
||||
f"No API base URL for provider '{cfg['provider']}'"
|
||||
)
|
||||
detail = "; ".join(parts) if parts else "LLM provider is not configured"
|
||||
raise LLMNotConfiguredError(
|
||||
f"{detail}. Configure it in Settings → AI Provider."
|
||||
@@ -364,9 +408,11 @@ class LLMService:
|
||||
"think": False,
|
||||
"options": {
|
||||
"temperature": temperature,
|
||||
# Allow up to 32K context so the model has room to think
|
||||
# AND produce output without hitting the 4K default limit.
|
||||
"num_ctx": 32768,
|
||||
# 8K context is sufficient for metadata enrichment
|
||||
# (prompt ~2-5K, output ~0.2-1K tokens). The old 32K
|
||||
# value was excessive for this use case and increased
|
||||
# Ollama VRAM usage unnecessarily.
|
||||
"num_ctx": 8192,
|
||||
},
|
||||
}
|
||||
if response_format is not None:
|
||||
@@ -480,11 +526,15 @@ class LLMService:
|
||||
temperature: float = 0.3,
|
||||
max_tokens: Optional[int] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Call the LLM and return parsed JSON.
|
||||
"""Call the LLM with ``response_format=json_object`` and return parsed JSON.
|
||||
|
||||
Sends ``response_format: {"type": "json_object"}`` when the provider
|
||||
supports it, and parses the response content as JSON. If parsing
|
||||
fails, retries once with a clarifying system message.
|
||||
``max_tokens`` is resolved in this order:
|
||||
1. Explicit caller-supplied ``max_tokens``
|
||||
2. Per-model ``limit.output`` from the model catalog
|
||||
3. A safe default of 4096 (sufficient for metadata enrichment)
|
||||
|
||||
If the response content is empty or not valid JSON, attempts
|
||||
:func:`_try_salvage_json` before raising.
|
||||
|
||||
Args:
|
||||
system_prompt: System-level instructions
|
||||
@@ -499,7 +549,7 @@ class LLMService:
|
||||
Raises:
|
||||
LLMNotConfiguredError: Provider not configured
|
||||
LLMRateLimitError: Rate limited
|
||||
LLMResponseError: JSON parse failure after retry
|
||||
LLMResponseError: Empty response or JSON parse failure
|
||||
"""
|
||||
|
||||
messages = [
|
||||
@@ -507,20 +557,66 @@ class LLMService:
|
||||
{"role": "user", "content": user_prompt},
|
||||
]
|
||||
|
||||
# First attempt with JSON mode.
|
||||
# Use a generous max_tokens so thinking-enabled models (e.g.
|
||||
# gemma4 via Ollama) have room to reason AND still emit content.
|
||||
effective_max = max_tokens or 131072
|
||||
result = await self.chat_completion(
|
||||
messages=messages,
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
response_format={"type": "json_object"},
|
||||
max_tokens=effective_max,
|
||||
)
|
||||
# Resolve max_tokens: caller override → catalog lookup → safe default
|
||||
if max_tokens is None:
|
||||
cfg = self._get_config()
|
||||
effective_max = _get_model_max_output(cfg["provider"], cfg["model"])
|
||||
else:
|
||||
effective_max = max_tokens
|
||||
if effective_max is None:
|
||||
effective_max = 4096
|
||||
|
||||
# Use json_schema (not json_object) for broader provider compatibility:
|
||||
# LM Studio and some other OpenAI-compatible servers reject
|
||||
# json_object but accept json_schema. {"type": "object"} is
|
||||
# functionally equivalent — it accepts any JSON object without
|
||||
# constraining specific fields.
|
||||
response_format = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "metadata",
|
||||
"schema": {"type": "object"},
|
||||
},
|
||||
}
|
||||
|
||||
try:
|
||||
parsed = json.loads(result["content"])
|
||||
result = await self.chat_completion(
|
||||
messages=messages,
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
response_format=response_format,
|
||||
max_tokens=effective_max,
|
||||
)
|
||||
except LLMResponseError as e:
|
||||
# Only fall back when the provider rejects the response_format
|
||||
# type value (e.g. "'response_format.type' must be..."). Avoid
|
||||
# catching unrelated 400 errors whose body happens to mention
|
||||
# "response_format" (e.g. "model does not support
|
||||
# response_format restrictions on this endpoint").
|
||||
if "'response_format.type'" not in str(e).lower():
|
||||
raise
|
||||
logger.info(
|
||||
"Provider rejected response_format, retrying without it. "
|
||||
"Falling back to prompt-only JSON mode. Error: %s",
|
||||
e,
|
||||
)
|
||||
result = await self.chat_completion(
|
||||
messages=messages,
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
response_format=None,
|
||||
max_tokens=effective_max,
|
||||
)
|
||||
|
||||
content = result.get("content", "") or ""
|
||||
if not content:
|
||||
raise LLMResponseError(
|
||||
"LLM returned empty content. "
|
||||
f"Raw response: {json.dumps(result)[:500]}"
|
||||
)
|
||||
|
||||
try:
|
||||
parsed = json.loads(content)
|
||||
logger.debug(
|
||||
"LLM raw content: %s",
|
||||
json.dumps(parsed, ensure_ascii=False)[:2000],
|
||||
@@ -529,64 +625,22 @@ class LLMService:
|
||||
except (json.JSONDecodeError, TypeError) as exc:
|
||||
logger.info(
|
||||
"LLM raw response (first 800 chars): %s",
|
||||
(result.get("content") or "")[:800],
|
||||
content[:800],
|
||||
)
|
||||
|
||||
# Last resort: attempt to salvage partial/truncated JSON
|
||||
salvaged = _try_salvage_json(content)
|
||||
if salvaged is not None:
|
||||
logger.warning(
|
||||
"LLM JSON parse failed on first attempt: %s. Retrying.", exc
|
||||
"LLM JSON salvaged from partial content (%d chars raw)",
|
||||
len(content),
|
||||
)
|
||||
return salvaged
|
||||
|
||||
# Retry WITHOUT response_format — some providers (Ollama with
|
||||
# thinking-enabled models like gemma4) may return empty content
|
||||
# when json_object mode is active. Fall back to a textual
|
||||
# instruction instead.
|
||||
previous_content = result.get("content", "") or ""
|
||||
retry_messages = messages + [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": previous_content or "(empty response)",
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": (
|
||||
"The previous response could not be parsed as JSON. "
|
||||
"Please respond with ONLY a valid JSON object, no "
|
||||
"markdown fences or extra text."
|
||||
),
|
||||
},
|
||||
]
|
||||
|
||||
result = await self.chat_completion(
|
||||
messages=retry_messages,
|
||||
model=model,
|
||||
temperature=0.0, # More deterministic for retry
|
||||
max_tokens=effective_max,
|
||||
raise LLMResponseError(
|
||||
f"LLM response could not be parsed as JSON: {content[:200]}"
|
||||
)
|
||||
|
||||
content = result.get("content", "") or ""
|
||||
if not content:
|
||||
raise LLMResponseError(
|
||||
"LLM response could not be parsed as JSON after retry: "
|
||||
f"Expecting value: line 1 column 1 (char 0)\n"
|
||||
f"Raw content: {content[:500]}"
|
||||
)
|
||||
|
||||
try:
|
||||
return json.loads(content)
|
||||
except (json.JSONDecodeError, TypeError) as parse_err:
|
||||
# Last resort: attempt to salvage partial JSON (closing unclosed
|
||||
# brackets/braces, truncating incomplete strings, etc.)
|
||||
salvaged = _try_salvage_json(content)
|
||||
if salvaged is not None:
|
||||
logger.warning(
|
||||
"LLM JSON salvaged from partial content (%d chars raw)",
|
||||
len(content),
|
||||
)
|
||||
return salvaged
|
||||
raise LLMResponseError(
|
||||
f"LLM response could not be parsed as JSON after retry: {parse_err}\n"
|
||||
f"Raw content: {content[:500]}"
|
||||
) from parse_err
|
||||
|
||||
|
||||
def _try_salvage_json(raw: str) -> Dict[str, Any] | None:
|
||||
"""Attempt to repair and parse a truncated JSON string.
|
||||
|
||||
@@ -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 []
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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
|
||||
|
||||
+43
-13
@@ -1,6 +1,7 @@
|
||||
import asyncio
|
||||
import time
|
||||
import logging
|
||||
import random
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
@@ -38,8 +39,8 @@ class ModelCache:
|
||||
|
||||
def __post_init__(self):
|
||||
self._lock = asyncio.Lock()
|
||||
# Cache for last sort: (sort_key, order) -> sorted list
|
||||
self._last_sort: Tuple[str, str] = (None, None)
|
||||
# Cache for last sort: (sort_key, order, seed) -> sorted list
|
||||
self._last_sort: Tuple[Optional[str], str, Optional[str]] = (None, "asc", None)
|
||||
self._last_sorted_data: List[Dict] = []
|
||||
self._normalize_raw_data()
|
||||
self.name_display_mode = self._normalize_display_mode(self.name_display_mode)
|
||||
@@ -203,9 +204,9 @@ class ModelCache:
|
||||
async def resort(self):
|
||||
"""Resort cached data according to last sort mode if set"""
|
||||
async with self._lock:
|
||||
if self._last_sort != (None, None):
|
||||
sort_key, order = self._last_sort
|
||||
sorted_data = self._sort_data(self.raw_data, sort_key, order)
|
||||
if self._last_sort[0] is not None:
|
||||
sort_key, order, seed = self._last_sort
|
||||
sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
|
||||
self._last_sorted_data = sorted_data
|
||||
# Update folder list
|
||||
# else: do nothing
|
||||
@@ -218,7 +219,7 @@ class ModelCache:
|
||||
self.folders = sorted(list(all_folders), key=lambda x: x.lower())
|
||||
self.rebuild_version_index()
|
||||
|
||||
def _sort_data(self, data: List[Dict], sort_key: str, order: str) -> List[Dict]:
|
||||
def _sort_data(self, data: List[Dict], sort_key: str, order: str, seed: Optional[str] = None) -> List[Dict]:
|
||||
"""Sort data by sort_key and order"""
|
||||
start_time = time.perf_counter()
|
||||
reverse = (order == 'desc')
|
||||
@@ -265,6 +266,13 @@ class ModelCache:
|
||||
),
|
||||
reverse=reverse
|
||||
)
|
||||
elif sort_key == 'random':
|
||||
# Random shuffle seeded for stable pagination: the same seed
|
||||
# always yields the same order, so successive page requests
|
||||
# stay consistent while browsing.
|
||||
rng = random.Random(seed or 'random')
|
||||
result = list(data)
|
||||
rng.shuffle(result)
|
||||
elif sort_key == 'versions_count':
|
||||
# Pre-dedup sort: fall back to name sort.
|
||||
# Actual re-sort by version_count happens in get_paginated_data after dedup.
|
||||
@@ -285,15 +293,16 @@ class ModelCache:
|
||||
logger.debug("ModelCache._sort_data(%s, %s) for %d items took %.3fs", sort_key, order, len(data), duration)
|
||||
return result
|
||||
|
||||
async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc') -> List[Dict]:
|
||||
async def get_sorted_data(self, sort_key: str = 'name', order: str = 'asc', seed: Optional[str] = None) -> List[Dict]:
|
||||
"""Get sorted data by sort_key and order, using cache if possible"""
|
||||
async with self._lock:
|
||||
if (sort_key, order) == self._last_sort:
|
||||
cache_key = (sort_key, order, seed)
|
||||
if cache_key == self._last_sort:
|
||||
return self._last_sorted_data
|
||||
|
||||
start_time = time.perf_counter()
|
||||
sorted_data = self._sort_data(self.raw_data, sort_key, order)
|
||||
self._last_sort = (sort_key, order)
|
||||
sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
|
||||
self._last_sort = cache_key
|
||||
self._last_sorted_data = sorted_data
|
||||
|
||||
duration = time.perf_counter() - start_time
|
||||
@@ -313,8 +322,8 @@ class ModelCache:
|
||||
self.name_display_mode = normalized
|
||||
|
||||
if self._last_sort[0] == 'name':
|
||||
sort_key, order = self._last_sort
|
||||
self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order)
|
||||
sort_key, order, seed = self._last_sort
|
||||
self._last_sorted_data = self._sort_data(self.raw_data, sort_key, order, seed)
|
||||
|
||||
async def update_preview_url(self, file_path: str, preview_url: str, preview_nsfw_level: int) -> bool:
|
||||
"""Update preview_url for a specific model in all cached data
|
||||
@@ -337,4 +346,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
|
||||
@@ -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()
|
||||
|
||||
@@ -48,6 +48,36 @@ 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.abspath()`` (NOT ``realpath``) to resolve ``..`` and ``.``
|
||||
while preserving symlinks — this keeps the check in business-path space.
|
||||
Skips when the scanner does not expose ``get_model_roots`` or the list
|
||||
is empty.
|
||||
"""
|
||||
|
||||
roots = None
|
||||
if hasattr(scanner, "get_model_roots"):
|
||||
try:
|
||||
roots = scanner.get_model_roots()
|
||||
except NotImplementedError:
|
||||
roots = None
|
||||
if not roots:
|
||||
return
|
||||
|
||||
resolved = os.path.abspath(os.path.normpath(file_path))
|
||||
|
||||
for root in roots:
|
||||
root_resolved = os.path.abspath(os.path.normpath(root))
|
||||
if resolved == root_resolved or resolved.startswith(root_resolved + os.sep):
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
f"{label} '{file_path}' is outside configured library directories"
|
||||
)
|
||||
|
||||
|
||||
class ModelLifecycleService:
|
||||
"""Co-ordinate destructive and mutating model operations."""
|
||||
|
||||
@@ -74,6 +104,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 +214,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 +263,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 +306,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 +319,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")
|
||||
|
||||
@@ -143,10 +143,18 @@ class ModelMetadataProvider(ABC):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
"""Fetch models owned by the specified user"""
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
"""Fetch one page of models owned by the specified user.
|
||||
|
||||
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
|
||||
or None when unsupported/failed. ``cursor`` continues a previous page.
|
||||
"""
|
||||
pass
|
||||
|
||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||
"""Published model count for the user; None when unsupported."""
|
||||
return None
|
||||
|
||||
class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
||||
"""Provider that uses Civitai API for metadata"""
|
||||
|
||||
@@ -175,8 +183,11 @@ class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
||||
return await self.client.get_model_version_info(version_id)
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
return await self.client.get_user_models(username)
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
return await self.client.get_user_models(username, cursor)
|
||||
|
||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||
return await self.client.get_creator_model_count(username)
|
||||
|
||||
class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
||||
"""Provider that uses CivArchive API for metadata"""
|
||||
@@ -196,7 +207,7 @@ class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
||||
return await self.client.get_model_version_info(version_id)
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
"""Not supported by CivArchive provider"""
|
||||
return None
|
||||
|
||||
@@ -347,7 +358,7 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
|
||||
version_data = await self._get_version_with_model_data(db, model_id, version_id)
|
||||
return version_data, None
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
"""Listing models by username is not supported for archive database"""
|
||||
return None
|
||||
|
||||
@@ -602,13 +613,14 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
||||
continue
|
||||
return None
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
for provider, label in self._iter_providers():
|
||||
try:
|
||||
result = await self._call_with_rate_limit(
|
||||
label,
|
||||
provider.get_user_models,
|
||||
username,
|
||||
cursor=cursor,
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
@@ -624,6 +636,19 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
||||
continue
|
||||
return None
|
||||
|
||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||
for provider, label in self._iter_providers():
|
||||
try:
|
||||
result = await provider.get_creator_model_count(username)
|
||||
if result is not None:
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.debug(
|
||||
"Provider %s failed for get_creator_model_count: %s", label, e
|
||||
)
|
||||
continue
|
||||
return None
|
||||
|
||||
def _iter_providers(self):
|
||||
return zip(self.providers, self._provider_labels)
|
||||
|
||||
@@ -704,13 +729,17 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
|
||||
version_id,
|
||||
)
|
||||
|
||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict]:
|
||||
return await self._rate_limit_helper.run(
|
||||
self._label,
|
||||
self._provider.get_user_models,
|
||||
username,
|
||||
cursor=cursor,
|
||||
)
|
||||
|
||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||
return await self._provider.get_creator_model_count(username)
|
||||
|
||||
class ModelMetadataProviderManager:
|
||||
"""Manager for selecting and using model metadata providers"""
|
||||
|
||||
@@ -776,10 +805,20 @@ class ModelMetadataProviderManager:
|
||||
except NotImplementedError:
|
||||
return None
|
||||
|
||||
async def get_user_models(self, username: str, provider_name: str = None) -> Optional[List[Dict]]:
|
||||
"""Fetch models owned by the specified user"""
|
||||
async def get_user_models(
|
||||
self,
|
||||
username: str,
|
||||
provider_name: str = None,
|
||||
cursor: Optional[str] = None,
|
||||
) -> Optional[Dict]:
|
||||
"""Fetch one page of models owned by the specified user"""
|
||||
provider = self._get_provider(provider_name)
|
||||
return await provider.get_user_models(username)
|
||||
return await provider.get_user_models(username, cursor)
|
||||
|
||||
async def get_creator_model_count(self, username: str, provider_name: str = None) -> Optional[int]:
|
||||
"""Best-effort published model count for the specified user"""
|
||||
provider = self._get_provider(provider_name)
|
||||
return await provider.get_creator_model_count(username)
|
||||
|
||||
def _get_provider(self, provider_name: str = None) -> ModelMetadataProvider:
|
||||
"""Get provider by name or default provider"""
|
||||
|
||||
@@ -85,6 +85,7 @@ class SortParams:
|
||||
|
||||
key: str
|
||||
order: str
|
||||
seed: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -116,7 +117,7 @@ class ModelCacheRepository:
|
||||
async def fetch_sorted(self, params: SortParams) -> List[Dict[str, Any]]:
|
||||
"""Fetch cached data pre-sorted according to ``params``."""
|
||||
cache = await self.get_cache()
|
||||
return await cache.get_sorted_data(params.key, params.order)
|
||||
return await cache.get_sorted_data(params.key, params.order, params.seed)
|
||||
|
||||
@staticmethod
|
||||
def parse_sort(sort_by: str) -> SortParams:
|
||||
@@ -132,10 +133,17 @@ class ModelCacheRepository:
|
||||
sort_key = sort_by.strip().lower() or "name"
|
||||
order = "asc"
|
||||
|
||||
if order not in ("asc", "desc"):
|
||||
seed = None
|
||||
if sort_key == "random":
|
||||
# Random sort: the portion after ':' is the shuffle seed.
|
||||
# A stable seed keeps paginated requests consistent; order is
|
||||
# meaningless for a random shuffle.
|
||||
seed = order if order and order not in ("asc", "desc") else None
|
||||
order = "asc"
|
||||
elif order not in ("asc", "desc"):
|
||||
order = "asc"
|
||||
|
||||
return SortParams(key=sort_key, order=order)
|
||||
return SortParams(key=sort_key, order=order, seed=seed)
|
||||
|
||||
|
||||
class ModelFilterSet:
|
||||
|
||||
+284
-11
@@ -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,
|
||||
@@ -922,6 +927,25 @@ class ModelScanner:
|
||||
# Update cache data
|
||||
self._cache.raw_data = [item for item in self._cache.raw_data if item['file_path'] not in missing_files]
|
||||
|
||||
dedup_removed = 0
|
||||
seen_paths: set = set()
|
||||
deduped: list = []
|
||||
for item in reversed(self._cache.raw_data):
|
||||
path = item.get('file_path', '')
|
||||
if path not in seen_paths:
|
||||
seen_paths.add(path)
|
||||
deduped.append(item)
|
||||
else:
|
||||
for tag in item.get('tags', []):
|
||||
if tag in self._tags_count:
|
||||
self._tags_count[tag] = max(0, self._tags_count[tag] - 1)
|
||||
if self._tags_count[tag] == 0:
|
||||
del self._tags_count[tag]
|
||||
dedup_removed += 1
|
||||
if dedup_removed > 0:
|
||||
self._cache.raw_data = list(reversed(deduped))
|
||||
total_removed += dedup_removed
|
||||
|
||||
# Resort cache if changes were made
|
||||
if total_added > 0 or total_removed > 0:
|
||||
# Update folders list
|
||||
@@ -1347,18 +1371,25 @@ class ModelScanner:
|
||||
# Update folder in metadata
|
||||
metadata_dict['folder'] = folder
|
||||
|
||||
# Add to cache
|
||||
self._cache.raw_data.append(metadata_dict)
|
||||
self._cache.add_to_version_index(metadata_dict)
|
||||
file_path = metadata_dict.get('file_path', '')
|
||||
if file_path:
|
||||
old_entries = [item for item in self._cache.raw_data if item.get('file_path') == file_path]
|
||||
for old_entry in old_entries:
|
||||
for tag in old_entry.get('tags', []):
|
||||
if tag in self._tags_count:
|
||||
self._tags_count[tag] = max(0, self._tags_count[tag] - 1)
|
||||
if self._tags_count[tag] == 0:
|
||||
del self._tags_count[tag]
|
||||
self._hash_index.remove_by_path(file_path)
|
||||
self._cache.raw_data = [item for item in self._cache.raw_data if item.get('file_path') != file_path]
|
||||
|
||||
for tag in metadata_dict.get('tags', []):
|
||||
self._tags_count[tag] = self._tags_count.get(tag, 0) + 1
|
||||
|
||||
self._cache.raw_data.append(metadata_dict)
|
||||
|
||||
# Resort cache data
|
||||
await self._cache.resort()
|
||||
|
||||
# Update folders list
|
||||
all_folders = set(self._cache.folders)
|
||||
all_folders.add(folder)
|
||||
self._cache.folders = sorted(list(all_folders), key=lambda x: x.lower())
|
||||
|
||||
# Update the hash index
|
||||
self._hash_index.add_entry(metadata_dict['sha256'], metadata_dict['file_path'])
|
||||
await self._persist_current_cache()
|
||||
@@ -1389,6 +1420,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 +1595,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[0] is not 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 +1859,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 +2000,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)
|
||||
|
||||
@@ -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 = ?",
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -21,7 +21,7 @@ from .checkpoint_scanner import CheckpointScanner
|
||||
from .settings_manager import get_settings_manager
|
||||
from .recipes.errors import RecipeNotFoundError
|
||||
from ..utils.civitai_utils import extract_civitai_image_id
|
||||
from ..utils.utils import calculate_recipe_fingerprint, fuzzy_match
|
||||
from ..utils.utils import calculate_recipe_fingerprint
|
||||
from natsort import natsorted
|
||||
import sys
|
||||
import re
|
||||
@@ -1020,13 +1020,16 @@ class RecipeScanner:
|
||||
|
||||
try:
|
||||
result = self._fts_index.search(search, fields)
|
||||
# Return None if empty to trigger fuzzy fallback
|
||||
# Empty FTS results may indicate query syntax issues or need for fuzzy matching
|
||||
# Return empty set for empty FTS results — do NOT fall back to
|
||||
# Python fuzzy matching, which freezes the server with 10k+ recipes.
|
||||
# FTS5 prefix matching with unicode61 tokenizer correctly handles
|
||||
# compound tokens (e.g. "illustrious" matches "path/illustrious/model").
|
||||
# If FTS returns nothing, there are genuinely no matching recipes.
|
||||
if not result:
|
||||
return None
|
||||
return set()
|
||||
return result
|
||||
except Exception as exc:
|
||||
logger.debug("FTS search failed, falling back to fuzzy search: %s", exc)
|
||||
logger.debug("FTS search failed, falling back to title-only search: %s", exc)
|
||||
return None
|
||||
|
||||
def _update_fts_index_for_recipe(
|
||||
@@ -2079,49 +2082,14 @@ class RecipeScanner:
|
||||
if str(item.get("id", "")) in fts_matching_ids
|
||||
]
|
||||
else:
|
||||
# Fallback to fuzzy_match (slower but always available)
|
||||
# Build the search predicate based on search options
|
||||
def matches_search(item):
|
||||
# Search in title if enabled
|
||||
if search_options.get("title", True):
|
||||
if fuzzy_match(str(item.get("title", "")), search):
|
||||
return True
|
||||
|
||||
# Search in tags if enabled
|
||||
if search_options.get("tags", True) and "tags" in item:
|
||||
for tag in item["tags"]:
|
||||
if fuzzy_match(tag, search):
|
||||
return True
|
||||
|
||||
# Search in lora file names if enabled
|
||||
if search_options.get("lora_name", True) and "loras" in item:
|
||||
for lora in item["loras"]:
|
||||
if fuzzy_match(str(lora.get("file_name", "")), search):
|
||||
return True
|
||||
|
||||
# Search in lora model names if enabled
|
||||
if search_options.get("lora_model", True) and "loras" in item:
|
||||
for lora in item["loras"]:
|
||||
if fuzzy_match(str(lora.get("modelName", "")), search):
|
||||
return True
|
||||
|
||||
# Search in prompt and negative_prompt if enabled
|
||||
if search_options.get("prompt", True) and "gen_params" in item:
|
||||
gen_params = item["gen_params"]
|
||||
if fuzzy_match(str(gen_params.get("prompt", "")), search):
|
||||
return True
|
||||
if fuzzy_match(
|
||||
str(gen_params.get("negative_prompt", "")), search
|
||||
):
|
||||
return True
|
||||
|
||||
# No match found
|
||||
return False
|
||||
|
||||
# Filter the data using the search predicate
|
||||
filtered_data = [
|
||||
item for item in filtered_data if matches_search(item)
|
||||
]
|
||||
# FTS index not yet built — return empty rather than
|
||||
# scanning 42k+ items in Python. The FTS background build
|
||||
# finishes in seconds; by the time a user navigates here
|
||||
# and types a search, it is already available.
|
||||
logger.debug(
|
||||
"FTS index not ready — search '%s' returning empty", search
|
||||
)
|
||||
filtered_data = []
|
||||
|
||||
# Apply additional filters
|
||||
if filters:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -122,6 +126,7 @@ class BulkMetadataRefreshUseCase:
|
||||
if sha256:
|
||||
model["sha256"] = sha256
|
||||
model["hash_status"] = "completed"
|
||||
hash_status = "completed"
|
||||
else:
|
||||
self._logger.error(f"Failed to calculate hash for {file_path}")
|
||||
failures.append({"name": model.get("model_name", file_path or "Unknown"), "error": "Failed to calculate hash"})
|
||||
@@ -144,6 +149,16 @@ class BulkMetadataRefreshUseCase:
|
||||
continue
|
||||
|
||||
await MetadataManager.hydrate_model_data(model)
|
||||
|
||||
# hydrate_model_data replaces model with .metadata.json content,
|
||||
# which may lack sha256. Restore from cache and persist the fix.
|
||||
if not model.get("sha256"):
|
||||
model["sha256"] = sha256
|
||||
model["hash_status"] = model.get("hash_status", hash_status)
|
||||
data_to_save = model.copy()
|
||||
data_to_save.pop("folder", None)
|
||||
await MetadataManager.save_metadata(file_path, data_to_save)
|
||||
|
||||
result, error_msg = await self._metadata_sync.fetch_and_update_model(
|
||||
sha256=model["sha256"],
|
||||
file_path=model["file_path"],
|
||||
|
||||
@@ -19,7 +19,7 @@ logger = logging.getLogger(__name__)
|
||||
_WILDCARD_PATTERN = re.compile(r"__([\w\s.\-+/*\\]+?)__")
|
||||
_OPTION_PATTERN = re.compile(r"{([^{}]*?)}")
|
||||
_TRIGGER_WORD_PATTERN = re.compile(r"^trigger_words\d+$")
|
||||
_WEIGHTED_OPTION_PATTERN = re.compile(r"^\s*([0-9.]+)::")
|
||||
_WEIGHTED_OPTION_PATTERN = re.compile(r"^\s*-?\d+(\.\d+)?::")
|
||||
_NUMERIC_PATTERN = re.compile(r"^-?\d+(\.\d+)?$")
|
||||
|
||||
|
||||
@@ -390,7 +390,7 @@ class WildcardService:
|
||||
) -> str | None:
|
||||
keyword = _normalize_wildcard_key(raw_key)
|
||||
if keyword in wildcard_dict:
|
||||
return rng.choice(wildcard_dict[keyword])
|
||||
return self._pick_weighted_or_plain(wildcard_dict[keyword], rng)
|
||||
|
||||
if "*" in keyword:
|
||||
regex_pattern = keyword.replace("*", ".*").replace("+", r"\+")
|
||||
@@ -400,7 +400,7 @@ class WildcardService:
|
||||
if compiled.match(key):
|
||||
aggregated.extend(values)
|
||||
if aggregated:
|
||||
return rng.choice(aggregated)
|
||||
return self._pick_weighted_or_plain(aggregated, rng)
|
||||
|
||||
if "/" not in keyword:
|
||||
fallback_keyword = _normalize_wildcard_key(f"*/{keyword}")
|
||||
@@ -409,6 +409,39 @@ class WildcardService:
|
||||
|
||||
return None
|
||||
|
||||
def _pick_weighted_or_plain(
|
||||
self, values: list[str], rng: random.Random
|
||||
) -> str:
|
||||
"""Pick a value from the list, respecting N::weight prefix if present.
|
||||
|
||||
When any value in the list uses the ``N::value`` weighted syntax with a
|
||||
weight different from 1, the pick uses weighted random selection. When
|
||||
no such weighting is present, a plain ``rng.choice`` is used (preserving
|
||||
backward compatibility for unweighted wildcard files).
|
||||
|
||||
In either case the ``N::`` prefix is always stripped from the returned
|
||||
value, matching the behaviour of ``{...}`` option groups.
|
||||
"""
|
||||
# Fast path: skip weighting logic entirely when no :: syntax exists
|
||||
if not any("::" in v for v in values):
|
||||
return rng.choice(values)
|
||||
|
||||
weighted_options: list[tuple[float, str]] = []
|
||||
for value in values:
|
||||
weight = 1.0
|
||||
parts = value.split("::", 1)
|
||||
if len(parts) == 2 and _is_numeric_string(parts[0].strip()):
|
||||
weight = float(parts[0].strip())
|
||||
weighted_options.append((weight, value))
|
||||
|
||||
any_weighted = any(w != 1.0 for w, _ in weighted_options)
|
||||
if any_weighted:
|
||||
picked = self._weighted_choice(weighted_options, rng)
|
||||
else:
|
||||
picked = rng.choice(values)
|
||||
|
||||
return self._strip_weight_prefix(picked)
|
||||
|
||||
|
||||
def is_trigger_words_input(name: str) -> bool:
|
||||
return bool(_TRIGGER_WORD_PATTERN.match(name))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -14,11 +14,16 @@ from ..services.service_registry import ServiceRegistry
|
||||
from ..utils.example_images_paths import (
|
||||
ExampleImagePathResolver,
|
||||
ensure_library_root_exists,
|
||||
get_example_images_root,
|
||||
is_hash_folder,
|
||||
uses_library_scoped_folders,
|
||||
)
|
||||
from ..utils.metadata_manager import MetadataManager
|
||||
from .example_images_processor import ExampleImagesProcessor
|
||||
from .example_images_metadata import MetadataUpdater
|
||||
from .example_images_metadata import (
|
||||
MetadataUpdater,
|
||||
update_cache_from_metadata,
|
||||
)
|
||||
from ..services.downloader import get_downloader
|
||||
from ..services.settings_manager import get_settings_manager
|
||||
|
||||
@@ -72,6 +77,7 @@ class _DownloadProgress(dict):
|
||||
refreshed_models=set(),
|
||||
failed_models=set(),
|
||||
reprocessed_models=set(),
|
||||
rate_limited_models=set(),
|
||||
)
|
||||
|
||||
def snapshot(self) -> dict:
|
||||
@@ -82,9 +88,17 @@ class _DownloadProgress(dict):
|
||||
snapshot["refreshed_models"] = list(self["refreshed_models"])
|
||||
snapshot["failed_models"] = list(self["failed_models"])
|
||||
snapshot["reprocessed_models"] = list(self.get("reprocessed_models", set()))
|
||||
snapshot["rate_limited_models"] = list(self.get("rate_limited_models", set()))
|
||||
return snapshot
|
||||
|
||||
|
||||
# When fewer candidates than this remain in check_pending_models, probe each
|
||||
# model folder directly (preserving legacy-folder migration semantics). Above
|
||||
# it, build a folder index with a single directory scan so libraries with
|
||||
# 100k+ models do not pay one syscall per candidate.
|
||||
_BULK_LOOKUP_THRESHOLD = 1000
|
||||
|
||||
|
||||
def _model_directory_has_files(path: str) -> bool:
|
||||
"""Return True when the provided directory exists and contains entries."""
|
||||
|
||||
@@ -101,6 +115,36 @@ def _model_directory_has_files(path: str) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _build_example_folder_index(output_dir: str) -> dict[str, bool]:
|
||||
"""Build a ``{hash: has_files}`` index for a library's example-image folders.
|
||||
|
||||
A single directory scan over the library root replaces ``O(candidates)``
|
||||
per-folder ``os.scandir`` calls, which is required for libraries with
|
||||
100k+ models. Each hash folder is classified by whether it contains any
|
||||
entries, matching the semantics of ``_model_directory_has_files``.
|
||||
"""
|
||||
|
||||
index: dict[str, bool] = {}
|
||||
if not output_dir or not os.path.isdir(output_dir):
|
||||
return index
|
||||
|
||||
try:
|
||||
with os.scandir(output_dir) as entries:
|
||||
for entry in entries:
|
||||
name = entry.name
|
||||
if not entry.is_dir() or not is_hash_folder(name):
|
||||
continue
|
||||
try:
|
||||
with os.scandir(entry.path) as subentries:
|
||||
index[name.lower()] = any(subentries)
|
||||
except OSError:
|
||||
index[name.lower()] = False
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
return index
|
||||
|
||||
|
||||
class DownloadManager:
|
||||
"""Manages downloading example images for models."""
|
||||
|
||||
@@ -153,13 +197,15 @@ class DownloadManager:
|
||||
# Step 3: Load progress file (I/O operation, done outside lock)
|
||||
processed_models = set()
|
||||
failed_models = set()
|
||||
rate_limited_models = set()
|
||||
|
||||
try:
|
||||
progress_file, processed_models, failed_models = await self._load_progress_file(output_dir)
|
||||
progress_file, processed_models, failed_models, rate_limited_models = await self._load_progress_file(output_dir)
|
||||
logger.debug(
|
||||
"Loaded previous progress, %s models already processed, %s models marked as failed",
|
||||
"Loaded previous progress, %s models already processed, %s models marked as failed, %s models rate-limited",
|
||||
len(processed_models),
|
||||
len(failed_models),
|
||||
len(rate_limited_models),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to load progress file: {e}")
|
||||
@@ -175,6 +221,7 @@ class DownloadManager:
|
||||
self._progress.reset()
|
||||
self._progress["processed_models"] = processed_models
|
||||
self._progress["failed_models"] = failed_models
|
||||
self._progress["rate_limited_models"] = rate_limited_models
|
||||
self._stop_requested = False
|
||||
self._progress["status"] = "running"
|
||||
self._progress["start_time"] = time.time()
|
||||
@@ -242,8 +289,8 @@ class DownloadManager:
|
||||
"status": self._progress.snapshot(),
|
||||
}
|
||||
|
||||
async def _load_progress_file(self, output_dir: str) -> tuple[str, set, set]:
|
||||
"""Load progress file from disk. Returns (progress_file_path, processed_models, failed_models).
|
||||
async def _load_progress_file(self, output_dir: str) -> tuple[str, set, set, set]:
|
||||
"""Load progress file from disk. Returns (progress_file_path, processed_models, failed_models, rate_limited_models).
|
||||
|
||||
This is a separate async method to allow running in executor to avoid blocking event loop.
|
||||
"""
|
||||
@@ -252,8 +299,12 @@ class DownloadManager:
|
||||
None, self._load_progress_file_sync, output_dir
|
||||
)
|
||||
|
||||
def _load_progress_file_sync(self, output_dir: str) -> tuple[str, set, set]:
|
||||
"""Synchronous implementation of progress file loading."""
|
||||
def _load_progress_file_sync(self, output_dir: str) -> tuple[str, set, set, set]:
|
||||
"""Synchronous implementation of progress file loading.
|
||||
|
||||
Returns:
|
||||
tuple: (progress_file_path, processed_models, failed_models, rate_limited_models)
|
||||
"""
|
||||
progress_file = os.path.join(output_dir, ".download_progress.json")
|
||||
progress_source = progress_file
|
||||
|
||||
@@ -289,6 +340,7 @@ class DownloadManager:
|
||||
|
||||
processed_models = set()
|
||||
failed_models = set()
|
||||
rate_limited_models = set()
|
||||
|
||||
if os.path.exists(progress_source):
|
||||
try:
|
||||
@@ -296,11 +348,11 @@ class DownloadManager:
|
||||
saved_progress = json.load(f)
|
||||
processed_models = set(saved_progress.get("processed_models", []))
|
||||
failed_models = set(saved_progress.get("failed_models", []))
|
||||
rate_limited_models = set(saved_progress.get("rate_limited_models", []))
|
||||
except Exception:
|
||||
# Return empty sets on error
|
||||
pass
|
||||
|
||||
return progress_file, processed_models, failed_models
|
||||
return progress_file, processed_models, failed_models, rate_limited_models
|
||||
|
||||
def _load_progress_sets_sync(self, progress_file: str) -> tuple[set, set]:
|
||||
"""Load only the processed and failed model sets from progress file.
|
||||
@@ -400,14 +452,49 @@ class DownloadManager:
|
||||
# Calculate pending count: check which models actually need processing.
|
||||
# A model is pending if it has a hash, is not already processed or known-failed,
|
||||
# and its folder doesn't exist or is empty.
|
||||
pending_hashes = set()
|
||||
for model_hash, model_name in all_models_with_hash:
|
||||
if model_hash not in processed_models and model_hash not in failed_models:
|
||||
candidate_hashes = [
|
||||
model_hash
|
||||
for model_hash, _ in all_models_with_hash
|
||||
if model_hash not in processed_models
|
||||
and model_hash not in failed_models
|
||||
]
|
||||
|
||||
pending_hashes: set[str] = set()
|
||||
# For small candidate counts the existing per-folder check is fine
|
||||
# and handles legacy folder migration.
|
||||
# For large libraries, scan the library root once and do set lookups.
|
||||
if len(candidate_hashes) <= _BULK_LOOKUP_THRESHOLD or not output_dir:
|
||||
for model_hash in candidate_hashes:
|
||||
model_dir = ExampleImagePathResolver.get_model_folder(
|
||||
model_hash, active_library
|
||||
)
|
||||
if not _model_directory_has_files(model_dir):
|
||||
pending_hashes.add(model_hash)
|
||||
else:
|
||||
folder_index = await asyncio.get_event_loop().run_in_executor(
|
||||
None, _build_example_folder_index, output_dir
|
||||
)
|
||||
# In multi-library mode, folders that have not been consolidated
|
||||
# into the library root yet (startup migration skipped, failed
|
||||
# move, or created at the legacy path afterwards) still live at
|
||||
# the legacy root/<hash> location. Only scan that root when at
|
||||
# least one candidate is missing from the library-root index, so
|
||||
# the fully-consolidated case does not pay an extra directory
|
||||
# pass on every call.
|
||||
if uses_library_scoped_folders() and any(
|
||||
not folder_index.get(model_hash, False)
|
||||
for model_hash in candidate_hashes
|
||||
):
|
||||
legacy_root = get_example_images_root()
|
||||
if legacy_root and legacy_root != output_dir:
|
||||
legacy_index = await asyncio.get_event_loop().run_in_executor(
|
||||
None, _build_example_folder_index, legacy_root
|
||||
)
|
||||
for hash_key, has_files in legacy_index.items():
|
||||
folder_index.setdefault(hash_key, has_files)
|
||||
for model_hash in candidate_hashes:
|
||||
if not folder_index.get(model_hash, False):
|
||||
pending_hashes.add(model_hash)
|
||||
|
||||
pending_count = len(pending_hashes)
|
||||
|
||||
@@ -732,11 +819,13 @@ class DownloadManager:
|
||||
success,
|
||||
is_stale,
|
||||
failed_images,
|
||||
rate_limited_images,
|
||||
) = await ExampleImagesProcessor.download_model_images_with_tracking(
|
||||
model_hash, model_name, images, model_dir, optimize, downloader
|
||||
)
|
||||
|
||||
failed_urls: Set[str] = set(failed_images)
|
||||
rate_limited_urls: Set[str] = set(rate_limited_images)
|
||||
|
||||
# If metadata is stale, try to refresh it
|
||||
if is_stale and model_hash not in self._progress["refreshed_models"]:
|
||||
@@ -760,6 +849,7 @@ class DownloadManager:
|
||||
success,
|
||||
_,
|
||||
additional_failed,
|
||||
additional_rate_limited,
|
||||
) = await ExampleImagesProcessor.download_model_images_with_tracking(
|
||||
model_hash,
|
||||
model_name,
|
||||
@@ -770,29 +860,50 @@ class DownloadManager:
|
||||
)
|
||||
|
||||
failed_urls.update(additional_failed)
|
||||
rate_limited_urls.update(additional_rate_limited)
|
||||
|
||||
self._progress["refreshed_models"].add(model_hash)
|
||||
|
||||
if failed_urls:
|
||||
# Separate permanent failures from rate-limited ones
|
||||
permanent_failures = failed_urls - rate_limited_urls
|
||||
|
||||
if permanent_failures:
|
||||
await self._remove_failed_images_from_metadata(
|
||||
model_hash,
|
||||
model_name,
|
||||
model_dir,
|
||||
failed_urls,
|
||||
permanent_failures,
|
||||
scanner,
|
||||
)
|
||||
|
||||
if failed_urls:
|
||||
if rate_limited_urls:
|
||||
self._progress["rate_limited_models"].add(model_hash)
|
||||
logger.warning(
|
||||
"%d example images for %s are rate-limited (429), will retry next time",
|
||||
len(rate_limited_urls),
|
||||
model_name,
|
||||
)
|
||||
# Clear failed_models so non-force runs can retry
|
||||
if force and model_hash in self._progress["failed_models"]:
|
||||
self._progress["failed_models"].discard(model_hash)
|
||||
logger.info(
|
||||
f"Removed {model_name} from failed_models after force retry with rate-limited images"
|
||||
)
|
||||
|
||||
if rate_limited_urls:
|
||||
# Don't mark as failed or fully processed — rate-limited
|
||||
# images will be retried next time.
|
||||
pass
|
||||
elif permanent_failures:
|
||||
self._progress["failed_models"].add(model_hash)
|
||||
self._progress["processed_models"].add(model_hash)
|
||||
logger.info(
|
||||
"Removed %s failed example images for %s",
|
||||
len(failed_urls),
|
||||
len(permanent_failures),
|
||||
model_name,
|
||||
)
|
||||
elif success:
|
||||
self._progress["processed_models"].add(model_hash)
|
||||
# Remove from failed_models if force mode enabled and model was previously failed
|
||||
if force and model_hash in self._progress["failed_models"]:
|
||||
self._progress["failed_models"].discard(model_hash)
|
||||
logger.info(
|
||||
@@ -850,6 +961,7 @@ class DownloadManager:
|
||||
"processed_models": list(self._progress["processed_models"]),
|
||||
"refreshed_models": list(self._progress["refreshed_models"]),
|
||||
"failed_models": list(self._progress["failed_models"]),
|
||||
"rate_limited_models": list(self._progress.get("rate_limited_models", set())),
|
||||
"completed": self._progress["completed"],
|
||||
"total": self._progress["total"],
|
||||
"last_update": time.time(),
|
||||
@@ -1155,11 +1267,13 @@ class DownloadManager:
|
||||
success,
|
||||
is_stale,
|
||||
failed_images,
|
||||
rate_limited_images,
|
||||
) = await ExampleImagesProcessor.download_model_images_with_tracking(
|
||||
model_hash, model_name, images, model_dir, optimize, downloader
|
||||
)
|
||||
|
||||
failed_urls: Set[str] = set(failed_images)
|
||||
rate_limited_urls: Set[str] = set(rate_limited_images)
|
||||
|
||||
# If metadata is stale, try to refresh it
|
||||
if is_stale and model_hash not in self._progress["refreshed_models"]:
|
||||
@@ -1183,6 +1297,7 @@ class DownloadManager:
|
||||
success,
|
||||
_,
|
||||
additional_failed_images,
|
||||
additional_rate_limited,
|
||||
) = await ExampleImagesProcessor.download_model_images_with_tracking(
|
||||
model_hash,
|
||||
model_name,
|
||||
@@ -1192,21 +1307,35 @@ class DownloadManager:
|
||||
downloader,
|
||||
)
|
||||
|
||||
# Combine failed images from both attempts
|
||||
failed_urls.update(additional_failed_images)
|
||||
rate_limited_urls.update(additional_rate_limited)
|
||||
|
||||
self._progress["refreshed_models"].add(model_hash)
|
||||
|
||||
# For forced downloads, remove failed images from metadata
|
||||
if failed_urls:
|
||||
# Separate permanent failures from rate-limited ones
|
||||
permanent_failures = failed_urls - rate_limited_urls
|
||||
|
||||
# Only remove permanently failed images from metadata
|
||||
if permanent_failures:
|
||||
await self._remove_failed_images_from_metadata(
|
||||
model_hash, model_name, model_dir, failed_urls, scanner
|
||||
model_hash, model_name, model_dir, permanent_failures, scanner
|
||||
)
|
||||
|
||||
# Mark as processed
|
||||
if (
|
||||
success or failed_urls
|
||||
): # Mark as processed if we successfully downloaded some images or removed failed ones
|
||||
if rate_limited_urls:
|
||||
self._progress["rate_limited_models"].add(model_hash)
|
||||
logger.warning(
|
||||
"%d example images for %s are rate-limited (429), will retry next time",
|
||||
len(rate_limited_urls),
|
||||
model_name,
|
||||
)
|
||||
|
||||
# Mark as processed only when no rate-limited images remain
|
||||
if rate_limited_urls:
|
||||
pass
|
||||
elif permanent_failures:
|
||||
self._progress["processed_models"].add(model_hash)
|
||||
self._progress["failed_models"].add(model_hash)
|
||||
elif success:
|
||||
self._progress["processed_models"].add(model_hash)
|
||||
|
||||
return True # Return True to indicate a remote download happened
|
||||
@@ -1229,15 +1358,20 @@ class DownloadManager:
|
||||
model_dir: str,
|
||||
failed_images: Iterable[str],
|
||||
scanner,
|
||||
error_type: str = "not_found",
|
||||
) -> None:
|
||||
"""Mark failed images in model metadata so they won't be retried."""
|
||||
"""Mark failed images in model metadata so they won't be retried.
|
||||
|
||||
Args:
|
||||
error_type: Reason string stored in the image's ``downloadError`` field
|
||||
(default ``"not_found"``).
|
||||
"""
|
||||
|
||||
failed_set: Set[str] = {url for url in failed_images if url}
|
||||
if not failed_set:
|
||||
return
|
||||
|
||||
try:
|
||||
# Get current model data
|
||||
model_data = await MetadataUpdater.get_updated_model(model_hash, scanner)
|
||||
if not model_data:
|
||||
logger.warning(
|
||||
@@ -1268,7 +1402,7 @@ class DownloadManager:
|
||||
continue
|
||||
|
||||
image["downloadFailed"] = True
|
||||
image.setdefault("downloadError", "not_found")
|
||||
image.setdefault("downloadError", error_type)
|
||||
logger.debug(
|
||||
"Marked example image %s for %s as failed due to missing remote asset",
|
||||
image_url,
|
||||
@@ -1286,8 +1420,8 @@ class DownloadManager:
|
||||
await MetadataManager.save_metadata(file_path, model_copy)
|
||||
|
||||
try:
|
||||
await scanner.update_single_model_cache(
|
||||
file_path, file_path, model_data
|
||||
await update_cache_from_metadata(
|
||||
scanner, file_path, model_copy
|
||||
)
|
||||
except AttributeError:
|
||||
logger.debug(
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
@@ -28,6 +29,31 @@ if TYPE_CHECKING: # pragma: no cover - import for type checkers only
|
||||
from ..services.settings_manager import SettingsManager
|
||||
|
||||
|
||||
async def update_cache_from_metadata(
|
||||
scanner: Any, file_path: str, metadata: Dict[str, Any]
|
||||
) -> bool:
|
||||
"""Update the scanner cache from a metadata dict using the in-place sync path.
|
||||
|
||||
``sync_cache_from_metadata`` patches the existing cache entry incrementally
|
||||
(tag/hash/version indexes, targeted single-row SQL update) and only resorts
|
||||
when a sort-key field changed. This avoids the ``O(n)`` full-list resort and
|
||||
full cache rewrite that ``update_single_model_cache`` performs on every call,
|
||||
which is critical for libraries with 100k+ models.
|
||||
|
||||
Falls back to the legacy full update when the scanner does not expose an
|
||||
async ``sync_cache_from_metadata`` method.
|
||||
|
||||
Returns:
|
||||
``True`` if the cache entry was updated, ``False`` otherwise.
|
||||
"""
|
||||
|
||||
sync_method = getattr(scanner, "sync_cache_from_metadata", None)
|
||||
if inspect.iscoroutinefunction(sync_method):
|
||||
return await sync_method(file_path, metadata)
|
||||
|
||||
return await scanner.update_single_model_cache(file_path, file_path, metadata)
|
||||
|
||||
|
||||
def _build_metadata_sync_service(settings_manager: "SettingsManager") -> MetadataSyncService:
|
||||
"""Construct a metadata sync service bound to the provided settings."""
|
||||
|
||||
@@ -103,8 +129,8 @@ class MetadataUpdater:
|
||||
progress['refreshed_models'].add(model_hash)
|
||||
|
||||
async def update_cache_func(old_path, new_path, metadata):
|
||||
return await scanner.update_single_model_cache(old_path, new_path, metadata)
|
||||
|
||||
return await update_cache_from_metadata(scanner, new_path, metadata)
|
||||
|
||||
await MetadataManager.hydrate_model_data(model_data)
|
||||
success, error = await _get_metadata_sync_service().fetch_and_update_model(
|
||||
sha256=model_hash,
|
||||
@@ -234,6 +260,7 @@ class MetadataUpdater:
|
||||
|
||||
# Save metadata to .metadata.json file
|
||||
file_path = model.get('file_path')
|
||||
model_copy: Optional[Dict[str, Any]] = None
|
||||
try:
|
||||
model_copy = model.copy()
|
||||
model_copy.pop('folder', None)
|
||||
@@ -241,14 +268,18 @@ class MetadataUpdater:
|
||||
logger.info(f"Saved metadata for {model.get('model_name')}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to save metadata for {model.get('model_name')}: {str(e)}")
|
||||
|
||||
# Save updated metadata to scanner cache
|
||||
success = await scanner.update_single_model_cache(file_path, file_path, model)
|
||||
if success:
|
||||
|
||||
# Save updated metadata to scanner cache. sync_cache_from_metadata
|
||||
# returns False both for "already in sync" and for actual failures,
|
||||
# so the cache sync result is deliberately not treated as an error;
|
||||
# the return value reflects whether the metadata was persisted.
|
||||
if file_path and model_copy is not None:
|
||||
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||
logger.info(f"Successfully updated metadata for {model.get('model_name')} with {len(images)} local examples")
|
||||
return True
|
||||
else:
|
||||
logger.warning(f"Failed to update metadata for {model.get('model_name')}")
|
||||
|
||||
logger.warning(f"Failed to update metadata for {model.get('model_name')}")
|
||||
return False
|
||||
|
||||
return False
|
||||
except Exception as e:
|
||||
@@ -336,6 +367,7 @@ class MetadataUpdater:
|
||||
|
||||
# Save metadata to .metadata.json file
|
||||
file_path = model_data.get('file_path')
|
||||
model_copy: Optional[Dict[str, Any]] = None
|
||||
if file_path:
|
||||
try:
|
||||
model_copy = model_data.copy()
|
||||
@@ -344,11 +376,11 @@ class MetadataUpdater:
|
||||
logger.info(f"Saved metadata for {model_data.get('model_name')}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to save metadata: {str(e)}")
|
||||
|
||||
|
||||
# Save updated metadata to scanner cache
|
||||
if file_path:
|
||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
||||
|
||||
if file_path and model_copy is not None:
|
||||
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||
|
||||
# Get regular images array (might be None)
|
||||
regular_images = civitai_data.get('images', [])
|
||||
|
||||
@@ -475,13 +507,19 @@ class MetadataUpdater:
|
||||
return False
|
||||
|
||||
model_folder = get_model_folder(model_hash)
|
||||
if not model_folder:
|
||||
if not model_folder or not os.path.isdir(model_folder):
|
||||
return False
|
||||
|
||||
civitai = getattr(metadata, "civitai", None)
|
||||
if not isinstance(civitai, dict):
|
||||
return False
|
||||
|
||||
# Read the directory listing once so every image entry reuses it.
|
||||
try:
|
||||
dir_entries = os.listdir(model_folder)
|
||||
except OSError:
|
||||
dir_entries = []
|
||||
|
||||
has_changes = False
|
||||
|
||||
custom_images = civitai.get("customImages")
|
||||
@@ -493,24 +531,15 @@ class MetadataUpdater:
|
||||
if not img_id:
|
||||
continue
|
||||
|
||||
if not os.path.isdir(model_folder):
|
||||
prefix = f"custom_{img_id}"
|
||||
found = any(
|
||||
f.startswith(prefix) and os.path.isfile(
|
||||
os.path.join(model_folder, f)
|
||||
)
|
||||
for f in dir_entries
|
||||
)
|
||||
if not found:
|
||||
stale.append(idx)
|
||||
else:
|
||||
found = False
|
||||
try:
|
||||
prefix = f"custom_{img_id}"
|
||||
for fname in os.listdir(model_folder):
|
||||
if fname.startswith(prefix) and os.path.isfile(
|
||||
os.path.join(model_folder, fname)
|
||||
):
|
||||
found = True
|
||||
break
|
||||
except OSError:
|
||||
stale.append(idx)
|
||||
continue
|
||||
|
||||
if not found:
|
||||
stale.append(idx)
|
||||
|
||||
if stale:
|
||||
for idx in reversed(stale):
|
||||
@@ -532,22 +561,9 @@ class MetadataUpdater:
|
||||
# is gone.
|
||||
continue
|
||||
|
||||
if not os.path.isdir(model_folder):
|
||||
prefix = f"image_{idx}."
|
||||
if not any(f.startswith(prefix) for f in dir_entries):
|
||||
stale.append(idx)
|
||||
else:
|
||||
found = False
|
||||
try:
|
||||
prefix = f"image_{idx}."
|
||||
for fname in os.listdir(model_folder):
|
||||
if fname.startswith(prefix):
|
||||
found = True
|
||||
break
|
||||
except OSError:
|
||||
stale.append(idx)
|
||||
continue
|
||||
|
||||
if not found:
|
||||
stale.append(idx)
|
||||
|
||||
if stale:
|
||||
for idx in reversed(stale):
|
||||
|
||||
@@ -3,11 +3,19 @@ import logging
|
||||
import os
|
||||
import re
|
||||
import json
|
||||
import shutil
|
||||
from ..services.settings_manager import get_settings_manager
|
||||
from ..services.service_registry import ServiceRegistry
|
||||
from ..utils.example_images_paths import iter_library_roots
|
||||
from ..utils.example_images_paths import (
|
||||
get_example_images_root,
|
||||
is_hash_folder,
|
||||
iter_library_roots,
|
||||
uses_library_scoped_folders,
|
||||
_library_folder_has_only_hash_dirs,
|
||||
)
|
||||
from ..utils.metadata_manager import MetadataManager
|
||||
from ..utils.example_images_processor import ExampleImagesProcessor
|
||||
from ..utils.example_images_metadata import update_cache_from_metadata
|
||||
from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -36,6 +44,90 @@ settings = _SettingsProxy()
|
||||
class ExampleImagesMigration:
|
||||
"""Handles migrations for example images naming conventions"""
|
||||
|
||||
@staticmethod
|
||||
def _consolidate_library_folders():
|
||||
"""Move hash folders from library-named subdirectories back to root.
|
||||
|
||||
When a user switches from multi-library mode back to single-library
|
||||
mode, example images previously stored under e.g.
|
||||
``<root>/default/<hash>/`` need to be moved back to
|
||||
``<root>/<hash>/``. Running this once at startup removes the need
|
||||
for ``get_model_folder()`` to perform directory scans on every
|
||||
request.
|
||||
"""
|
||||
if uses_library_scoped_folders():
|
||||
return
|
||||
|
||||
root = get_example_images_root()
|
||||
if not root or not os.path.isdir(root):
|
||||
return
|
||||
|
||||
moved: list[str] = []
|
||||
cleaned: list[str] = []
|
||||
|
||||
try:
|
||||
for entry in os.listdir(root):
|
||||
# Fast regex checks first — no filesystem I/O.
|
||||
if is_hash_folder(entry) or entry == "_deleted":
|
||||
continue
|
||||
|
||||
entry_path = os.path.join(root, entry)
|
||||
if not os.path.isdir(entry_path):
|
||||
continue
|
||||
if not _library_folder_has_only_hash_dirs(entry_path):
|
||||
continue
|
||||
|
||||
try:
|
||||
for hash_entry in os.listdir(entry_path):
|
||||
hash_path = os.path.join(entry_path, hash_entry)
|
||||
if not os.path.isdir(hash_path) or not is_hash_folder(hash_entry):
|
||||
continue
|
||||
target = os.path.join(root, hash_entry)
|
||||
if not os.path.exists(target):
|
||||
try:
|
||||
shutil.move(hash_path, target)
|
||||
moved.append(hash_entry)
|
||||
except (OSError, shutil.Error) as exc:
|
||||
logger.error(
|
||||
"Failed to move '%s' → '%s': %s",
|
||||
hash_path, target, exc,
|
||||
)
|
||||
except OSError as exc:
|
||||
logger.error(
|
||||
"Failed to list library subdirectory '%s': %s",
|
||||
entry_path, exc,
|
||||
)
|
||||
|
||||
try:
|
||||
remaining = os.listdir(entry_path)
|
||||
except OSError:
|
||||
remaining = []
|
||||
if not remaining:
|
||||
try:
|
||||
os.rmdir(entry_path)
|
||||
cleaned.append(entry)
|
||||
except OSError as exc:
|
||||
logger.debug(
|
||||
"Could not remove empty library dir '%s': %s",
|
||||
entry_path, exc,
|
||||
)
|
||||
except OSError as exc:
|
||||
logger.error(
|
||||
"Failed to list example images root during consolidation: %s",
|
||||
exc,
|
||||
)
|
||||
|
||||
if moved:
|
||||
logger.info(
|
||||
"Consolidated %d example image folder(s) to root",
|
||||
len(moved),
|
||||
)
|
||||
if cleaned:
|
||||
logger.info(
|
||||
"Removed %d empty library directories",
|
||||
len(cleaned),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def check_and_run_migrations():
|
||||
"""Check if migrations are needed and run them in background"""
|
||||
@@ -44,6 +136,10 @@ class ExampleImagesMigration:
|
||||
logger.debug("No example images path configured or path doesn't exist, skipping migrations")
|
||||
return
|
||||
|
||||
# Run library-to-root consolidation once at startup so the hot
|
||||
# path (get_model_folder) stays a pure-path computation.
|
||||
ExampleImagesMigration._consolidate_library_folders()
|
||||
|
||||
for library_name, library_path in iter_library_roots():
|
||||
if not library_path or not os.path.exists(library_path):
|
||||
continue
|
||||
@@ -326,7 +422,7 @@ class ExampleImagesMigration:
|
||||
await MetadataManager.save_metadata(file_path, model_copy)
|
||||
|
||||
# Update scanner cache
|
||||
await scanner.update_single_model_cache(file_path, file_path, model_metadata)
|
||||
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||
|
||||
updated_models += 1
|
||||
except Exception as e:
|
||||
|
||||
@@ -83,7 +83,12 @@ def ensure_library_root_exists(library_name: Optional[str] = None) -> str:
|
||||
|
||||
|
||||
def get_model_folder(model_hash: str, library_name: Optional[str] = None) -> str:
|
||||
"""Return the folder path for a model's example images."""
|
||||
"""Return the folder path for a model's example images.
|
||||
|
||||
Multi-library ↔ single-library consolidation is handled once at startup by
|
||||
``ExampleImagesMigration._consolidate_library_folders`` — this function is a
|
||||
pure path computation on the hot path (no directory scans).
|
||||
"""
|
||||
|
||||
if not model_hash:
|
||||
return ""
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
@@ -8,7 +9,7 @@ from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
||||
from ..services.service_registry import ServiceRegistry
|
||||
from ..services.settings_manager import get_settings_manager
|
||||
from ..utils.example_images_paths import get_model_folder, get_model_relative_path
|
||||
from .example_images_metadata import MetadataUpdater
|
||||
from .example_images_metadata import MetadataUpdater, update_cache_from_metadata
|
||||
from ..utils.metadata_manager import MetadataManager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -194,16 +195,22 @@ class ExampleImagesProcessor:
|
||||
|
||||
return model_success, False # (success, is_metadata_stale)
|
||||
|
||||
@staticmethod
|
||||
def _extract_retry_after(error_message: str) -> int:
|
||||
if not error_message:
|
||||
return 60
|
||||
match = re.search(r"retry after (\d+)s", str(error_message))
|
||||
if match:
|
||||
return max(1, int(match.group(1)))
|
||||
return 60
|
||||
|
||||
@staticmethod
|
||||
async def download_model_images_with_tracking(model_hash, model_name, model_images, model_dir, optimize, downloader):
|
||||
"""Download images for a single model with tracking of failed image URLs
|
||||
|
||||
Returns:
|
||||
tuple: (success, is_stale_metadata, failed_images) - whether download was successful, whether metadata is stale, list of failed image URLs
|
||||
"""
|
||||
model_success = True
|
||||
failed_images = []
|
||||
|
||||
rate_limited_images = []
|
||||
any_successful_download = False
|
||||
|
||||
for i, image in enumerate(model_images):
|
||||
image_url = image.get('url')
|
||||
if not image_url:
|
||||
@@ -221,64 +228,110 @@ class ExampleImagesProcessor:
|
||||
original_url = image_url
|
||||
if optimize and 'civitai.com' in image_url:
|
||||
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
||||
|
||||
# Download the file first to determine the actual file type
|
||||
try:
|
||||
logger.debug(f"Downloading media file {i} for {model_name}")
|
||||
|
||||
# Download using the unified downloader with headers
|
||||
success, content, headers = await downloader.download_to_memory(
|
||||
|
||||
async def _attempt_download() -> tuple:
|
||||
logger.debug("Downloading media file %s for %s", i, model_name)
|
||||
return await downloader.download_to_memory(
|
||||
image_url,
|
||||
use_auth=False, # Example images don't need auth
|
||||
return_headers=True
|
||||
use_auth=False,
|
||||
return_headers=True,
|
||||
)
|
||||
|
||||
|
||||
try:
|
||||
success, content, headers = await _attempt_download()
|
||||
|
||||
if success:
|
||||
# Determine file extension from content or headers
|
||||
media_ext = ExampleImagesProcessor._get_file_extension_from_content_or_headers(
|
||||
content, headers, original_url, image.get("type")
|
||||
)
|
||||
|
||||
# Check if the detected file type is supported
|
||||
is_image = media_ext in SUPPORTED_MEDIA_EXTENSIONS['images']
|
||||
is_video = media_ext in SUPPORTED_MEDIA_EXTENSIONS['videos']
|
||||
|
||||
|
||||
if not (is_image or is_video):
|
||||
logger.debug(f"Skipping unsupported file type: {media_ext}")
|
||||
logger.debug("Skipping unsupported file type: %s", media_ext)
|
||||
continue
|
||||
|
||||
# Use 0-based indexing with the detected extension
|
||||
|
||||
save_filename = f"image_{i}{media_ext}"
|
||||
save_path = os.path.join(model_dir, save_filename)
|
||||
|
||||
# Check if already downloaded
|
||||
|
||||
if os.path.exists(save_path):
|
||||
logger.debug(f"File already exists: {save_path}")
|
||||
logger.debug("File already exists: %s", save_path)
|
||||
continue
|
||||
|
||||
# Save the file
|
||||
|
||||
with open(save_path, 'wb') as f:
|
||||
f.write(content)
|
||||
|
||||
any_successful_download = True
|
||||
|
||||
elif ExampleImagesProcessor._is_not_found_error(content):
|
||||
error_msg = f"Failed to download file: {image_url}, status code: 404 - Model metadata might be stale"
|
||||
logger.warning(error_msg)
|
||||
model_success = False # Mark the model as failed due to 404 error
|
||||
failed_images.append(image_url) # Track failed URL
|
||||
# Return early to trigger metadata refresh attempt
|
||||
return False, True, failed_images # (success, is_metadata_stale, failed_images)
|
||||
model_success = False
|
||||
failed_images.append(image_url)
|
||||
return False, True, failed_images, rate_limited_images
|
||||
|
||||
elif "Rate limited (429)" in str(content):
|
||||
max_attempts = 3
|
||||
for attempt in range(1, max_attempts + 1):
|
||||
wait = ExampleImagesProcessor._extract_retry_after(str(content)) * (2 ** (attempt - 1))
|
||||
logger.warning(
|
||||
"Rate limited (429) for %s, retry %d/%d after %ds",
|
||||
image_url, attempt, max_attempts, wait,
|
||||
)
|
||||
await asyncio.sleep(wait)
|
||||
|
||||
success, content, headers = await _attempt_download()
|
||||
if success:
|
||||
media_ext = ExampleImagesProcessor._get_file_extension_from_content_or_headers(
|
||||
content, headers, original_url, image.get("type")
|
||||
)
|
||||
is_image = media_ext in SUPPORTED_MEDIA_EXTENSIONS['images']
|
||||
is_video = media_ext in SUPPORTED_MEDIA_EXTENSIONS['videos']
|
||||
|
||||
if not (is_image or is_video):
|
||||
logger.debug("Skipping unsupported file type: %s", media_ext)
|
||||
break
|
||||
|
||||
save_filename = f"image_{i}{media_ext}"
|
||||
save_path = os.path.join(model_dir, save_filename)
|
||||
if os.path.exists(save_path):
|
||||
logger.debug("File already exists: %s", save_path)
|
||||
break
|
||||
|
||||
with open(save_path, 'wb') as f:
|
||||
f.write(content)
|
||||
any_successful_download = True
|
||||
break
|
||||
elif "Rate limited (429)" in str(content):
|
||||
continue
|
||||
elif ExampleImagesProcessor._is_not_found_error(content):
|
||||
logger.warning("Failed to download file: %s, status code: 404", image_url)
|
||||
model_success = False
|
||||
failed_images.append(image_url)
|
||||
break
|
||||
else:
|
||||
logger.warning("Failed to download file: %s, error: %s", image_url, content)
|
||||
model_success = False
|
||||
failed_images.append(image_url)
|
||||
break
|
||||
else:
|
||||
logger.warning(
|
||||
"Giving up on %s after %d retries due to rate limiting",
|
||||
image_url, max_attempts,
|
||||
)
|
||||
rate_limited_images.append(image_url)
|
||||
model_success = False
|
||||
else:
|
||||
error_msg = f"Failed to download file: {image_url}, error: {content}"
|
||||
logger.warning(error_msg)
|
||||
model_success = False # Mark the model as failed
|
||||
failed_images.append(image_url) # Track failed URL
|
||||
model_success = False
|
||||
failed_images.append(image_url)
|
||||
except Exception as e:
|
||||
error_msg = f"Error downloading file {image_url}: {str(e)}"
|
||||
logger.error(error_msg)
|
||||
model_success = False # Mark the model as failed
|
||||
failed_images.append(image_url) # Track failed URL
|
||||
|
||||
return model_success, False, failed_images # (success, is_metadata_stale, failed_images)
|
||||
model_success = False
|
||||
failed_images.append(image_url)
|
||||
|
||||
return any_successful_download or model_success, False, failed_images, rate_limited_images
|
||||
|
||||
@staticmethod
|
||||
async def process_local_examples(model_file_path, model_file_name, model_name, model_dir, optimize):
|
||||
@@ -591,7 +644,7 @@ class ExampleImagesProcessor:
|
||||
}, status=500)
|
||||
|
||||
# Update cache
|
||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
||||
await update_cache_from_metadata(scanner, file_path, model_data)
|
||||
|
||||
# Get regular images array (might be None)
|
||||
regular_images = civitai_data.get('images', [])
|
||||
@@ -706,7 +759,7 @@ class ExampleImagesProcessor:
|
||||
model_copy = model_data.copy()
|
||||
model_copy.pop('folder', None)
|
||||
await MetadataManager.save_metadata(file_path, model_copy)
|
||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
||||
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
|
||||
@@ -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
|
||||
|
||||
+54
-1
@@ -1,7 +1,7 @@
|
||||
from difflib import SequenceMatcher
|
||||
import os
|
||||
import re
|
||||
from typing import Dict
|
||||
from typing import Any, Dict, List, Optional
|
||||
from ..services.service_registry import ServiceRegistry
|
||||
from ..config import config
|
||||
from ..services.settings_manager import get_settings_manager
|
||||
@@ -294,6 +294,53 @@ def _format_model_name_for_comfyui(file_path: str, model_roots: list) -> str:
|
||||
return os.path.basename(file_path)
|
||||
|
||||
|
||||
def model_patcher_to_name(model_patcher: Any) -> Optional[str]:
|
||||
"""Extract a ComfyUI-style model name from a MODEL (ModelPatcher) object.
|
||||
|
||||
Core ComfyUI loaders record the absolute weight file path on the patcher's
|
||||
``cached_patcher_init`` attribute:
|
||||
- load_checkpoint_guess_config -> (fn, (ckpt_path, ...), index)
|
||||
- load_diffusion_model -> (fn, (unet_path, model_options))
|
||||
Patcher clones (LoRA loaders, model merges, ...) preserve the attribute,
|
||||
so the name is recoverable anywhere downstream of a core loader — including
|
||||
from LoRA Manager's own loaders (CheckpointLoaderLM / UNETLoaderLM), which
|
||||
call the same core load functions.
|
||||
|
||||
The absolute path is converted to the ComfyUI-style relative name used by
|
||||
the metadata pipeline (covering standard ComfyUI roots and LoRA Manager
|
||||
extra folder paths).
|
||||
|
||||
Returns None when the path cannot be recovered (e.g. third-party loaders
|
||||
that never set ``cached_patcher_init``).
|
||||
"""
|
||||
init = getattr(model_patcher, "cached_patcher_init", None)
|
||||
if not isinstance(init, (tuple, list)) or len(init) < 2:
|
||||
return None
|
||||
args = init[1]
|
||||
abs_path = args[0] if args else None
|
||||
if not isinstance(abs_path, str) or not abs_path:
|
||||
return None
|
||||
return _abs_model_path_to_name(abs_path)
|
||||
|
||||
|
||||
def _abs_model_path_to_name(abs_path: str) -> str:
|
||||
"""Convert an absolute model path to a ComfyUI-style relative name.
|
||||
|
||||
Tries standard ComfyUI model roots plus LoRA Manager extra folder paths;
|
||||
falls back to the bare filename.
|
||||
"""
|
||||
try:
|
||||
roots: List[str] = list(config.base_models_roots or [])
|
||||
roots.extend(config.extra_checkpoints_roots or [])
|
||||
roots.extend(config.extra_unet_roots or [])
|
||||
formatted = _format_model_name_for_comfyui(abs_path, roots)
|
||||
if formatted:
|
||||
return formatted
|
||||
except Exception:
|
||||
pass
|
||||
return os.path.basename(abs_path)
|
||||
|
||||
|
||||
def fuzzy_match(text: str, pattern: str, threshold: float = 0.85) -> bool:
|
||||
"""
|
||||
Check if text matches pattern using fuzzy matching.
|
||||
@@ -488,6 +535,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
@@ -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.2.0"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = [
|
||||
"aiohttp",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -151,6 +151,7 @@ body.modal-open {
|
||||
.support-section,
|
||||
.changelog-section,
|
||||
.update-info,
|
||||
.update-channels,
|
||||
.info-item,
|
||||
.path-preview {
|
||||
background: var(--surface-subtle);
|
||||
|
||||
@@ -577,13 +577,14 @@
|
||||
border: 1px solid var(--border-color);
|
||||
border-radius: var(--border-radius-sm);
|
||||
cursor: pointer;
|
||||
transition: var(--transition-base);
|
||||
transition: var(--transition-base), box-shadow var(--transition-fast), transform var(--transition-fast);
|
||||
background: var(--bg-color);
|
||||
}
|
||||
|
||||
.file-option:hover {
|
||||
border-color: var(--lora-accent);
|
||||
box-shadow: var(--shadow-sm);
|
||||
box-shadow: var(--shadow-md);
|
||||
transform: translateY(-1px);
|
||||
}
|
||||
|
||||
.file-option.selected {
|
||||
@@ -698,10 +699,25 @@
|
||||
color: var(--lora-accent);
|
||||
}
|
||||
|
||||
/* Batch Preview List */
|
||||
/* BUG 1 FIX: Single scrollbar — modal-content becomes a flex column so the
|
||||
batch preview step can flex; the list scrolls instead of the modal-content. */
|
||||
#downloadModal .modal-content {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
}
|
||||
|
||||
#batchPreviewStep {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
min-height: 0;
|
||||
flex: 1;
|
||||
}
|
||||
|
||||
/* Batch Preview List — no max-height; flexes inside #batchPreviewStep */
|
||||
.batch-preview-list {
|
||||
max-height: 400px;
|
||||
flex: 1;
|
||||
overflow-y: auto;
|
||||
min-height: 0;
|
||||
margin: var(--space-2) 0;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
@@ -859,6 +875,8 @@
|
||||
position: sticky;
|
||||
top: 0;
|
||||
z-index: 1;
|
||||
backdrop-filter: blur(8px);
|
||||
-webkit-backdrop-filter: blur(8px);
|
||||
}
|
||||
|
||||
.batch-preview-select-all input[type="checkbox"] {
|
||||
@@ -884,3 +902,100 @@
|
||||
[data-theme="dark"] .batch-preview-select-all {
|
||||
background: var(--lora-surface);
|
||||
}
|
||||
|
||||
/* FEATURE 2: HF repo grouping — collapsible groups by repo */
|
||||
.batch-preview-group {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
background: var(--surface-base);
|
||||
}
|
||||
|
||||
.batch-preview-group-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
padding: 10px 12px;
|
||||
background: var(--color-accent-subtle);
|
||||
border-bottom: 1px solid var(--color-accent-border);
|
||||
cursor: pointer;
|
||||
user-select: none;
|
||||
transition: background var(--transition-fast);
|
||||
}
|
||||
|
||||
.batch-preview-group-header:hover {
|
||||
background: oklch(var(--color-accent-l) var(--color-accent-c) var(--color-accent-h) / 0.18);
|
||||
}
|
||||
|
||||
.batch-preview-group-toggle {
|
||||
width: 14px;
|
||||
font-size: 0.75em;
|
||||
color: var(--text-color);
|
||||
opacity: 0.7;
|
||||
transition: transform var(--transition-fast);
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
.batch-preview-group-toggle.expanded {
|
||||
transform: rotate(90deg);
|
||||
}
|
||||
|
||||
.batch-preview-group-name {
|
||||
flex: 1;
|
||||
min-width: 0;
|
||||
font-weight: 600;
|
||||
color: var(--text-color);
|
||||
white-space: nowrap;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
font-size: 0.95em;
|
||||
}
|
||||
|
||||
.batch-preview-group-count {
|
||||
font-size: 0.8em;
|
||||
color: var(--text-color);
|
||||
opacity: 0.7;
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
.batch-preview-group-select-all {
|
||||
width: 18px;
|
||||
height: 18px;
|
||||
cursor: pointer;
|
||||
accent-color: var(--lora-accent);
|
||||
flex-shrink: 0;
|
||||
padding: 0;
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
.batch-preview-group-body {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 1px;
|
||||
background: var(--border-color);
|
||||
overflow: hidden;
|
||||
max-height: 0;
|
||||
opacity: 0;
|
||||
transition: max-height 0.35s ease, opacity 0.2s ease;
|
||||
}
|
||||
|
||||
.batch-preview-group-body.expanded {
|
||||
opacity: 1;
|
||||
max-height: 9999px; /* rest state: content visible; JS inline style overrides during transitions */
|
||||
}
|
||||
|
||||
/* Dark theme overrides for group styles */
|
||||
[data-theme="dark"] .batch-preview-group {
|
||||
background: var(--surface-base);
|
||||
}
|
||||
|
||||
[data-theme="dark"] .batch-preview-group-header {
|
||||
background: var(--color-accent-subtle);
|
||||
}
|
||||
|
||||
[data-theme="dark"] .batch-preview-group-header:hover {
|
||||
background: oklch(var(--color-accent-l) var(--color-accent-c) var(--color-accent-h) / 0.22);
|
||||
}
|
||||
|
||||
[data-theme="dark"] .batch-preview-group-body {
|
||||
background: var(--border-color);
|
||||
}
|
||||
|
||||
@@ -21,18 +21,22 @@
|
||||
margin-bottom: 4px;
|
||||
}
|
||||
|
||||
.input-group {
|
||||
#relinkCivitaiModal .input-group,
|
||||
#linkHfModal .input-group {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
margin-bottom: var(--space-2);
|
||||
}
|
||||
|
||||
.input-group label {
|
||||
#relinkCivitaiModal .input-group label,
|
||||
#linkHfModal .input-group label {
|
||||
margin-bottom: var(--space-1);
|
||||
font-weight: 500;
|
||||
}
|
||||
|
||||
.input-group input {
|
||||
#relinkCivitaiModal .input-group input,
|
||||
#linkHfModal .input-group input {
|
||||
width: auto;
|
||||
padding: 8px 12px;
|
||||
border-radius: var(--border-radius-xs);
|
||||
border: 1px solid var(--border-color);
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -93,15 +93,13 @@
|
||||
.update-content {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: var(--space-3);
|
||||
gap: var(--space-2);
|
||||
}
|
||||
|
||||
.update-info {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
border-radius: var(--border-radius-sm);
|
||||
padding: var(--space-3);
|
||||
}
|
||||
|
||||
.update-info .version-info {
|
||||
@@ -175,7 +173,6 @@
|
||||
border: 1px solid var(--lora-border);
|
||||
border-radius: var(--border-radius-sm);
|
||||
padding: var(--space-2);
|
||||
margin: var(--space-2) 0;
|
||||
}
|
||||
|
||||
[data-theme="dark"] .update-progress {
|
||||
@@ -233,11 +230,6 @@
|
||||
}
|
||||
|
||||
/* Changelog section */
|
||||
.changelog-section {
|
||||
border-radius: var(--border-radius-sm);
|
||||
padding: var(--space-3);
|
||||
}
|
||||
|
||||
.changelog-section h3 {
|
||||
margin-top: 0;
|
||||
margin-bottom: var(--space-2);
|
||||
@@ -349,6 +341,131 @@
|
||||
text-decoration: underline;
|
||||
}
|
||||
|
||||
/* Channel Toggle */
|
||||
.update-channels {
|
||||
}
|
||||
|
||||
.channels-label {
|
||||
font-size: 0.9em;
|
||||
color: var(--text-color);
|
||||
opacity: 0.8;
|
||||
margin-bottom: 8px;
|
||||
}
|
||||
|
||||
.channel-toggle {
|
||||
display: flex;
|
||||
gap: 0;
|
||||
background: var(--lora-surface);
|
||||
border-radius: 8px;
|
||||
padding: 3px;
|
||||
width: fit-content;
|
||||
}
|
||||
|
||||
.channel-btn {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
padding: 8px 20px;
|
||||
border: none;
|
||||
border-radius: 6px;
|
||||
background: transparent;
|
||||
color: var(--text-secondary, #999);
|
||||
cursor: pointer;
|
||||
font-size: 0.9em;
|
||||
font-weight: 500;
|
||||
transition: all 0.2s ease;
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
.channel-btn:hover {
|
||||
color: var(--text-primary, #ddd);
|
||||
background: rgba(255, 255, 255, 0.04);
|
||||
}
|
||||
|
||||
.channel-btn.active {
|
||||
background: var(--lora-accent, #4285F4);
|
||||
color: #fff;
|
||||
box-shadow: 0 1px 3px rgba(0, 0, 0, 0.2);
|
||||
}
|
||||
|
||||
.channel-btn.active i {
|
||||
color: #fff;
|
||||
}
|
||||
|
||||
.channel-btn i {
|
||||
font-size: 0.85em;
|
||||
}
|
||||
|
||||
/* Channel Switch Confirmation Overlay */
|
||||
.channel-switch-overlay {
|
||||
position: fixed;
|
||||
inset: 0;
|
||||
background: rgba(0, 0, 0, 0.6);
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
z-index: 10000;
|
||||
backdrop-filter: blur(2px);
|
||||
}
|
||||
|
||||
.channel-switch-dialog {
|
||||
background: var(--lora-surface);
|
||||
border: 1px solid var(--border-color, rgba(255, 255, 255, 0.1));
|
||||
border-radius: 12px;
|
||||
padding: 28px 32px;
|
||||
max-width: 420px;
|
||||
width: 90%;
|
||||
box-shadow: 0 8px 32px rgba(0, 0, 0, 0.4);
|
||||
}
|
||||
|
||||
.channel-switch-dialog h3 {
|
||||
margin: 0 0 12px;
|
||||
font-size: 1.1em;
|
||||
color: var(--text-primary, #eee);
|
||||
}
|
||||
|
||||
.channel-switch-dialog p {
|
||||
margin: 0 0 24px;
|
||||
font-size: 0.9em;
|
||||
color: var(--text-secondary, #aaa);
|
||||
line-height: 1.6;
|
||||
}
|
||||
|
||||
.channel-switch-actions {
|
||||
display: flex;
|
||||
justify-content: flex-end;
|
||||
gap: 10px;
|
||||
}
|
||||
|
||||
.channel-switch-cancel {
|
||||
padding: 8px 18px;
|
||||
border: 1px solid var(--border-color, rgba(255, 255, 255, 0.1));
|
||||
border-radius: 6px;
|
||||
background: transparent;
|
||||
color: var(--text-secondary, #aaa);
|
||||
cursor: pointer;
|
||||
font-size: 0.9em;
|
||||
}
|
||||
|
||||
.channel-switch-cancel:hover {
|
||||
background: rgba(255, 255, 255, 0.04);
|
||||
}
|
||||
|
||||
.channel-switch-confirm {
|
||||
padding: 8px 18px;
|
||||
border: none;
|
||||
border-radius: 6px;
|
||||
background: var(--lora-accent, #4285F4);
|
||||
color: #fff;
|
||||
cursor: pointer;
|
||||
font-size: 0.9em;
|
||||
font-weight: 500;
|
||||
}
|
||||
|
||||
.channel-switch-confirm:hover {
|
||||
opacity: 0.9;
|
||||
}
|
||||
|
||||
/* Update preferences section */
|
||||
.update-preferences {
|
||||
border-top: 1px solid var(--lora-border);
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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`,
|
||||
|
||||
@@ -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();
|
||||
|
||||
|
||||
@@ -391,6 +391,15 @@ export class BulkContextMenu extends BaseContextMenu {
|
||||
`Enriching metadata for ${modelPaths.length} models...`
|
||||
);
|
||||
|
||||
function cleanupCallbacks() {
|
||||
const pIdx = agentManager.progressCallbacks.indexOf(onProgress);
|
||||
if (pIdx >= 0) agentManager.progressCallbacks.splice(pIdx, 1);
|
||||
const cIdx = agentManager.completeCallbacks.indexOf(onComplete);
|
||||
if (cIdx >= 0) agentManager.completeCallbacks.splice(cIdx, 1);
|
||||
const eIdx = agentManager.errorCallbacks.indexOf(onError);
|
||||
if (eIdx >= 0) agentManager.errorCallbacks.splice(eIdx, 1);
|
||||
}
|
||||
|
||||
const onProgress = (data) => {
|
||||
if (data.status === 'processing' && data.current_path && data.updated_data && Object.keys(data.updated_data).length > 0) {
|
||||
if (state.virtualScroller?.updateSingleItem) {
|
||||
@@ -404,36 +413,37 @@ export class BulkContextMenu extends BaseContextMenu {
|
||||
agentManager.onProgress(onProgress);
|
||||
|
||||
const onComplete = (data) => {
|
||||
const pIdx = agentManager.progressCallbacks.indexOf(onProgress);
|
||||
if (pIdx >= 0) agentManager.progressCallbacks.splice(pIdx, 1);
|
||||
const cIdx = agentManager.completeCallbacks.indexOf(onComplete);
|
||||
if (cIdx >= 0) agentManager.completeCallbacks.splice(cIdx, 1);
|
||||
cleanupCallbacks();
|
||||
|
||||
if (data.status === 'completed') {
|
||||
if (state.bulkMode) bulkManager.toggleBulkMode();
|
||||
progressUI.complete(data.summary || 'Enrich complete');
|
||||
showToast(
|
||||
'toast.agent.enrichComplete',
|
||||
{ summary: data.summary || 'Done' },
|
||||
'success'
|
||||
);
|
||||
} else if (data.status === 'error') {
|
||||
state.loadingManager.hide();
|
||||
showToast(
|
||||
'toast.agent.enrichFailed',
|
||||
{ error: data.error || 'Unknown error' },
|
||||
'error'
|
||||
);
|
||||
}
|
||||
};
|
||||
agentManager.onComplete(onComplete);
|
||||
|
||||
const onError = (data) => {
|
||||
cleanupCallbacks();
|
||||
if (state.bulkMode) bulkManager.toggleBulkMode();
|
||||
state.loadingManager.hide();
|
||||
showToast(
|
||||
'toast.agent.enrichFailed',
|
||||
{ error: data.error || 'Unknown error' },
|
||||
'error'
|
||||
);
|
||||
};
|
||||
agentManager.onError(onError);
|
||||
|
||||
try {
|
||||
await agentManager.executeSkill('enrich_hf_metadata', modelPaths);
|
||||
} catch (error) {
|
||||
const pIdx = agentManager.progressCallbacks.indexOf(onProgress);
|
||||
if (pIdx >= 0) agentManager.progressCallbacks.splice(pIdx, 1);
|
||||
const cIdx = agentManager.completeCallbacks.indexOf(onComplete);
|
||||
if (cIdx >= 0) agentManager.completeCallbacks.splice(cIdx, 1);
|
||||
cleanupCallbacks();
|
||||
if (state.bulkMode) bulkManager.toggleBulkMode();
|
||||
state.loadingManager.hide();
|
||||
showToast(
|
||||
'toast.agent.enrichFailed',
|
||||
|
||||
@@ -32,6 +32,9 @@ export class LoraContextMenu extends BaseContextMenu {
|
||||
if (!enrichItem) return;
|
||||
const hasHfUrl = !!card.dataset.hf_url;
|
||||
enrichItem.classList.toggle('disabled', !hasHfUrl);
|
||||
enrichItem.title = hasHfUrl
|
||||
? ''
|
||||
: 'Link this model to a HuggingFace repo first (Link Model \u2192 Link to HuggingFace)';
|
||||
}
|
||||
|
||||
handleMenuAction(action, menuItem) {
|
||||
@@ -99,6 +102,15 @@ export class LoraContextMenu extends BaseContextMenu {
|
||||
'Enriching metadata with AI...'
|
||||
);
|
||||
|
||||
function cleanupCallbacks() {
|
||||
const pIdx = agentManager.progressCallbacks.indexOf(onProgress);
|
||||
if (pIdx >= 0) agentManager.progressCallbacks.splice(pIdx, 1);
|
||||
const cIdx = agentManager.completeCallbacks.indexOf(onComplete);
|
||||
if (cIdx >= 0) agentManager.completeCallbacks.splice(cIdx, 1);
|
||||
const eIdx = agentManager.errorCallbacks.indexOf(onError);
|
||||
if (eIdx >= 0) agentManager.errorCallbacks.splice(eIdx, 1);
|
||||
}
|
||||
|
||||
const onProgress = (data) => {
|
||||
if (data.status === 'processing' && data.current_path && data.updated_data && Object.keys(data.updated_data).length > 0) {
|
||||
if (state.virtualScroller?.updateSingleItem) {
|
||||
@@ -112,28 +124,26 @@ export class LoraContextMenu extends BaseContextMenu {
|
||||
agentManager.onProgress(onProgress);
|
||||
|
||||
const onComplete = (data) => {
|
||||
const pIdx = agentManager.progressCallbacks.indexOf(onProgress);
|
||||
if (pIdx >= 0) agentManager.progressCallbacks.splice(pIdx, 1);
|
||||
const cIdx = agentManager.completeCallbacks.indexOf(onComplete);
|
||||
if (cIdx >= 0) agentManager.completeCallbacks.splice(cIdx, 1);
|
||||
cleanupCallbacks();
|
||||
|
||||
if (data.status === 'completed') {
|
||||
progressUI.complete(data.summary || 'Enrich complete');
|
||||
showToast('toast.agent.enrichComplete', { summary: data.summary || 'Done' }, 'success');
|
||||
} else if (data.status === 'error') {
|
||||
state.loadingManager.hide();
|
||||
showToast('toast.agent.enrichFailed', { error: data.error || 'Unknown error' }, 'error');
|
||||
}
|
||||
};
|
||||
agentManager.onComplete(onComplete);
|
||||
|
||||
const onError = (data) => {
|
||||
cleanupCallbacks();
|
||||
state.loadingManager.hide();
|
||||
showToast('toast.agent.enrichFailed', { error: data.error || 'Unknown error' }, 'error');
|
||||
};
|
||||
agentManager.onError(onError);
|
||||
|
||||
try {
|
||||
await agentManager.executeSkill('enrich_hf_metadata', [filePath]);
|
||||
} catch (error) {
|
||||
const pIdx = agentManager.progressCallbacks.indexOf(onProgress);
|
||||
if (pIdx >= 0) agentManager.progressCallbacks.splice(pIdx, 1);
|
||||
const cIdx = agentManager.completeCallbacks.indexOf(onComplete);
|
||||
if (cIdx >= 0) agentManager.completeCallbacks.splice(cIdx, 1);
|
||||
cleanupCallbacks();
|
||||
state.loadingManager.hide();
|
||||
showToast('toast.agent.enrichFailed', { error: error.message }, 'error');
|
||||
}
|
||||
@@ -142,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');
|
||||
}
|
||||
|
||||
@@ -187,6 +187,74 @@ export const ModelContextMenuMixin = {
|
||||
setTimeout(() => urlInput.focus(), 50);
|
||||
},
|
||||
|
||||
// HuggingFace linking methods
|
||||
showLinkHfModal() {
|
||||
const filePath = this.currentCard.dataset.filepath;
|
||||
if (!filePath) return;
|
||||
|
||||
const confirmBtn = document.getElementById('confirmLinkHfBtn');
|
||||
const urlInput = document.getElementById('hfModelUrl');
|
||||
const errorDiv = document.getElementById('hfModelUrlError');
|
||||
|
||||
if (this._boundLinkHfHandler) {
|
||||
confirmBtn.removeEventListener('click', this._boundLinkHfHandler);
|
||||
}
|
||||
|
||||
this._boundLinkHfHandler = async () => {
|
||||
const hfUrl = urlInput.value.trim();
|
||||
if (!hfUrl) {
|
||||
errorDiv.textContent = 'Please enter a HuggingFace repository URL.';
|
||||
return;
|
||||
}
|
||||
|
||||
const hfPattern = /^https?:\/\/huggingface\.co\/([^/]+\/[^/]+)\/?$/;
|
||||
if (!hfPattern.test(hfUrl)) {
|
||||
errorDiv.textContent = 'Invalid URL format. Expected: https://huggingface.co/user/repo';
|
||||
return;
|
||||
}
|
||||
|
||||
errorDiv.textContent = '';
|
||||
modalManager.closeModal('linkHfModal');
|
||||
|
||||
try {
|
||||
state.loadingManager.showSimpleLoading('Linking to HuggingFace...');
|
||||
|
||||
const response = await fetch('/api/lm/set-hf-url', {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ file_path: filePath, hf_url: hfUrl }),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errData = await response.json().catch(() => ({}));
|
||||
throw new Error(errData.error || `Request failed: ${response.statusText}`);
|
||||
}
|
||||
|
||||
const data = await response.json();
|
||||
if (data.success) {
|
||||
showToast('toast.contextMenu.linkHfSuccess', {}, 'success');
|
||||
await this.resetAndReload();
|
||||
} else {
|
||||
throw new Error(data.error || 'Failed to link model');
|
||||
}
|
||||
} catch (error) {
|
||||
console.error('Error linking model to HuggingFace:', error);
|
||||
showToast('toast.contextMenu.linkHfFailed', { message: error.message }, 'error');
|
||||
} finally {
|
||||
state.loadingManager.hide();
|
||||
}
|
||||
};
|
||||
|
||||
confirmBtn.addEventListener('click', this._boundLinkHfHandler);
|
||||
|
||||
urlInput.value = '';
|
||||
errorDiv.textContent = '';
|
||||
|
||||
modalManager.showModal('linkHfModal');
|
||||
|
||||
setTimeout(() => urlInput.focus(), 50);
|
||||
},
|
||||
|
||||
extractModelVersionId(url) {
|
||||
return extractCivitaiModelUrlParts(url);
|
||||
},
|
||||
@@ -295,6 +363,9 @@ export const ModelContextMenuMixin = {
|
||||
case 'relink-civitai':
|
||||
this.showRelinkCivitaiModal();
|
||||
return true;
|
||||
case 'link-hf':
|
||||
this.showLinkHfModal();
|
||||
return true;
|
||||
case 'set-nsfw':
|
||||
this.showNSFWLevelSelector(null, null, this.currentCard);
|
||||
return true;
|
||||
|
||||
@@ -260,8 +260,9 @@ export class RecipeContextMenu extends BaseContextMenu {
|
||||
strength: lora.strength || 1.0,
|
||||
|
||||
// Model identifiers
|
||||
modelId: lora.modelId || lora.model_id || civitaiInfo.modelId,
|
||||
hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash,
|
||||
modelVersionId: civitaiInfo.id || lora.modelVersionId,
|
||||
id: civitaiInfo.id || lora.modelVersionId,
|
||||
|
||||
// Metadata
|
||||
thumbnailUrl: civitaiInfo.images?.[0]?.url || '',
|
||||
|
||||
@@ -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">
|
||||
|
||||
@@ -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) {
|
||||
@@ -1421,6 +1421,7 @@ class RecipeModal {
|
||||
strength: lora.strength || 1.0,
|
||||
|
||||
// Model identifiers
|
||||
modelId: lora.modelId || lora.model_id || civitaiInfo.modelId,
|
||||
hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash,
|
||||
id: civitaiInfo.id || lora.modelVersionId,
|
||||
|
||||
@@ -1606,7 +1607,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">
|
||||
|
||||
@@ -108,10 +108,20 @@ export class PageControls {
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
if (sortSelect) {
|
||||
initSortDropdown(sortSelect);
|
||||
sortSelect.value = this.pageState.sortBy;
|
||||
this.applySortToSelect(this.pageState.sortBy);
|
||||
sortSelect.addEventListener('change', async (e) => {
|
||||
this.pageState.sortBy = e.target.value;
|
||||
this.saveSortPreference(e.target.value);
|
||||
let value = e.target.value;
|
||||
if (value.startsWith('random')) {
|
||||
// Every pick of Random reshuffles the list: generate a
|
||||
// fresh seed so the backend keeps a stable order across
|
||||
// paginated requests.
|
||||
value = this._randomizeSortValue();
|
||||
}
|
||||
this.pageState.sortBy = value;
|
||||
this.saveSortPreference(value);
|
||||
// Reset the seeded Random option when switching away from
|
||||
// Random, or re-apply the fresh seed when picking it again.
|
||||
this.applySortToSelect(value);
|
||||
await this.resetAndReload();
|
||||
});
|
||||
}
|
||||
@@ -312,6 +322,44 @@ export class PageControls {
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Apply a sort value to the native sort <select>, keeping the Random
|
||||
* option's value in sync when the persisted value carries a seed
|
||||
* (e.g. "random:abc123"). Must be used instead of assigning
|
||||
* sortSelect.value directly whenever the value may be a seeded random
|
||||
* sort, otherwise the native select has no matching option.
|
||||
* @param {string} sortValue - Sort value like "name:asc" or "random:<seed>"
|
||||
*/
|
||||
applySortToSelect(sortValue) {
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
if (!sortSelect) return;
|
||||
const randomOpt = sortSelect.querySelector('option[value="random"], option[value^="random:"]');
|
||||
if (randomOpt) {
|
||||
randomOpt.value = String(sortValue).startsWith('random') ? sortValue : 'random';
|
||||
}
|
||||
sortSelect.value = sortValue;
|
||||
}
|
||||
|
||||
/**
|
||||
* Generate a fresh seeded random sort value ("random:<seed>") and keep
|
||||
* the native <select> in sync so its value matches the persisted sort
|
||||
* string and the dropdown shows the selected label.
|
||||
* @returns {string} The new sort value, e.g. "random:abc123xyz"
|
||||
*/
|
||||
_randomizeSortValue() {
|
||||
const seed = Math.random().toString(36).slice(2, 12);
|
||||
const value = `random:${seed}`;
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
if (sortSelect) {
|
||||
const randomOpt = sortSelect.querySelector('option[value="random"], option[value^="random:"]');
|
||||
if (randomOpt) {
|
||||
randomOpt.value = value;
|
||||
}
|
||||
sortSelect.value = value;
|
||||
}
|
||||
return value;
|
||||
}
|
||||
|
||||
/**
|
||||
* Load sort preference from storage
|
||||
*/
|
||||
@@ -326,10 +374,7 @@ export class PageControls {
|
||||
// Handle legacy format conversion
|
||||
const convertedSort = this.convertLegacySortFormat(savedSort);
|
||||
this.pageState.sortBy = convertedSort;
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
if (sortSelect) {
|
||||
sortSelect.value = convertedSort;
|
||||
}
|
||||
this.applySortToSelect(convertedSort);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -523,9 +568,9 @@ export class PageControls {
|
||||
this.pageState.sortBy = restoredSort;
|
||||
this.saveSortPreference(restoredSort);
|
||||
this._removeVlmSortOption();
|
||||
this.applySortToSelect(restoredSort);
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
if (sortSelect) {
|
||||
sortSelect.value = restoredSort;
|
||||
sortSelect.disabled = false;
|
||||
}
|
||||
}
|
||||
@@ -575,10 +620,7 @@ export class PageControls {
|
||||
const savedGroupedSort = getStorageItem(groupedKey);
|
||||
if (savedGroupedSort) {
|
||||
this.pageState.sortBy = savedGroupedSort;
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
if (sortSelect) {
|
||||
sortSelect.value = savedGroupedSort;
|
||||
}
|
||||
this.applySortToSelect(savedGroupedSort);
|
||||
}
|
||||
} else {
|
||||
// Leaving group mode: persist current sort for next time, restore non-group sort
|
||||
@@ -586,10 +628,7 @@ export class PageControls {
|
||||
const savedNormalSort = getStorageItem(`${this.pageType}_sort`);
|
||||
if (savedNormalSort) {
|
||||
this.pageState.sortBy = savedNormalSort;
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
if (sortSelect) {
|
||||
sortSelect.value = savedNormalSort;
|
||||
}
|
||||
this.applySortToSelect(savedNormalSort);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -874,7 +913,7 @@ export class PageControls {
|
||||
}
|
||||
|
||||
if (sortSelect) {
|
||||
sortSelect.value = this.pageState.sortBy;
|
||||
this.applySortToSelect(this.pageState.sortBy);
|
||||
}
|
||||
if (searchInput) {
|
||||
searchInput.value = this.pageState.filters?.search || '';
|
||||
|
||||
@@ -96,7 +96,16 @@ export function initSortDropdown(select) {
|
||||
};
|
||||
|
||||
const choose = (value) => {
|
||||
if (select.value === value) return;
|
||||
if (select.value === value) {
|
||||
// Re-picking the already-selected option is normally a no-op,
|
||||
// matching native <select> behavior. The seeded Random sort is
|
||||
// the exception: clicking it again should reshuffle, so let the
|
||||
// change handler (PageControls) generate a fresh seed.
|
||||
if (String(value).startsWith('random')) {
|
||||
select.dispatchEvent(new Event('change', { bubbles: true }));
|
||||
}
|
||||
return;
|
||||
}
|
||||
select.value = value;
|
||||
select.dispatchEvent(new Event('change', { bubbles: true }));
|
||||
};
|
||||
@@ -277,9 +286,10 @@ export function initSortDropdown(select) {
|
||||
}
|
||||
|
||||
// Rebuild the menu when <option>s change (VLM adds/removes a temporary
|
||||
// option at runtime).
|
||||
// option at runtime, and the seeded Random sort option gets a new value
|
||||
// attribute each time it is picked).
|
||||
const observer = new MutationObserver(() => buildMenu());
|
||||
observer.observe(select, { childList: true });
|
||||
observer.observe(select, { childList: true, subtree: true, attributes: true, attributeFilter: ['value'] });
|
||||
|
||||
buildMenu();
|
||||
group.dataset.sortReady = '1';
|
||||
|
||||
@@ -489,6 +489,12 @@ export function createModelCard(model, modelType) {
|
||||
const modelId = civitaiData?.modelId ?? civitaiData?.model_id;
|
||||
if (modelId !== undefined && modelId !== null && modelId !== '') {
|
||||
card.dataset.modelId = modelId;
|
||||
} else if (model.hf_url) {
|
||||
// For HF-only models, derive a group key from hf_url for version grouping
|
||||
const match = model.hf_url.match(/https?:\/\/huggingface\.co\/([^/]+\/[^/]+)/);
|
||||
if (match) {
|
||||
card.dataset.modelId = 'hf:' + match[1];
|
||||
}
|
||||
}
|
||||
|
||||
// LoRA specific data
|
||||
@@ -643,7 +649,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 ?
|
||||
|
||||
@@ -473,7 +473,14 @@ export async function showModelModal(model, modelType) {
|
||||
const loadingExamplesText = translate('modals.model.loading.examples', {}, 'Loading examples...');
|
||||
|
||||
const loadingVersionsText = translate('modals.model.loading.versions', {}, 'Loading versions...');
|
||||
const civitaiModelId = modelWithFullData.civitai?.modelId || '';
|
||||
// Use CivitAI modelId, or derive HF group key for HF-only models
|
||||
let civitaiModelId = modelWithFullData.civitai?.modelId || '';
|
||||
if (!civitaiModelId && modelWithFullData.hf_url) {
|
||||
const match = modelWithFullData.hf_url.match(/https?:\/\/huggingface\.co\/([^/]+\/[^/]+)/);
|
||||
if (match) {
|
||||
civitaiModelId = 'hf:' + match[1];
|
||||
}
|
||||
}
|
||||
const civitaiVersionId = modelWithFullData.civitai?.id || '';
|
||||
const navAriaLabel = translate('modals.model.navigation.label', {}, 'Model navigation');
|
||||
const previousTitle = translate('modals.model.navigation.previousWithShortcut', {}, 'Previous model (←)');
|
||||
@@ -885,7 +892,8 @@ function setupEventHandlers(filePath, modelType) {
|
||||
case 'view-creator':
|
||||
const username = target.dataset.username;
|
||||
if (username) {
|
||||
window.open(`https://civitai.com/user/${username}`, '_blank');
|
||||
const host = state.global.settings.civitai_host || 'civitai.com';
|
||||
window.open(`https://${host}/user/${username}`, '_blank');
|
||||
}
|
||||
break;
|
||||
case 'open-file-location':
|
||||
|
||||
@@ -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>
|
||||
`;
|
||||
}
|
||||
@@ -950,6 +950,26 @@ export function initVersionsTab({
|
||||
renderErrorState(container, translate('modals.model.versions.missingModelId', {}, 'This model is missing a Civitai model id.'));
|
||||
return;
|
||||
}
|
||||
// HF group keys (e.g. "hf:user/repo") are not real CivitAI model IDs —
|
||||
// skip the remote API call and show a helpful message instead.
|
||||
const isHfGroupKey = typeof modelId === 'string' && modelId.startsWith('hf:');
|
||||
if (isHfGroupKey) {
|
||||
controller.isLoading = false;
|
||||
controller.hasLoaded = true;
|
||||
controller.record = null;
|
||||
const hfMsg = translate(
|
||||
'modals.model.versions.hfGroupInfo',
|
||||
{},
|
||||
'This is a HuggingFace model group. Open the library to see all versions in the grid.'
|
||||
);
|
||||
container.innerHTML = `
|
||||
<div class="versions-empty-state">
|
||||
<i class="fas fa-info-circle"></i>
|
||||
<p>${escapeHtml(hfMsg)}</p>
|
||||
</div>
|
||||
`;
|
||||
return;
|
||||
}
|
||||
if (controller.hasLoaded && !forceRefresh) {
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -586,6 +586,7 @@ export function initMediaControlHandlers(container) {
|
||||
const imageMetaRaw = this.dataset.imageMeta;
|
||||
const imageUrl = this.dataset.imageUrl;
|
||||
const imageNsfw = this.dataset.imageNsfw;
|
||||
const imgId = this.dataset.imgId || '';
|
||||
const localPath = this.dataset.localPath || '';
|
||||
const showcaseSection = this.closest('.showcase-section');
|
||||
const modelHash = showcaseSection ? showcaseSection.dataset.modelHash : '';
|
||||
@@ -613,6 +614,7 @@ export function initMediaControlHandlers(container) {
|
||||
meta: imageMeta,
|
||||
url: imageUrl,
|
||||
nsfwLevel: imageNsfw ? parseInt(imageNsfw, 10) : undefined,
|
||||
id: imgId || undefined,
|
||||
},
|
||||
model_hash: modelHash,
|
||||
model_name: modelName || modelHash,
|
||||
|
||||
@@ -213,8 +213,8 @@ function renderMediaItem(img, index, exampleFiles) {
|
||||
const model = meta.Model || '';
|
||||
const steps = meta.steps || '';
|
||||
const sampler = meta.sampler || '';
|
||||
const cfgScale = meta.cfgScale || '';
|
||||
const clipSkip = meta.clipSkip || '';
|
||||
const cfgScale = meta.cfg_scale || meta.cfgScale || '';
|
||||
const clipSkip = meta.clip_skip || meta.clipSkip || '';
|
||||
|
||||
// Check if we have any meaningful generation parameters
|
||||
const hasParams = seed || model || steps || sampler || cfgScale || clipSkip;
|
||||
@@ -245,6 +245,7 @@ function renderMediaItem(img, index, exampleFiles) {
|
||||
data-image-url="${img.url || ''}"
|
||||
data-image-nsfw="${img.nsfwLevel ?? ''}"
|
||||
data-image-id="${cdnImageId}"
|
||||
data-img-id="${img.id || ''}"
|
||||
data-local-path="${localFile ? localFile.path : ''}">
|
||||
<i class="fas fa-book-open"></i>
|
||||
</button>
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import { modalManager } from './ModalManager.js';
|
||||
import { showToast } from '../utils/uiHelpers.js';
|
||||
import { showToast, setupAutoNewlineOnPaste } from '../utils/uiHelpers.js';
|
||||
import { translate } from '../utils/i18nHelpers.js';
|
||||
import { WS_ENDPOINTS } from '../api/apiConfig.js';
|
||||
import { getStorageItem, setStorageItem } from '../utils/storageHelpers.js';
|
||||
@@ -43,6 +43,9 @@ export class BatchImportManager {
|
||||
setStorageItem('batch_import_skip_no_metadata', e.target.checked);
|
||||
});
|
||||
}
|
||||
|
||||
// Auto-append newline after pasting a URL in the batch URL input
|
||||
setupAutoNewlineOnPaste('batchUrlInput');
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user