mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-16 02:33:21 -03:00
Compare commits
61 Commits
v1.1.8
..
186ef4da78
| Author | SHA1 | Date | |
|---|---|---|---|
| 186ef4da78 | |||
| dc674098e7 | |||
| 9087b4b07c | |||
| 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 |
@@ -137,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
|
||||
|
||||
|
||||
+10
@@ -17,6 +17,8 @@ try: # pragma: no cover - import fallback for pytest collection
|
||||
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
|
||||
@@ -62,6 +64,12 @@ except (
|
||||
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 = {
|
||||
@@ -83,6 +91,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
LoraCyclerLM.NAME: LoraCyclerLM,
|
||||
LoraInfoLM.NAME: LoraInfoLM,
|
||||
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
|
||||
CreateHookLoraLM.NAME: CreateHookLoraLM,
|
||||
MetadataOverwriteLM.NAME: MetadataOverwriteLM,
|
||||
}
|
||||
|
||||
WEB_DIRECTORY = "./web/comfyui"
|
||||
|
||||
+313
-291
File diff suppressed because it is too large
Load Diff
+2227
-2205
File diff suppressed because it is too large
Load Diff
+23
-1
@@ -714,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",
|
||||
@@ -771,6 +773,8 @@
|
||||
"deleteAll": "Delete Selected",
|
||||
"downloadMissingLoras": "Download Missing LoRAs",
|
||||
"downloadExamples": "Download Example Images",
|
||||
"downloadMissingExamples": "Download Missing",
|
||||
"reprocessExamples": "Re-process All",
|
||||
"clear": "Clear Selection",
|
||||
"skipMetadataRefreshCount": "Skip ({count} models)",
|
||||
"resumeMetadataRefreshCount": "Resume ({count} models)",
|
||||
@@ -806,6 +810,8 @@
|
||||
"sendToWorkflowReplace": "Send to Workflow (Replace)",
|
||||
"openExamples": "Open Examples Folder",
|
||||
"downloadExamples": "Download Example Images",
|
||||
"downloadMissingExamples": "Download Missing",
|
||||
"reprocessExamples": "Re-process All",
|
||||
"replacePreview": "Replace Preview",
|
||||
"setContentRating": "Set Content Rating",
|
||||
"moveToFolder": "Move to Folder",
|
||||
@@ -1548,6 +1554,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?"
|
||||
},
|
||||
@@ -1751,6 +1758,12 @@
|
||||
"checkingMessage": "Please wait while we check for the latest version.",
|
||||
"showNotifications": "Show update notifications",
|
||||
"latestBadge": "Latest",
|
||||
"latestMain": "Latest main",
|
||||
"channel": "Update Channel",
|
||||
"channels": {
|
||||
"release": "Release",
|
||||
"nightly": "Nightly"
|
||||
},
|
||||
"updateProgress": {
|
||||
"preparing": "Preparing update...",
|
||||
"installing": "Installing update...",
|
||||
@@ -1771,6 +1784,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.",
|
||||
|
||||
+2227
-2205
File diff suppressed because it is too large
Load Diff
+2227
-2205
File diff suppressed because it is too large
Load Diff
+2227
-2205
File diff suppressed because it is too large
Load Diff
+2227
-2205
File diff suppressed because it is too large
Load Diff
+2227
-2205
File diff suppressed because it is too large
Load Diff
+2227
-2205
File diff suppressed because it is too large
Load Diff
+2227
-2205
File diff suppressed because it is too large
Load Diff
+2227
-2205
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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)
|
||||
@@ -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,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:
|
||||
|
||||
@@ -1562,6 +1562,11 @@ class SettingsHandler:
|
||||
{"success": False, "error": validation_error}
|
||||
)
|
||||
|
||||
if key == "update_channel" and value not in ("release", "nightly"):
|
||||
return web.json_response(
|
||||
{"success": False, "error": "update_channel must be 'release' or 'nightly'"}
|
||||
)
|
||||
|
||||
if value == "__DELETE__" and key in (
|
||||
"proxy_username",
|
||||
"proxy_password",
|
||||
@@ -2585,6 +2590,8 @@ class ModelLibraryHandler:
|
||||
status=400,
|
||||
)
|
||||
|
||||
cursor = request.query.get("cursor")
|
||||
|
||||
metadata_provider = await self._metadata_provider_factory()
|
||||
if not metadata_provider:
|
||||
return web.json_response(
|
||||
@@ -2593,7 +2600,7 @@ class ModelLibraryHandler:
|
||||
)
|
||||
|
||||
try:
|
||||
models = await metadata_provider.get_user_models(username)
|
||||
result = await metadata_provider.get_user_models(username, cursor)
|
||||
except NotImplementedError:
|
||||
return web.json_response(
|
||||
{
|
||||
@@ -2603,14 +2610,35 @@ class ModelLibraryHandler:
|
||||
status=501,
|
||||
)
|
||||
|
||||
if models is None:
|
||||
if result is None:
|
||||
return web.json_response(
|
||||
{"success": False, "error": "Failed to fetch user models"},
|
||||
status=502,
|
||||
)
|
||||
|
||||
if isinstance(result, dict):
|
||||
models = result.get("items")
|
||||
next_cursor = result.get("nextCursor")
|
||||
else:
|
||||
# Defensive: tolerate providers that still return a raw list
|
||||
models = result
|
||||
next_cursor = None
|
||||
|
||||
if not isinstance(models, list):
|
||||
models = []
|
||||
if next_cursor is not None and not isinstance(next_cursor, str):
|
||||
next_cursor = str(next_cursor)
|
||||
|
||||
estimated_total = None
|
||||
if cursor is None:
|
||||
get_count = getattr(metadata_provider, "get_creator_model_count", None)
|
||||
if get_count is not None:
|
||||
try:
|
||||
estimated_total = await get_count(username)
|
||||
except Exception: # best-effort only
|
||||
estimated_total = None
|
||||
if not isinstance(estimated_total, int):
|
||||
estimated_total = None
|
||||
|
||||
lora_scanner = await self._service_registry.get_lora_scanner()
|
||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||
@@ -2630,6 +2658,7 @@ class ModelLibraryHandler:
|
||||
versions: list[dict] = []
|
||||
history_service = await self._get_download_history_service()
|
||||
model_ids: list[int] = []
|
||||
model_count = 0
|
||||
for model in models:
|
||||
try:
|
||||
model_ids.append(int(model.get("id")))
|
||||
@@ -2663,6 +2692,8 @@ class ModelLibraryHandler:
|
||||
if model_type not in normalized_allowed_types:
|
||||
continue
|
||||
|
||||
model_count += 1
|
||||
|
||||
scanner = type_scanner_map.get(model_type)
|
||||
if scanner is None:
|
||||
return web.json_response(
|
||||
@@ -2728,7 +2759,15 @@ class ModelLibraryHandler:
|
||||
)
|
||||
|
||||
return web.json_response(
|
||||
{"success": True, "username": username, "versions": versions}
|
||||
{
|
||||
"success": True,
|
||||
"username": username,
|
||||
"versions": versions,
|
||||
"modelCount": model_count,
|
||||
"nextCursor": next_cursor,
|
||||
"hasMore": next_cursor is not None,
|
||||
"estimatedTotal": estimated_total,
|
||||
}
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
|
||||
|
||||
@@ -394,12 +394,14 @@ class ModelListingHandler:
|
||||
)
|
||||
|
||||
# View-local-versions filter: show all local versions of a specific model
|
||||
# Accepts either a CivitAI modelId (int) or a HF group key like "hf:user/repo"
|
||||
civitai_model_id = request.query.get("civitai_model_id")
|
||||
if civitai_model_id is not None:
|
||||
try:
|
||||
civitai_model_id = int(civitai_model_id)
|
||||
except (TypeError, ValueError):
|
||||
civitai_model_id = None
|
||||
# Keep as string — could be an HF group key (e.g. "hf:user/repo")
|
||||
pass
|
||||
|
||||
return {
|
||||
"page": page,
|
||||
@@ -537,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
|
||||
@@ -544,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,
|
||||
@@ -566,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:
|
||||
|
||||
+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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1389,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)
|
||||
|
||||
@@ -1827,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)
|
||||
@@ -1842,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,
|
||||
@@ -74,6 +74,9 @@ class DownloadQueueService:
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_dh_completed ON download_history(completed_at DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_dh_status ON download_history(status);
|
||||
"""
|
||||
|
||||
_CREATE_UNIQUE_INDEX = """
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_dh_download_id
|
||||
ON download_history(download_id) WHERE download_id IS NOT NULL;
|
||||
"""
|
||||
@@ -115,10 +118,39 @@ class DownloadQueueService:
|
||||
if self._schema_initialized:
|
||||
return
|
||||
with self._connect() as conn:
|
||||
conn.executescript(self._SCHEMA)
|
||||
conn.executescript(self._SCHEMA_TABLES)
|
||||
|
||||
# Creating the unique index on download_history.download_id can
|
||||
# fail if pre-existing rows have duplicate values (e.g. from a
|
||||
# previous version that lacked the index). Deduplicate first so
|
||||
# that the migration does not crash on startup.
|
||||
if not self._index_exists(conn, "idx_dh_download_id"):
|
||||
self._remove_duplicate_download_ids(conn)
|
||||
conn.executescript(self._CREATE_UNIQUE_INDEX)
|
||||
|
||||
conn.commit()
|
||||
self._schema_initialized = True
|
||||
|
||||
@staticmethod
|
||||
def _index_exists(conn: sqlite3.Connection, name: str) -> bool:
|
||||
return conn.execute(
|
||||
"SELECT 1 FROM sqlite_master WHERE type='index' AND name=?",
|
||||
(name,),
|
||||
).fetchone() is not None
|
||||
|
||||
@staticmethod
|
||||
def _remove_duplicate_download_ids(conn: sqlite3.Connection) -> None:
|
||||
conn.execute("""
|
||||
DELETE FROM download_history
|
||||
WHERE id NOT IN (
|
||||
SELECT MIN(id)
|
||||
FROM download_history
|
||||
WHERE download_id IS NOT NULL
|
||||
GROUP BY download_id
|
||||
)
|
||||
AND download_id IS NOT NULL
|
||||
""")
|
||||
|
||||
def get_database_path(self) -> str:
|
||||
"""Return the resolved database file path."""
|
||||
return self._db_path
|
||||
|
||||
@@ -566,18 +566,52 @@ class LLMService:
|
||||
if effective_max is None:
|
||||
effective_max = 4096
|
||||
|
||||
result = await self.chat_completion(
|
||||
messages=messages,
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
response_format={"type": "json_object"},
|
||||
max_tokens=effective_max,
|
||||
)
|
||||
# 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:
|
||||
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 in json_object mode. "
|
||||
"LLM returned empty content. "
|
||||
f"Raw response: {json.dumps(result)[:500]}"
|
||||
)
|
||||
|
||||
|
||||
+21
-12
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -927,6 +927,25 @@ class ModelScanner:
|
||||
# Update cache data
|
||||
self._cache.raw_data = [item for item in self._cache.raw_data if item['file_path'] not in missing_files]
|
||||
|
||||
dedup_removed = 0
|
||||
seen_paths: set = set()
|
||||
deduped: list = []
|
||||
for item in reversed(self._cache.raw_data):
|
||||
path = item.get('file_path', '')
|
||||
if path not in seen_paths:
|
||||
seen_paths.add(path)
|
||||
deduped.append(item)
|
||||
else:
|
||||
for tag in item.get('tags', []):
|
||||
if tag in self._tags_count:
|
||||
self._tags_count[tag] = max(0, self._tags_count[tag] - 1)
|
||||
if self._tags_count[tag] == 0:
|
||||
del self._tags_count[tag]
|
||||
dedup_removed += 1
|
||||
if dedup_removed > 0:
|
||||
self._cache.raw_data = list(reversed(deduped))
|
||||
total_removed += dedup_removed
|
||||
|
||||
# Resort cache if changes were made
|
||||
if total_added > 0 or total_removed > 0:
|
||||
# Update folders list
|
||||
@@ -1352,18 +1371,25 @@ class ModelScanner:
|
||||
# Update folder in metadata
|
||||
metadata_dict['folder'] = folder
|
||||
|
||||
# Add to cache
|
||||
self._cache.raw_data.append(metadata_dict)
|
||||
self._cache.add_to_version_index(metadata_dict)
|
||||
file_path = metadata_dict.get('file_path', '')
|
||||
if file_path:
|
||||
old_entries = [item for item in self._cache.raw_data if item.get('file_path') == file_path]
|
||||
for old_entry in old_entries:
|
||||
for tag in old_entry.get('tags', []):
|
||||
if tag in self._tags_count:
|
||||
self._tags_count[tag] = max(0, self._tags_count[tag] - 1)
|
||||
if self._tags_count[tag] == 0:
|
||||
del self._tags_count[tag]
|
||||
self._hash_index.remove_by_path(file_path)
|
||||
self._cache.raw_data = [item for item in self._cache.raw_data if item.get('file_path') != file_path]
|
||||
|
||||
for tag in metadata_dict.get('tags', []):
|
||||
self._tags_count[tag] = self._tags_count.get(tag, 0) + 1
|
||||
|
||||
self._cache.raw_data.append(metadata_dict)
|
||||
|
||||
# Resort cache data
|
||||
await self._cache.resort()
|
||||
|
||||
# Update folders list
|
||||
all_folders = set(self._cache.folders)
|
||||
all_folders.add(folder)
|
||||
self._cache.folders = sorted(list(all_folders), key=lambda x: x.lower())
|
||||
|
||||
# Update the hash index
|
||||
self._hash_index.add_entry(metadata_dict['sha256'], metadata_dict['file_path'])
|
||||
await self._persist_current_cache()
|
||||
@@ -1394,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)
|
||||
|
||||
@@ -1723,7 +1752,7 @@ class ModelScanner:
|
||||
# ---- Conditional resort (only when sort-key fields changed) ----
|
||||
need_resort = False
|
||||
_last = cache._last_sort
|
||||
sort_key: Optional[str] = _last[0] if _last != (None, None) else None
|
||||
sort_key: Optional[str] = _last[0] if _last[0] is not None else None
|
||||
if sort_key == "name":
|
||||
if (
|
||||
old_model_name != desired_entry.get("model_name", "")
|
||||
@@ -1971,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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -126,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"})
|
||||
@@ -148,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
|
||||
|
||||
@@ -87,6 +92,13 @@ class _DownloadProgress(dict):
|
||||
return snapshot
|
||||
|
||||
|
||||
# When fewer candidates than this remain in check_pending_models, probe each
|
||||
# model folder directly (preserving legacy-folder migration semantics). Above
|
||||
# it, build a folder index with a single directory scan so libraries with
|
||||
# 100k+ models do not pay one syscall per candidate.
|
||||
_BULK_LOOKUP_THRESHOLD = 1000
|
||||
|
||||
|
||||
def _model_directory_has_files(path: str) -> bool:
|
||||
"""Return True when the provided directory exists and contains entries."""
|
||||
|
||||
@@ -103,6 +115,36 @@ def _model_directory_has_files(path: str) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _build_example_folder_index(output_dir: str) -> dict[str, bool]:
|
||||
"""Build a ``{hash: has_files}`` index for a library's example-image folders.
|
||||
|
||||
A single directory scan over the library root replaces ``O(candidates)``
|
||||
per-folder ``os.scandir`` calls, which is required for libraries with
|
||||
100k+ models. Each hash folder is classified by whether it contains any
|
||||
entries, matching the semantics of ``_model_directory_has_files``.
|
||||
"""
|
||||
|
||||
index: dict[str, bool] = {}
|
||||
if not output_dir or not os.path.isdir(output_dir):
|
||||
return index
|
||||
|
||||
try:
|
||||
with os.scandir(output_dir) as entries:
|
||||
for entry in entries:
|
||||
name = entry.name
|
||||
if not entry.is_dir() or not is_hash_folder(name):
|
||||
continue
|
||||
try:
|
||||
with os.scandir(entry.path) as subentries:
|
||||
index[name.lower()] = any(subentries)
|
||||
except OSError:
|
||||
index[name.lower()] = False
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
return index
|
||||
|
||||
|
||||
class DownloadManager:
|
||||
"""Manages downloading example images for models."""
|
||||
|
||||
@@ -130,6 +172,7 @@ class DownloadManager:
|
||||
model_types = data.get("model_types", ["lora", "checkpoint"])
|
||||
delay = float(data.get("delay", 0.2))
|
||||
force = data.get("force", False)
|
||||
model_hashes = data.get("model_hashes", [])
|
||||
|
||||
# Step 2: Validate configuration (fast lookup)
|
||||
settings_manager = get_settings_manager()
|
||||
@@ -199,6 +242,7 @@ class DownloadManager:
|
||||
delay,
|
||||
active_library,
|
||||
force,
|
||||
model_hashes,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -410,14 +454,49 @@ class DownloadManager:
|
||||
# Calculate pending count: check which models actually need processing.
|
||||
# A model is pending if it has a hash, is not already processed or known-failed,
|
||||
# and its folder doesn't exist or is empty.
|
||||
pending_hashes = set()
|
||||
for model_hash, model_name in all_models_with_hash:
|
||||
if model_hash not in processed_models and model_hash not in failed_models:
|
||||
candidate_hashes = [
|
||||
model_hash
|
||||
for model_hash, _ in all_models_with_hash
|
||||
if model_hash not in processed_models
|
||||
and model_hash not in failed_models
|
||||
]
|
||||
|
||||
pending_hashes: set[str] = set()
|
||||
# For small candidate counts the existing per-folder check is fine
|
||||
# and handles legacy folder migration.
|
||||
# For large libraries, scan the library root once and do set lookups.
|
||||
if len(candidate_hashes) <= _BULK_LOOKUP_THRESHOLD or not output_dir:
|
||||
for model_hash in candidate_hashes:
|
||||
model_dir = ExampleImagePathResolver.get_model_folder(
|
||||
model_hash, active_library
|
||||
)
|
||||
if not _model_directory_has_files(model_dir):
|
||||
pending_hashes.add(model_hash)
|
||||
else:
|
||||
folder_index = await asyncio.get_event_loop().run_in_executor(
|
||||
None, _build_example_folder_index, output_dir
|
||||
)
|
||||
# In multi-library mode, folders that have not been consolidated
|
||||
# into the library root yet (startup migration skipped, failed
|
||||
# move, or created at the legacy path afterwards) still live at
|
||||
# the legacy root/<hash> location. Only scan that root when at
|
||||
# least one candidate is missing from the library-root index, so
|
||||
# the fully-consolidated case does not pay an extra directory
|
||||
# pass on every call.
|
||||
if uses_library_scoped_folders() and any(
|
||||
not folder_index.get(model_hash, False)
|
||||
for model_hash in candidate_hashes
|
||||
):
|
||||
legacy_root = get_example_images_root()
|
||||
if legacy_root and legacy_root != output_dir:
|
||||
legacy_index = await asyncio.get_event_loop().run_in_executor(
|
||||
None, _build_example_folder_index, legacy_root
|
||||
)
|
||||
for hash_key, has_files in legacy_index.items():
|
||||
folder_index.setdefault(hash_key, has_files)
|
||||
for model_hash in candidate_hashes:
|
||||
if not folder_index.get(model_hash, False):
|
||||
pending_hashes.add(model_hash)
|
||||
|
||||
pending_count = len(pending_hashes)
|
||||
|
||||
@@ -500,8 +579,9 @@ class DownloadManager:
|
||||
delay,
|
||||
library_name,
|
||||
force: bool = False,
|
||||
model_hashes: list[str] | None = None,
|
||||
):
|
||||
"""Download example images for all models."""
|
||||
"""Download example images for all models (or only the given hashes)."""
|
||||
|
||||
downloader = await get_downloader()
|
||||
|
||||
@@ -529,6 +609,18 @@ class DownloadManager:
|
||||
if model.get("sha256"):
|
||||
all_models.append((scanner_type, model, scanner))
|
||||
|
||||
# Restrict to the requested hashes when provided (empty = all models).
|
||||
# Explicit targets are a directed user request, so previously failed
|
||||
# models are retried instead of skipped.
|
||||
explicit_targets = bool(model_hashes)
|
||||
if model_hashes:
|
||||
hash_set = {h.lower() for h in model_hashes}
|
||||
all_models = [
|
||||
(scanner_type, model, scanner)
|
||||
for scanner_type, model, scanner in all_models
|
||||
if model.get("sha256", "").lower() in hash_set
|
||||
]
|
||||
|
||||
# Update total count
|
||||
self._progress["total"] = len(all_models)
|
||||
logger.debug(f"Found {self._progress['total']} models to process")
|
||||
@@ -552,6 +644,7 @@ class DownloadManager:
|
||||
downloader,
|
||||
library_name,
|
||||
force,
|
||||
explicit_targets,
|
||||
)
|
||||
|
||||
# Update progress
|
||||
@@ -648,6 +741,7 @@ class DownloadManager:
|
||||
downloader,
|
||||
library_name,
|
||||
force: bool = False,
|
||||
explicit_targets: bool = False,
|
||||
):
|
||||
"""Process a single model download."""
|
||||
|
||||
@@ -670,8 +764,9 @@ class DownloadManager:
|
||||
self._progress["current_model"] = f"{model_name} ({model_hash[:8]})"
|
||||
await self._broadcast_progress(status="running")
|
||||
|
||||
# Skip if already in failed models (unless force mode is enabled)
|
||||
if not force and model_hash in self._progress["failed_models"]:
|
||||
# Skip if already in failed models (unless force mode is enabled or
|
||||
# the model was explicitly targeted by hash)
|
||||
if not force and not explicit_targets and model_hash in self._progress["failed_models"]:
|
||||
logger.debug(f"Skipping known failed model: {model_name}")
|
||||
return False
|
||||
|
||||
@@ -680,30 +775,34 @@ class DownloadManager:
|
||||
)
|
||||
existing_files = _model_directory_has_files(model_dir)
|
||||
|
||||
# Skip if already processed AND directory exists with files
|
||||
if model_hash in self._progress["processed_models"]:
|
||||
if existing_files:
|
||||
logger.debug(f"Skipping already processed model: {model_name}")
|
||||
# Model-level guard: a populated folder counts as done. Explicitly
|
||||
# targeted models bypass it so the per-image existence pre-check can
|
||||
# fill individual gaps without re-fetching existing files.
|
||||
if not explicit_targets:
|
||||
# Skip if already processed AND directory exists with files
|
||||
if model_hash in self._progress["processed_models"]:
|
||||
if existing_files:
|
||||
logger.debug(f"Skipping already processed model: {model_name}")
|
||||
return False
|
||||
|
||||
logger.debug(
|
||||
"Model %s (%s) marked as processed but folder empty or missing, reprocessing triggered",
|
||||
model_name,
|
||||
model_hash,
|
||||
)
|
||||
# Track that we are reprocessing this model for summary logging
|
||||
self._progress["reprocessed_models"].add(model_hash)
|
||||
# Remove from processed models since we need to reprocess
|
||||
self._progress["processed_models"].discard(model_hash)
|
||||
|
||||
if existing_files and model_hash not in self._progress["processed_models"]:
|
||||
logger.debug(
|
||||
"Model folder already populated for %s, marking as processed without download",
|
||||
model_name,
|
||||
)
|
||||
self._progress["processed_models"].add(model_hash)
|
||||
return False
|
||||
|
||||
logger.debug(
|
||||
"Model %s (%s) marked as processed but folder empty or missing, reprocessing triggered",
|
||||
model_name,
|
||||
model_hash,
|
||||
)
|
||||
# Track that we are reprocessing this model for summary logging
|
||||
self._progress["reprocessed_models"].add(model_hash)
|
||||
# Remove from processed models since we need to reprocess
|
||||
self._progress["processed_models"].discard(model_hash)
|
||||
|
||||
if existing_files and model_hash not in self._progress["processed_models"]:
|
||||
logger.debug(
|
||||
"Model folder already populated for %s, marking as processed without download",
|
||||
model_name,
|
||||
)
|
||||
self._progress["processed_models"].add(model_hash)
|
||||
return False
|
||||
|
||||
if not model_dir:
|
||||
logger.warning(
|
||||
"Unable to resolve example images folder for model %s (%s)",
|
||||
@@ -807,7 +906,7 @@ class DownloadManager:
|
||||
model_name,
|
||||
)
|
||||
# Clear failed_models so non-force runs can retry
|
||||
if force and model_hash in self._progress["failed_models"]:
|
||||
if (force or explicit_targets) and model_hash in self._progress["failed_models"]:
|
||||
self._progress["failed_models"].discard(model_hash)
|
||||
logger.info(
|
||||
f"Removed {model_name} from failed_models after force retry with rate-limited images"
|
||||
@@ -827,7 +926,7 @@ class DownloadManager:
|
||||
)
|
||||
elif success:
|
||||
self._progress["processed_models"].add(model_hash)
|
||||
if force and model_hash in self._progress["failed_models"]:
|
||||
if (force or explicit_targets) and model_hash in self._progress["failed_models"]:
|
||||
self._progress["failed_models"].discard(model_hash)
|
||||
logger.info(
|
||||
f"Removed {model_name} from failed_models after successful force retry"
|
||||
@@ -1343,8 +1442,8 @@ class DownloadManager:
|
||||
await MetadataManager.save_metadata(file_path, model_copy)
|
||||
|
||||
try:
|
||||
await scanner.update_single_model_cache(
|
||||
file_path, file_path, model_data
|
||||
await update_cache_from_metadata(
|
||||
scanner, file_path, model_copy
|
||||
)
|
||||
except AttributeError:
|
||||
logger.debug(
|
||||
|
||||
@@ -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 ""
|
||||
@@ -113,35 +118,6 @@ def get_model_folder(model_hash: str, library_name: Optional[str] = None) -> str
|
||||
exc,
|
||||
)
|
||||
return legacy_folder
|
||||
elif not os.path.exists(resolved_folder):
|
||||
# Reverse migration: when consolidating from multi-library to
|
||||
# single-library mode (e.g. after "default" was cleaned up), look
|
||||
# for existing example images inside library-named subdirectories
|
||||
# and bring them back to the root level.
|
||||
root = get_example_images_root()
|
||||
if root:
|
||||
try:
|
||||
for entry in os.listdir(root):
|
||||
entry_path = os.path.join(root, entry)
|
||||
if not os.path.isdir(entry_path):
|
||||
continue
|
||||
if is_hash_folder(entry) or entry == "_deleted":
|
||||
continue
|
||||
if not _library_folder_has_only_hash_dirs(entry_path):
|
||||
continue
|
||||
legacy = os.path.join(entry_path, normalized_hash)
|
||||
if os.path.exists(legacy):
|
||||
shutil.move(legacy, resolved_folder)
|
||||
logger.info(
|
||||
"Consolidated example images from '%s' to '%s'",
|
||||
legacy, resolved_folder,
|
||||
)
|
||||
break
|
||||
except OSError as exc:
|
||||
logger.error(
|
||||
"Failed to consolidate example images during "
|
||||
"library merge: %s", exc,
|
||||
)
|
||||
|
||||
return resolved_folder
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ from ..utils.constants import SUPPORTED_MEDIA_EXTENSIONS
|
||||
from ..services.service_registry import ServiceRegistry
|
||||
from ..services.settings_manager import get_settings_manager
|
||||
from ..utils.example_images_paths import get_model_folder, get_model_relative_path
|
||||
from .example_images_metadata import MetadataUpdater
|
||||
from .example_images_metadata import MetadataUpdater, update_cache_from_metadata
|
||||
from ..utils.metadata_manager import MetadataManager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -113,6 +113,26 @@ class ExampleImagesProcessor:
|
||||
message = str(error).lower()
|
||||
return '404' in message or 'file not found' in message
|
||||
|
||||
@staticmethod
|
||||
def _example_image_file_exists(model_dir: str, index: int, media_type_hint: str | None = None) -> bool:
|
||||
"""Return True when the file that would be written for a media index already exists.
|
||||
|
||||
The final filename (``image_{index}{extension}``) depends on the downloaded
|
||||
content, so the extension cannot be known ahead of time. The post-download
|
||||
check skips the write when the exact target file exists; this pre-check
|
||||
approximates that with the candidate extensions for the media type (videos
|
||||
only when the metadata hints at a video) so the network request is avoided
|
||||
for files that already exist on disk.
|
||||
"""
|
||||
if media_type_hint == "video":
|
||||
extensions = SUPPORTED_MEDIA_EXTENSIONS['videos']
|
||||
else:
|
||||
extensions = SUPPORTED_MEDIA_EXTENSIONS['images']
|
||||
return any(
|
||||
os.path.exists(os.path.join(model_dir, f"image_{index}{ext}"))
|
||||
for ext in extensions
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def download_model_images(model_hash, model_name, model_images, model_dir, optimize, downloader):
|
||||
"""Download images for a single model
|
||||
@@ -139,7 +159,12 @@ class ExampleImagesProcessor:
|
||||
original_url = image_url
|
||||
if optimize and 'civitai.com' in image_url:
|
||||
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
||||
|
||||
|
||||
# Skip the download when the file already exists on disk
|
||||
if ExampleImagesProcessor._example_image_file_exists(model_dir, i, image.get("type")):
|
||||
logger.debug("File already exists, skipping download for %s", image_url)
|
||||
continue
|
||||
|
||||
# Download the file first to determine the actual file type
|
||||
try:
|
||||
logger.debug(f"Downloading media file {i} for {model_name}")
|
||||
@@ -229,6 +254,11 @@ class ExampleImagesProcessor:
|
||||
if optimize and 'civitai.com' in image_url:
|
||||
image_url = ExampleImagesProcessor.get_civitai_optimized_url(image_url)
|
||||
|
||||
# Skip the download when the file already exists on disk
|
||||
if ExampleImagesProcessor._example_image_file_exists(model_dir, i, image.get("type")):
|
||||
logger.debug("File already exists, skipping download for %s", image_url)
|
||||
continue
|
||||
|
||||
async def _attempt_download() -> tuple:
|
||||
logger.debug("Downloading media file %s for %s", i, model_name)
|
||||
return await downloader.download_to_memory(
|
||||
@@ -644,7 +674,7 @@ class ExampleImagesProcessor:
|
||||
}, status=500)
|
||||
|
||||
# Update cache
|
||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
||||
await update_cache_from_metadata(scanner, file_path, model_data)
|
||||
|
||||
# Get regular images array (might be None)
|
||||
regular_images = civitai_data.get('images', [])
|
||||
@@ -759,7 +789,7 @@ class ExampleImagesProcessor:
|
||||
model_copy = model_data.copy()
|
||||
model_copy.pop('folder', None)
|
||||
await MetadataManager.save_metadata(file_path, model_copy)
|
||||
await scanner.update_single_model_cache(file_path, file_path, model_data)
|
||||
await update_cache_from_metadata(scanner, file_path, model_copy)
|
||||
|
||||
return web.json_response({
|
||||
'success': True,
|
||||
|
||||
+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.8"
|
||||
version = "1.2.0"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = [
|
||||
"aiohttp",
|
||||
|
||||
@@ -151,6 +151,7 @@ body.modal-open {
|
||||
.support-section,
|
||||
.changelog-section,
|
||||
.update-info,
|
||||
.update-channels,
|
||||
.info-item,
|
||||
.path-preview {
|
||||
background: var(--surface-subtle);
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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`,
|
||||
@@ -188,7 +184,8 @@ export const DOWNLOAD_ENDPOINTS = {
|
||||
downloadGet: '/api/lm/download-model-get',
|
||||
cancelGet: '/api/lm/cancel-download-get',
|
||||
progress: '/api/lm/download-progress',
|
||||
exampleImages: '/api/lm/force-download-example-images' // New endpoint for downloading example images
|
||||
exampleImages: '/api/lm/force-download-example-images', // Re-process example images ignoring previous status
|
||||
exampleImagesMissing: '/api/lm/download-example-images' // Download only missing example images
|
||||
};
|
||||
|
||||
// Hugging Face API endpoints
|
||||
|
||||
@@ -1641,7 +1641,7 @@ export class BaseModelApiClient {
|
||||
}
|
||||
}
|
||||
|
||||
async downloadExampleImages(modelHashes, modelTypes = null) {
|
||||
async downloadExampleImages(modelHashes, modelTypes = null, { force = true } = {}) {
|
||||
let ws = null;
|
||||
|
||||
await state.loadingManager.showWithProgress(async (loading) => {
|
||||
@@ -1700,8 +1700,13 @@ export class BaseModelApiClient {
|
||||
// Determine optimize setting
|
||||
const optimize = state.global?.settings?.optimize_example_images ?? true;
|
||||
|
||||
// force=false routes to the regular endpoint, which skips already-processed models
|
||||
const endpoint = force
|
||||
? DOWNLOAD_ENDPOINTS.exampleImages
|
||||
: DOWNLOAD_ENDPOINTS.exampleImagesMissing;
|
||||
|
||||
// Make the API request to start the download process
|
||||
const response = await fetch(DOWNLOAD_ENDPOINTS.exampleImages, {
|
||||
const response = await fetch(endpoint, {
|
||||
method: 'POST',
|
||||
headers: {
|
||||
'Content-Type': 'application/json'
|
||||
@@ -1710,6 +1715,7 @@ export class BaseModelApiClient {
|
||||
model_hashes: modelHashes,
|
||||
output_dir: outputDir,
|
||||
optimize: optimize,
|
||||
force: force,
|
||||
model_types: modelTypes || [this.apiConfig.config.singularName]
|
||||
})
|
||||
});
|
||||
|
||||
@@ -137,11 +137,10 @@ export class BulkContextMenu extends BaseContextMenu {
|
||||
downloadMissingLorasItem.style.display = currentModelType === 'recipes' ? 'flex' : 'none';
|
||||
}
|
||||
|
||||
const downloadExampleImagesItem = this.menu.querySelector('[data-action="download-example-images"]');
|
||||
if (downloadExampleImagesItem) {
|
||||
const downloadExampleImagesSubmenu = this.menu.querySelector('[data-has-submenu="download-example-images"]');
|
||||
if (downloadExampleImagesSubmenu) {
|
||||
// Show on model pages (loras, checkpoints, embeddings), hide on recipes
|
||||
const modelPages = ['loras', 'checkpoints', 'embeddings'];
|
||||
downloadExampleImagesItem.style.display = modelPages.includes(currentModelType) ? 'flex' : 'none';
|
||||
downloadExampleImagesSubmenu.style.display = ['loras', 'checkpoints', 'embeddings'].includes(currentModelType) ? 'flex' : 'none';
|
||||
}
|
||||
|
||||
const skipMetadataRefreshItem = this.menu.querySelector('[data-action="skip-metadata-refresh"]');
|
||||
@@ -294,8 +293,11 @@ export class BulkContextMenu extends BaseContextMenu {
|
||||
case 'download-missing-loras':
|
||||
this.handleDownloadMissingLoras();
|
||||
break;
|
||||
case 'download-missing-example-images':
|
||||
this.handleDownloadExampleImages({ force: false });
|
||||
break;
|
||||
case 'download-example-images':
|
||||
this.handleDownloadExampleImages();
|
||||
this.handleDownloadExampleImages({ force: true });
|
||||
break;
|
||||
case 'clear':
|
||||
bulkManager.clearSelection();
|
||||
@@ -340,7 +342,7 @@ export class BulkContextMenu extends BaseContextMenu {
|
||||
await bulkMissingLoraDownloadManager.downloadMissingLoras(selectedRecipes);
|
||||
}
|
||||
|
||||
async handleDownloadExampleImages() {
|
||||
async handleDownloadExampleImages({ force = true } = {}) {
|
||||
if (state.selectedModels.size === 0) {
|
||||
return;
|
||||
}
|
||||
@@ -361,7 +363,7 @@ export class BulkContextMenu extends BaseContextMenu {
|
||||
|
||||
try {
|
||||
const apiClient = getModelApiClient();
|
||||
await apiClient.downloadExampleImages([...hashes]);
|
||||
await apiClient.downloadExampleImages([...hashes], null, { force });
|
||||
} catch (error) {
|
||||
console.error('Bulk download example images failed:', error);
|
||||
}
|
||||
|
||||
@@ -347,7 +347,10 @@ export const ModelContextMenuMixin = {
|
||||
openExampleImagesFolder(this.currentCard.dataset.sha256);
|
||||
return true;
|
||||
case 'download-examples':
|
||||
this.downloadExampleImages();
|
||||
this.downloadExampleImages(false);
|
||||
return true;
|
||||
case 'download-examples-force':
|
||||
this.downloadExampleImages(true);
|
||||
return true;
|
||||
case 'civitai':
|
||||
if (this.currentCard.dataset.from_civitai === 'true') {
|
||||
@@ -378,7 +381,7 @@ export const ModelContextMenuMixin = {
|
||||
},
|
||||
|
||||
// Download example images method
|
||||
async downloadExampleImages() {
|
||||
async downloadExampleImages(force = false) {
|
||||
const modelHash = this.currentCard.dataset.sha256;
|
||||
if (!modelHash) {
|
||||
showToast('toast.contextMenu.missingHash', {}, 'error');
|
||||
@@ -387,7 +390,7 @@ export const ModelContextMenuMixin = {
|
||||
|
||||
try {
|
||||
const apiClient = getModelApiClient();
|
||||
await apiClient.downloadExampleImages([modelHash]);
|
||||
await apiClient.downloadExampleImages([modelHash], null, { force });
|
||||
} catch (error) {
|
||||
console.error('Error downloading example images:', error);
|
||||
}
|
||||
|
||||
@@ -260,8 +260,9 @@ export class RecipeContextMenu extends BaseContextMenu {
|
||||
strength: lora.strength || 1.0,
|
||||
|
||||
// Model identifiers
|
||||
modelId: lora.modelId || lora.model_id || civitaiInfo.modelId,
|
||||
hash: modelFile?.hashes?.SHA256?.toLowerCase() || lora.hash,
|
||||
modelVersionId: civitaiInfo.id || lora.modelVersionId,
|
||||
id: civitaiInfo.id || lora.modelVersionId,
|
||||
|
||||
// Metadata
|
||||
thumbnailUrl: civitaiInfo.images?.[0]?.url || '',
|
||||
|
||||
@@ -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,
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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':
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -27,6 +27,8 @@ export class BulkManager {
|
||||
|
||||
// Drag detection properties
|
||||
this.dragThreshold = 5; // Pixels to move before considering it a drag
|
||||
this.dragDelayMs = 100; // Minimum hold time before a drag is treated as a marquee
|
||||
this.minMarqueeSize = 10; // Minimum drag box (px) before a marquee counts as a selection
|
||||
this.mouseDownTime = 0;
|
||||
this.mouseDownPosition = { x: 0, y: 0 };
|
||||
|
||||
@@ -88,7 +90,7 @@ export class BulkManager {
|
||||
moveAll: true,
|
||||
autoOrganize: false,
|
||||
deleteAll: true,
|
||||
setContentRating: false,
|
||||
setContentRating: true,
|
||||
skipMetadataRefresh: false,
|
||||
setFavorite: true,
|
||||
unfavorite: true,
|
||||
@@ -173,6 +175,19 @@ export class BulkManager {
|
||||
});
|
||||
|
||||
eventManager.addHandler('mousemove', 'bulkManager-marquee-move', (e) => {
|
||||
// Only track marquee/drag while the left button is physically held.
|
||||
// mouseup can be missed (release outside the window, focus loss, driver quirks),
|
||||
// so mousemove must verify the button state itself instead of relying on it.
|
||||
if (!(e.buttons & 1)) {
|
||||
if (this.isMarqueeActive) {
|
||||
this.endMarqueeSelection(e);
|
||||
} else {
|
||||
this.mouseDownTime = 0;
|
||||
this.isDragging = false;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
if (this.isMarqueeActive) {
|
||||
this.lastClientX = e.clientX;
|
||||
this.lastClientY = e.clientY;
|
||||
@@ -184,7 +199,10 @@ export class BulkManager {
|
||||
const dy = e.clientY - this.mouseDownPosition.y;
|
||||
const distance = Math.sqrt(dx * dx + dy * dy);
|
||||
|
||||
if (distance >= this.dragThreshold) {
|
||||
// Require both enough movement AND enough hold time so quick
|
||||
// click jitter from micro-movement input devices is not a marquee.
|
||||
const heldTime = Date.now() - this.mouseDownTime;
|
||||
if (heldTime >= this.dragDelayMs && distance >= this.dragThreshold) {
|
||||
this.isDragging = true;
|
||||
this.startMarqueeSelection(e, true);
|
||||
}
|
||||
@@ -1510,14 +1528,18 @@ export class BulkManager {
|
||||
let failureCount = 0;
|
||||
|
||||
try {
|
||||
const apiClient = getModelApiClient();
|
||||
const isRecipesPage = state.currentPageType === 'recipes';
|
||||
for (const filePath of targets) {
|
||||
if (cancelled) {
|
||||
showToast('toast.api.operationCancelled', {}, 'info');
|
||||
break;
|
||||
}
|
||||
try {
|
||||
await apiClient.saveModelMetadata(filePath, { preview_nsfw_level: level });
|
||||
if (isRecipesPage) {
|
||||
await updateRecipeMetadata(filePath, { preview_nsfw_level: level });
|
||||
} else {
|
||||
await getModelApiClient().saveModelMetadata(filePath, { preview_nsfw_level: level });
|
||||
}
|
||||
successCount++;
|
||||
} catch (error) {
|
||||
failureCount++;
|
||||
@@ -1958,9 +1980,31 @@ export class BulkManager {
|
||||
// Remove visual feedback class
|
||||
document.body.classList.remove('marquee-selecting');
|
||||
|
||||
// Compute the actual drag box size in document coordinates, matching how
|
||||
// updateMarqueeSelectionFromPosition tracks the rectangle. Client-space
|
||||
// size would wrongly flag auto-scroll marquees (tiny pointer movement,
|
||||
// large document-space box) as accidental clicks.
|
||||
const container = document.querySelector('.page-content');
|
||||
const scrollX = container?.scrollLeft || 0;
|
||||
const scrollY = container?.scrollTop || 0;
|
||||
const dragWidth = Math.abs((e.clientX + scrollX) - this.marqueeStartDoc.x);
|
||||
const dragHeight = Math.abs((e.clientY + scrollY) - this.marqueeStartDoc.y);
|
||||
const isTinyMarquee = dragWidth < this.minMarqueeSize && dragHeight < this.minMarqueeSize;
|
||||
|
||||
// Get selection count
|
||||
const selectionCount = state.selectedModels.size;
|
||||
|
||||
// A tiny box (e.g. click jitter that happened to graze a card) is treated
|
||||
// as an accidental click: undo any selection and leave bulk mode.
|
||||
if (isTinyMarquee) {
|
||||
this.clearSelection();
|
||||
if (state.bulkMode) {
|
||||
this.toggleBulkMode();
|
||||
}
|
||||
this.initialSelectedModels.clear();
|
||||
return;
|
||||
}
|
||||
|
||||
// If no models were selected, exit bulk mode
|
||||
if (selectionCount === 0) {
|
||||
if (state.bulkMode) {
|
||||
|
||||
@@ -158,6 +158,7 @@ export class DownloadManager {
|
||||
this.modelVersionId = null;
|
||||
this.source = null;
|
||||
this.selectedFile = null;
|
||||
this._isDiffusionModel = false;
|
||||
|
||||
this.selectedFolder = '';
|
||||
this.batchModels = [];
|
||||
@@ -787,24 +788,40 @@ export class DownloadManager {
|
||||
async proceedToLocationContent() {
|
||||
|
||||
try {
|
||||
// Fetch model roots
|
||||
const rootsData = await this.apiClient.fetchModelRoots();
|
||||
const _isDiffusionModel = this.selectedFile
|
||||
? (this.selectedFile.type === 'UNet' || this.selectedFile.type === 'Diffusion Model')
|
||||
: (this.currentVersion?.files || []).some(
|
||||
f => f.type === 'UNet' || f.type === 'Diffusion Model'
|
||||
);
|
||||
this._isDiffusionModel = _isDiffusionModel;
|
||||
|
||||
let rootsData;
|
||||
if (this._isDiffusionModel && this.apiClient.modelType === 'checkpoints') {
|
||||
rootsData = await this.apiClient.fetchModelRoots('diffusion_model');
|
||||
} else {
|
||||
rootsData = await this.apiClient.fetchModelRoots();
|
||||
}
|
||||
const modelRoot = document.getElementById('modelRoot');
|
||||
modelRoot.innerHTML = rootsData.roots.map(root =>
|
||||
`<option value="${root}">${root}</option>`
|
||||
).join('');
|
||||
|
||||
// Set default root if available
|
||||
const singularType = this.apiClient.modelType.replace(/s$/, '');
|
||||
const singularType = this._isDiffusionModel
|
||||
? 'unet'
|
||||
: this.apiClient.modelType.replace(/s$/, '');
|
||||
const defaultRootKey = `default_${singularType}_root`;
|
||||
const defaultRoot = state.global.settings[defaultRootKey];
|
||||
console.log(`Default root for ${this.apiClient.modelType}:`, defaultRoot);
|
||||
console.log(`Default root for ${singularType}:`, defaultRoot);
|
||||
console.log('Available roots:', rootsData.roots);
|
||||
if (defaultRoot && rootsData.roots.includes(defaultRoot)) {
|
||||
console.log(`Setting default root: ${defaultRoot}`);
|
||||
modelRoot.value = defaultRoot;
|
||||
}
|
||||
|
||||
const subtypeDisplay = this._isDiffusionModel ? 'Diffusion Model' : this.apiClient.apiConfig.config.displayName;
|
||||
document.getElementById('modelRootLabel').textContent =
|
||||
translate('modals.download.selectTypeRoot', { type: subtypeDisplay });
|
||||
|
||||
// Set autocomplete="off" on folderPath input
|
||||
const folderPathInput = document.getElementById('folderPath');
|
||||
if (folderPathInput) {
|
||||
@@ -1776,13 +1793,15 @@ export class DownloadManager {
|
||||
const modelRoot = document.getElementById('modelRoot').value;
|
||||
const config = this.apiClient.apiConfig.config;
|
||||
|
||||
let fullPath = modelRoot || translate('modals.download.selectTypeRoot', { type: config.displayName });
|
||||
const subtypeDisplay = this._isDiffusionModel ? 'Diffusion Model' : config.displayName;
|
||||
let fullPath = modelRoot || translate('modals.download.selectTypeRoot', { type: subtypeDisplay });
|
||||
|
||||
if (modelRoot) {
|
||||
if (this.useDefaultPath) {
|
||||
// Show actual template path
|
||||
try {
|
||||
const singularType = this.apiClient.modelType.replace(/s$/, '');
|
||||
const singularType = this._isDiffusionModel
|
||||
? 'unet'
|
||||
: this.apiClient.modelType.replace(/s$/, '');
|
||||
const templates = state.global.settings.download_path_templates;
|
||||
const template = templates[singularType];
|
||||
fullPath += `/${template}`;
|
||||
|
||||
@@ -729,10 +729,12 @@ export class FilterManager {
|
||||
const pageState = getCurrentPageState();
|
||||
const storageKey = `${this.currentPage}_filters`;
|
||||
|
||||
// Save filters to localStorage (exclude EMPTY_WILDCARD_MARKER)
|
||||
// Save filters to localStorage (exclude EMPTY_WILDCARD_MARKER and transient search)
|
||||
const filtersSnapshot = this.cloneFilters();
|
||||
// Don't persist EMPTY_WILDCARD_MARKER - it's a runtime-only marker
|
||||
filtersSnapshot.baseModel = filtersSnapshot.baseModel.filter(m => m !== EMPTY_WILDCARD_MARKER);
|
||||
// Don't persist search - it's transient and managed by SearchManager
|
||||
delete filtersSnapshot.search;
|
||||
setStorageItem(storageKey, filtersSnapshot);
|
||||
|
||||
// Update state with current filters
|
||||
@@ -984,6 +986,7 @@ export class FilterManager {
|
||||
}
|
||||
|
||||
cloneFilters() {
|
||||
const pageState = getCurrentPageState();
|
||||
return {
|
||||
...this.filters,
|
||||
baseModel: [...(this.filters.baseModel || [])],
|
||||
@@ -991,7 +994,8 @@ export class FilterManager {
|
||||
autoTags: { ...(this.filters.autoTags || {}) },
|
||||
license: { ...(this.filters.license || {}) },
|
||||
modelTypes: [...(this.filters.modelTypes || [])],
|
||||
tagLogic: this.filters.tagLogic || 'any'
|
||||
tagLogic: this.filters.tagLogic || 'any',
|
||||
search: pageState?.filters?.search ?? ''
|
||||
};
|
||||
}
|
||||
|
||||
|
||||
@@ -1517,11 +1517,20 @@ export class SettingsManager {
|
||||
return data;
|
||||
}
|
||||
|
||||
async loadLoraRoots() {
|
||||
try {
|
||||
const defaultLoraRootSelect = document.getElementById('defaultLoraRoot');
|
||||
if (!defaultLoraRootSelect) return;
|
||||
showNoRootsPlaceholder(select) {
|
||||
select.innerHTML = '';
|
||||
const option = document.createElement('option');
|
||||
option.value = '';
|
||||
option.textContent = translate('settings.folderSettings.noDefault', {}, 'No Default');
|
||||
select.appendChild(option);
|
||||
select.disabled = true;
|
||||
}
|
||||
|
||||
async loadLoraRoots() {
|
||||
const defaultLoraRootSelect = document.getElementById('defaultLoraRoot');
|
||||
if (!defaultLoraRootSelect) return;
|
||||
|
||||
try {
|
||||
// Fetch lora roots
|
||||
const response = await fetch('/api/lm/loras/roots');
|
||||
if (!response.ok) {
|
||||
@@ -1530,10 +1539,12 @@ export class SettingsManager {
|
||||
|
||||
const data = await response.json();
|
||||
if (!data.roots || data.roots.length === 0) {
|
||||
throw new Error('No LoRA roots found');
|
||||
this.showNoRootsPlaceholder(defaultLoraRootSelect);
|
||||
return;
|
||||
}
|
||||
|
||||
defaultLoraRootSelect.innerHTML = '';
|
||||
defaultLoraRootSelect.disabled = false;
|
||||
|
||||
// Add options for each root
|
||||
data.roots.forEach(root => {
|
||||
@@ -1548,15 +1559,16 @@ export class SettingsManager {
|
||||
|
||||
} catch (error) {
|
||||
console.error('Error loading LoRA roots:', error);
|
||||
this.showNoRootsPlaceholder(defaultLoraRootSelect);
|
||||
showToast('toast.settings.loraRootsFailed', { message: error.message }, 'error');
|
||||
}
|
||||
}
|
||||
|
||||
async loadCheckpointRoots() {
|
||||
try {
|
||||
const defaultCheckpointRootSelect = document.getElementById('defaultCheckpointRoot');
|
||||
if (!defaultCheckpointRootSelect) return;
|
||||
const defaultCheckpointRootSelect = document.getElementById('defaultCheckpointRoot');
|
||||
if (!defaultCheckpointRootSelect) return;
|
||||
|
||||
try {
|
||||
// Fetch checkpoint roots (checkpoint paths only, not unet)
|
||||
const response = await fetch('/api/lm/checkpoints/checkpoints_roots');
|
||||
if (!response.ok) {
|
||||
@@ -1565,10 +1577,12 @@ export class SettingsManager {
|
||||
|
||||
const data = await response.json();
|
||||
if (!data.roots || data.roots.length === 0) {
|
||||
throw new Error('No checkpoint roots found');
|
||||
this.showNoRootsPlaceholder(defaultCheckpointRootSelect);
|
||||
return;
|
||||
}
|
||||
|
||||
defaultCheckpointRootSelect.innerHTML = '';
|
||||
defaultCheckpointRootSelect.disabled = false;
|
||||
|
||||
// Add options for each root
|
||||
data.roots.forEach(root => {
|
||||
@@ -1583,15 +1597,16 @@ export class SettingsManager {
|
||||
|
||||
} catch (error) {
|
||||
console.error('Error loading checkpoint roots:', error);
|
||||
this.showNoRootsPlaceholder(defaultCheckpointRootSelect);
|
||||
showToast('toast.settings.checkpointRootsFailed', { message: error.message }, 'error');
|
||||
}
|
||||
}
|
||||
|
||||
async loadUnetRoots() {
|
||||
try {
|
||||
const defaultUnetRootSelect = document.getElementById('defaultUnetRoot');
|
||||
if (!defaultUnetRootSelect) return;
|
||||
const defaultUnetRootSelect = document.getElementById('defaultUnetRoot');
|
||||
if (!defaultUnetRootSelect) return;
|
||||
|
||||
try {
|
||||
// Fetch unet roots (diffusion model paths only)
|
||||
const response = await fetch('/api/lm/checkpoints/unet_roots');
|
||||
if (!response.ok) {
|
||||
@@ -1600,10 +1615,12 @@ export class SettingsManager {
|
||||
|
||||
const data = await response.json();
|
||||
if (!data.roots || data.roots.length === 0) {
|
||||
throw new Error('No diffusion model roots found');
|
||||
this.showNoRootsPlaceholder(defaultUnetRootSelect);
|
||||
return;
|
||||
}
|
||||
|
||||
defaultUnetRootSelect.innerHTML = '';
|
||||
defaultUnetRootSelect.disabled = false;
|
||||
|
||||
// Add options for each root
|
||||
data.roots.forEach(root => {
|
||||
@@ -1618,15 +1635,16 @@ export class SettingsManager {
|
||||
|
||||
} catch (error) {
|
||||
console.error('Error loading diffusion model roots:', error);
|
||||
this.showNoRootsPlaceholder(defaultUnetRootSelect);
|
||||
showToast('toast.settings.unetRootsFailed', { message: error.message }, 'error');
|
||||
}
|
||||
}
|
||||
|
||||
async loadEmbeddingRoots() {
|
||||
try {
|
||||
const defaultEmbeddingRootSelect = document.getElementById('defaultEmbeddingRoot');
|
||||
if (!defaultEmbeddingRootSelect) return;
|
||||
const defaultEmbeddingRootSelect = document.getElementById('defaultEmbeddingRoot');
|
||||
if (!defaultEmbeddingRootSelect) return;
|
||||
|
||||
try {
|
||||
// Fetch embedding roots
|
||||
const response = await fetch('/api/lm/embeddings/roots');
|
||||
if (!response.ok) {
|
||||
@@ -1635,10 +1653,12 @@ export class SettingsManager {
|
||||
|
||||
const data = await response.json();
|
||||
if (!data.roots || data.roots.length === 0) {
|
||||
throw new Error('No embedding roots found');
|
||||
this.showNoRootsPlaceholder(defaultEmbeddingRootSelect);
|
||||
return;
|
||||
}
|
||||
|
||||
defaultEmbeddingRootSelect.innerHTML = '';
|
||||
defaultEmbeddingRootSelect.disabled = false;
|
||||
|
||||
// Add options for each root
|
||||
data.roots.forEach(root => {
|
||||
@@ -1653,6 +1673,7 @@ export class SettingsManager {
|
||||
|
||||
} catch (error) {
|
||||
console.error('Error loading embedding roots:', error);
|
||||
this.showNoRootsPlaceholder(defaultEmbeddingRootSelect);
|
||||
showToast('toast.settings.embeddingRootsFailed', { message: error.message }, 'error');
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,11 +1,12 @@
|
||||
import { modalManager } from './ModalManager.js';
|
||||
import {
|
||||
getStorageItem,
|
||||
setStorageItem,
|
||||
getStoredVersionInfo,
|
||||
import {
|
||||
getStorageItem,
|
||||
setStorageItem,
|
||||
getStoredVersionInfo,
|
||||
setStoredVersionInfo,
|
||||
isVersionMatch
|
||||
} from '../utils/storageHelpers.js';
|
||||
import { state } from '../state/index.js';
|
||||
import { bannerService } from './BannerService.js';
|
||||
import { translate } from '../utils/i18nHelpers.js';
|
||||
|
||||
@@ -24,7 +25,11 @@ export class UpdateService {
|
||||
this.updateNotificationsEnabled = getStorageItem('show_update_notifications', true);
|
||||
this.lastCheckTime = parseInt(getStorageItem('last_update_check') || '0');
|
||||
this.isUpdating = false;
|
||||
this.nightlyMode = getStorageItem('nightly_updates', false);
|
||||
this.channelMode = null;
|
||||
this.hasGit = false;
|
||||
this.nightlyNotifyDate = getStorageItem('nightly_notify_date', '');
|
||||
this.nightlyBadgeShown = false;
|
||||
this.progressKeepVisible = false;
|
||||
this.currentVersionInfo = null;
|
||||
this.versionMismatch = false;
|
||||
this.activeNotificationTab = 'updates';
|
||||
@@ -49,43 +54,180 @@ export class UpdateService {
|
||||
updateBtn.addEventListener('click', () => this.performUpdate());
|
||||
}
|
||||
|
||||
// Register event listener for nightly update toggle
|
||||
const nightlyCheckbox = document.getElementById('nightlyUpdateToggle');
|
||||
if (nightlyCheckbox) {
|
||||
nightlyCheckbox.checked = this.nightlyMode;
|
||||
nightlyCheckbox.addEventListener('change', (e) => {
|
||||
this.nightlyMode = e.target.checked;
|
||||
setStorageItem('nightly_updates', e.target.checked);
|
||||
this.updateNightlyWarning();
|
||||
this.updateModalContent();
|
||||
// Re-check for updates when switching channels
|
||||
this.manualCheckForUpdates();
|
||||
});
|
||||
this.updateNightlyWarning();
|
||||
}
|
||||
this.wireChannelButtons();
|
||||
|
||||
this.setupNotificationCenter();
|
||||
window.addEventListener('lm:banner-history-updated', this.handleBannerHistoryUpdated);
|
||||
this.updateTabBadges();
|
||||
|
||||
// Perform update check if needed
|
||||
this.checkForUpdates().then(() => {
|
||||
// Ensure badges are updated after checking
|
||||
this.updateBadgeVisibility();
|
||||
this.checkVersionInfo().then(() => {
|
||||
this.checkForUpdates().then(() => {
|
||||
this.updateBadgeVisibility();
|
||||
});
|
||||
});
|
||||
|
||||
// Immediately update modal content with current values (even if from default)
|
||||
this.updateModalContent();
|
||||
|
||||
// Check version info for mismatch after loading basic info
|
||||
this.checkVersionInfo();
|
||||
}
|
||||
|
||||
updateNightlyWarning() {
|
||||
const warning = document.getElementById('nightlyWarning');
|
||||
if (warning) {
|
||||
warning.style.display = this.nightlyMode ? 'flex' : 'none';
|
||||
wireChannelButtons() {
|
||||
const releaseBtn = document.getElementById('channelRelease');
|
||||
const nightlyBtn = document.getElementById('channelNightly');
|
||||
if (releaseBtn) {
|
||||
releaseBtn.addEventListener('click', () => this.switchChannel('release'));
|
||||
}
|
||||
if (nightlyBtn) {
|
||||
nightlyBtn.addEventListener('click', () => this.switchChannel('nightly'));
|
||||
}
|
||||
}
|
||||
|
||||
async switchChannel(channel) {
|
||||
if (channel === this.channelMode) {
|
||||
return;
|
||||
}
|
||||
if (this.isUpdating) {
|
||||
return;
|
||||
}
|
||||
if (!this.hasGit && channel === 'nightly') {
|
||||
const confirmed = await this._confirmChannelSwitch(
|
||||
'update.channelSwitch.nightlyTitle',
|
||||
'update.channelSwitch.nightlyMessage'
|
||||
);
|
||||
if (!confirmed) return;
|
||||
}
|
||||
if (this.hasGit && channel === 'release') {
|
||||
const confirmed = await this._confirmChannelSwitch(
|
||||
'update.channelSwitch.releaseTitle',
|
||||
'update.channelSwitch.releaseMessage'
|
||||
);
|
||||
if (!confirmed) return;
|
||||
}
|
||||
|
||||
try {
|
||||
this.isUpdating = true;
|
||||
this.showUpdateProgress(true);
|
||||
this.updateProgress(10, translate('update.channelSwitch.switching', { channel }));
|
||||
|
||||
const response = await fetch('/api/lm/switch-channel', {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ channel })
|
||||
});
|
||||
|
||||
const data = await response.json();
|
||||
|
||||
if (data.success) {
|
||||
this.channelMode = channel;
|
||||
// Persist channel preference to settings.json
|
||||
fetch('/api/lm/settings', {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ update_channel: channel })
|
||||
}).then(r => {
|
||||
if (!r.ok) console.warn('Failed to persist update channel:', r.status);
|
||||
}).catch(e => console.warn('Failed to persist update channel:', e));
|
||||
await this.checkForUpdates({ force: true });
|
||||
this.updateModalContent();
|
||||
this.updateChannelUI();
|
||||
this._showSwitchCompleteMessage(data.new_version);
|
||||
this.progressKeepVisible = true;
|
||||
} else {
|
||||
throw new Error(data.error || translate('update.channelSwitch.failed'));
|
||||
}
|
||||
} catch (error) {
|
||||
console.error('Channel switch failed:', error);
|
||||
this.updateProgress(0, translate('update.channelSwitch.failed'));
|
||||
} finally {
|
||||
if (this.progressKeepVisible) {
|
||||
this.isUpdating = false;
|
||||
this.progressKeepVisible = false;
|
||||
} else {
|
||||
setTimeout(() => {
|
||||
this.showUpdateProgress(false);
|
||||
this.isUpdating = false;
|
||||
}, 2000);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
updateChannelUI() {
|
||||
const releaseBtn = document.getElementById('channelRelease');
|
||||
const nightlyBtn = document.getElementById('channelNightly');
|
||||
|
||||
if (releaseBtn) {
|
||||
releaseBtn.classList.toggle('active', this.channelMode === 'release');
|
||||
}
|
||||
if (nightlyBtn) {
|
||||
nightlyBtn.classList.toggle('active', this.channelMode === 'nightly');
|
||||
}
|
||||
}
|
||||
|
||||
_resolveChannelFromSettings() {
|
||||
const stored = state?.global?.settings?.update_channel;
|
||||
if (stored === 'nightly' || stored === 'release') {
|
||||
return stored;
|
||||
}
|
||||
if (!this.hasGit) {
|
||||
return 'release';
|
||||
}
|
||||
if (this.gitInfo?.branch === 'detached') {
|
||||
return 'release';
|
||||
}
|
||||
return 'nightly';
|
||||
}
|
||||
|
||||
async _confirmChannelSwitch(titleKey, messageKey) {
|
||||
return new Promise((resolve) => {
|
||||
const title = translate(titleKey);
|
||||
const message = translate(messageKey);
|
||||
const cancelText = translate('common.cancel');
|
||||
const confirmText = translate('common.confirm');
|
||||
|
||||
const overlay = document.createElement('div');
|
||||
overlay.className = 'channel-switch-overlay';
|
||||
overlay.innerHTML = `
|
||||
<div class="channel-switch-dialog">
|
||||
<h3>${title}</h3>
|
||||
<p>${message}</p>
|
||||
<div class="channel-switch-actions">
|
||||
<button class="secondary-btn channel-switch-cancel">${cancelText}</button>
|
||||
<button class="primary-btn channel-switch-confirm">${confirmText}</button>
|
||||
</div>
|
||||
</div>
|
||||
`;
|
||||
|
||||
const dismiss = (result) => {
|
||||
document.removeEventListener('keydown', onKeydown);
|
||||
overlay.remove();
|
||||
resolve(result);
|
||||
};
|
||||
|
||||
const onKeydown = (e) => {
|
||||
if (e.key === 'Escape') {
|
||||
e.stopPropagation();
|
||||
e.preventDefault();
|
||||
dismiss(false);
|
||||
}
|
||||
};
|
||||
|
||||
document.addEventListener('keydown', onKeydown, { capture: true });
|
||||
|
||||
overlay.addEventListener('click', (e) => {
|
||||
if (e.target === overlay) {
|
||||
dismiss(false);
|
||||
}
|
||||
});
|
||||
|
||||
overlay.querySelector('.channel-switch-cancel').addEventListener('click', () => {
|
||||
dismiss(false);
|
||||
});
|
||||
|
||||
overlay.querySelector('.channel-switch-confirm').addEventListener('click', () => {
|
||||
dismiss(true);
|
||||
});
|
||||
|
||||
document.body.appendChild(overlay);
|
||||
});
|
||||
}
|
||||
|
||||
setupNotificationCenter() {
|
||||
@@ -355,6 +497,18 @@ export class UpdateService {
|
||||
}
|
||||
|
||||
async checkForUpdates({ force = false } = {}) {
|
||||
let needsMigration = false;
|
||||
if (this.channelMode === null) {
|
||||
const stored = state?.global?.settings?.update_channel;
|
||||
if (stored === 'nightly' || stored === 'release') {
|
||||
this.channelMode = stored;
|
||||
} else if (!this.hasGit) {
|
||||
this.channelMode = 'release';
|
||||
needsMigration = true;
|
||||
}
|
||||
// hasGit=true with no stored value: wait for gitInfo.branch
|
||||
}
|
||||
|
||||
if (!force && !this.updateNotificationsEnabled) {
|
||||
return;
|
||||
}
|
||||
@@ -373,7 +527,8 @@ export class UpdateService {
|
||||
|
||||
try {
|
||||
// Call backend API to check for updates with nightly flag
|
||||
const response = await fetch(`/api/lm/check-updates?nightly=${this.nightlyMode}`);
|
||||
const nightly = (this.channelMode ?? (this.hasGit ? 'nightly' : 'release')) === 'nightly';
|
||||
const response = await fetch(`/api/lm/check-updates?nightly=${nightly}`);
|
||||
const data = await response.json();
|
||||
|
||||
if (data.success) {
|
||||
@@ -381,17 +536,35 @@ export class UpdateService {
|
||||
this.latestVersion = data.latest_version || "v0.0.0";
|
||||
this.updateInfo = data;
|
||||
this.gitInfo = data.git_info || this.gitInfo;
|
||||
|
||||
// Explicitly set update availability based on version comparison
|
||||
this.updateAvailable = this.isNewerVersion(this.latestVersion, this.currentVersion);
|
||||
|
||||
// Update last check time
|
||||
this.hasGit = data.has_git || false;
|
||||
|
||||
if (needsMigration || this.channelMode === null) {
|
||||
this.channelMode = this._resolveChannelFromSettings();
|
||||
if (state?.global?.settings) {
|
||||
state.global.settings.update_channel = this.channelMode;
|
||||
}
|
||||
fetch('/api/lm/settings', {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ update_channel: this.channelMode })
|
||||
}).then(r => {
|
||||
if (!r.ok) console.warn('Failed to persist update channel:', r.status);
|
||||
}).catch(e => console.warn('Failed to persist update channel:', e));
|
||||
}
|
||||
|
||||
this.updateAvailable = data.update_available;
|
||||
|
||||
// Nightly channel: surface the update badge at most once per calendar day.
|
||||
if (this.updateAvailable && this.channelMode === 'nightly' && this.nightlyNotifyDate !== this._getTodayKey()) {
|
||||
this._markNightlyNotified();
|
||||
}
|
||||
|
||||
this.lastCheckTime = now;
|
||||
setStorageItem('last_update_check', now.toString());
|
||||
|
||||
// Update UI
|
||||
|
||||
this.updateBadgeVisibility();
|
||||
this.updateModalContent();
|
||||
this.updateChannelUI();
|
||||
|
||||
console.log("Update check complete:", {
|
||||
currentVersion: this.currentVersion,
|
||||
@@ -435,6 +608,28 @@ export class UpdateService {
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
_getTodayKey() {
|
||||
const now = new Date();
|
||||
const month = String(now.getMonth() + 1).padStart(2, '0');
|
||||
const day = String(now.getDate()).padStart(2, '0');
|
||||
return `${now.getFullYear()}-${month}-${day}`;
|
||||
}
|
||||
|
||||
_isNightlyBadgeAllowed() {
|
||||
if (this.channelMode !== 'nightly') {
|
||||
return true;
|
||||
}
|
||||
// Keep the badge visible for the rest of the session once shown, but do
|
||||
// not show it again on later sessions within the same calendar day.
|
||||
return this.nightlyNotifyDate !== this._getTodayKey() || this.nightlyBadgeShown;
|
||||
}
|
||||
|
||||
_markNightlyNotified() {
|
||||
this.nightlyNotifyDate = this._getTodayKey();
|
||||
this.nightlyBadgeShown = true;
|
||||
setStorageItem('nightly_notify_date', this.nightlyNotifyDate);
|
||||
}
|
||||
|
||||
updateBadgeVisibility() {
|
||||
const updateToggle = document.querySelector('.update-toggle');
|
||||
@@ -443,9 +638,12 @@ export class UpdateService {
|
||||
? bannerService.getUnreadBannerCount()
|
||||
: 0;
|
||||
|
||||
// Force updating badges visibility based on current state
|
||||
const shouldShowUpdate = this.updateNotificationsEnabled && this.updateAvailable && this._isNightlyBadgeAllowed();
|
||||
|
||||
if (updateToggle) {
|
||||
let tooltipKey = 'header.actions.notifications';
|
||||
if (this.updateNotificationsEnabled && this.updateAvailable) {
|
||||
if (shouldShowUpdate) {
|
||||
tooltipKey = 'update.updateAvailable';
|
||||
} else if (unreadBanners > 0) {
|
||||
tooltipKey = 'update.tabs.messages';
|
||||
@@ -453,8 +651,6 @@ export class UpdateService {
|
||||
updateToggle.title = translate(tooltipKey);
|
||||
}
|
||||
|
||||
// Force updating badges visibility based on current state
|
||||
const shouldShowUpdate = this.updateNotificationsEnabled && this.updateAvailable;
|
||||
const shouldShow = shouldShowUpdate || unreadBanners > 0;
|
||||
|
||||
if (updateBadge) {
|
||||
@@ -482,8 +678,31 @@ export class UpdateService {
|
||||
|
||||
if (currentVersionEl) currentVersionEl.textContent = this.currentVersion;
|
||||
|
||||
const newVersionLabel = modal.querySelector('.new-version .label');
|
||||
if (newVersionLabel) {
|
||||
newVersionLabel.textContent = (this.updateInfo?.nightly)
|
||||
? `${translate('update.latestMain')}:`
|
||||
: `${translate('update.newVersion')}:`;
|
||||
}
|
||||
|
||||
if (newVersionEl) {
|
||||
newVersionEl.textContent = this.latestVersion;
|
||||
if (this.updateInfo?.nightly) {
|
||||
const behind = this.updateInfo.behind_by || 0;
|
||||
const remoteHash = this.latestVersion.replace('main-', '');
|
||||
const localHash = this.gitInfo.short_hash || '';
|
||||
const date = this.updateInfo.commit_date || '';
|
||||
const datePart = date ? ` · ${date}` : '';
|
||||
|
||||
if (behind > 0) {
|
||||
newVersionEl.textContent = `${behind} commit${behind !== 1 ? 's' : ''} behind main (${remoteHash}${datePart})`;
|
||||
} else if (localHash !== remoteHash) {
|
||||
newVersionEl.textContent = `Behind main (${remoteHash}${datePart})`;
|
||||
} else {
|
||||
newVersionEl.textContent = `Up to date (${remoteHash}${datePart})`;
|
||||
}
|
||||
} else {
|
||||
newVersionEl.textContent = this.latestVersion;
|
||||
}
|
||||
}
|
||||
|
||||
// Update update button state
|
||||
@@ -599,8 +818,12 @@ export class UpdateService {
|
||||
// Update GitHub link to point to the specific release if available
|
||||
const githubLink = modal.querySelector('.update-link');
|
||||
if (githubLink && this.latestVersion) {
|
||||
const versionTag = this.latestVersion.replace(/^v/, '');
|
||||
githubLink.href = `https://github.com/willmiao/ComfyUI-Lora-Manager/releases/tag/v${versionTag}`;
|
||||
if (this.updateInfo?.nightly) {
|
||||
githubLink.href = 'https://github.com/willmiao/ComfyUI-Lora-Manager/commits/main';
|
||||
} else {
|
||||
const versionTag = this.latestVersion.replace(/^v/, '');
|
||||
githubLink.href = `https://github.com/willmiao/ComfyUI-Lora-Manager/releases/tag/v${versionTag}`;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -623,7 +846,7 @@ export class UpdateService {
|
||||
'Content-Type': 'application/json'
|
||||
},
|
||||
body: JSON.stringify({
|
||||
nightly: this.nightlyMode
|
||||
nightly: this.channelMode === 'nightly'
|
||||
})
|
||||
});
|
||||
|
||||
@@ -698,7 +921,26 @@ export class UpdateService {
|
||||
progressText.textContent = text;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
_showSwitchCompleteMessage(version) {
|
||||
this.showUpdateProgress(true);
|
||||
this.updateProgress(100, '');
|
||||
const progressText = document.getElementById('updateProgressText');
|
||||
if (progressText) {
|
||||
progressText.innerHTML = `
|
||||
<div style="text-align: center; color: var(--lora-success);">
|
||||
<i class="fas fa-check-circle" style="margin-right: 8px;"></i>
|
||||
${translate('update.completion.successMessage', { version })}
|
||||
<br><br>
|
||||
<div style="opacity: 0.95; color: var(--lora-error); font-size: 1em;">
|
||||
${translate('update.completion.restartMessage')}<br>
|
||||
${translate('update.completion.reloadMessage')}
|
||||
</div>
|
||||
</div>
|
||||
`;
|
||||
}
|
||||
}
|
||||
|
||||
showUpdateCompleteMessage(newVersion) {
|
||||
const modal = document.getElementById('updateModal');
|
||||
if (!modal) return;
|
||||
@@ -771,6 +1013,7 @@ export class UpdateService {
|
||||
|
||||
// Update the modal content immediately with current data
|
||||
this.updateModalContent();
|
||||
this.updateChannelUI();
|
||||
this.renderRecentBanners();
|
||||
|
||||
// Show the modal with current data
|
||||
@@ -801,8 +1044,8 @@ export class UpdateService {
|
||||
|
||||
if (data.success) {
|
||||
this.currentVersionInfo = data.version;
|
||||
|
||||
// Check if version matches stored version
|
||||
this.hasGit = data.has_git || false;
|
||||
|
||||
this.versionMismatch = !isVersionMatch(this.currentVersionInfo);
|
||||
|
||||
if (this.versionMismatch) {
|
||||
|
||||
@@ -3,6 +3,7 @@ import { translate } from '../../utils/i18nHelpers.js';
|
||||
import { getModelApiClient } from '../../api/modelApiFactory.js';
|
||||
import { MODEL_TYPES } from '../../api/apiConfig.js';
|
||||
import { getStorageItem } from '../../utils/storageHelpers.js';
|
||||
import { state } from '../../state/index.js';
|
||||
|
||||
export class DownloadManager {
|
||||
constructor(importManager) {
|
||||
@@ -125,11 +126,25 @@ export class DownloadManager {
|
||||
showToast('toast.recipes.nameSaved', { name: this.importManager.recipeName }, 'success');
|
||||
}
|
||||
|
||||
// Close modal
|
||||
modalManager.closeModal('importModal');
|
||||
|
||||
// Refresh the recipe
|
||||
window.recipeManager.loadRecipes(true);
|
||||
if (isDownloadOnly && state.virtualScroller) {
|
||||
const recipeId = this.importManager.recipeId;
|
||||
try {
|
||||
const detailRes = await fetch(`/api/lm/recipe/${encodeURIComponent(recipeId)}`);
|
||||
if (detailRes.ok) {
|
||||
const updated = await detailRes.json();
|
||||
state.virtualScroller.updateSingleItem(updated.file_path, updated);
|
||||
} else {
|
||||
throw new Error(`API returned ${detailRes.status}`);
|
||||
}
|
||||
} catch (e) {
|
||||
console.warn('Failed to update recipe card in-place, falling back to reload:', e);
|
||||
await window.recipeManager.loadRecipes({ resetPage: true, preserveScroll: true });
|
||||
}
|
||||
} else {
|
||||
window.recipeManager.loadRecipes({ resetPage: true, preserveScroll: true });
|
||||
}
|
||||
|
||||
} catch (error) {
|
||||
console.error('Error:', error);
|
||||
|
||||
@@ -333,6 +333,7 @@ export const PATH_TEMPLATE_PLACEHOLDERS = [
|
||||
export const DEFAULT_PATH_TEMPLATES = {
|
||||
lora: '{base_model}/{first_tag}',
|
||||
checkpoint: '{base_model}',
|
||||
unet: '{base_model}',
|
||||
embedding: '{first_tag}'
|
||||
};
|
||||
|
||||
@@ -369,21 +370,24 @@ export function getMatureBlurThreshold(settings = {}) {
|
||||
export const NODE_TYPES = {
|
||||
LORA_LOADER: 1,
|
||||
LORA_STACKER: 2,
|
||||
WAN_VIDEO_LORA_SELECT: 3
|
||||
WAN_VIDEO_LORA_SELECT: 3,
|
||||
HOOK_LORA: 4
|
||||
};
|
||||
|
||||
// Node type names to IDs mapping
|
||||
export const NODE_TYPE_NAMES = {
|
||||
"Lora Loader (LoraManager)": NODE_TYPES.LORA_LOADER,
|
||||
"Lora Stacker (LoraManager)": NODE_TYPES.LORA_STACKER,
|
||||
"WanVideo Lora Select (LoraManager)": NODE_TYPES.WAN_VIDEO_LORA_SELECT
|
||||
"WanVideo Lora Select (LoraManager)": NODE_TYPES.WAN_VIDEO_LORA_SELECT,
|
||||
"Create Hook LoRA (LoraManager)": NODE_TYPES.HOOK_LORA
|
||||
};
|
||||
|
||||
// Node type icons
|
||||
export const NODE_TYPE_ICONS = {
|
||||
[NODE_TYPES.LORA_LOADER]: "fas fa-l",
|
||||
[NODE_TYPES.LORA_STACKER]: "fas fa-s",
|
||||
[NODE_TYPES.WAN_VIDEO_LORA_SELECT]: "fas fa-w"
|
||||
[NODE_TYPES.WAN_VIDEO_LORA_SELECT]: "fas fa-w",
|
||||
[NODE_TYPES.HOOK_LORA]: "fas fa-h"
|
||||
};
|
||||
|
||||
// Default ComfyUI node color when bgcolor is null
|
||||
|
||||
@@ -32,7 +32,12 @@
|
||||
<div class="context-menu-separator menu-section-break"></div>
|
||||
<!-- Media / Preview -->
|
||||
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
|
||||
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }}</div>
|
||||
<div class="context-menu-item has-submenu" data-has-submenu="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }} <i class="fas fa-chevron-right submenu-arrow"></i>
|
||||
<div class="context-submenu">
|
||||
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadMissingExamples') }}</div>
|
||||
<div class="context-menu-item" data-action="download-examples-force"><i class="fas fa-redo-alt"></i> {{ t('loras.contextMenu.reprocessExamples') }}</div>
|
||||
</div>
|
||||
</div>
|
||||
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
|
||||
<div class="context-menu-separator menu-section-break"></div>
|
||||
<!-- Attributes -->
|
||||
|
||||
@@ -44,8 +44,18 @@
|
||||
<div class="context-menu-item" data-action="preview">
|
||||
<i class="fas fa-folder-open"></i> <span>{{ t('loras.contextMenu.openExamples') }}</span>
|
||||
</div>
|
||||
<div class="context-menu-item" data-action="download-examples">
|
||||
<i class="fas fa-download"></i> <span>{{ t('loras.contextMenu.downloadExamples') }}</span>
|
||||
<div class="context-menu-item has-submenu" data-has-submenu="download-examples">
|
||||
<i class="fas fa-download"></i>
|
||||
<span>{{ t('loras.contextMenu.downloadExamples') }}</span>
|
||||
<i class="fas fa-chevron-right submenu-arrow"></i>
|
||||
<div class="context-submenu">
|
||||
<div class="context-menu-item" data-action="download-examples">
|
||||
<i class="fas fa-download"></i> <span>{{ t('loras.contextMenu.downloadMissingExamples') }}</span>
|
||||
</div>
|
||||
<div class="context-menu-item" data-action="download-examples-force">
|
||||
<i class="fas fa-redo-alt"></i> <span>{{ t('loras.contextMenu.reprocessExamples') }}</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<div class="context-menu-item" data-action="replace-preview">
|
||||
<i class="fas fa-image"></i> <span>{{ t('loras.contextMenu.replacePreview') }}</span>
|
||||
@@ -136,8 +146,18 @@
|
||||
</div>
|
||||
<div class="context-menu-section" data-section="download">
|
||||
<div class="context-menu-section-header">{{ t('loras.bulkOperations.sections.download') }}</div>
|
||||
<div class="context-menu-item" data-action="download-example-images">
|
||||
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadExamples') }}</span>
|
||||
<div class="context-menu-item has-submenu" data-has-submenu="download-example-images">
|
||||
<i class="fas fa-download"></i>
|
||||
<span>{{ t('loras.bulkOperations.downloadExamples') }}</span>
|
||||
<i class="fas fa-chevron-right submenu-arrow"></i>
|
||||
<div class="context-submenu">
|
||||
<div class="context-menu-item" data-action="download-missing-example-images">
|
||||
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingExamples') }}</span>
|
||||
</div>
|
||||
<div class="context-menu-item" data-action="download-example-images">
|
||||
<i class="fas fa-redo-alt"></i> <span>{{ t('loras.bulkOperations.reprocessExamples') }}</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<div class="context-menu-item" data-action="download-missing-loras">
|
||||
<i class="fas fa-download"></i> <span>{{ t('loras.bulkOperations.downloadMissingLoras') }}</span>
|
||||
|
||||
@@ -48,6 +48,11 @@
|
||||
<option value="versions_count:asc">{{ t('loras.controls.sort.versionsCountAsc', default='Fewest versions first') }}</option>
|
||||
</optgroup>
|
||||
{% endif %}
|
||||
{% if page_id != 'recipes' %}
|
||||
<optgroup label="{{ t('loras.controls.sort.random', default='Random') }}">
|
||||
<option value="random">{{ t('loras.controls.sort.randomAction', default='Randomize (shuffle)') }}</option>
|
||||
</optgroup>
|
||||
{% endif %}
|
||||
{% if page_id == 'recipes' %}
|
||||
<optgroup label="{{ t('recipes.controls.sort.lorasCount') }}">
|
||||
<option value="loras_count:desc">{{ t('recipes.controls.sort.lorasCountDesc') }}</option>
|
||||
|
||||
@@ -19,6 +19,20 @@
|
||||
<div class="notification-panels">
|
||||
<div class="notification-panel active" id="updatesPanel" role="tabpanel" aria-labelledby="updatesTab" aria-hidden="false" tabindex="0" data-notification-panel="updates">
|
||||
<div class="update-content">
|
||||
|
||||
<!-- Channel Selector -->
|
||||
<div class="update-channels" id="updateChannels">
|
||||
<div class="channels-label">{{ t('update.channel') }}</div>
|
||||
<div class="channel-toggle">
|
||||
<button type="button" class="channel-btn" data-channel="release" id="channelRelease">
|
||||
<i class="fas fa-tag"></i> {{ t('update.channels.release') }}
|
||||
</button>
|
||||
<button type="button" class="channel-btn" data-channel="nightly" id="channelNightly">
|
||||
<i class="fas fa-moon"></i> {{ t('update.channels.nightly') }}
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="update-info">
|
||||
<div class="version-info">
|
||||
<div class="current-version">
|
||||
|
||||
@@ -32,7 +32,12 @@
|
||||
<div class="context-menu-separator menu-section-break"></div>
|
||||
<!-- Media / Preview -->
|
||||
<div class="context-menu-item" data-action="preview"><i class="fas fa-folder-open"></i> {{ t('loras.contextMenu.openExamples') }}</div>
|
||||
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }}</div>
|
||||
<div class="context-menu-item has-submenu" data-has-submenu="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadExamples') }} <i class="fas fa-chevron-right submenu-arrow"></i>
|
||||
<div class="context-submenu">
|
||||
<div class="context-menu-item" data-action="download-examples"><i class="fas fa-download"></i> {{ t('loras.contextMenu.downloadMissingExamples') }}</div>
|
||||
<div class="context-menu-item" data-action="download-examples-force"><i class="fas fa-redo-alt"></i> {{ t('loras.contextMenu.reprocessExamples') }}</div>
|
||||
</div>
|
||||
</div>
|
||||
<div class="context-menu-item" data-action="replace-preview"><i class="fas fa-image"></i> {{ t('loras.contextMenu.replacePreview') }}</div>
|
||||
<div class="context-menu-separator menu-section-break"></div>
|
||||
<!-- Attributes -->
|
||||
|
||||
@@ -85,6 +85,7 @@ sys.modules['comfy.utils'] = comfy_mock.utils
|
||||
sys.modules['comfy.sd'] = comfy_mock.sd
|
||||
sys.modules['comfy.model_management'] = comfy_mock.model_management
|
||||
sys.modules['comfy.comfy_types'] = comfy_mock.comfy_types
|
||||
sys.modules['comfy.hooks'] = MockModule("comfy.hooks")
|
||||
|
||||
execution_mock = MockModule("execution")
|
||||
execution_mock.PromptExecutor = mock.MagicMock()
|
||||
|
||||
@@ -2155,4 +2155,35 @@ describe('Interaction-level regression coverage', () => {
|
||||
excludedItem.dispatchEvent(new Event('click', { bubbles: true }));
|
||||
expect(window.pageControls.enterExcludedView).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it('routes single-model example downloads to missing-only and force paths', async () => {
|
||||
document.body.innerHTML = `
|
||||
<div id="loraContextMenu" class="context-menu">
|
||||
<div class="context-menu-item has-submenu" data-has-submenu="download-examples">
|
||||
<div class="context-submenu">
|
||||
<div class="context-menu-item" data-action="download-examples"></div>
|
||||
<div class="context-menu-item" data-action="download-examples-force"></div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
`;
|
||||
|
||||
const { LoraContextMenu } = await import('../../../static/js/components/ContextMenu/LoraContextMenu.js');
|
||||
const contextMenu = new LoraContextMenu();
|
||||
|
||||
const card = document.createElement('div');
|
||||
card.className = 'model-card';
|
||||
card.dataset.filepath = '/models/test.safetensors';
|
||||
card.dataset.sha256 = 'abc123hash';
|
||||
document.body.appendChild(card);
|
||||
|
||||
contextMenu.showMenu(100, 100, card);
|
||||
|
||||
document.querySelector('[data-action="download-examples"]').dispatchEvent(new Event('click', { bubbles: true }));
|
||||
expect(downloadExampleImagesApiMock).toHaveBeenCalledWith(['abc123hash'], null, { force: false });
|
||||
|
||||
contextMenu.showMenu(100, 100, card);
|
||||
document.querySelector('[data-action="download-examples-force"]').dispatchEvent(new Event('click', { bubbles: true }));
|
||||
expect(downloadExampleImagesApiMock).toHaveBeenCalledWith(['abc123hash'], null, { force: true });
|
||||
});
|
||||
});
|
||||
|
||||
@@ -0,0 +1,221 @@
|
||||
import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest';
|
||||
|
||||
const resetAndReloadMock = vi.fn();
|
||||
const getModelApiClientMock = vi.fn();
|
||||
|
||||
vi.mock('../../../static/js/api/modelApiFactory.js', () => ({
|
||||
getModelApiClient: getModelApiClientMock,
|
||||
resetAndReload: resetAndReloadMock,
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
|
||||
showToast: vi.fn(),
|
||||
openCivitaiByMetadata: vi.fn(),
|
||||
updatePanelPositions: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/managers/DownloadManager.js', () => ({
|
||||
downloadManager: { showDownloadModal: vi.fn() },
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/components/SidebarManager.js', () => ({
|
||||
sidebarManager: {
|
||||
setHostPageControls: vi.fn(),
|
||||
initialize: vi.fn(async () => {}),
|
||||
refresh: vi.fn(async () => {}),
|
||||
cleanup: vi.fn(),
|
||||
isInitialized: false,
|
||||
},
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/components/alphabet/index.js', () => ({
|
||||
createAlphabetBar: vi.fn(() => ({ destroy: vi.fn() })),
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/utils/updateCheckHelpers.js', () => ({
|
||||
performModelUpdateCheck: vi.fn(async () => ({ status: 'success', displayName: 'LoRA', records: [] })),
|
||||
}));
|
||||
|
||||
beforeEach(() => {
|
||||
vi.resetModules();
|
||||
vi.clearAllMocks();
|
||||
localStorage.clear();
|
||||
sessionStorage.clear();
|
||||
|
||||
resetAndReloadMock.mockResolvedValue(undefined);
|
||||
getModelApiClientMock.mockReturnValue({});
|
||||
|
||||
global.fetch = vi.fn().mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => ({ success: true, base_models: [] }),
|
||||
});
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
delete window.bulkManager;
|
||||
delete window.modelDuplicatesManager;
|
||||
delete global.fetch;
|
||||
});
|
||||
|
||||
function renderControlsDom(pageKey) {
|
||||
document.body.dataset.page = pageKey;
|
||||
document.body.innerHTML = `
|
||||
<div class="controls">
|
||||
<div id="excludedViewBanner" class="excluded-view-banner hidden">
|
||||
<button id="excludedViewBackBtn">Back</button>
|
||||
</div>
|
||||
<div class="actions">
|
||||
<div class="action-buttons">
|
||||
<div class="control-group">
|
||||
<select id="sortSelect">
|
||||
<option value="name:asc">Name Asc</option>
|
||||
<option value="name:desc">Name Desc</option>
|
||||
<option value="random">Randomize (shuffle)</option>
|
||||
</select>
|
||||
</div>
|
||||
<div class="control-group dropdown-group">
|
||||
<button data-action="refresh" class="dropdown-main"></button>
|
||||
<button class="dropdown-toggle"></button>
|
||||
<div class="dropdown-menu">
|
||||
<div class="dropdown-item" data-action="full-rebuild"></div>
|
||||
</div>
|
||||
</div>
|
||||
<div class="control-group">
|
||||
<button data-action="fetch"></button>
|
||||
</div>
|
||||
<div class="control-group">
|
||||
<button data-action="download"></button>
|
||||
</div>
|
||||
<div class="control-group">
|
||||
<button data-action="bulk"></button>
|
||||
</div>
|
||||
<div class="control-group">
|
||||
<button data-action="find-duplicates"></button>
|
||||
</div>
|
||||
<div class="control-group">
|
||||
<button id="favoriteFilterBtn" class="favorite-filter"></button>
|
||||
</div>
|
||||
<div class="control-group dropdown-group update-filter-group">
|
||||
<button id="updateFilterBtn" class="dropdown-main update-filter" aria-busy="false">
|
||||
<span>Updates</span>
|
||||
</button>
|
||||
<button id="updateFilterMenuToggle" class="dropdown-toggle"></button>
|
||||
<div class="dropdown-menu">
|
||||
<div id="checkUpdatesMenuItem" class="dropdown-item" data-action="check-updates">
|
||||
<span>Check updates</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<div id="customFilterIndicator" class="control-group hidden">
|
||||
<div class="filter-active">
|
||||
<span class="customFilterText" title=""></span>
|
||||
<i class="fas fa-times-circle clear-filter"></i>
|
||||
</div>
|
||||
</div>
|
||||
<div id="breadcrumbContainer"></div>
|
||||
<div id="duplicatesBanner" style="display: none;"></div>
|
||||
<div class="alphabet-bar-container"></div>
|
||||
`;
|
||||
}
|
||||
|
||||
async function createControls() {
|
||||
const stateModule = await import('../../../static/js/state/index.js');
|
||||
stateModule.initPageState('loras');
|
||||
const { LorasControls } = await import('../../../static/js/components/controls/LorasControls.js');
|
||||
return { stateModule, controls: new LorasControls() };
|
||||
}
|
||||
|
||||
describe('Random sort option', () => {
|
||||
it('generates a seeded sort value when Random is picked', async () => {
|
||||
renderControlsDom('loras');
|
||||
const { controls } = await createControls();
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||
|
||||
sortSelect.value = 'random';
|
||||
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||
await Promise.resolve();
|
||||
|
||||
expect(controls.pageState.sortBy).toMatch(/^random:[a-z0-9]+$/);
|
||||
expect(localStorage.getItem('lora_manager_loras_sort')).toBe(controls.pageState.sortBy);
|
||||
expect(randomOpt.value).toBe(controls.pageState.sortBy);
|
||||
expect(sortSelect.value).toBe(controls.pageState.sortBy);
|
||||
expect(resetAndReloadMock).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it('reshuffles with a fresh seed every time Random is picked again', async () => {
|
||||
renderControlsDom('loras');
|
||||
const { controls } = await createControls();
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||
|
||||
// First pick
|
||||
sortSelect.value = 'random';
|
||||
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||
await Promise.resolve();
|
||||
const firstSeed = controls.pageState.sortBy;
|
||||
|
||||
// Second pick: the option now carries the seeded value, like a menu click
|
||||
sortSelect.value = randomOpt.value;
|
||||
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||
await Promise.resolve();
|
||||
|
||||
expect(controls.pageState.sortBy).toMatch(/^random:[a-z0-9]+$/);
|
||||
expect(controls.pageState.sortBy).not.toBe(firstSeed);
|
||||
});
|
||||
|
||||
it('restores a persisted seeded random sort on load', async () => {
|
||||
renderControlsDom('loras');
|
||||
const savedSort = 'random:persistedseed';
|
||||
localStorage.setItem('lora_manager_loras_sort', savedSort);
|
||||
|
||||
const { controls } = await createControls();
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
|
||||
expect(controls.pageState.sortBy).toBe(savedSort);
|
||||
expect(sortSelect.value).toBe(savedSort);
|
||||
expect(sortSelect.querySelector('option[value="random:persistedseed"]')).not.toBeNull();
|
||||
});
|
||||
|
||||
it('applies a non-random sort back to the plain random option', async () => {
|
||||
renderControlsDom('loras');
|
||||
const { controls } = await createControls();
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||
|
||||
// Seed a random sort, then switch to a normal sort
|
||||
sortSelect.value = 'random';
|
||||
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||
await Promise.resolve();
|
||||
controls.applySortToSelect('name:desc');
|
||||
|
||||
expect(sortSelect.value).toBe('name:desc');
|
||||
expect(randomOpt.value).toBe('random');
|
||||
});
|
||||
|
||||
it('resets the seeded option when switching away from Random via the dropdown change handler', async () => {
|
||||
renderControlsDom('loras');
|
||||
const { controls } = await createControls();
|
||||
const sortSelect = document.getElementById('sortSelect');
|
||||
const randomOpt = sortSelect.querySelector('option[value="random"]');
|
||||
|
||||
// Pick Random: the option is now seeded
|
||||
sortSelect.value = 'random';
|
||||
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||
await Promise.resolve();
|
||||
expect(randomOpt.value).toMatch(/^random:[a-z0-9]+$/);
|
||||
|
||||
// Switch to a non-random sort through the change handler (as a menu
|
||||
// click does); the option must go back to the plain "random" value
|
||||
sortSelect.value = 'name:desc';
|
||||
sortSelect.dispatchEvent(new Event('change', { bubbles: true }));
|
||||
await Promise.resolve();
|
||||
|
||||
expect(controls.pageState.sortBy).toBe('name:desc');
|
||||
expect(sortSelect.value).toBe('name:desc');
|
||||
expect(randomOpt.value).toBe('random');
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,68 @@
|
||||
import { describe, it, beforeEach, expect } from 'vitest';
|
||||
import { initSortDropdown } from '../../../static/js/components/controls/SortDropdown.js';
|
||||
|
||||
function renderSortDropdownDom() {
|
||||
document.body.innerHTML = `
|
||||
<div class="sort-dropdown-group">
|
||||
<select id="sortSelect">
|
||||
<option value="name:asc">Name Asc</option>
|
||||
<option value="name:desc">Name Desc</option>
|
||||
<option value="random" selected>Randomize (shuffle)</option>
|
||||
</select>
|
||||
<button class="sort-trigger" type="button">
|
||||
<span class="sort-trigger__label"></span>
|
||||
</button>
|
||||
<div class="sort-dropdown-menu"></div>
|
||||
</div>
|
||||
`;
|
||||
return {
|
||||
select: document.getElementById('sortSelect'),
|
||||
menu: document.querySelector('.sort-dropdown-menu'),
|
||||
label: document.querySelector('.sort-trigger__label'),
|
||||
};
|
||||
}
|
||||
|
||||
describe('SortDropdown menu sync', () => {
|
||||
let select;
|
||||
let menu;
|
||||
let label;
|
||||
|
||||
beforeEach(() => {
|
||||
({ select, menu, label } = renderSortDropdownDom());
|
||||
initSortDropdown(select);
|
||||
});
|
||||
|
||||
it('rebuilds the menu and highlights the selected item when an option value attribute changes', async () => {
|
||||
// The seeded Random option gets a new value each time it is picked.
|
||||
// The select's value getter follows the selected option's new value.
|
||||
const randomOpt = select.querySelector('option[value="random"]');
|
||||
randomOpt.value = 'random:abc123';
|
||||
await Promise.resolve();
|
||||
|
||||
const items = [...menu.querySelectorAll('.sort-option')];
|
||||
expect(items.map((el) => el.dataset.value)).toContain('random:abc123');
|
||||
const seededItem = items.find((el) => el.dataset.value === 'random:abc123');
|
||||
expect(seededItem.classList.contains('is-selected')).toBe(true);
|
||||
expect(label.textContent).toBe('Randomize (shuffle)');
|
||||
});
|
||||
|
||||
it('drops the stale seeded item and re-selects the plain random item when the option is reset', async () => {
|
||||
const randomOpt = select.querySelector('option[value="random"]');
|
||||
randomOpt.value = 'random:abc123';
|
||||
await Promise.resolve();
|
||||
|
||||
// The rebuild must have happened: the seeded item is in the menu
|
||||
const seededItems = [...menu.querySelectorAll('.sort-option')]
|
||||
.filter((el) => el.dataset.value === 'random:abc123');
|
||||
expect(seededItems).toHaveLength(1);
|
||||
|
||||
// PageControls resets the option to "random" when switching away
|
||||
randomOpt.value = 'random';
|
||||
await Promise.resolve();
|
||||
|
||||
const items = [...menu.querySelectorAll('.sort-option')];
|
||||
expect(items.map((el) => el.dataset.value)).not.toContain('random:abc123');
|
||||
const randomItem = items.find((el) => el.dataset.value === 'random');
|
||||
expect(randomItem.classList.contains('is-selected')).toBe(true);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,195 @@
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
|
||||
const { APP_MODULE, EXTENSION_MODULE, appMock, registeredExtensions } =
|
||||
vi.hoisted(() => {
|
||||
const registeredExtensions = [];
|
||||
const appMock = {
|
||||
configuringGraph: false,
|
||||
registerExtension: (ext) => registeredExtensions.push(ext),
|
||||
};
|
||||
return {
|
||||
APP_MODULE: new URL("../../../scripts/app.js", import.meta.url).pathname,
|
||||
EXTENSION_MODULE: new URL(
|
||||
"../../../web/comfyui/lora_stack_dynamic_inputs.js",
|
||||
import.meta.url
|
||||
).pathname,
|
||||
appMock,
|
||||
registeredExtensions,
|
||||
};
|
||||
});
|
||||
|
||||
vi.mock(APP_MODULE, () => ({
|
||||
app: appMock,
|
||||
}));
|
||||
|
||||
describe("Lora Stack Combiner dynamic inputs", () => {
|
||||
let extension;
|
||||
|
||||
beforeEach(async () => {
|
||||
vi.resetModules();
|
||||
registeredExtensions.length = 0;
|
||||
appMock.configuringGraph = false;
|
||||
await import(EXTENSION_MODULE);
|
||||
extension = registeredExtensions.find(
|
||||
(ext) => ext.name === "Comfy.LoraManager.LoraStackCombiner"
|
||||
);
|
||||
expect(extension).toBeDefined();
|
||||
});
|
||||
|
||||
function createNodeType() {
|
||||
const nodeType = { prototype: {} };
|
||||
extension.beforeRegisterNodeDef(
|
||||
nodeType,
|
||||
{ name: "Lora Stack Combiner (LoraManager)" },
|
||||
appMock
|
||||
);
|
||||
return nodeType;
|
||||
}
|
||||
|
||||
function createNode(inputs = []) {
|
||||
const node = {
|
||||
comfyClass: "Lora Stack Combiner (LoraManager)",
|
||||
inputs: inputs.map((name) => ({ name, type: "LORA_STACK" })),
|
||||
addInput: vi.fn(function (name, type, opts) {
|
||||
this.inputs.push({ name, type, ...opts });
|
||||
}),
|
||||
removeInput: vi.fn(function (index) {
|
||||
this.inputs.splice(index, 1);
|
||||
}),
|
||||
};
|
||||
return node;
|
||||
}
|
||||
|
||||
function makeLinkInfo() {
|
||||
return { id: 999, origin_id: 1, target_id: 2 };
|
||||
}
|
||||
|
||||
it("adds a third input when the last slot gets connected", () => {
|
||||
const nodeType = createNodeType();
|
||||
const node = createNode(["lora_stack1", "lora_stack2"]);
|
||||
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||
|
||||
node.onConnectionsChange(1, 1, true, makeLinkInfo());
|
||||
|
||||
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||
"lora_stack1",
|
||||
"lora_stack2",
|
||||
"lora_stack3",
|
||||
]);
|
||||
});
|
||||
|
||||
it("does not add an input when a non-last slot gets connected", () => {
|
||||
const nodeType = createNodeType();
|
||||
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
|
||||
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||
|
||||
node.onConnectionsChange(1, 0, true, makeLinkInfo());
|
||||
|
||||
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||
"lora_stack1",
|
||||
"lora_stack2",
|
||||
"lora_stack3",
|
||||
]);
|
||||
});
|
||||
|
||||
it("removes a disconnected middle slot and renumbers", () => {
|
||||
// Simulates a real LiteGraph disconnect event: it fires only for slots that
|
||||
// had a link, and input.link has already been cleared before the event fires.
|
||||
const nodeType = createNodeType();
|
||||
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
|
||||
node.inputs[0].link = 11;
|
||||
node.inputs[1].link = null; // slot 2 was just disconnected
|
||||
node.inputs[2].link = 13;
|
||||
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||
|
||||
node.onConnectionsChange(1, 1, false, makeLinkInfo());
|
||||
|
||||
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||
"lora_stack1",
|
||||
"lora_stack2",
|
||||
]);
|
||||
});
|
||||
|
||||
it("keeps the last slot when it is disconnected", () => {
|
||||
const nodeType = createNodeType();
|
||||
const node = createNode(["lora_stack1", "lora_stack2", "lora_stack3"]);
|
||||
node.inputs[0].link = 11;
|
||||
node.inputs[1].link = 12;
|
||||
node.inputs[2].link = null; // last slot was just disconnected
|
||||
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||
|
||||
node.onConnectionsChange(1, 2, false, makeLinkInfo());
|
||||
|
||||
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||
"lora_stack1",
|
||||
"lora_stack2",
|
||||
"lora_stack3",
|
||||
]);
|
||||
expect(node.removeInput).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("keeps at least two inputs when disconnecting", () => {
|
||||
const nodeType = createNodeType();
|
||||
const node = createNode(["lora_stack1", "lora_stack2"]);
|
||||
node.inputs[0].link = 11;
|
||||
node.inputs[1].link = null; // slot 2 was just disconnected
|
||||
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||
|
||||
node.onConnectionsChange(1, 1, false, makeLinkInfo());
|
||||
|
||||
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||
"lora_stack1",
|
||||
"lora_stack2",
|
||||
]);
|
||||
expect(node.removeInput).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("does nothing while the graph is being configured", () => {
|
||||
appMock.configuringGraph = true;
|
||||
const nodeType = createNodeType();
|
||||
const node = createNode(["lora_stack1", "lora_stack2"]);
|
||||
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||
|
||||
node.onConnectionsChange(1, 1, true, makeLinkInfo());
|
||||
|
||||
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||
"lora_stack1",
|
||||
"lora_stack2",
|
||||
]);
|
||||
expect(node.addInput).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("leaves legacy lora_stack_a/b inputs untouched", () => {
|
||||
const nodeType = createNodeType();
|
||||
const node = createNode(["lora_stack_a", "lora_stack_b"]);
|
||||
node.onConnectionsChange = nodeType.prototype.onConnectionsChange;
|
||||
|
||||
node.onConnectionsChange(1, 0, true, makeLinkInfo());
|
||||
|
||||
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||
"lora_stack_a",
|
||||
"lora_stack_b",
|
||||
]);
|
||||
expect(node.addInput).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("ensures two numbered inputs exist on creation", () => {
|
||||
const node = createNode([]);
|
||||
extension.nodeCreated(node, {});
|
||||
|
||||
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||
"lora_stack1",
|
||||
"lora_stack2",
|
||||
]);
|
||||
});
|
||||
|
||||
it("does not add numbered inputs to legacy workflows", () => {
|
||||
const node = createNode(["lora_stack_a", "lora_stack_b"]);
|
||||
extension.nodeCreated(node, {});
|
||||
|
||||
expect(node.inputs.map((input) => input.name)).toEqual([
|
||||
"lora_stack_a",
|
||||
"lora_stack_b",
|
||||
]);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,133 @@
|
||||
import { describe, it, beforeEach, expect, vi } from 'vitest';
|
||||
|
||||
const showToastMock = vi.fn();
|
||||
const translateMock = vi.fn((key, params, fallback) => (typeof fallback === 'string' ? fallback : key));
|
||||
const getNSFWLevelNameMock = vi.fn((level) => {
|
||||
if (level >= 16) return 'XXX';
|
||||
if (level >= 8) return 'X';
|
||||
if (level >= 4) return 'R';
|
||||
if (level >= 2) return 'PG13';
|
||||
if (level >= 1) return 'PG';
|
||||
return 'Unknown';
|
||||
});
|
||||
|
||||
const loadingManagerStub = {
|
||||
showSimpleLoading: vi.fn(),
|
||||
showCancelButton: vi.fn(),
|
||||
hide: vi.fn(),
|
||||
};
|
||||
|
||||
const stateStub = {
|
||||
currentPageType: 'recipes',
|
||||
bulkMode: false,
|
||||
selectedModels: new Set(),
|
||||
loadingManager: loadingManagerStub,
|
||||
virtualScroller: { updateSingleItem: vi.fn() },
|
||||
global: { settings: {} },
|
||||
};
|
||||
|
||||
const saveModelMetadataMock = vi.fn();
|
||||
const getModelApiClientMock = vi.fn(() => ({ saveModelMetadata: saveModelMetadataMock }));
|
||||
const updateRecipeMetadataMock = vi.fn(() => Promise.resolve({ success: true }));
|
||||
|
||||
vi.mock('../../../static/js/state/index.js', () => ({
|
||||
state: stateStub,
|
||||
getCurrentPageState: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/utils/uiHelpers.js', () => ({
|
||||
showToast: showToastMock,
|
||||
copyToClipboard: vi.fn(),
|
||||
sendLoraToWorkflow: vi.fn(),
|
||||
sendEmbeddingToWorkflow: vi.fn(),
|
||||
buildLoraSyntax: vi.fn(),
|
||||
getNSFWLevelName: getNSFWLevelNameMock,
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/api/modelApiFactory.js', () => ({
|
||||
getModelApiClient: getModelApiClientMock,
|
||||
resetAndReload: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/api/recipeApi.js', () => ({
|
||||
RecipeSidebarApiClient: class {},
|
||||
updateRecipeMetadata: updateRecipeMetadataMock,
|
||||
extractRecipeId: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/api/apiConfig.js', () => ({
|
||||
MODEL_TYPES: { LORA: 'loras', CHECKPOINT: 'checkpoints', EMBEDDING: 'embeddings' },
|
||||
MODEL_CONFIG: {},
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/managers/ModalManager.js', () => ({
|
||||
modalManager: { showModal: vi.fn(), closeModal: vi.fn() },
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/components/shared/ModelCard.js', () => ({
|
||||
updateCardsForBulkMode: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/utils/i18nHelpers.js', () => ({
|
||||
translate: translateMock,
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/utils/priorityTagHelpers.js', () => ({
|
||||
getPriorityTagSuggestions: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock('../../../static/js/components/shared/NsfwLevelSelector.js', () => ({
|
||||
getNsfwLevelSelector: vi.fn(),
|
||||
}));
|
||||
|
||||
describe('BulkManager bulk content rating', () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
stateStub.currentPageType = 'recipes';
|
||||
stateStub.bulkMode = false;
|
||||
stateStub.selectedModels.clear();
|
||||
saveModelMetadataMock.mockResolvedValue(undefined);
|
||||
updateRecipeMetadataMock.mockResolvedValue({ success: true });
|
||||
});
|
||||
|
||||
async function createBulkManager() {
|
||||
const { BulkManager } = await import('../../../static/js/managers/BulkManager.js');
|
||||
return new BulkManager();
|
||||
}
|
||||
|
||||
it('exposes the content rating action on the recipes page action config', async () => {
|
||||
const bulk = await createBulkManager();
|
||||
expect(bulk.actionConfig.recipes.setContentRating).toBe(true);
|
||||
});
|
||||
|
||||
it('persists the rating through the recipe API when on the recipes page', async () => {
|
||||
const bulk = await createBulkManager();
|
||||
stateStub.currentPageType = 'recipes';
|
||||
stateStub.selectedModels.add('/recipes/test.webp');
|
||||
|
||||
const ok = await bulk.setBulkContentRating(4, ['/recipes/test.webp']);
|
||||
|
||||
expect(ok).toBe(true);
|
||||
expect(updateRecipeMetadataMock).toHaveBeenCalledWith('/recipes/test.webp', { preview_nsfw_level: 4 });
|
||||
expect(updateRecipeMetadataMock).toHaveBeenCalledTimes(1);
|
||||
expect(saveModelMetadataMock).not.toHaveBeenCalled();
|
||||
expect(showToastMock).toHaveBeenCalledWith(
|
||||
'toast.models.bulkContentRatingSet',
|
||||
{ count: 1, level: 'R' },
|
||||
'success'
|
||||
);
|
||||
});
|
||||
|
||||
it('persists the rating through the model API on model pages', async () => {
|
||||
const bulk = await createBulkManager();
|
||||
stateStub.currentPageType = 'loras';
|
||||
stateStub.selectedModels.add('/models/test.safetensors');
|
||||
|
||||
const ok = await bulk.setBulkContentRating(8, ['/models/test.safetensors']);
|
||||
|
||||
expect(ok).toBe(true);
|
||||
expect(saveModelMetadataMock).toHaveBeenCalledWith('/models/test.safetensors', { preview_nsfw_level: 8 });
|
||||
expect(saveModelMetadataMock).toHaveBeenCalledTimes(1);
|
||||
expect(updateRecipeMetadataMock).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,186 @@
|
||||
import { describe, it, beforeEach, afterEach, expect, vi } from 'vitest';
|
||||
import { state } from '../../../static/js/state/index.js';
|
||||
import { MODEL_TYPES } from '../../../static/js/api/apiConfig.js';
|
||||
import { eventManager } from '../../../static/js/utils/EventManager.js';
|
||||
import { BulkManager } from '../../../static/js/managers/BulkManager.js';
|
||||
|
||||
function fire(type, init = {}) {
|
||||
return new MouseEvent(type, { bubbles: true, cancelable: true, ...init });
|
||||
}
|
||||
|
||||
describe('BulkManager marquee guards', () => {
|
||||
beforeEach(() => {
|
||||
vi.useFakeTimers();
|
||||
// jsdom may not provide requestAnimationFrame; stub it so the auto-scroll loop is a no-op.
|
||||
window.requestAnimationFrame = vi.fn();
|
||||
window.cancelAnimationFrame = vi.fn();
|
||||
|
||||
eventManager.cleanup();
|
||||
state.currentPageType = MODEL_TYPES.LORA;
|
||||
state.bulkMode = false;
|
||||
state.selectedModels.clear();
|
||||
|
||||
document.body.innerHTML = '<div class="page-content"></div>';
|
||||
const pageContent = document.querySelector('.page-content');
|
||||
pageContent.getBoundingClientRect = () => ({
|
||||
top: 0,
|
||||
left: 0,
|
||||
right: 1000,
|
||||
bottom: 1000,
|
||||
width: 1000,
|
||||
height: 1000,
|
||||
x: 0,
|
||||
y: 0,
|
||||
toJSON: () => ({}),
|
||||
});
|
||||
pageContent.scrollBy = vi.fn();
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
eventManager.cleanup();
|
||||
vi.useRealTimers();
|
||||
document.body.innerHTML = '';
|
||||
});
|
||||
|
||||
function createBulkManager() {
|
||||
const bulk = new BulkManager();
|
||||
bulk.initialize();
|
||||
return bulk;
|
||||
}
|
||||
|
||||
it('never starts a marquee when the left button is not held', () => {
|
||||
const bulk = createBulkManager();
|
||||
const pageContent = document.querySelector('.page-content');
|
||||
|
||||
pageContent.dispatchEvent(fire('mousedown', { button: 0, clientX: 10, clientY: 10 }));
|
||||
document.dispatchEvent(fire('mousemove', { buttons: 0, clientX: 50, clientY: 50 }));
|
||||
|
||||
expect(bulk.mouseDownTime).toBe(0);
|
||||
expect(bulk.isMarqueeActive).toBe(false);
|
||||
expect(state.bulkMode).toBe(false);
|
||||
expect(document.querySelector('.marquee-selection')).toBeNull();
|
||||
});
|
||||
|
||||
it('requires holding the left button for the drag delay before starting a marquee', () => {
|
||||
const bulk = createBulkManager();
|
||||
const pageContent = document.querySelector('.page-content');
|
||||
|
||||
pageContent.dispatchEvent(fire('mousedown', { button: 0, clientX: 10, clientY: 10 }));
|
||||
|
||||
// Fast movement: far enough, but too soon after mousedown.
|
||||
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 30, clientY: 10 }));
|
||||
expect(state.bulkMode).toBe(false);
|
||||
expect(bulk.isMarqueeActive).toBe(false);
|
||||
|
||||
// Once the hold time has elapsed, the same drag qualifies.
|
||||
vi.advanceTimersByTime(100);
|
||||
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 35, clientY: 12 }));
|
||||
expect(state.bulkMode).toBe(true);
|
||||
expect(bulk.isMarqueeActive).toBe(true);
|
||||
expect(document.querySelector('.marquee-selection')).not.toBeNull();
|
||||
});
|
||||
|
||||
it('ends an active marquee if the left button is released without a mouseup event', () => {
|
||||
const bulk = createBulkManager();
|
||||
bulk.mouseDownPosition = { x: 10, y: 10 };
|
||||
bulk.startMarqueeSelection({}, true);
|
||||
expect(state.bulkMode).toBe(true);
|
||||
expect(document.querySelector('.marquee-selection')).not.toBeNull();
|
||||
|
||||
// No mouseup was dispatched; a plain move with the button released finalizes it.
|
||||
document.dispatchEvent(fire('mousemove', { buttons: 0, clientX: 50, clientY: 50 }));
|
||||
|
||||
expect(bulk.isMarqueeActive).toBe(false);
|
||||
expect(document.querySelector('.marquee-selection')).toBeNull();
|
||||
expect(state.bulkMode).toBe(false); // zero selected -> auto-exit
|
||||
});
|
||||
|
||||
it('treats a tiny marquee as an accidental click: clears selection and exits bulk mode', () => {
|
||||
const bulk = createBulkManager();
|
||||
const card = document.createElement('div');
|
||||
card.className = 'model-card selected';
|
||||
card.dataset.filepath = '/models/test.safetensors';
|
||||
document.body.appendChild(card);
|
||||
state.selectedModels.add('/models/test.safetensors');
|
||||
|
||||
bulk.mouseDownPosition = { x: 100, y: 100 };
|
||||
bulk.startMarqueeSelection({}, true);
|
||||
expect(state.bulkMode).toBe(true);
|
||||
|
||||
bulk.endMarqueeSelection({ clientX: 103, clientY: 104 });
|
||||
|
||||
expect(state.bulkMode).toBe(false);
|
||||
expect(state.selectedModels.size).toBe(0);
|
||||
expect(card.classList.contains('selected')).toBe(false);
|
||||
});
|
||||
|
||||
it('keeps selection and bulk mode when the marquee is large enough', () => {
|
||||
const bulk = createBulkManager();
|
||||
const card = document.createElement('div');
|
||||
card.className = 'model-card selected';
|
||||
card.dataset.filepath = '/models/test.safetensors';
|
||||
document.body.appendChild(card);
|
||||
state.selectedModels.add('/models/test.safetensors');
|
||||
|
||||
bulk.mouseDownPosition = { x: 100, y: 100 };
|
||||
bulk.startMarqueeSelection({}, true);
|
||||
|
||||
bulk.endMarqueeSelection({ clientX: 130, clientY: 140 });
|
||||
|
||||
expect(state.bulkMode).toBe(true);
|
||||
expect(state.selectedModels.has('/models/test.safetensors')).toBe(true);
|
||||
expect(card.classList.contains('selected')).toBe(true);
|
||||
});
|
||||
|
||||
it('keeps auto-scroll marquee selections when the pointer only moved a few pixels', () => {
|
||||
const bulk = createBulkManager();
|
||||
const pageContent = document.querySelector('.page-content');
|
||||
|
||||
// Card just below the press point in document coordinates.
|
||||
const card = document.createElement('div');
|
||||
card.className = 'model-card';
|
||||
card.dataset.filepath = '/models/off-screen.safetensors';
|
||||
card.getBoundingClientRect = () => ({
|
||||
top: 950,
|
||||
left: 400,
|
||||
right: 600,
|
||||
bottom: 1050,
|
||||
width: 200,
|
||||
height: 100,
|
||||
x: 400,
|
||||
y: 950,
|
||||
toJSON: () => ({}),
|
||||
});
|
||||
document.body.appendChild(card);
|
||||
|
||||
pageContent.dispatchEvent(fire('mousedown', { button: 0, clientX: 500, clientY: 900 }));
|
||||
vi.advanceTimersByTime(100);
|
||||
|
||||
// Small pointer move: enough to start the marquee, but under minMarqueeSize.
|
||||
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 506, clientY: 906 }));
|
||||
expect(bulk.isMarqueeActive).toBe(true);
|
||||
|
||||
// Auto-scroll grows the document-space box while the pointer stays nearly still.
|
||||
pageContent.scrollTop = 200;
|
||||
card.getBoundingClientRect = () => ({
|
||||
top: 750,
|
||||
left: 400,
|
||||
right: 600,
|
||||
bottom: 850,
|
||||
width: 200,
|
||||
height: 100,
|
||||
x: 400,
|
||||
y: 750,
|
||||
toJSON: () => ({}),
|
||||
});
|
||||
document.dispatchEvent(fire('mousemove', { buttons: 1, clientX: 506, clientY: 906 }));
|
||||
|
||||
expect(state.selectedModels.has('/models/off-screen.safetensors')).toBe(true);
|
||||
|
||||
// Release: the client-space box is tiny, but the document-space box is not.
|
||||
document.dispatchEvent(fire('mouseup', { button: 0, clientX: 506, clientY: 906 }));
|
||||
|
||||
expect(state.selectedModels.has('/models/off-screen.safetensors')).toBe(true);
|
||||
expect(state.bulkMode).toBe(true);
|
||||
});
|
||||
});
|
||||
@@ -106,6 +106,118 @@ afterEach(() => {
|
||||
});
|
||||
});
|
||||
|
||||
describe('SettingsManager root selects', () => {
|
||||
const rootCases = [
|
||||
{
|
||||
method: 'loadLoraRoots',
|
||||
selectId: 'defaultLoraRoot',
|
||||
endpoint: '/api/lm/loras/roots',
|
||||
errorKey: 'toast.settings.loraRootsFailed',
|
||||
},
|
||||
{
|
||||
method: 'loadCheckpointRoots',
|
||||
selectId: 'defaultCheckpointRoot',
|
||||
endpoint: '/api/lm/checkpoints/checkpoints_roots',
|
||||
errorKey: 'toast.settings.checkpointRootsFailed',
|
||||
},
|
||||
{
|
||||
method: 'loadUnetRoots',
|
||||
selectId: 'defaultUnetRoot',
|
||||
endpoint: '/api/lm/checkpoints/unet_roots',
|
||||
errorKey: 'toast.settings.unetRootsFailed',
|
||||
},
|
||||
{
|
||||
method: 'loadEmbeddingRoots',
|
||||
selectId: 'defaultEmbeddingRoot',
|
||||
endpoint: '/api/lm/embeddings/roots',
|
||||
errorKey: 'toast.settings.embeddingRootsFailed',
|
||||
},
|
||||
];
|
||||
|
||||
const appendRootSelect = (id) => {
|
||||
const select = document.createElement('select');
|
||||
select.id = id;
|
||||
document.body.appendChild(select);
|
||||
return select;
|
||||
};
|
||||
|
||||
it.each(rootCases)(
|
||||
'populates the $method select with roots and keeps it enabled',
|
||||
async ({ method, selectId, endpoint }) => {
|
||||
const manager = createManager();
|
||||
const select = appendRootSelect(selectId);
|
||||
select.disabled = true;
|
||||
|
||||
global.fetch = vi.fn().mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => ({
|
||||
success: true,
|
||||
roots: ['/models/root-a', '/models/root-b'],
|
||||
}),
|
||||
});
|
||||
|
||||
await manager[method]();
|
||||
|
||||
expect(global.fetch).toHaveBeenCalledWith(endpoint);
|
||||
expect(Array.from(select.options).map(option => option.value)).toEqual([
|
||||
'/models/root-a',
|
||||
'/models/root-b',
|
||||
]);
|
||||
expect(select.disabled).toBe(false);
|
||||
expect(showToast).not.toHaveBeenCalled();
|
||||
}
|
||||
);
|
||||
|
||||
it.each(rootCases)(
|
||||
'shows a placeholder and no error toast when $method has empty roots',
|
||||
async ({ method, selectId, endpoint }) => {
|
||||
const manager = createManager();
|
||||
const select = appendRootSelect(selectId);
|
||||
|
||||
global.fetch = vi.fn().mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => ({
|
||||
success: true,
|
||||
roots: [],
|
||||
}),
|
||||
});
|
||||
|
||||
await manager[method]();
|
||||
|
||||
expect(global.fetch).toHaveBeenCalledWith(endpoint);
|
||||
expect(select.options).toHaveLength(1);
|
||||
expect(select.options[0].value).toBe('');
|
||||
expect(select.options[0].textContent).toBe('No Default');
|
||||
expect(select.disabled).toBe(true);
|
||||
expect(showToast).not.toHaveBeenCalled();
|
||||
}
|
||||
);
|
||||
|
||||
it.each(rootCases)(
|
||||
'shows an error toast when the $method roots request fails',
|
||||
async ({ method, selectId, errorKey }) => {
|
||||
const manager = createManager();
|
||||
const select = appendRootSelect(selectId);
|
||||
|
||||
global.fetch = vi.fn().mockResolvedValue({
|
||||
ok: false,
|
||||
status: 500,
|
||||
});
|
||||
|
||||
await manager[method]();
|
||||
|
||||
expect(select.options).toHaveLength(1);
|
||||
expect(select.options[0].value).toBe('');
|
||||
expect(select.disabled).toBe(true);
|
||||
expect(showToast).toHaveBeenCalledWith(
|
||||
errorKey,
|
||||
expect.objectContaining({ message: expect.any(String) }),
|
||||
'error',
|
||||
);
|
||||
}
|
||||
);
|
||||
});
|
||||
|
||||
describe('SettingsManager library controls', () => {
|
||||
it('loads libraries and populates the select', async () => {
|
||||
const manager = createManager();
|
||||
|
||||
@@ -1,12 +1,26 @@
|
||||
import { describe, beforeEach, afterEach, expect, it, vi } from 'vitest';
|
||||
import { UpdateService } from '../../../static/js/managers/UpdateService.js';
|
||||
import { state } from '../../../static/js/state/index.js';
|
||||
|
||||
function createFetchResponse(payload) {
|
||||
return {
|
||||
json: vi.fn().mockResolvedValue(payload)
|
||||
json: vi.fn().mockResolvedValue(payload),
|
||||
ok: true,
|
||||
};
|
||||
}
|
||||
|
||||
function stubSettingsUpdateChannel(channel) {
|
||||
state.global = state.global || {};
|
||||
state.global.settings = state.global.settings || {};
|
||||
state.global.settings.update_channel = channel;
|
||||
}
|
||||
|
||||
function clearSettingsUpdateChannel() {
|
||||
if (state.global?.settings) {
|
||||
delete state.global.settings.update_channel;
|
||||
}
|
||||
}
|
||||
|
||||
describe('UpdateService passive checks', () => {
|
||||
let service;
|
||||
let fetchMock;
|
||||
@@ -16,10 +30,13 @@ describe('UpdateService passive checks', () => {
|
||||
success: true,
|
||||
current_version: 'v1.0.0',
|
||||
latest_version: 'v1.0.0',
|
||||
git_info: { short_hash: 'abc123' }
|
||||
git_info: { short_hash: 'abc123' },
|
||||
has_git: true,
|
||||
}));
|
||||
global.fetch = fetchMock;
|
||||
|
||||
stubSettingsUpdateChannel('release');
|
||||
|
||||
service = new UpdateService();
|
||||
service.updateNotificationsEnabled = false;
|
||||
service.lastCheckTime = 0;
|
||||
@@ -28,6 +45,7 @@ describe('UpdateService passive checks', () => {
|
||||
|
||||
afterEach(() => {
|
||||
delete global.fetch;
|
||||
clearSettingsUpdateChannel();
|
||||
});
|
||||
|
||||
it('skips passive update checks when notifications are disabled', async () => {
|
||||
@@ -43,3 +61,106 @@ describe('UpdateService passive checks', () => {
|
||||
expect(fetchMock).toHaveBeenCalledWith('/api/lm/check-updates?nightly=false');
|
||||
});
|
||||
});
|
||||
|
||||
describe('UpdateService nightly notification throttling', () => {
|
||||
let fetchMock;
|
||||
let updateToggle;
|
||||
let updateBadge;
|
||||
|
||||
function stubUpdateBadgeDom() {
|
||||
updateToggle = document.createElement('div');
|
||||
updateToggle.className = 'update-toggle';
|
||||
updateBadge = document.createElement('span');
|
||||
updateBadge.className = 'update-badge';
|
||||
updateToggle.appendChild(updateBadge);
|
||||
document.body.appendChild(updateToggle);
|
||||
|
||||
vi.spyOn(document, 'querySelector').mockImplementation((selector) => {
|
||||
if (selector === '.update-toggle') return updateToggle;
|
||||
if (selector === '.update-toggle .update-badge') return updateBadge;
|
||||
return null;
|
||||
});
|
||||
}
|
||||
|
||||
function makeUpdateResponse(channel) {
|
||||
return {
|
||||
success: true,
|
||||
current_version: 'v1.0.0',
|
||||
latest_version: channel === 'nightly' ? 'main-abc1234' : 'v1.1.0',
|
||||
update_available: true,
|
||||
git_info: { short_hash: 'abc123' },
|
||||
has_git: true,
|
||||
nightly: channel === 'nightly',
|
||||
changelog: ['test: change'],
|
||||
releases: [],
|
||||
behind_by: 3,
|
||||
commit_date: '2026-07-31',
|
||||
};
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
fetchMock = vi.fn().mockResolvedValue(createFetchResponse(makeUpdateResponse('release')));
|
||||
global.fetch = fetchMock;
|
||||
stubUpdateBadgeDom();
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
vi.restoreAllMocks();
|
||||
delete global.fetch;
|
||||
});
|
||||
|
||||
it('shows the nightly badge once and keeps it visible for the session', async () => {
|
||||
stubSettingsUpdateChannel('nightly');
|
||||
fetchMock.mockResolvedValue(createFetchResponse(makeUpdateResponse('nightly')));
|
||||
|
||||
const service = new UpdateService();
|
||||
service.updateNotificationsEnabled = true;
|
||||
|
||||
await service.checkForUpdates({ force: true });
|
||||
|
||||
expect(service.updateAvailable).toBe(true);
|
||||
expect(service.nightlyBadgeShown).toBe(true);
|
||||
expect(service.nightlyNotifyDate).toBe(service._getTodayKey());
|
||||
expect(updateBadge.classList.contains('visible')).toBe(true);
|
||||
|
||||
// A repeated check within the same session keeps the badge visible.
|
||||
await service.checkForUpdates({ force: true });
|
||||
expect(updateBadge.classList.contains('visible')).toBe(true);
|
||||
});
|
||||
|
||||
it('suppresses the nightly badge on a later session in the same day', async () => {
|
||||
stubSettingsUpdateChannel('nightly');
|
||||
fetchMock.mockResolvedValue(createFetchResponse(makeUpdateResponse('nightly')));
|
||||
|
||||
const firstService = new UpdateService();
|
||||
firstService.updateNotificationsEnabled = true;
|
||||
await firstService.checkForUpdates({ force: true });
|
||||
expect(updateBadge.classList.contains('visible')).toBe(true);
|
||||
|
||||
// Simulate a fresh page session on the same calendar day.
|
||||
const secondService = new UpdateService();
|
||||
secondService.updateNotificationsEnabled = true;
|
||||
await secondService.checkForUpdates({ force: true });
|
||||
|
||||
expect(secondService.updateAvailable).toBe(true);
|
||||
expect(secondService.nightlyBadgeShown).toBe(false);
|
||||
expect(updateBadge.classList.contains('visible')).toBe(false);
|
||||
});
|
||||
|
||||
it('is not affected by the daily limit on the release channel', async () => {
|
||||
stubSettingsUpdateChannel('release');
|
||||
fetchMock.mockResolvedValue(createFetchResponse(makeUpdateResponse('release')));
|
||||
|
||||
const firstService = new UpdateService();
|
||||
firstService.updateNotificationsEnabled = true;
|
||||
await firstService.checkForUpdates({ force: true });
|
||||
expect(updateBadge.classList.contains('visible')).toBe(true);
|
||||
|
||||
const secondService = new UpdateService();
|
||||
secondService.updateNotificationsEnabled = true;
|
||||
await secondService.checkForUpdates({ force: true });
|
||||
|
||||
expect(secondService.updateAvailable).toBe(true);
|
||||
expect(updateBadge.classList.contains('visible')).toBe(true);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -30,10 +30,10 @@ def test_metadata_hook_installs_and_traces_execution(monkeypatch, metadata_regis
|
||||
|
||||
calls = []
|
||||
|
||||
def record_stub(self, node_id, class_type, inputs, outputs):
|
||||
def record_stub(self, node_id, class_type, inputs, outputs, return_types=None):
|
||||
calls.append(("record", node_id, class_type, inputs))
|
||||
|
||||
def update_stub(self, node_id, class_type, outputs):
|
||||
def update_stub(self, node_id, class_type, outputs, return_types=None):
|
||||
calls.append(("update", node_id, class_type, outputs))
|
||||
|
||||
monkeypatch.setattr(MetadataRegistry, "record_node_execution", record_stub)
|
||||
@@ -820,3 +820,227 @@ def test_lora_manager_checkpoint_and_unet_loaders_extract_models(metadata_regist
|
||||
"type": "checkpoint",
|
||||
"node_id": "unet_node",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# MetadataOverwriteExtractor & overwrite merge tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
from py.metadata_collector.constants import OVERWRITE, METADATA_OVERWRITE_FIELDS
|
||||
from py.metadata_collector.node_extractors import MetadataOverwriteExtractor
|
||||
|
||||
|
||||
def test_metadata_overwrite_extractor_stores_truthy_values(metadata_registry):
|
||||
"""Extractor should store truthy inputs under the OVERWRITE category."""
|
||||
metadata_registry.start_collection("prompt-ow")
|
||||
metadata = metadata_registry.prompt_metadata["prompt-ow"]
|
||||
|
||||
inputs = {
|
||||
"prompt": "a beautiful landscape",
|
||||
"negative_prompt": "",
|
||||
"seed": 42,
|
||||
"steps": 0,
|
||||
"cfg_scale": 7.5,
|
||||
"sampler": "",
|
||||
"scheduler": "",
|
||||
"model": "myModel.safetensors",
|
||||
"loras": "<lora:detail:0.8>",
|
||||
"size": "1024x768",
|
||||
"clip_skip": 0,
|
||||
"additional_data": '{"Copyright": "CC0"}',
|
||||
}
|
||||
|
||||
MetadataOverwriteExtractor.extract("ow-1", inputs, None, metadata)
|
||||
|
||||
assert OVERWRITE in metadata
|
||||
assert "ow-1" in metadata[OVERWRITE]
|
||||
params = metadata[OVERWRITE]["ow-1"]["parameters"]
|
||||
|
||||
# Truthy values stored
|
||||
assert params["prompt"] == "a beautiful landscape"
|
||||
assert params["seed"] == 42
|
||||
assert params["cfg_scale"] == 7.5
|
||||
assert params["model"] == "myModel.safetensors"
|
||||
assert params["loras"] == "<lora:detail:0.8>"
|
||||
assert params["size"] == "1024x768"
|
||||
assert params["additional_data"] == '{"Copyright": "CC0"}'
|
||||
|
||||
# Falsy values NOT stored
|
||||
assert "negative_prompt" not in params
|
||||
assert "steps" not in params
|
||||
assert "sampler" not in params
|
||||
assert "scheduler" not in params
|
||||
# clip_skip=0 is now stored (0 != sentinel -25) — wired 0 is valid
|
||||
assert params["clip_skip"] == 0
|
||||
|
||||
metadata_registry.clear_metadata()
|
||||
|
||||
|
||||
def test_metadata_overwrite_extractor_empty_inputs(metadata_registry):
|
||||
"""Extractor with all-falsy inputs should NOT create OVERWRITE category."""
|
||||
metadata_registry.start_collection("prompt-ow2")
|
||||
metadata = metadata_registry.prompt_metadata["prompt-ow2"]
|
||||
|
||||
from py.metadata_collector.constants import CLIP_SKIP_SENTINEL
|
||||
|
||||
inputs = {key: "" for key in METADATA_OVERWRITE_FIELDS}
|
||||
inputs.update({"seed": 0, "steps": 0, "cfg_scale": 0.0, "clip_skip": CLIP_SKIP_SENTINEL})
|
||||
|
||||
MetadataOverwriteExtractor.extract("ow-2", inputs, None, metadata)
|
||||
|
||||
# start_collection pre-creates empty dicts for all categories,
|
||||
# but no node should have populated OVERWRITE with any data
|
||||
assert not metadata[OVERWRITE]
|
||||
|
||||
metadata_registry.clear_metadata()
|
||||
|
||||
|
||||
def test_extract_generation_params_applies_overwrite(metadata_registry, populated_registry, monkeypatch):
|
||||
"""overwrite values should replace inferred params in extract_generation_params."""
|
||||
import py.metadata_collector.metadata_processor as mp
|
||||
|
||||
monkeypatch.setattr(mp, "standalone_mode", False)
|
||||
|
||||
metadata = populated_registry["metadata"]
|
||||
registry_obj = populated_registry["registry"]
|
||||
|
||||
# Simulate the MetadataOverwriteLM node having been executed with overwrite values
|
||||
registry_obj.start_collection("promptA")
|
||||
# Re-populate with the same data (start_collection resets)
|
||||
registry_obj.set_current_prompt(populated_registry["prompt"])
|
||||
metadata2 = registry_obj.prompt_metadata["promptA"]
|
||||
|
||||
# Inject overwrite data into metadata
|
||||
metadata2[OVERWRITE] = {
|
||||
"ow-1": {
|
||||
"parameters": {
|
||||
"seed": 777,
|
||||
"additional_data": '{"AuthorURL": "https://civitai.com/user/foo"}',
|
||||
},
|
||||
"node_id": "ow-1",
|
||||
}
|
||||
}
|
||||
# Copy other categories from original populated metadata
|
||||
for cat in ("models", "prompts", "sampling", "loras", "size", "images"):
|
||||
if cat in metadata:
|
||||
metadata2[cat] = metadata[cat]
|
||||
metadata2["execution_order"] = metadata["execution_order"]
|
||||
|
||||
params = MetadataProcessor.extract_generation_params(metadata2, id="vae")
|
||||
|
||||
# Overwritten values
|
||||
assert params["seed"] == 777
|
||||
assert params["additional_data"] == '{"AuthorURL": "https://civitai.com/user/foo"}'
|
||||
|
||||
# Inferred values still present (not overwritten)
|
||||
assert params["prompt"] == "A castle on a hill"
|
||||
assert params["cfg_scale"] == 7.5
|
||||
assert params["checkpoint"] == "model.safetensors"
|
||||
|
||||
registry_obj.clear_metadata()
|
||||
|
||||
|
||||
def test_extract_generation_params_overwrite_falsy_skipped(metadata_registry, populated_registry, monkeypatch):
|
||||
"""Overwrite entries with falsy values should NOT replace inferred params."""
|
||||
import py.metadata_collector.metadata_processor as mp
|
||||
|
||||
monkeypatch.setattr(mp, "standalone_mode", False)
|
||||
|
||||
metadata = populated_registry["metadata"]
|
||||
registry_obj = populated_registry["registry"]
|
||||
|
||||
registry_obj.start_collection("promptA")
|
||||
registry_obj.set_current_prompt(populated_registry["prompt"])
|
||||
metadata2 = registry_obj.prompt_metadata["promptA"]
|
||||
|
||||
# Inject overwrite with falsy values (except clip_skip=0 which is now
|
||||
# treated as a valid wired input thanks to the -25 sentinel)
|
||||
metadata2[OVERWRITE] = {
|
||||
"ow-1": {
|
||||
"parameters": {
|
||||
"seed": 0,
|
||||
"steps": 0,
|
||||
"cfg_scale": 0.0,
|
||||
"prompt": "",
|
||||
"clip_skip": 0,
|
||||
},
|
||||
"node_id": "ow-1",
|
||||
}
|
||||
}
|
||||
for cat in ("models", "prompts", "sampling", "loras", "size", "images"):
|
||||
if cat in metadata:
|
||||
metadata2[cat] = metadata[cat]
|
||||
metadata2["execution_order"] = metadata["execution_order"]
|
||||
|
||||
params = MetadataProcessor.extract_generation_params(metadata2, id="vae")
|
||||
|
||||
# Falsy overwrites should NOT have replaced inferred values
|
||||
assert params["prompt"] == "A castle on a hill"
|
||||
assert params["cfg_scale"] == 7.5
|
||||
|
||||
# clip_skip=0 is a valid wired value (not the -25 sentinel) — should be applied
|
||||
assert params["clip_skip"] == 0
|
||||
|
||||
registry_obj.clear_metadata()
|
||||
|
||||
|
||||
def test_fill_missing_metadata_skips_overwrite_for_bypassed_node(metadata_registry):
|
||||
"""Bypassed (mode=4) node should not have OVERWRITE filled from cache."""
|
||||
metadata_registry.start_collection("prompt-bypass")
|
||||
|
||||
# Simulate a previous execution that cached overwrite data
|
||||
metadata_registry.record_node_execution(
|
||||
"ow-1",
|
||||
"MetadataOverwriteLM",
|
||||
{"seed": 99, "prompt": "test", "steps": 0, "cfg_scale": 0.0,
|
||||
"negative_prompt": "", "sampler": "", "scheduler": "", "model": "",
|
||||
"loras": "", "size": "", "clip_skip": 0, "additional_data": ""},
|
||||
None,
|
||||
)
|
||||
|
||||
# Now start a new prompt where the node is bypassed (mode=4)
|
||||
metadata_registry.start_collection("prompt-bypass-2")
|
||||
original_prompt = {
|
||||
"ow-1": {"class_type": "MetadataOverwriteLM", "inputs": {}, "mode": 4},
|
||||
}
|
||||
metadata_registry.set_current_prompt(
|
||||
SimpleNamespace(original_prompt=original_prompt)
|
||||
)
|
||||
|
||||
metadata = metadata_registry.get_metadata("prompt-bypass-2")
|
||||
|
||||
# The overwrite data should NOT be present (node was bypassed, not
|
||||
# a cache hit — it should not inherit previous execution's overwrite)
|
||||
assert "ow-1" not in metadata.get(OVERWRITE, {})
|
||||
|
||||
metadata_registry.clear_metadata()
|
||||
|
||||
|
||||
def test_fill_missing_metadata_fills_overwrite_for_muted_node(metadata_registry):
|
||||
"""Muted (mode=2) node should also not have OVERWRITE filled from cache."""
|
||||
metadata_registry.start_collection("prompt-mute")
|
||||
|
||||
# Simulate a previous execution that cached overwrite data
|
||||
metadata_registry.record_node_execution(
|
||||
"ow-1",
|
||||
"MetadataOverwriteLM",
|
||||
{"seed": 88, "prompt": "test2", "steps": 0, "cfg_scale": 0.0,
|
||||
"negative_prompt": "", "sampler": "", "scheduler": "", "model": "",
|
||||
"loras": "", "size": "", "clip_skip": 0, "additional_data": ""},
|
||||
None,
|
||||
)
|
||||
|
||||
# Start a new prompt where the node is muted (mode=2)
|
||||
metadata_registry.start_collection("prompt-mute-2")
|
||||
original_prompt = {
|
||||
"ow-1": {"class_type": "MetadataOverwriteLM", "inputs": {}, "mode": 2},
|
||||
}
|
||||
metadata_registry.set_current_prompt(
|
||||
SimpleNamespace(original_prompt=original_prompt)
|
||||
)
|
||||
|
||||
metadata = metadata_registry.get_metadata("prompt-mute-2")
|
||||
|
||||
assert "ow-1" not in metadata.get(OVERWRITE, {})
|
||||
|
||||
metadata_registry.clear_metadata()
|
||||
|
||||
@@ -1,4 +1,11 @@
|
||||
from py.nodes.lora_stack_combiner import LoraStackCombinerLM
|
||||
import types
|
||||
|
||||
import pytest
|
||||
|
||||
from py.nodes.lora_stack_combiner import (
|
||||
LoraStackCombinerLM,
|
||||
_LoraStackOptionalInputs,
|
||||
)
|
||||
|
||||
|
||||
def test_combine_stacks_preserves_order():
|
||||
@@ -49,3 +56,109 @@ def test_combine_stacks_allows_duplicate_entries():
|
||||
(combined_stack,) = node.combine_stacks([duplicate_entry], [duplicate_entry])
|
||||
|
||||
assert combined_stack == [duplicate_entry, duplicate_entry]
|
||||
|
||||
|
||||
def test_combine_stacks_returns_empty_when_both_unconnected():
|
||||
node = LoraStackCombinerLM()
|
||||
|
||||
(combined_stack,) = node.combine_stacks()
|
||||
|
||||
assert combined_stack == []
|
||||
|
||||
|
||||
def test_combine_stacks_returns_other_when_one_unconnected():
|
||||
node = LoraStackCombinerLM()
|
||||
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||
|
||||
(combined_stack_a,) = node.combine_stacks(lora_stack1=stack_a)
|
||||
(combined_stack_b,) = node.combine_stacks(lora_stack2=stack_a)
|
||||
|
||||
assert combined_stack_a == stack_a
|
||||
assert combined_stack_b == stack_a
|
||||
|
||||
|
||||
def test_combine_stacks_with_dynamic_third_slot():
|
||||
node = LoraStackCombinerLM()
|
||||
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
|
||||
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
|
||||
|
||||
(combined_stack,) = node.combine_stacks(
|
||||
lora_stack1=stack_a, lora_stack2=stack_b, lora_stack3=stack_c
|
||||
)
|
||||
|
||||
assert combined_stack == stack_a + stack_b + stack_c
|
||||
|
||||
|
||||
def test_combine_stacks_orders_by_slot_number_not_call_order():
|
||||
node = LoraStackCombinerLM()
|
||||
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
|
||||
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
|
||||
|
||||
(combined_stack,) = node.combine_stacks(
|
||||
lora_stack3=stack_c, lora_stack2=stack_b, lora_stack1=stack_a
|
||||
)
|
||||
|
||||
assert combined_stack == stack_a + stack_b + stack_c
|
||||
|
||||
|
||||
def test_combine_stacks_accepts_only_dynamic_slot():
|
||||
node = LoraStackCombinerLM()
|
||||
stack_c = [("folder/c.safetensors", 1.0, 0.9)]
|
||||
|
||||
(combined_stack,) = node.combine_stacks(lora_stack3=stack_c)
|
||||
|
||||
assert combined_stack == stack_c
|
||||
|
||||
|
||||
def test_combine_stacks_handles_legacy_input_names():
|
||||
node = LoraStackCombinerLM()
|
||||
stack_a = [("folder/a.safetensors", 0.7, 0.6)]
|
||||
stack_b = [("folder/b.safetensors", 0.8, 0.8)]
|
||||
|
||||
(combined_stack,) = node.combine_stacks(lora_stack_a=stack_a, lora_stack_b=stack_b)
|
||||
|
||||
assert combined_stack == stack_a + stack_b
|
||||
|
||||
|
||||
def test_input_types_exposes_two_default_slots():
|
||||
input_types = LoraStackCombinerLM.INPUT_TYPES()
|
||||
|
||||
assert set(input_types["optional"]) == {"lora_stack1", "lora_stack2"}
|
||||
assert input_types["optional"]["lora_stack1"][0] == "LORA_STACK"
|
||||
assert input_types["optional"]["lora_stack2"][0] == "LORA_STACK"
|
||||
|
||||
|
||||
def test_input_types_recognizes_dynamic_slots_from_get_input_info(monkeypatch):
|
||||
frames = [None, None, types.SimpleNamespace(function="get_input_info")]
|
||||
monkeypatch.setattr(
|
||||
"py.nodes.lora_stack_combiner.inspect.stack", lambda: frames
|
||||
)
|
||||
|
||||
input_types = LoraStackCombinerLM.INPUT_TYPES()
|
||||
optional = input_types["optional"]
|
||||
|
||||
assert "lora_stack3" in optional
|
||||
assert optional["lora_stack3"][0] == "LORA_STACK"
|
||||
assert "lora_stack25" in optional
|
||||
assert optional["lora_stack25"][0] == "LORA_STACK"
|
||||
|
||||
|
||||
def test_lora_stack_optional_inputs_proxy():
|
||||
proxy = _LoraStackOptionalInputs({"lora_stack1": ("LORA_STACK", {})})
|
||||
|
||||
assert "lora_stack1" in proxy
|
||||
assert "lora_stack2" in proxy
|
||||
assert "lora_stack10" in proxy
|
||||
assert "lora_stack_a" in proxy
|
||||
assert "lora_stack" not in proxy
|
||||
assert "lora_stacka" not in proxy
|
||||
assert "lora_stack_1" not in proxy
|
||||
assert "text" not in proxy
|
||||
|
||||
assert proxy["lora_stack1"][0] == "LORA_STACK"
|
||||
assert proxy["lora_stack5"][0] == "LORA_STACK"
|
||||
|
||||
with pytest.raises(KeyError):
|
||||
proxy["not_a_stack"]
|
||||
|
||||
@@ -59,7 +59,7 @@ def test_save_image_defaults_to_writing_png_metadata(monkeypatch, tmp_path):
|
||||
|
||||
image_path = tmp_path / "sample_00001_.png"
|
||||
with Image.open(image_path) as img:
|
||||
assert img.info["parameters"] == "prompt text\nSeed: 123"
|
||||
assert img.info["parameters"] == "prompt text\nSeed: 123, Version: ComfyUI"
|
||||
|
||||
|
||||
def test_save_image_skips_png_parameters_when_metadata_disabled_and_keeps_workflow(
|
||||
@@ -363,3 +363,102 @@ def test_save_image_as_recipe_writes_recipe_without_async_scanner_calls(
|
||||
assert recipe["gen_params"] == {"prompt": "prompt text", "seed": 123}
|
||||
assert scanner._json_path_map[recipe["id"]] == os.path.normpath(str(recipe_files[0]))
|
||||
assert scanner.fts_updates == [(recipe["id"], "add")]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests for webp_method and jpeg_subsampling parameters
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _capture_save_kwargs(monkeypatch):
|
||||
"""Monkeypatch Image.Image.save to capture kwargs while still saving to disk."""
|
||||
real_save = Image.Image.save
|
||||
captured_kwargs = {}
|
||||
|
||||
def _fake_save(self, fp, *args, **kwargs):
|
||||
captured_kwargs.update(kwargs)
|
||||
return real_save(self, fp, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(Image.Image, "save", _fake_save)
|
||||
return captured_kwargs
|
||||
|
||||
|
||||
def test_webp_method_default_passed_to_pillow_save(monkeypatch, tmp_path):
|
||||
_configure_save_paths(monkeypatch, tmp_path)
|
||||
_configure_metadata(monkeypatch, {"prompt": "test", "seed": 1})
|
||||
captured = _capture_save_kwargs(monkeypatch)
|
||||
|
||||
node = SaveImageLM()
|
||||
node.save_images([_make_image()], "ComfyUI", "webp", id="node-1")
|
||||
|
||||
assert "method" in captured
|
||||
assert captured["method"] == 6
|
||||
|
||||
|
||||
def test_webp_method_custom_value_passed_to_pillow_save(monkeypatch, tmp_path):
|
||||
_configure_save_paths(monkeypatch, tmp_path)
|
||||
_configure_metadata(monkeypatch, {"prompt": "test", "seed": 1})
|
||||
captured = _capture_save_kwargs(monkeypatch)
|
||||
|
||||
node = SaveImageLM()
|
||||
node.save_images(
|
||||
[_make_image()], "ComfyUI", "webp", id="node-1", webp_method=3
|
||||
)
|
||||
|
||||
assert captured["method"] == 3
|
||||
|
||||
|
||||
def test_jpeg_subsampling_default_passed_to_pillow_save(monkeypatch, tmp_path):
|
||||
_configure_save_paths(monkeypatch, tmp_path)
|
||||
_configure_metadata(monkeypatch, {"prompt": "test", "seed": 1})
|
||||
captured = _capture_save_kwargs(monkeypatch)
|
||||
|
||||
node = SaveImageLM()
|
||||
node.save_images([_make_image()], "ComfyUI", "jpeg", id="node-1")
|
||||
|
||||
assert "subsampling" in captured
|
||||
assert captured["subsampling"] == 0
|
||||
|
||||
|
||||
def test_jpeg_subsampling_custom_value_passed_to_pillow_save(monkeypatch, tmp_path):
|
||||
_configure_save_paths(monkeypatch, tmp_path)
|
||||
_configure_metadata(monkeypatch, {"prompt": "test", "seed": 1})
|
||||
captured = _capture_save_kwargs(monkeypatch)
|
||||
|
||||
node = SaveImageLM()
|
||||
node.save_images(
|
||||
[_make_image()], "ComfyUI", "jpeg", id="node-1", jpeg_subsampling=1
|
||||
)
|
||||
|
||||
assert captured["subsampling"] == 1
|
||||
|
||||
|
||||
class TestParameterDefaultConsistency:
|
||||
"""Verify defaults match across INPUT_TYPES, save_images(), and process_image()."""
|
||||
|
||||
def test_webp_method_defaults_are_consistent(self):
|
||||
input_types = SaveImageLM.INPUT_TYPES()
|
||||
optional = input_types["optional"]
|
||||
|
||||
assert optional["webp_method"][1]["default"] == 6
|
||||
assert SaveImageLM.save_images.__defaults__[4] == 6 # positional: webp_method=6 is at index 4
|
||||
assert SaveImageLM.process_image.__defaults__[6] == 6
|
||||
|
||||
def test_jpeg_subsampling_defaults_are_consistent(self):
|
||||
input_types = SaveImageLM.INPUT_TYPES()
|
||||
optional = input_types["optional"]
|
||||
|
||||
assert optional["jpeg_subsampling"][1]["default"] == 0
|
||||
assert SaveImageLM.save_images.__defaults__[5] == 0
|
||||
assert SaveImageLM.process_image.__defaults__[7] == 0
|
||||
|
||||
|
||||
def test_png_does_not_pass_webp_method_or_jpeg_subsampling(monkeypatch, tmp_path):
|
||||
_configure_save_paths(monkeypatch, tmp_path)
|
||||
_configure_metadata(monkeypatch, {"prompt": "test", "seed": 1})
|
||||
captured = _capture_save_kwargs(monkeypatch)
|
||||
|
||||
node = SaveImageLM()
|
||||
node.save_images([_make_image()], "ComfyUI", "png", id="node-1")
|
||||
|
||||
assert "method" not in captured
|
||||
assert "subsampling" not in captured
|
||||
|
||||
@@ -900,18 +900,28 @@ class FakeMetadataProvider:
|
||||
async def get_model_versions(self, _model_id):
|
||||
return {"modelVersions": [], "name": "", "type": "lora"}
|
||||
|
||||
async def get_user_models(self, _username):
|
||||
return []
|
||||
async def get_user_models(self, _username, cursor=None):
|
||||
return {"items": [], "nextCursor": None}
|
||||
|
||||
async def get_creator_model_count(self, _username):
|
||||
return None
|
||||
|
||||
|
||||
class FakeUserModelsProvider(FakeMetadataProvider):
|
||||
def __init__(self, models):
|
||||
def __init__(self, models, next_cursor=None, estimated_total=None):
|
||||
self.models = models
|
||||
self.next_cursor = next_cursor
|
||||
self.estimated_total = estimated_total
|
||||
self.received_usernames: list[str] = []
|
||||
self.received_cursors: list = []
|
||||
|
||||
async def get_user_models(self, username):
|
||||
async def get_user_models(self, username, cursor=None):
|
||||
self.received_usernames.append(username)
|
||||
return self.models
|
||||
self.received_cursors.append(cursor)
|
||||
return {"items": self.models, "nextCursor": self.next_cursor}
|
||||
|
||||
async def get_creator_model_count(self, _username):
|
||||
return self.estimated_total
|
||||
|
||||
|
||||
async def fake_metadata_provider_factory():
|
||||
@@ -1286,6 +1296,88 @@ async def test_get_civitai_user_models_requires_username():
|
||||
assert "username" in payload["error"].lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_civitai_user_models_returns_pagination_fields():
|
||||
models = [
|
||||
{
|
||||
"id": 1,
|
||||
"name": "Model A",
|
||||
"type": "LORA",
|
||||
"tags": [],
|
||||
"modelVersions": [
|
||||
{"id": 100, "name": "v1", "images": [{"url": "http://example.com/a.jpg"}]},
|
||||
],
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"name": "Unsupported",
|
||||
"type": "Other",
|
||||
"modelVersions": [{"id": 200, "name": "v1"}],
|
||||
},
|
||||
]
|
||||
|
||||
provider = FakeUserModelsProvider(models, next_cursor="cursor-token", estimated_total=2140)
|
||||
|
||||
async def provider_factory():
|
||||
return provider
|
||||
|
||||
handler = ModelLibraryHandler(
|
||||
ServiceRegistryAdapter(
|
||||
get_lora_scanner=fake_scanner_factory,
|
||||
get_checkpoint_scanner=fake_scanner_factory,
|
||||
get_embedding_scanner=fake_scanner_factory,
|
||||
get_downloaded_version_history_service=fake_download_history_service_factory,
|
||||
),
|
||||
metadata_provider_factory=provider_factory,
|
||||
)
|
||||
|
||||
response = await handler.get_civitai_user_models(
|
||||
FakeRequest(query={"username": "pixel"})
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
# modelCount only counts models surviving the type filter
|
||||
assert payload["modelCount"] == 1
|
||||
assert payload["nextCursor"] == "cursor-token"
|
||||
assert payload["hasMore"] is True
|
||||
# first page includes the estimated total
|
||||
assert payload["estimatedTotal"] == 2140
|
||||
assert provider.received_cursors == [None]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_civitai_user_models_passes_cursor_and_omits_estimate():
|
||||
provider = FakeUserModelsProvider([], next_cursor=None, estimated_total=999)
|
||||
|
||||
async def provider_factory():
|
||||
return provider
|
||||
|
||||
handler = ModelLibraryHandler(
|
||||
ServiceRegistryAdapter(
|
||||
get_lora_scanner=fake_scanner_factory,
|
||||
get_checkpoint_scanner=fake_scanner_factory,
|
||||
get_embedding_scanner=fake_scanner_factory,
|
||||
get_downloaded_version_history_service=fake_download_history_service_factory,
|
||||
),
|
||||
metadata_provider_factory=provider_factory,
|
||||
)
|
||||
|
||||
response = await handler.get_civitai_user_models(
|
||||
FakeRequest(query={"username": "pixel", "cursor": "opaque-token"})
|
||||
)
|
||||
payload = json.loads(response.text)
|
||||
|
||||
assert response.status == 200
|
||||
assert payload["success"] is True
|
||||
assert payload["nextCursor"] is None
|
||||
assert payload["hasMore"] is False
|
||||
# cursor requests must not include the estimated total
|
||||
assert payload["estimatedTotal"] is None
|
||||
assert provider.received_cursors == ["opaque-token"]
|
||||
|
||||
|
||||
def test_ensure_handler_mapping_caches_result():
|
||||
call_records = []
|
||||
|
||||
|
||||
@@ -1,10 +1,33 @@
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
from aiohttp import ClientError
|
||||
from aiohttp import web
|
||||
import pytest
|
||||
|
||||
from py.routes import update_routes
|
||||
|
||||
|
||||
def _fake_request(body=None, query_params=None):
|
||||
from multidict import MultiDict
|
||||
|
||||
q = MultiDict(query_params or {})
|
||||
|
||||
req = type("Req", (), {
|
||||
"has_body": body is not None,
|
||||
"match_info": {},
|
||||
"rel_url": type("U", (), {"query": q})(),
|
||||
"query": q,
|
||||
"app": {},
|
||||
})()
|
||||
|
||||
async def _json():
|
||||
return body or {}
|
||||
|
||||
req.json = _json
|
||||
return req
|
||||
|
||||
|
||||
class OfflineDownloader:
|
||||
async def make_request(self, *_, **__):
|
||||
return False, "Cannot connect to host"
|
||||
@@ -53,10 +76,12 @@ async def test_get_nightly_version_network_error_logs_warning(monkeypatch, caplo
|
||||
caplog.set_level(logging.WARNING)
|
||||
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(RaisingDownloader()))
|
||||
|
||||
version, changelog = await update_routes.UpdateRoutes._get_nightly_version()
|
||||
version, changelog, behind_by, commit_date = await update_routes.UpdateRoutes._get_nightly_version()
|
||||
|
||||
assert version == "main"
|
||||
assert changelog == []
|
||||
assert behind_by == 0
|
||||
assert commit_date == ""
|
||||
assert "Unable to reach GitHub for nightly version" in caplog.text
|
||||
assert "Traceback" not in caplog.text
|
||||
|
||||
@@ -236,3 +261,305 @@ async def test_perform_git_update_stable_preserves_user_dirs(monkeypatch, tmp_pa
|
||||
clean_args = clean_calls[0][1]
|
||||
for name in update_routes._PRESERVE_DIRS:
|
||||
assert name in clean_args, f"{name} missing from git clean excludes (stable)"
|
||||
|
||||
def test_init_git_repo_creates_valid_repo(tmp_path, monkeypatch):
|
||||
if not shutil.which("git"):
|
||||
pytest.skip("git executable not found")
|
||||
|
||||
plugin_root = tmp_path / "plugin"
|
||||
plugin_root.mkdir()
|
||||
(plugin_root / ".tracking").write_text("pyproject.toml")
|
||||
(plugin_root / "settings.json").write_text('{"some": "value"}')
|
||||
|
||||
try:
|
||||
success, version = update_routes.UpdateRoutes._init_git_repo(str(plugin_root))
|
||||
except Exception as e:
|
||||
pytest.skip(f"Network unavailable for git fetch: {e}")
|
||||
|
||||
assert success is True
|
||||
assert version.startswith("main-")
|
||||
assert len(version) > len("main-")
|
||||
assert (plugin_root / ".git").is_dir()
|
||||
assert not (plugin_root / ".tracking").exists()
|
||||
assert (plugin_root / "settings.json").exists()
|
||||
assert (plugin_root / "pyproject.toml").exists()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_switch_channel_invalid_channel_returns_error():
|
||||
req = _fake_request({"channel": "bad_channel"})
|
||||
resp = await update_routes.UpdateRoutes.switch_channel(req)
|
||||
|
||||
data = _raw_body(resp)
|
||||
assert not data["success"]
|
||||
assert "Invalid channel" in data["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_switch_channel_to_nightly_without_git_inits_repo(monkeypatch, tmp_path):
|
||||
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
|
||||
routes_file.parent.mkdir(parents=True)
|
||||
routes_file.write_text("")
|
||||
monkeypatch.setattr(update_routes, "__file__", str(routes_file))
|
||||
monkeypatch.setattr(update_routes, "ensure_settings_file", lambda logger: str(tmp_path / "settings.json"))
|
||||
monkeypatch.setattr(
|
||||
update_routes.UpdateRoutes,
|
||||
"_init_git_repo",
|
||||
staticmethod(lambda plugin_root: (True, "main-fedcba9")),
|
||||
)
|
||||
|
||||
req = _fake_request({"channel": "nightly"})
|
||||
resp = await update_routes.UpdateRoutes.switch_channel(req)
|
||||
data = _raw_body(resp)
|
||||
|
||||
assert data["success"] is True
|
||||
assert data["channel"] == "nightly"
|
||||
assert data["new_version"] == "main-fedcba9"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_switch_channel_to_nightly_with_git_calls_git_update(monkeypatch, tmp_path):
|
||||
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
|
||||
routes_file.parent.mkdir(parents=True)
|
||||
routes_file.write_text("")
|
||||
monkeypatch.setattr(update_routes, "__file__", str(routes_file))
|
||||
monkeypatch.setattr(update_routes, "ensure_settings_file", lambda logger: str(tmp_path / "settings.json"))
|
||||
|
||||
(tmp_path / ".git").mkdir()
|
||||
|
||||
async def _fake_git_update(*args, **kwargs):
|
||||
return True, "main-1111111"
|
||||
|
||||
monkeypatch.setattr(
|
||||
update_routes.UpdateRoutes, "_perform_git_update", _fake_git_update
|
||||
)
|
||||
|
||||
req = _fake_request({"channel": "nightly"})
|
||||
resp = await update_routes.UpdateRoutes.switch_channel(req)
|
||||
data = _raw_body(resp)
|
||||
|
||||
assert data["success"] is True
|
||||
assert data["channel"] == "nightly"
|
||||
assert data["new_version"] == "main-1111111"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_switch_channel_to_release_with_git_calls_git_update(monkeypatch, tmp_path):
|
||||
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
|
||||
routes_file.parent.mkdir(parents=True)
|
||||
routes_file.write_text("")
|
||||
monkeypatch.setattr(update_routes, "__file__", str(routes_file))
|
||||
monkeypatch.setattr(update_routes, "ensure_settings_file", lambda logger: str(tmp_path / "settings.json"))
|
||||
|
||||
(tmp_path / ".git").mkdir()
|
||||
|
||||
async def _fake_git_update(*args, **kwargs):
|
||||
return True, "v9.9.9"
|
||||
|
||||
monkeypatch.setattr(
|
||||
update_routes.UpdateRoutes, "_perform_git_update", _fake_git_update
|
||||
)
|
||||
|
||||
req = _fake_request({"channel": "release"})
|
||||
resp = await update_routes.UpdateRoutes.switch_channel(req)
|
||||
data = _raw_body(resp)
|
||||
|
||||
assert data["success"] is True
|
||||
assert data["channel"] == "release"
|
||||
assert data["new_version"] == "v9.9.9"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_switch_channel_to_release_without_git_still_downloads_zip(monkeypatch, tmp_path):
|
||||
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
|
||||
routes_file.parent.mkdir(parents=True)
|
||||
routes_file.write_text("")
|
||||
monkeypatch.setattr(update_routes, "__file__", str(routes_file))
|
||||
monkeypatch.setattr(update_routes, "ensure_settings_file", lambda logger: str(tmp_path / "settings.json"))
|
||||
|
||||
async def _fake_zip(*args, **kwargs):
|
||||
return True, "v2.0.0"
|
||||
|
||||
monkeypatch.setattr(
|
||||
update_routes.UpdateRoutes, "_download_and_replace_zip", _fake_zip
|
||||
)
|
||||
|
||||
req = _fake_request({"channel": "release"})
|
||||
resp = await update_routes.UpdateRoutes.switch_channel(req)
|
||||
data = _raw_body(resp)
|
||||
|
||||
assert data["success"] is True
|
||||
assert data["channel"] == "release"
|
||||
assert data["new_version"] == "v2.0.0"
|
||||
|
||||
|
||||
class _NightlyDownloader:
|
||||
"""Returns a fake main-branch commit AND a compare response."""
|
||||
|
||||
commit_sha = "7777777"
|
||||
commit_msg = "test: add nightly feature"
|
||||
commit_date = "2026-07-27T12:00:00Z"
|
||||
behind_by = 5
|
||||
|
||||
async def make_request(self, method, url, **kwargs):
|
||||
if "/compare/" in url:
|
||||
return True, {"behind_by": self.behind_by}
|
||||
return True, {
|
||||
"sha": self.commit_sha,
|
||||
"commit": {
|
||||
"message": self.commit_msg,
|
||||
"committer": {"date": self.commit_date},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_nightly_version_parses_behind_by(monkeypatch):
|
||||
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(_NightlyDownloader()))
|
||||
|
||||
version, changelog, behind_by, commit_date = await update_routes.UpdateRoutes._get_nightly_version(
|
||||
local_hash="abc1234"
|
||||
)
|
||||
|
||||
assert version == "main-7777777"
|
||||
assert behind_by == 5
|
||||
assert commit_date == "2026-07-27"
|
||||
assert len(changelog) == 1
|
||||
assert changelog[0] == "test: add nightly feature"
|
||||
|
||||
|
||||
class _AheadCompareDownloader:
|
||||
"""Fake compare API response with status='ahead' (main is ahead of local)."""
|
||||
|
||||
commit_sha = "9999999"
|
||||
commit_msg = "latest commit"
|
||||
commit_date = "2026-07-28T00:00:00Z"
|
||||
ahead_by = 3
|
||||
|
||||
async def make_request(self, method, url, **kwargs):
|
||||
if "/compare/" in url:
|
||||
return True, {"status": "ahead", "ahead_by": self.ahead_by, "behind_by": 0}
|
||||
return True, {
|
||||
"sha": self.commit_sha,
|
||||
"commit": {
|
||||
"message": self.commit_msg,
|
||||
"committer": {"date": self.commit_date},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_nightly_version_reads_ahead_by_when_ahead(monkeypatch):
|
||||
"""compare/{local}...main returns status='ahead' → read ahead_by."""
|
||||
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(_AheadCompareDownloader()))
|
||||
|
||||
version, changelog, behind_by, commit_date = await update_routes.UpdateRoutes._get_nightly_version(
|
||||
local_hash="oldhash"
|
||||
)
|
||||
|
||||
assert version == "main-9999999"
|
||||
assert behind_by == 3
|
||||
assert commit_date == "2026-07-28"
|
||||
|
||||
|
||||
class _DivergedCompareDownloader:
|
||||
"""Fake compare API response with status='diverged' (both have unique commits)."""
|
||||
|
||||
commit_sha = "aaaaaaa"
|
||||
commit_msg = "diverged test"
|
||||
commit_date = "2026-07-29T00:00:00Z"
|
||||
|
||||
async def make_request(self, method, url, **kwargs):
|
||||
if "/compare/" in url:
|
||||
return True, {"status": "diverged", "ahead_by": 5, "behind_by": 2}
|
||||
return True, {
|
||||
"sha": self.commit_sha,
|
||||
"commit": {
|
||||
"message": self.commit_msg,
|
||||
"committer": {"date": self.commit_date},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_nightly_version_reads_ahead_by_when_diverged(monkeypatch):
|
||||
"""compare/{local}...main returns status='diverged' → read ahead_by (remote ahead)."""
|
||||
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(_DivergedCompareDownloader()))
|
||||
|
||||
version, changelog, behind_by, commit_date = await update_routes.UpdateRoutes._get_nightly_version(
|
||||
local_hash="divhash"
|
||||
)
|
||||
|
||||
assert behind_by == 5
|
||||
|
||||
|
||||
class _CheckUpdatesDownloader:
|
||||
"""Fake downloader returning both a release list and a nightly commit + compare."""
|
||||
|
||||
commit_sha = "8888888"
|
||||
commit_date = "2026-07-28T00:00:00Z"
|
||||
|
||||
async def make_request(self, method, url, **kwargs):
|
||||
if "/releases" in url:
|
||||
return True, [
|
||||
{
|
||||
"tag_name": "v3.0.0",
|
||||
"body": "- Feature A\n- Feature B",
|
||||
"published_at": "2026-07-20T00:00:00Z",
|
||||
}
|
||||
]
|
||||
if "/compare/" in url:
|
||||
return True, {"behind_by": 3}
|
||||
return True, {
|
||||
"sha": self.commit_sha + "0" * 33,
|
||||
"commit": {
|
||||
"message": "latest commit",
|
||||
"committer": {"date": self.commit_date},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_updates_nightly_response_includes_behind_and_date(monkeypatch, tmp_path):
|
||||
monkeypatch.setattr(update_routes, "get_downloader", lambda: _stub_downloader(_CheckUpdatesDownloader()))
|
||||
|
||||
monkeypatch.setattr(
|
||||
update_routes.UpdateRoutes,
|
||||
"_get_local_version",
|
||||
staticmethod(lambda: "v1.0.0"),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
update_routes.UpdateRoutes,
|
||||
"_get_git_info",
|
||||
staticmethod(lambda: {
|
||||
"commit_hash": "abc1234",
|
||||
"short_hash": "abc1234",
|
||||
"branch": "main",
|
||||
"commit_date": "2026-01-01",
|
||||
}),
|
||||
)
|
||||
|
||||
routes_file = tmp_path / "py" / "routes" / "update_routes.py"
|
||||
routes_file.parent.mkdir(parents=True)
|
||||
routes_file.write_text("")
|
||||
monkeypatch.setattr(update_routes, "__file__", str(routes_file))
|
||||
(tmp_path / ".git").mkdir()
|
||||
|
||||
req = _fake_request(query_params={"nightly": "true"})
|
||||
resp = await update_routes.UpdateRoutes.check_updates(req)
|
||||
data = _raw_body(resp)
|
||||
|
||||
assert data["success"] is True
|
||||
assert data["nightly"] is True
|
||||
assert data["has_git"] is True
|
||||
assert data["behind_by"] == 3
|
||||
assert data["commit_date"] == "2026-07-28"
|
||||
assert data["latest_version"] == "main-8888888"
|
||||
assert isinstance(data["releases"], list)
|
||||
assert len(data["releases"]) == 1
|
||||
assert data["releases"][0]["version"] == "v3.0.0"
|
||||
|
||||
|
||||
def _raw_body(response):
|
||||
import json
|
||||
return json.loads(response._body.decode())
|
||||
|
||||
@@ -183,7 +183,7 @@ class FakeCache:
|
||||
def __init__(self, items):
|
||||
self.items = list(items)
|
||||
|
||||
async def get_sorted_data(self, sort_key, order):
|
||||
async def get_sorted_data(self, sort_key, order, seed=None):
|
||||
if sort_key == "name":
|
||||
data = sorted(self.items, key=lambda x: x["model_name"].lower())
|
||||
if order == "desc":
|
||||
@@ -1252,3 +1252,69 @@ async def test_get_model_civitai_url_falls_back_when_host_setting_is_not_a_strin
|
||||
"model_id": "123",
|
||||
"version_id": "456",
|
||||
}
|
||||
|
||||
|
||||
class TestHfGroupKey:
|
||||
"""Tests for _extract_hf_group_key and _extract_group_key."""
|
||||
|
||||
# --- _extract_hf_group_key ---
|
||||
|
||||
def test_hf_group_key_valid_url(self):
|
||||
"""Standard HF URL returns hf:user/repo."""
|
||||
item = {"hf_url": "https://huggingface.co/unsloth/qwen-edit"}
|
||||
assert BaseModelService._extract_hf_group_key(item) == "hf:unsloth/qwen-edit"
|
||||
|
||||
def test_hf_group_key_url_with_subpath(self):
|
||||
"""URL with subpath still extracts just owner/repo."""
|
||||
item = {"hf_url": "https://huggingface.co/user/repo/resolve/main/file.safetensors"}
|
||||
assert BaseModelService._extract_hf_group_key(item) == "hf:user/repo"
|
||||
|
||||
def test_hf_group_key_empty_url(self):
|
||||
"""Empty hf_url returns None."""
|
||||
assert BaseModelService._extract_hf_group_key({"hf_url": ""}) is None
|
||||
|
||||
def test_hf_group_key_no_url(self):
|
||||
"""Missing hf_url key returns None."""
|
||||
assert BaseModelService._extract_hf_group_key({}) is None
|
||||
|
||||
def test_hf_group_key_none_url(self):
|
||||
"""None hf_url returns None."""
|
||||
assert BaseModelService._extract_hf_group_key({"hf_url": None}) is None
|
||||
|
||||
def test_hf_group_key_invalid_url(self):
|
||||
"""Malformed HF URL returns None."""
|
||||
assert BaseModelService._extract_hf_group_key({"hf_url": "not-a-url"}) is None
|
||||
assert BaseModelService._extract_hf_group_key({"hf_url": "https://example.com"}) is None
|
||||
|
||||
# --- _extract_group_key ---
|
||||
|
||||
def test_group_key_civitai_only(self):
|
||||
"""CivitAI modelId returned as int."""
|
||||
item = {"civitai": {"modelId": 123}}
|
||||
assert BaseModelService._extract_group_key(item) == 123
|
||||
|
||||
def test_group_key_hf_only(self):
|
||||
"""HF-only item returns hf:user/repo string."""
|
||||
item = {"hf_url": "https://huggingface.co/user/repo"}
|
||||
assert BaseModelService._extract_group_key(item) == "hf:user/repo"
|
||||
|
||||
def test_group_key_civitai_preferred(self):
|
||||
"""CivitAI modelId takes precedence over hf_url."""
|
||||
item = {
|
||||
"civitai": {"modelId": 456},
|
||||
"hf_url": "https://huggingface.co/other/repo",
|
||||
}
|
||||
assert BaseModelService._extract_group_key(item) == 456
|
||||
|
||||
def test_group_key_neither(self):
|
||||
"""No CivitAI or HF returns None."""
|
||||
assert BaseModelService._extract_group_key({}) is None
|
||||
assert BaseModelService._extract_group_key({"some": "data"}) is None
|
||||
|
||||
def test_group_key_civitai_none_model_id(self):
|
||||
"""civitai.modelId=None falls through to HF."""
|
||||
item = {
|
||||
"civitai": {"modelId": None},
|
||||
"hf_url": "https://huggingface.co/user/repo",
|
||||
}
|
||||
assert BaseModelService._extract_group_key(item) == "hf:user/repo"
|
||||
|
||||
@@ -363,6 +363,148 @@ async def test_check_pending_models_handles_corrupted_progress_file(
|
||||
assert result["pending_count"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("tmp_path")
|
||||
async def test_check_pending_models_uses_bulk_folder_index_for_large_libraries(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
tmp_path,
|
||||
settings_manager,
|
||||
):
|
||||
"""For >1000 candidates the pre-check scans the library root once instead of
|
||||
probing every folder individually."""
|
||||
|
||||
ws_manager = RecordingWebSocketManager()
|
||||
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||
|
||||
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||
|
||||
# 1500 unprocessed models triggers the bulk lookup path
|
||||
models = [
|
||||
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
|
||||
for i in range(1500)
|
||||
]
|
||||
|
||||
# Create folders with files for the first 500 models
|
||||
for i in range(500):
|
||||
model_dir = tmp_path / f"{i:064x}"
|
||||
model_dir.mkdir()
|
||||
(model_dir / "image_0.png").write_text("data")
|
||||
|
||||
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
|
||||
|
||||
per_model_checks = 0
|
||||
|
||||
def counting_model_directory_has_files(path: str) -> bool:
|
||||
nonlocal per_model_checks
|
||||
per_model_checks += 1
|
||||
return False
|
||||
|
||||
monkeypatch.setattr(
|
||||
download_module,
|
||||
"_model_directory_has_files",
|
||||
counting_model_directory_has_files,
|
||||
)
|
||||
|
||||
result = await manager.check_pending_models(["lora"])
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["total_models"] == 1500
|
||||
assert result["pending_count"] == 1000
|
||||
assert result["needs_download"] is True
|
||||
# The per-folder check should not be used once we cross the threshold.
|
||||
assert per_model_checks == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("tmp_path")
|
||||
async def test_check_pending_models_uses_per_folder_check_for_small_candidate_sets(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
tmp_path,
|
||||
settings_manager,
|
||||
):
|
||||
"""For <=1000 candidates the pre-check keeps the accurate per-folder path."""
|
||||
|
||||
ws_manager = RecordingWebSocketManager()
|
||||
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||
|
||||
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||
|
||||
models = [
|
||||
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
|
||||
for i in range(500)
|
||||
]
|
||||
|
||||
# Create folders with files for the first 200 models
|
||||
for i in range(200):
|
||||
model_dir = tmp_path / f"{i:064x}"
|
||||
model_dir.mkdir()
|
||||
(model_dir / "image_0.png").write_text("data")
|
||||
|
||||
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
|
||||
|
||||
per_model_checks = 0
|
||||
original_has_files = download_module._model_directory_has_files
|
||||
|
||||
def counting_model_directory_has_files(path: str) -> bool:
|
||||
nonlocal per_model_checks
|
||||
per_model_checks += 1
|
||||
return original_has_files(path)
|
||||
|
||||
monkeypatch.setattr(
|
||||
download_module,
|
||||
"_model_directory_has_files",
|
||||
counting_model_directory_has_files,
|
||||
)
|
||||
|
||||
result = await manager.check_pending_models(["lora"])
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["total_models"] == 500
|
||||
assert result["pending_count"] == 300
|
||||
assert result["needs_download"] is True
|
||||
# Per-folder path should run once per candidate.
|
||||
assert per_model_checks == 500
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("tmp_path")
|
||||
async def test_check_pending_models_bulk_index_includes_legacy_folders(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
tmp_path,
|
||||
settings_manager,
|
||||
):
|
||||
"""In multi-library mode the bulk index also scans the legacy root so models
|
||||
whose folders have not been consolidated yet are not reported pending."""
|
||||
|
||||
ws_manager = RecordingWebSocketManager()
|
||||
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||
|
||||
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(tmp_path))
|
||||
monkeypatch.setitem(settings_manager.settings, "libraries", {"default": {}, "extra": {}})
|
||||
monkeypatch.setitem(settings_manager.settings, "active_library", "extra")
|
||||
|
||||
# 1500 unprocessed models triggers the bulk lookup path
|
||||
models = [
|
||||
{"sha256": f"{i:064x}", "model_name": f"Model {i}"}
|
||||
for i in range(1500)
|
||||
]
|
||||
|
||||
# Folders live at the LEGACY root/<hash> path (not yet consolidated)
|
||||
for i in range(500):
|
||||
model_dir = tmp_path / f"{i:064x}"
|
||||
model_dir.mkdir()
|
||||
(model_dir / "image_0.png").write_text("data")
|
||||
|
||||
_patch_scanners(monkeypatch, lora_scanner=StubScanner(models))
|
||||
|
||||
result = await manager.check_pending_models(["lora"])
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["total_models"] == 1500
|
||||
assert result["pending_count"] == 1000
|
||||
assert result["needs_download"] is True
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def settings_manager():
|
||||
return get_settings_manager()
|
||||
|
||||
@@ -35,9 +35,11 @@ class DummyDownloader:
|
||||
def reset_singletons():
|
||||
CivitaiClient._instance = None
|
||||
ModelMetadataProviderManager._instance = None
|
||||
civitai_client_module._creator_model_count_cache.clear()
|
||||
yield
|
||||
CivitaiClient._instance = None
|
||||
ModelMetadataProviderManager._instance = None
|
||||
civitai_client_module._creator_model_count_cache.clear()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -622,3 +624,162 @@ async def test_get_image_info_handles_invalid_id(monkeypatch, downloader, caplog
|
||||
|
||||
assert result is None
|
||||
assert "Invalid image ID format" in caplog.text
|
||||
|
||||
|
||||
async def test_get_user_models_requests_first_page_with_stable_params(downloader):
|
||||
request_calls = []
|
||||
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
request_calls.append({"method": method, "url": url, "kwargs": kwargs})
|
||||
return True, {
|
||||
"items": [
|
||||
{
|
||||
"id": 1,
|
||||
"modelVersions": [
|
||||
{"id": 100, "images": [{"meta": {"comfy": {"x": 1}}}]}
|
||||
],
|
||||
}
|
||||
],
|
||||
"metadata": {"nextCursor": "next-token"},
|
||||
}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
result = await client.get_user_models("pixel")
|
||||
|
||||
assert result is not None
|
||||
assert result["nextCursor"] == "next-token"
|
||||
assert len(result["items"]) == 1
|
||||
# comfy metadata is still stripped
|
||||
assert "comfy" not in result["items"][0]["modelVersions"][0]["images"][0]["meta"]
|
||||
|
||||
call = request_calls[0]
|
||||
assert call["method"] == "GET"
|
||||
assert call["url"] == "https://civitai.red/api/v1/models"
|
||||
params = call["kwargs"]["params"]
|
||||
assert params["username"] == "pixel"
|
||||
assert params["nsfw"] == "true"
|
||||
assert params["limit"] == 100
|
||||
assert params["sort"] == "Newest"
|
||||
assert params["period"] == "AllTime"
|
||||
assert "cursor" not in params
|
||||
|
||||
|
||||
async def test_get_user_models_passes_cursor_and_stringifies_next_cursor(downloader):
|
||||
request_calls = []
|
||||
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
request_calls.append(kwargs)
|
||||
return True, {"items": [], "metadata": {"nextCursor": 12345}}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
result = await client.get_user_models("pixel", cursor="opaque-token")
|
||||
|
||||
assert request_calls[0]["params"]["cursor"] == "opaque-token"
|
||||
assert result == {"items": [], "nextCursor": "12345"}
|
||||
|
||||
|
||||
async def test_get_user_models_without_next_cursor_returns_none_cursor(downloader):
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
return True, {"items": [{"id": 1, "modelVersions": []}], "metadata": {}}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
result = await client.get_user_models("pixel")
|
||||
|
||||
assert result == {"items": [{"id": 1, "modelVersions": []}], "nextCursor": None}
|
||||
|
||||
|
||||
async def test_get_user_models_failure_returns_none(downloader):
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
return False, "500 server error"
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
result = await client.get_user_models("pixel")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
async def test_get_creator_model_count_matches_exact_username(downloader):
|
||||
request_calls = []
|
||||
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
request_calls.append({"url": url, "kwargs": kwargs})
|
||||
return True, {
|
||||
"items": [
|
||||
{"username": "pixelart", "modelCount": 5},
|
||||
{"username": "Pixel", "modelCount": 2140},
|
||||
]
|
||||
}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
count = await client.get_creator_model_count("pixel")
|
||||
|
||||
assert count == 2140
|
||||
assert request_calls[0]["url"] == "https://civitai.red/api/v1/creators"
|
||||
assert request_calls[0]["kwargs"]["params"] == {"query": "pixel", "limit": 10}
|
||||
|
||||
|
||||
async def test_get_creator_model_count_without_exact_match_returns_none(downloader):
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
return True, {"items": [{"username": "pixelart", "modelCount": 5}]}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
count = await client.get_creator_model_count("pixel")
|
||||
|
||||
assert count is None
|
||||
|
||||
|
||||
async def test_get_creator_model_count_caches_results(downloader):
|
||||
request_count = 0
|
||||
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
nonlocal request_count
|
||||
request_count += 1
|
||||
return True, {"items": [{"username": "pixel", "modelCount": 42}]}
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
|
||||
assert await client.get_creator_model_count("pixel") == 42
|
||||
# case-insensitive cache key, second call served from cache
|
||||
assert await client.get_creator_model_count("Pixel") == 42
|
||||
assert request_count == 1
|
||||
|
||||
|
||||
async def test_get_creator_model_count_caches_failures(downloader):
|
||||
request_count = 0
|
||||
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
nonlocal request_count
|
||||
request_count += 1
|
||||
return False, "500 server error"
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
|
||||
assert await client.get_creator_model_count("pixel") is None
|
||||
assert await client.get_creator_model_count("pixel") is None
|
||||
assert request_count == 1
|
||||
|
||||
|
||||
async def test_get_creator_model_count_never_raises(downloader):
|
||||
async def fake_make_request(method, url, use_auth=True, **kwargs):
|
||||
return True, "unexpected non-dict payload"
|
||||
|
||||
downloader.make_request = fake_make_request
|
||||
|
||||
client = await CivitaiClient.get_instance()
|
||||
assert await client.get_creator_model_count("pixel") is None
|
||||
|
||||
@@ -1189,6 +1189,109 @@ def test_relative_path_sanitizes_model_and_version_placeholders():
|
||||
assert relative_path == "Fancy_Model/Version_One"
|
||||
|
||||
|
||||
def test_relative_path_empty_first_tag_fallback():
|
||||
"""Test that empty first_tag falls back to 'no tags'."""
|
||||
manager = DownloadManager()
|
||||
settings_manager = get_settings_manager()
|
||||
settings_manager.settings["download_path_templates"]["lora"] = (
|
||||
"{base_model}/{first_tag}"
|
||||
)
|
||||
|
||||
version_info = {
|
||||
"baseModel": "SDXL",
|
||||
"model": {"name": "Test Model", "tags": []},
|
||||
"creator": {"username": "Author"},
|
||||
}
|
||||
|
||||
relative_path = manager._calculate_relative_path(version_info, "lora")
|
||||
|
||||
assert relative_path == "SDXL/no tags"
|
||||
|
||||
|
||||
def test_relative_path_empty_base_model_and_first_tag():
|
||||
"""Test that empty base_model + empty first_tag does NOT produce a leading slash."""
|
||||
manager = DownloadManager()
|
||||
settings_manager = get_settings_manager()
|
||||
settings_manager.settings["download_path_templates"]["lora"] = (
|
||||
"{base_model}/{first_tag}"
|
||||
)
|
||||
|
||||
version_info = {
|
||||
"baseModel": "",
|
||||
"model": {"name": "Test Model", "tags": []},
|
||||
"creator": {"username": "Author"},
|
||||
}
|
||||
|
||||
relative_path = manager._calculate_relative_path(version_info, "lora")
|
||||
|
||||
assert not relative_path.startswith("/")
|
||||
assert relative_path == "no tags"
|
||||
|
||||
|
||||
def test_relative_path_sanitizes_double_slashes():
|
||||
"""Test that empty placeholder substitutions don't produce double slashes."""
|
||||
manager = DownloadManager()
|
||||
settings_manager = get_settings_manager()
|
||||
settings_manager.settings["download_path_templates"]["lora"] = (
|
||||
"{base_model}/{first_tag}/{author}"
|
||||
)
|
||||
|
||||
version_info = {
|
||||
"baseModel": "SDXL",
|
||||
"model": {"name": "Test Model", "tags": []},
|
||||
"creator": {"username": "Author"},
|
||||
}
|
||||
|
||||
relative_path = manager._calculate_relative_path(version_info, "lora")
|
||||
|
||||
assert "//" not in relative_path
|
||||
assert relative_path == "SDXL/no tags/Author"
|
||||
|
||||
|
||||
def test_download_containment_accepts_symlink_save_dir(tmp_path):
|
||||
"""Verify the download path containment check (download_manager.py:1395-1397)
|
||||
accepts save directories reached through user-created symlinks inside the
|
||||
library root — reproducing the symlink scenario from issue #1028."""
|
||||
# Library root with a symlink subdirectory pointing to an external drive
|
||||
lora_root = tmp_path / "loras"
|
||||
lora_root.mkdir()
|
||||
|
||||
external_drive = tmp_path / "external" / "models"
|
||||
external_drive.mkdir(parents=True)
|
||||
|
||||
symlink = lora_root / "Krea 2"
|
||||
symlink.symlink_to(str(external_drive))
|
||||
|
||||
# Simulate a download: base_save_dir = library root,
|
||||
# relative_path = "Krea 2/concept/NewModel"
|
||||
base_save_dir = str(lora_root)
|
||||
save_dir = os.path.join(base_save_dir, "Krea 2", "concept", "NewModel")
|
||||
|
||||
# Replicate the exact containment check from download_manager.py
|
||||
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
|
||||
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
|
||||
|
||||
# Must NOT be rejected — symlinks are legitimate business paths
|
||||
assert resolved_dir.startswith(base_dir + os.sep)
|
||||
|
||||
|
||||
def test_download_containment_rejects_dot_dot_traversal(tmp_path):
|
||||
"""Verify the download path containment check still blocks ``..`` traversal
|
||||
after the realpath → abspath change."""
|
||||
lora_root = tmp_path / "loras"
|
||||
lora_root.mkdir()
|
||||
|
||||
base_save_dir = str(lora_root)
|
||||
save_dir = os.path.join(base_save_dir, "..", "..", "etc", "passwd")
|
||||
|
||||
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
|
||||
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
|
||||
|
||||
# Must be rejected — dot-dot escapes the library root
|
||||
assert not resolved_dir.startswith(base_dir + os.sep)
|
||||
assert resolved_dir != base_dir
|
||||
|
||||
|
||||
def test_distribute_preview_to_entries_moves_and_copies(tmp_path):
|
||||
"""Test that preview distribution moves file to first entry and copies to others."""
|
||||
manager = DownloadManager()
|
||||
|
||||
@@ -26,6 +26,7 @@ class StubScanner:
|
||||
|
||||
def __init__(self, models: list[dict]) -> None:
|
||||
self._cache = SimpleNamespace(raw_data=models)
|
||||
self.sync_calls: list[tuple[str, dict]] = []
|
||||
|
||||
async def get_cached_data(self):
|
||||
return self._cache
|
||||
@@ -38,6 +39,14 @@ class StubScanner:
|
||||
break
|
||||
return True
|
||||
|
||||
async def sync_cache_from_metadata(self, file_path: str, metadata: dict) -> bool:
|
||||
self.sync_calls.append((file_path, metadata))
|
||||
for index, model in enumerate(self._cache.raw_data):
|
||||
if model.get("file_path") == metadata.get("file_path"):
|
||||
self._cache.raw_data[index] = metadata
|
||||
break
|
||||
return True
|
||||
|
||||
|
||||
def _patch_scanner(monkeypatch: pytest.MonkeyPatch, scanner: StubScanner) -> None:
|
||||
async def _get_lora_scanner(cls):
|
||||
@@ -520,7 +529,8 @@ async def test_not_found_example_images_are_cleaned(
|
||||
|
||||
model_dir = images_root / model_hash
|
||||
model_dir.mkdir(parents=True, exist_ok=True)
|
||||
(model_dir / "image_0.png").write_bytes(b"first")
|
||||
# Pre-existing file collides with the valid image index (1) so the
|
||||
# pre-download existence check must skip it without a network request
|
||||
(model_dir / "image_1.png").write_bytes(b"second")
|
||||
|
||||
async def fake_process_local_examples(*_args, **_kwargs):
|
||||
@@ -588,6 +598,9 @@ async def test_not_found_example_images_are_cleaned(
|
||||
assert missing_url in downloader.calls
|
||||
assert manager._progress["failed_models"] == {model_hash}
|
||||
assert model_hash in manager._progress["processed_models"]
|
||||
assert scanner.sync_calls
|
||||
assert len(scanner.sync_calls) == 1
|
||||
assert scanner.sync_calls[0][0] == str(model_path)
|
||||
|
||||
remaining_images = model_metadata["civitai"]["images"]
|
||||
assert remaining_images == [
|
||||
@@ -596,11 +609,188 @@ async def test_not_found_example_images_are_cleaned(
|
||||
]
|
||||
|
||||
files = sorted(p.name for p in model_dir.iterdir())
|
||||
assert files == ["image_0.png", "image_1.png"]
|
||||
assert (model_dir / "image_0.png").read_bytes() == b"first"
|
||||
assert files == ["image_1.png"]
|
||||
assert (model_dir / "image_1.png").read_bytes() == b"second"
|
||||
|
||||
|
||||
async def test_failed_models_retried_when_explicitly_targeted(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
tmp_path,
|
||||
settings_manager,
|
||||
):
|
||||
ws_manager = RecordingWebSocketManager()
|
||||
manager = download_module.DownloadManager(ws_manager=ws_manager)
|
||||
|
||||
images_root = tmp_path / "examples"
|
||||
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(images_root))
|
||||
|
||||
model_hash = "a" * 64
|
||||
model_path = tmp_path / "model.safetensors"
|
||||
model_path.write_text("data", encoding="utf-8")
|
||||
|
||||
model_metadata = {
|
||||
"sha256": model_hash,
|
||||
"model_name": "Failed Example",
|
||||
"file_path": str(model_path),
|
||||
"file_name": "model.safetensors",
|
||||
"civitai": {"images": [{"url": "https://example.com/valid.png"}]},
|
||||
}
|
||||
|
||||
scanner = StubScanner([model_metadata.copy()])
|
||||
_patch_scanner(monkeypatch, scanner)
|
||||
|
||||
# Persist a previous failure so the skip path is exercised
|
||||
images_root.mkdir(parents=True, exist_ok=True)
|
||||
(images_root / ".download_progress.json").write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"failed_models": [model_hash],
|
||||
"processed_models": [],
|
||||
"rate_limited_models": [],
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
async def fake_process_local_examples(*_args, **_kwargs):
|
||||
return False
|
||||
|
||||
async def fake_get_updated_model(model_hash_arg, _scanner):
|
||||
return model_metadata
|
||||
|
||||
class DownloaderStub:
|
||||
def __init__(self):
|
||||
self.calls: list[str] = []
|
||||
|
||||
async def download_to_memory(self, url, *_args, **_kwargs):
|
||||
self.calls.append(url)
|
||||
return True, b"\x89PNG\r\n\x1a\n", {"content-type": "image/png"}
|
||||
|
||||
downloader = DownloaderStub()
|
||||
|
||||
async def fake_get_downloader():
|
||||
return downloader
|
||||
|
||||
monkeypatch.setattr(
|
||||
download_module.ExampleImagesProcessor,
|
||||
"process_local_examples",
|
||||
staticmethod(fake_process_local_examples),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
download_module.MetadataUpdater,
|
||||
"get_updated_model",
|
||||
staticmethod(fake_get_updated_model),
|
||||
)
|
||||
monkeypatch.setattr(download_module, "get_downloader", fake_get_downloader)
|
||||
|
||||
# Without explicit hashes the previously failed model is skipped
|
||||
skipped_manager = download_module.DownloadManager(ws_manager=RecordingWebSocketManager())
|
||||
result = await skipped_manager.start_download({"model_types": ["lora"], "delay": 0})
|
||||
assert result["success"] is True
|
||||
if skipped_manager._download_task is not None:
|
||||
await asyncio.wait_for(skipped_manager._download_task, timeout=1)
|
||||
assert downloader.calls == []
|
||||
|
||||
# With explicit hashes the previously failed model is retried and cleared
|
||||
result = await manager.start_download(
|
||||
{"model_types": ["lora"], "delay": 0, "model_hashes": [model_hash]}
|
||||
)
|
||||
assert result["success"] is True
|
||||
if manager._download_task is not None:
|
||||
await asyncio.wait_for(manager._download_task, timeout=1)
|
||||
assert downloader.calls == ["https://example.com/valid.png"]
|
||||
assert manager._progress["failed_models"] == set()
|
||||
assert model_hash in manager._progress["processed_models"]
|
||||
|
||||
|
||||
async def test_explicit_targets_fill_partial_example_gaps(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
tmp_path,
|
||||
settings_manager,
|
||||
):
|
||||
ws_manager = RecordingWebSocketManager()
|
||||
|
||||
images_root = tmp_path / "examples"
|
||||
monkeypatch.setitem(settings_manager.settings, "example_images_path", str(images_root))
|
||||
|
||||
model_hash = "b" * 64
|
||||
model_path = tmp_path / "model.safetensors"
|
||||
model_path.write_text("data", encoding="utf-8")
|
||||
|
||||
model_metadata = {
|
||||
"sha256": model_hash,
|
||||
"model_name": "Partial Example",
|
||||
"file_path": str(model_path),
|
||||
"file_name": "model.safetensors",
|
||||
"civitai": {
|
||||
"images": [
|
||||
{"url": "https://example.com/first.png"},
|
||||
{"url": "https://example.com/second.png"},
|
||||
]
|
||||
},
|
||||
}
|
||||
|
||||
scanner = StubScanner([model_metadata.copy()])
|
||||
_patch_scanner(monkeypatch, scanner)
|
||||
|
||||
# Simulate a partially populated folder: index 0 already downloaded
|
||||
model_dir = images_root / model_hash
|
||||
model_dir.mkdir(parents=True, exist_ok=True)
|
||||
(model_dir / "image_0.png").write_bytes(b"existing")
|
||||
|
||||
async def fake_process_local_examples(*_args, **_kwargs):
|
||||
return False
|
||||
|
||||
async def fake_get_updated_model(model_hash_arg, _scanner):
|
||||
return model_metadata
|
||||
|
||||
class DownloaderStub:
|
||||
def __init__(self):
|
||||
self.calls: list[str] = []
|
||||
|
||||
async def download_to_memory(self, url, *_args, **_kwargs):
|
||||
self.calls.append(url)
|
||||
return True, b"\x89PNG\r\n\x1a\n", {"content-type": "image/png"}
|
||||
|
||||
downloader = DownloaderStub()
|
||||
|
||||
async def fake_get_downloader():
|
||||
return downloader
|
||||
|
||||
monkeypatch.setattr(
|
||||
download_module.ExampleImagesProcessor,
|
||||
"process_local_examples",
|
||||
staticmethod(fake_process_local_examples),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
download_module.MetadataUpdater,
|
||||
"get_updated_model",
|
||||
staticmethod(fake_get_updated_model),
|
||||
)
|
||||
monkeypatch.setattr(download_module, "get_downloader", fake_get_downloader)
|
||||
|
||||
# Untargeted run treats the populated folder as done
|
||||
untargeted = download_module.DownloadManager(ws_manager=RecordingWebSocketManager())
|
||||
result = await untargeted.start_download({"model_types": ["lora"], "delay": 0})
|
||||
assert result["success"] is True
|
||||
if untargeted._download_task is not None:
|
||||
await asyncio.wait_for(untargeted._download_task, timeout=1)
|
||||
assert downloader.calls == []
|
||||
|
||||
# Explicitly targeted run fills only the missing index, skipping the
|
||||
# existing file without a network request
|
||||
targeted = download_module.DownloadManager(ws_manager=ws_manager)
|
||||
result = await targeted.start_download(
|
||||
{"model_types": ["lora"], "delay": 0, "model_hashes": [model_hash]}
|
||||
)
|
||||
assert result["success"] is True
|
||||
if targeted._download_task is not None:
|
||||
await asyncio.wait_for(targeted._download_task, timeout=1)
|
||||
assert downloader.calls == ["https://example.com/second.png"]
|
||||
assert (model_dir / "image_1.png").exists()
|
||||
assert (model_dir / "image_0.png").read_bytes() == b"existing"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def settings_manager():
|
||||
return get_settings_manager()
|
||||
|
||||
@@ -243,6 +243,56 @@ class TestLLMServiceChatCompletionJson:
|
||||
|
||||
assert result == {"key": "value"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_json_falls_back_on_response_format_rejection(
|
||||
self, llm_service,
|
||||
):
|
||||
"""Retry without response_format when provider rejects it (HTTP 400)."""
|
||||
error_response = MockResponse(
|
||||
400,
|
||||
text_data=(
|
||||
'{"error":"\'response_format.type\' must be '
|
||||
'\'json_schema\' or \'text\'"}'
|
||||
),
|
||||
)
|
||||
success_response = MockResponse(
|
||||
200,
|
||||
json_data={
|
||||
"choices": [{"message": {"content": '{"key": "value"}'}}],
|
||||
"usage": {},
|
||||
"model": "local-model",
|
||||
},
|
||||
)
|
||||
|
||||
call_index = 0
|
||||
|
||||
class FallbackMockSession:
|
||||
def __init__(self):
|
||||
self.last_url = None
|
||||
self.last_json = None
|
||||
|
||||
def post(self, url, json=None, headers=None):
|
||||
nonlocal call_index
|
||||
self.last_url = url
|
||||
self.last_json = json
|
||||
call_index += 1
|
||||
return error_response if call_index == 1 else success_response
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
pass
|
||||
|
||||
with mock.patch("aiohttp.ClientSession", return_value=FallbackMockSession()):
|
||||
result = await llm_service.chat_completion_json(
|
||||
system_prompt="You are helpful.",
|
||||
user_prompt="Return JSON.",
|
||||
)
|
||||
|
||||
assert result == {"key": "value"}
|
||||
assert call_index == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_json_raises_on_non_json(self, llm_service):
|
||||
# Non-JSON content raises LLMResponseError (salvage also fails)
|
||||
|
||||
@@ -1,13 +1,181 @@
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from py.services.model_lifecycle_service import ModelLifecycleService
|
||||
from py.services.model_lifecycle_service import ModelLifecycleService, _require_path_in_library_roots
|
||||
from py.utils.metadata_manager import MetadataManager
|
||||
from py.utils.models import LoraMetadata
|
||||
|
||||
|
||||
class ScannerWithRoots:
|
||||
def __init__(self, roots):
|
||||
self._roots = list(roots)
|
||||
|
||||
def get_model_roots(self):
|
||||
return self._roots
|
||||
|
||||
|
||||
class TestRequirePathInLibraryRoots:
|
||||
def test_accepts_path_within_root(self, tmp_path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
model = root / "model.safetensors"
|
||||
model.write_text("")
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
_require_path_in_library_roots(str(model), scanner)
|
||||
|
||||
def test_rejects_path_outside_roots(self, tmp_path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
outside = tmp_path / "outside" / "model.safetensors"
|
||||
outside.parent.mkdir(parents=True)
|
||||
outside.write_text("")
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
with pytest.raises(ValueError, match="outside configured library"):
|
||||
_require_path_in_library_roots(str(outside), scanner)
|
||||
|
||||
def test_passes_when_no_roots_configured(self, tmp_path):
|
||||
f = tmp_path / "model.safetensors"
|
||||
f.write_text("")
|
||||
|
||||
scanner = ScannerWithRoots([])
|
||||
_require_path_in_library_roots(str(f), scanner)
|
||||
|
||||
def test_accepts_path_matching_root_exactly(self, tmp_path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
_require_path_in_library_roots(str(root), scanner)
|
||||
|
||||
def test_accepts_symlink_within_root(self, tmp_path):
|
||||
"""Symlinks under a configured root are legitimate business paths
|
||||
and should be accepted — containment works on business-path space,
|
||||
not resolved physical paths."""
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
|
||||
outside_dir = tmp_path / "outside"
|
||||
outside_dir.mkdir()
|
||||
outside_file = outside_dir / "escaped.safetensors"
|
||||
outside_file.write_text("")
|
||||
|
||||
symlink = root / "link.safetensors"
|
||||
symlink.symlink_to(outside_file)
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
# Symlink path is under root in business-path space → accepted
|
||||
_require_path_in_library_roots(str(symlink), scanner)
|
||||
|
||||
def test_rejects_dot_dot_traversal(self, tmp_path):
|
||||
"""Verify that ``..`` components are still resolved and blocked —
|
||||
``abspath`` normalises dot-dot but does not resolve symlinks."""
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
|
||||
# A path that traverses up out of the root via ..
|
||||
escaped = os.path.join(str(root), "..", "..", "etc", "passwd")
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
with pytest.raises(ValueError, match="outside configured library"):
|
||||
_require_path_in_library_roots(escaped, scanner)
|
||||
|
||||
|
||||
class ScannerForDelete:
|
||||
def __init__(self, raw_data, roots, model_type="lora"):
|
||||
self.model_type = model_type
|
||||
self.cache = DummyCache(raw_data)
|
||||
self._hash_index = DummyHashIndex()
|
||||
self._roots = list(roots)
|
||||
self._persist_calls = []
|
||||
|
||||
def get_model_roots(self):
|
||||
return self._roots
|
||||
|
||||
async def get_cached_data(self):
|
||||
return self.cache
|
||||
|
||||
async def _persist_current_cache(self):
|
||||
self._persist_calls.append(True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_model_rejects_path_outside_roots(tmp_path: Path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
model = root / "model.safetensors"
|
||||
model.write_bytes(b"data")
|
||||
|
||||
scanner = ScannerForDelete(
|
||||
raw_data=[{"file_path": str(model)}],
|
||||
roots=[str(root)],
|
||||
)
|
||||
service = ModelLifecycleService(
|
||||
scanner=scanner,
|
||||
metadata_manager=DummyMetadataManager({"civitai": {"modelId": 1}}),
|
||||
metadata_loader=lambda x: {},
|
||||
)
|
||||
# Path within root should work (model file exists)
|
||||
result = await service.delete_model(str(model))
|
||||
assert result["success"] is True
|
||||
|
||||
# Path outside root should be rejected
|
||||
outside = tmp_path / "outside.safetensors"
|
||||
outside.write_bytes(b"data")
|
||||
scanner2 = ScannerForDelete(
|
||||
raw_data=[],
|
||||
roots=[str(root)],
|
||||
)
|
||||
service2 = ModelLifecycleService(
|
||||
scanner=scanner2,
|
||||
metadata_manager=DummyMetadataManager({}),
|
||||
metadata_loader=lambda x: {},
|
||||
)
|
||||
with pytest.raises(ValueError, match="outside configured library"):
|
||||
await service2.delete_model(str(outside))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rename_model_rejects_path_outside_roots(tmp_path: Path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
service = ModelLifecycleService(
|
||||
scanner=scanner,
|
||||
metadata_manager=DummyMetadataManager({}),
|
||||
metadata_loader=lambda x: {},
|
||||
)
|
||||
outside = tmp_path / "outside.safetensors"
|
||||
outside.write_bytes(b"data")
|
||||
|
||||
with pytest.raises(ValueError, match="outside configured library"):
|
||||
await service.rename_model(file_path=str(outside), new_file_name="new_name")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_delete_rejects_any_path_outside_roots(tmp_path: Path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
model_ok = root / "model.safetensors"
|
||||
model_ok.write_bytes(b"data")
|
||||
outside = tmp_path / "outside.safetensors"
|
||||
outside.write_bytes(b"data")
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
service = ModelLifecycleService(
|
||||
scanner=scanner,
|
||||
metadata_manager=DummyMetadataManager({}),
|
||||
metadata_loader=lambda x: {},
|
||||
)
|
||||
with pytest.raises(ValueError, match="outside configured library"):
|
||||
await service.bulk_delete_models([str(model_ok), str(outside)])
|
||||
|
||||
|
||||
class DummyCache:
|
||||
def __init__(self, raw_data):
|
||||
self.raw_data = raw_data
|
||||
|
||||
@@ -884,7 +884,7 @@ async def test_sync_cache_conditional_resort_skipped(tmp_path: Path, monkeypatch
|
||||
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
|
||||
)
|
||||
await scanner._cache.resort()
|
||||
scanner._cache._last_sort = ("name", "asc") # name sort is active
|
||||
scanner._cache._last_sort = ("name", "asc", None) # name sort is active
|
||||
scanner._tags_count = {"alpha": 1}
|
||||
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
|
||||
|
||||
@@ -935,7 +935,7 @@ async def test_sync_cache_conditional_resort_triggered(tmp_path: Path, monkeypat
|
||||
raw_data=[dict(entry)], folders=[], name_display_mode="model_name"
|
||||
)
|
||||
await scanner._cache.resort()
|
||||
scanner._cache._last_sort = ("name", "asc")
|
||||
scanner._cache._last_sort = ("name", "asc", None)
|
||||
scanner._tags_count = {"alpha": 1}
|
||||
scanner._hash_index.add_entry("abc123", "/m/a.safetensors")
|
||||
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Tests for sort parsing and the seeded random sort mode."""
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from py.services.model_cache import ModelCache
|
||||
from py.services.model_query import ModelCacheRepository, SortParams
|
||||
|
||||
|
||||
def _make_cache(items):
|
||||
return ModelCache(
|
||||
raw_data=[
|
||||
{
|
||||
"file_path": f"/models/{name}.safetensors",
|
||||
"file_name": f"{name}.safetensors",
|
||||
"model_name": name,
|
||||
"folder": "",
|
||||
"size": 100,
|
||||
"modified": 0.0,
|
||||
}
|
||||
for name in items
|
||||
],
|
||||
folders=[],
|
||||
)
|
||||
|
||||
|
||||
class TestParseSort:
|
||||
def test_random_with_seed(self):
|
||||
params = ModelCacheRepository.parse_sort("random:abc123")
|
||||
assert params == SortParams(key="random", order="asc", seed="abc123")
|
||||
|
||||
def test_random_without_seed(self):
|
||||
params = ModelCacheRepository.parse_sort("random")
|
||||
assert params == SortParams(key="random", order="asc", seed=None)
|
||||
|
||||
def test_random_empty_seed_falls_back_to_none(self):
|
||||
params = ModelCacheRepository.parse_sort("random:")
|
||||
assert params.seed is None
|
||||
|
||||
def test_regular_sorts_unaffected(self):
|
||||
params = ModelCacheRepository.parse_sort("name:desc")
|
||||
assert params == SortParams(key="name", order="desc", seed=None)
|
||||
|
||||
|
||||
class TestRandomShuffle:
|
||||
@pytest.mark.asyncio
|
||||
async def test_same_seed_yields_same_order(self):
|
||||
cache = _make_cache(["a", "b", "c", "d", "e"])
|
||||
await asyncio.sleep(0) # allow background resort task to run
|
||||
|
||||
first = await cache.get_sorted_data("random", "asc", "seed1")
|
||||
second = await cache.get_sorted_data("random", "asc", "seed1")
|
||||
|
||||
assert [item["model_name"] for item in first] == [
|
||||
item["model_name"] for item in second
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_different_seeds_yield_different_orders(self):
|
||||
cache = _make_cache([f"m{i}" for i in range(20)])
|
||||
await asyncio.sleep(0)
|
||||
|
||||
first = await cache.get_sorted_data("random", "asc", "seed-a")
|
||||
second = await cache.get_sorted_data("random", "asc", "seed-b")
|
||||
|
||||
assert [item["model_name"] for item in first] != [
|
||||
item["model_name"] for item in second
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shuffle_is_a_permutation(self):
|
||||
cache = _make_cache(["a", "b", "c", "d", "e"])
|
||||
await asyncio.sleep(0)
|
||||
|
||||
shuffled = await cache.get_sorted_data("random", "asc", "seed")
|
||||
|
||||
assert sorted(item["model_name"] for item in shuffled) == [
|
||||
"a",
|
||||
"b",
|
||||
"c",
|
||||
"d",
|
||||
"e",
|
||||
]
|
||||
assert len({item["file_path"] for item in shuffled}) == 5
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_seed_is_stable(self):
|
||||
cache = _make_cache(["a", "b", "c", "d", "e"])
|
||||
await asyncio.sleep(0)
|
||||
|
||||
first = await cache.get_sorted_data("random", "asc")
|
||||
second = await cache.get_sorted_data("random", "asc")
|
||||
|
||||
assert [item["model_name"] for item in first] == [
|
||||
item["model_name"] for item in second
|
||||
]
|
||||
@@ -77,7 +77,15 @@ def recipe_scanner(tmp_path: Path, monkeypatch):
|
||||
monkeypatch.setattr(config, "loras_roots", [str(tmp_path)])
|
||||
stub = StubLoraScanner()
|
||||
scanner = RecipeScanner(lora_scanner=stub)
|
||||
asyncio.run(scanner.refresh_cache(force=True))
|
||||
|
||||
async def _init():
|
||||
await scanner.refresh_cache(force=True)
|
||||
# Wait for FTS index build to finish — asyncio.run()
|
||||
# cancels background tasks on return, so we must await it here.
|
||||
if scanner._fts_index_task:
|
||||
await scanner._fts_index_task
|
||||
|
||||
asyncio.run(_init())
|
||||
yield scanner, stub
|
||||
RecipeScanner._instance = None
|
||||
settings_manager_module.reset_settings_manager()
|
||||
|
||||
@@ -139,3 +139,122 @@ def test_contains_dynamic_syntax_detects_wildcards_and_options():
|
||||
assert contains_dynamic_syntax("__flower__") is True
|
||||
assert contains_dynamic_syntax("{red|blue}") is True
|
||||
assert contains_dynamic_syntax("{2$$, $$red|blue|green}") is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _pick_weighted_or_plain
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_pick_weighted_or_plain_plain_values(monkeypatch, tmp_path):
|
||||
"""Plain values without :: are picked via rng.choice (fast path)."""
|
||||
service, _ = _make_service(monkeypatch, tmp_path)
|
||||
|
||||
import random
|
||||
rng = random.Random(42)
|
||||
|
||||
result = service._pick_weighted_or_plain(["red", "green", "blue"], rng)
|
||||
assert result in {"red", "green", "blue"}
|
||||
assert "::" not in result
|
||||
|
||||
|
||||
def test_pick_weighted_or_plain_deterministic_with_seed(monkeypatch, tmp_path):
|
||||
"""Same seed produces the same result for plain values."""
|
||||
service, _ = _make_service(monkeypatch, tmp_path)
|
||||
|
||||
import random
|
||||
first = service._pick_weighted_or_plain(["a", "b", "c"], random.Random(99))
|
||||
second = service._pick_weighted_or_plain(["a", "b", "c"], random.Random(99))
|
||||
assert first == second
|
||||
|
||||
|
||||
def test_pick_weighted_or_plain_weighted_values(monkeypatch, tmp_path):
|
||||
"""Weighted values use weighted selection and strip the N:: prefix."""
|
||||
service, _ = _make_service(monkeypatch, tmp_path)
|
||||
|
||||
import random
|
||||
values = ["3::apple", "1::banana"]
|
||||
results = {"apple": 0, "banana": 0}
|
||||
for seed in range(4000):
|
||||
result = service._pick_weighted_or_plain(values, random.Random(seed))
|
||||
assert result in results, f"Unexpected result: {result!r}"
|
||||
assert "::" not in result
|
||||
results[result] += 1
|
||||
|
||||
total = results["apple"] + results["banana"]
|
||||
# 3:1 weight → apple ≈ 75%, banana ≈ 25%
|
||||
assert 2700 < results["apple"] < 3300, f"apple count out of range: {results['apple']}"
|
||||
assert 700 < results["banana"] < 1300, f"banana count out of range: {results['banana']}"
|
||||
|
||||
|
||||
def test_pick_weighted_or_plain_weight_one_values(monkeypatch, tmp_path):
|
||||
"""Values with explicit 1:: prefix have prefix stripped but are not weighted."""
|
||||
service, _ = _make_service(monkeypatch, tmp_path)
|
||||
|
||||
import random
|
||||
# All weights are 1.0 → no actual weighting, but :: prefix is stripped
|
||||
values = ["1::foo", "1::bar"]
|
||||
rng = random.Random(42)
|
||||
results = {service._pick_weighted_or_plain(values, rng) for _ in range(200)}
|
||||
assert results == {"foo", "bar"}
|
||||
# Ensure the prefix is always stripped
|
||||
for result in results:
|
||||
assert "::" not in result
|
||||
|
||||
|
||||
def test_pick_weighted_or_plain_mixed_weighted_and_plain(monkeypatch, tmp_path):
|
||||
"""Mixed list with some weighted and some unweighted values."""
|
||||
service, _ = _make_service(monkeypatch, tmp_path)
|
||||
|
||||
import random
|
||||
values = ["5::x", "y", "z"] # x has weight 5, y/z have default weight 1
|
||||
results = {"x": 0, "y": 0, "z": 0}
|
||||
for seed in range(4000):
|
||||
result = service._pick_weighted_or_plain(values, random.Random(seed))
|
||||
assert result in results
|
||||
assert "::" not in result
|
||||
results[result] += 1
|
||||
|
||||
# x (5) vs combined y+z (1+1=2) → ~71% / ~29%
|
||||
x_pct = results["x"] / sum(results.values())
|
||||
assert 0.65 < x_pct < 0.78, f"x proportion out of range: {x_pct:.3f}"
|
||||
|
||||
|
||||
def test_pick_weighted_or_plain_invalid_weight_prefix(monkeypatch, tmp_path):
|
||||
"""Invalid numeric prefix (e.g. 1.2.3) is NOT treated as a weight and
|
||||
the prefix is NOT stripped, matching the updated strict regex."""
|
||||
service, _ = _make_service(monkeypatch, tmp_path)
|
||||
|
||||
import random
|
||||
rng = random.Random(42)
|
||||
|
||||
# "1.2.3::a" is not a valid number → treated as plain text value
|
||||
result = service._pick_weighted_or_plain(["1.2.3::a", "b"], rng)
|
||||
# It should keep the full text including :: because the prefix isn't a
|
||||
# valid numeric weight according to the strict regex
|
||||
assert result == "1.2.3::a" or result == "b"
|
||||
|
||||
|
||||
def test_pick_weighted_or_plain_glob_aggregation(monkeypatch, tmp_path):
|
||||
"""Weighted wildcard resolution through glob aggregation (__*__)."""
|
||||
service, wildcards_dir = _make_service(monkeypatch, tmp_path)
|
||||
wildcards_dir.mkdir()
|
||||
|
||||
(wildcards_dir / "animals").mkdir()
|
||||
(wildcards_dir / "animals" / "cat.txt").write_text("3::tabby\n1::persian\n", encoding="utf-8")
|
||||
(wildcards_dir / "animals" / "dog.txt").write_text("retriever\npoodle\n", encoding="utf-8")
|
||||
|
||||
# __animals/*__ aggregates all values across both files
|
||||
# Weighted values should have :: stripped
|
||||
results = {"tabby": 0, "persian": 0, "retriever": 0, "poodle": 0}
|
||||
for seed in range(4000):
|
||||
expanded = service.expand_text("__animals/*__", seed=seed)
|
||||
assert expanded in results, f"Unexpected result: {expanded!r}"
|
||||
assert "::" not in expanded
|
||||
results[expanded] += 1
|
||||
|
||||
# tabby (3) vs persian (1) → ~75% / ~25% within the cat subset
|
||||
cat_total = results["tabby"] + results["persian"]
|
||||
if cat_total > 0:
|
||||
tabby_pct = results["tabby"] / cat_total
|
||||
assert 0.65 < tabby_pct < 0.85, f"tabby proportion out of range: {tabby_pct:.3f}"
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user