mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-06 22:10:14 -03:00
Compare commits
16 Commits
v1.1.9
...
f49b4ba4db
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f49b4ba4db | ||
|
|
84e708328b | ||
|
|
125bed3f09 | ||
|
|
077e70169d | ||
|
|
e6dc169a05 | ||
|
|
f34c02756d | ||
|
|
1e4c315481 | ||
|
|
a8283a0d00 | ||
|
|
55896669fc | ||
|
|
e341e0b9d2 | ||
|
|
e6538c83bb | ||
|
|
92e1285ea5 | ||
|
|
2aabd1d90e | ||
|
|
7b8b778f83 | ||
|
|
7c8dc57d55 | ||
|
|
fe95fae5f2 |
@@ -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
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ try: # pragma: no cover - import fallback for pytest collection
|
||||
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
|
||||
@@ -66,6 +67,9 @@ except (
|
||||
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 = {
|
||||
@@ -88,6 +92,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
LoraInfoLM.NAME: LoraInfoLM,
|
||||
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
|
||||
CreateHookLoraLM.NAME: CreateHookLoraLM,
|
||||
MetadataOverwriteLM.NAME: MetadataOverwriteLM,
|
||||
}
|
||||
|
||||
WEB_DIRECTORY = "./web/comfyui"
|
||||
|
||||
@@ -9,6 +9,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,21 @@ 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 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,7 @@ 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, METADATA_OVERWRITE_FIELDS
|
||||
|
||||
|
||||
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
||||
@@ -31,11 +31,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 +1221,32 @@ 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 = {}
|
||||
for key in METADATA_OVERWRITE_FIELDS:
|
||||
value = inputs.get(key)
|
||||
if value: # truthy — only overwrite when user provided a real value
|
||||
overwrite_params[key] = value
|
||||
|
||||
if overwrite_params:
|
||||
metadata.setdefault(OVERWRITE, {})
|
||||
metadata[OVERWRITE][node_id] = {
|
||||
"parameters": overwrite_params,
|
||||
"node_id": node_id,
|
||||
}
|
||||
|
||||
|
||||
# Registry of node-specific extractors
|
||||
# Keys are node class names
|
||||
NODE_EXTRACTORS = {
|
||||
@@ -1221,5 +1314,7 @@ NODE_EXTRACTORS = {
|
||||
"CFGGuider": CFGGuiderExtractor, # Add CFGGuider
|
||||
# Image
|
||||
"VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor
|
||||
# Metadata overwrite
|
||||
"MetadataOverwriteLM": MetadataOverwriteExtractor,
|
||||
# Add other nodes as needed
|
||||
}
|
||||
|
||||
157
py/nodes/metadata_overwrite.py
Normal file
157
py/nodes/metadata_overwrite.py
Normal file
@@ -0,0 +1,157 @@
|
||||
"""Metadata Overwrite node — allows users to manually specify generation parameters
|
||||
that override the automatically collected/inferred metadata.
|
||||
|
||||
All inputs have falsy defaults: only truthy (non-empty / non-zero) values
|
||||
will overwrite the corresponding field in the final metadata.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ..metadata_collector.constants import METADATA_OVERWRITE_FIELDS
|
||||
|
||||
|
||||
class MetadataOverwriteLM:
|
||||
NAME = "Metadata Overwrite (LoraManager)"
|
||||
CATEGORY = "Lora Manager/utils"
|
||||
DESCRIPTION = (
|
||||
"Manually specify generation parameters to override automatically collected "
|
||||
"metadata. Only filled/connected inputs will take effect — empty defaults "
|
||||
"are ignored."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||
return {
|
||||
"optional": {
|
||||
"prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Positive prompt. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"negative_prompt": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": "Negative prompt. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"seed": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 0xFFFFFFFFFFFFFFFF,
|
||||
"control_after_generate": False,
|
||||
"tooltip": "Seed value. Only overwrites when > 0.",
|
||||
},
|
||||
),
|
||||
"steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 10000,
|
||||
"tooltip": "Number of steps. Only overwrites when > 0.",
|
||||
},
|
||||
),
|
||||
"cfg_scale": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0,
|
||||
"tooltip": "CFG scale. Only overwrites when > 0.",
|
||||
},
|
||||
),
|
||||
"sampler": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Sampler name. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"scheduler": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": "Scheduler name. Only overwrites when non-empty.",
|
||||
},
|
||||
),
|
||||
"model": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"The checkpoint or diffusion model (UNet) used "
|
||||
"for generation. Only overwrites when non-empty."
|
||||
),
|
||||
},
|
||||
),
|
||||
"loras": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"multiline": True,
|
||||
"tooltip": (
|
||||
"LoRA syntax, e.g. <lora:name:strength> "
|
||||
"or <lora:name:model_strength:clip_strength>, "
|
||||
"separated by spaces. Only overwrites when non-empty."
|
||||
),
|
||||
},
|
||||
),
|
||||
"size": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"tooltip": (
|
||||
"Image size in WIDTHxHEIGHT format (e.g. 512x768). "
|
||||
"Only overwrites when non-empty."
|
||||
),
|
||||
},
|
||||
),
|
||||
"clip_skip": (
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": -24,
|
||||
"max": 24,
|
||||
"tooltip": "Clip skip. Only overwrites when non-zero.",
|
||||
},
|
||||
),
|
||||
"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-falsy input values into a metadata dict.
|
||||
|
||||
Only values that are truthy (non-empty string, non-zero number)
|
||||
are included — matching the overwrite logic in the metadata pipeline.
|
||||
"""
|
||||
result: dict[str, Any] = {}
|
||||
for key in METADATA_OVERWRITE_FIELDS:
|
||||
value = kwargs.get(key)
|
||||
if value:
|
||||
result[key] = value
|
||||
return (result,)
|
||||
@@ -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,184 @@ 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_display = self._get_civitai_sampler_name(sampler, scheduler)
|
||||
|
||||
# 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 = []
|
||||
params: list[str] = []
|
||||
if steps is not None:
|
||||
params.append(f"Steps: {steps}")
|
||||
if sampler_display:
|
||||
params.append(f"Sampler: {sampler_display}")
|
||||
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:
|
||||
try:
|
||||
cs = int(clip_skip)
|
||||
if cs != 0:
|
||||
params.append(f"Clip skip: {abs(cs)}")
|
||||
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=(',', ':'))}"
|
||||
)
|
||||
|
||||
# Add standard parameters in the correct order
|
||||
if "steps" in metadata_dict:
|
||||
add_param_if_not_none(params, "Steps", metadata_dict.get("steps"))
|
||||
|
||||
# Combine sampler and scheduler information
|
||||
sampler_name = None
|
||||
scheduler_name = None
|
||||
|
||||
if "sampler" in metadata_dict:
|
||||
sampler = metadata_dict.get("sampler")
|
||||
# Convert ComfyUI sampler names to user-friendly names
|
||||
sampler_mapping = {
|
||||
"euler": "Euler",
|
||||
"euler_ancestral": "Euler a",
|
||||
"dpm_2": "DPM2",
|
||||
"dpm_2_ancestral": "DPM2 a",
|
||||
"heun": "Heun",
|
||||
"dpm_fast": "DPM fast",
|
||||
"dpm_adaptive": "DPM adaptive",
|
||||
"lms": "LMS",
|
||||
"dpmpp_2s_ancestral": "DPM++ 2S a",
|
||||
"dpmpp_sde": "DPM++ SDE",
|
||||
"dpmpp_sde_gpu": "DPM++ SDE",
|
||||
"dpmpp_2m": "DPM++ 2M",
|
||||
"dpmpp_2m_sde": "DPM++ 2M SDE",
|
||||
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
|
||||
"ddim": "DDIM",
|
||||
}
|
||||
sampler_name = sampler_mapping.get(sampler, sampler)
|
||||
|
||||
if "scheduler" in metadata_dict:
|
||||
scheduler = metadata_dict.get("scheduler")
|
||||
scheduler_mapping = {
|
||||
"normal": "Simple",
|
||||
"karras": "Karras",
|
||||
"exponential": "Exponential",
|
||||
"sgm_uniform": "SGM Uniform",
|
||||
"sgm_quadratic": "SGM Quadratic",
|
||||
}
|
||||
scheduler_name = scheduler_mapping.get(scheduler, scheduler)
|
||||
|
||||
# Add combined sampler and scheduler information
|
||||
if sampler_name:
|
||||
if scheduler_name:
|
||||
params.append(f"Sampler: {sampler_name} {scheduler_name}")
|
||||
else:
|
||||
params.append(f"Sampler: {sampler_name}")
|
||||
|
||||
# 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 +777,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 +833,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 +927,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 +958,8 @@ class SaveImageLM:
|
||||
extra_pnginfo,
|
||||
lossless_webp,
|
||||
quality,
|
||||
webp_method,
|
||||
jpeg_subsampling,
|
||||
embed_workflow,
|
||||
save_with_metadata,
|
||||
add_counter_to_filename,
|
||||
|
||||
@@ -537,6 +537,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 +545,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 +593,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:
|
||||
|
||||
@@ -1392,8 +1392,8 @@ class DownloadManager:
|
||||
base_save_dir = save_dir
|
||||
save_dir = os.path.join(save_dir, relative_path)
|
||||
# Security: validate path containment after joining
|
||||
resolved_dir = os.path.realpath(os.path.normpath(save_dir))
|
||||
base_dir = os.path.realpath(os.path.normpath(base_save_dir))
|
||||
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",
|
||||
|
||||
@@ -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]}"
|
||||
)
|
||||
|
||||
|
||||
@@ -51,9 +51,10 @@ async def delete_model_artifacts(
|
||||
def _require_path_in_library_roots(file_path: str, scanner, *, label: str = "path") -> None:
|
||||
"""Raise ``ValueError`` if *file_path* is not inside a configured model root.
|
||||
|
||||
Uses ``os.path.realpath()`` to resolve symlinks before comparing,
|
||||
so symlink-based escapes are also caught. Skips when the scanner
|
||||
does not expose ``get_model_roots`` or the list is empty.
|
||||
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
|
||||
@@ -65,10 +66,10 @@ def _require_path_in_library_roots(file_path: str, scanner, *, label: str = "pat
|
||||
if not roots:
|
||||
return
|
||||
|
||||
resolved = os.path.realpath(os.path.normpath(file_path))
|
||||
resolved = os.path.abspath(os.path.normpath(file_path))
|
||||
|
||||
for root in roots:
|
||||
root_resolved = os.path.realpath(os.path.normpath(root))
|
||||
root_resolved = os.path.abspath(os.path.normpath(root))
|
||||
if resolved == root_resolved or resolved.startswith(root_resolved + os.sep):
|
||||
return
|
||||
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -885,7 +885,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':
|
||||
|
||||
@@ -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,220 @@ 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
|
||||
assert "clip_skip" not in params
|
||||
|
||||
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"]
|
||||
|
||||
inputs = {key: "" for key in METADATA_OVERWRITE_FIELDS}
|
||||
inputs.update({"seed": 0, "steps": 0, "cfg_scale": 0.0, "clip_skip": 0})
|
||||
|
||||
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
|
||||
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
|
||||
|
||||
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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1248,6 +1248,50 @@ def test_relative_path_sanitizes_double_slashes():
|
||||
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()
|
||||
|
||||
@@ -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,4 +1,5 @@
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
@@ -51,11 +52,12 @@ class TestRequirePathInLibraryRoots:
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
_require_path_in_library_roots(str(root), scanner)
|
||||
|
||||
def test_rejects_symlink_escape(self, tmp_path):
|
||||
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()
|
||||
model = root / "model.safetensors"
|
||||
model.write_text("")
|
||||
|
||||
outside_dir = tmp_path / "outside"
|
||||
outside_dir.mkdir()
|
||||
@@ -65,9 +67,22 @@ class TestRequirePathInLibraryRoots:
|
||||
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(str(symlink), scanner)
|
||||
_require_path_in_library_roots(escaped, scanner)
|
||||
|
||||
|
||||
class ScannerForDelete:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -37,22 +37,28 @@ app.registerExtension({
|
||||
|
||||
// Handle broadcast mode (for Desktop/non-browser support)
|
||||
if (numericNodeId === -1) {
|
||||
// Find all Lora Loader nodes in the current graph
|
||||
const loraLoaderNodes = getAllGraphNodes(app.graph)
|
||||
// Find all compatible nodes in the current graph
|
||||
const compatibleClasses = new Set([
|
||||
"Lora Loader (LoraManager)",
|
||||
"Lora Stacker (LoraManager)",
|
||||
"WanVideo Lora Select (LoraManager)",
|
||||
"Create Hook LoRA (LoraManager)",
|
||||
]);
|
||||
const targetNodes = getAllGraphNodes(app.graph)
|
||||
.map(({ node }) => node)
|
||||
.filter((node) => node?.comfyClass === "Lora Loader (LoraManager)");
|
||||
.filter((node) => compatibleClasses.has(node?.comfyClass));
|
||||
|
||||
// Update each Lora Loader node found
|
||||
if (loraLoaderNodes.length > 0) {
|
||||
loraLoaderNodes.forEach((node) => {
|
||||
// Update each node found
|
||||
if (targetNodes.length > 0) {
|
||||
targetNodes.forEach((node) => {
|
||||
this.updateNodeLoraCode(node, loraCode, mode);
|
||||
});
|
||||
console.log(
|
||||
`Updated ${loraLoaderNodes.length} Lora Loader nodes in broadcast mode`
|
||||
`Updated ${targetNodes.length} nodes in broadcast mode`
|
||||
);
|
||||
} else {
|
||||
console.warn(
|
||||
"No Lora Loader nodes found in the workflow for broadcast update"
|
||||
"No compatible LoRA nodes found in the workflow for broadcast update"
|
||||
);
|
||||
}
|
||||
|
||||
@@ -65,10 +71,11 @@ app.registerExtension({
|
||||
!node ||
|
||||
(node.comfyClass !== "Lora Loader (LoraManager)" &&
|
||||
node.comfyClass !== "Lora Stacker (LoraManager)" &&
|
||||
node.comfyClass !== "WanVideo Lora Select (LoraManager)")
|
||||
node.comfyClass !== "WanVideo Lora Select (LoraManager)" &&
|
||||
node.comfyClass !== "Create Hook LoRA (LoraManager)")
|
||||
) {
|
||||
console.warn(
|
||||
"Node not found or not a LoraLoader:",
|
||||
"Node not found or not a compatible LoRA node:",
|
||||
graphId ?? "root",
|
||||
nodeId
|
||||
);
|
||||
|
||||
@@ -751,7 +751,11 @@ export function addLorasWidget(node, name, opts, callback) {
|
||||
}
|
||||
}
|
||||
|
||||
renderLoras(widgetValue, widget);
|
||||
// Skip DOM re-render during drag to preserve pointer capture and event listeners.
|
||||
// The strength inputs are updated directly via the pointermove handler instead.
|
||||
if (!widget.__dragActive) {
|
||||
renderLoras(widgetValue, widget);
|
||||
}
|
||||
},
|
||||
hideOnZoom: true,
|
||||
selectOn: ['click', 'focus']
|
||||
|
||||
@@ -37,17 +37,18 @@ export function handleStrengthDrag(name, initialStrength, initialX, event, widge
|
||||
syncClipStrengthIfCollapsed(lorasData[loraIndex]);
|
||||
}
|
||||
|
||||
// Update the widget value only if updateWidget flag is true
|
||||
// This allows us to update inputs directly during drag without triggering re-render
|
||||
if (updateWidget) {
|
||||
widget.value = formatLoraValue(lorasData);
|
||||
}
|
||||
// Always write back to widget.value to persist the mutation.
|
||||
// During drag (updateWidget=false), setValue skips renderLoras via __dragActive flag,
|
||||
// so the DOM survives and pointer capture is preserved.
|
||||
widget.value = formatLoraValue(lorasData);
|
||||
|
||||
// Force re-render via callback only if updateWidget is true
|
||||
// Only fire callback on the final commit, not during drag
|
||||
if (updateWidget && widget.callback) {
|
||||
widget.callback(widget.value);
|
||||
}
|
||||
}
|
||||
|
||||
return newStrength;
|
||||
}
|
||||
|
||||
// Function to handle proportional strength adjustment for all LoRAs via header dragging
|
||||
@@ -90,12 +91,11 @@ export function handleAllStrengthsDrag(initialStrengths, initialX, event, widget
|
||||
lorasData[index].clipStrength = Number(newClipStrength);
|
||||
});
|
||||
|
||||
// Update widget value only if updateWidget flag is true
|
||||
if (updateWidget) {
|
||||
widget.value = formatLoraValue(lorasData);
|
||||
}
|
||||
// Always write back to widget.value to persist mutations.
|
||||
// During drag (updateWidget=false), setValue skips renderLoras via __dragActive flag.
|
||||
widget.value = formatLoraValue(lorasData);
|
||||
|
||||
// Force re-render via callback only if updateWidget is true
|
||||
// Only fire callback on the final commit, not during drag
|
||||
if (updateWidget && widget.callback) {
|
||||
widget.callback(widget.value);
|
||||
}
|
||||
@@ -149,6 +149,13 @@ export function initDrag(
|
||||
activePointerId = e.pointerId;
|
||||
currentDragElement = e.currentTarget;
|
||||
|
||||
// Suppress renderLoras in setValue during drag so the DOM survives.
|
||||
// The getter creates a new array on every read, so mutations to a
|
||||
// parsed copy are lost unless we write back through widget.value.
|
||||
// Writing back would normally trigger a full DOM re-render via setValue,
|
||||
// destroying pointer capture. __dragActive tells setValue to skip the render.
|
||||
widget.__dragActive = true;
|
||||
|
||||
// Capture pointer to receive all subsequent events regardless of stopPropagation
|
||||
const target = e.currentTarget;
|
||||
target.setPointerCapture(e.pointerId);
|
||||
@@ -181,17 +188,12 @@ export function initDrag(
|
||||
}
|
||||
|
||||
// Call the strength adjustment function without updating widget.value during drag
|
||||
handleStrengthDrag(name, initialStrength, initialX, e, widget, isClipStrength, false);
|
||||
const newStrength = handleStrengthDrag(name, initialStrength, initialX, e, widget, isClipStrength, false);
|
||||
|
||||
// Update strength input directly instead of re-rendering to avoid losing event listeners
|
||||
const strengthInput = currentDragElement.querySelector('.lm-lora-strength-input');
|
||||
if (strengthInput) {
|
||||
const lorasData = parseLoraValue(widget.value);
|
||||
const loraData = lorasData.find(l => l.name === name);
|
||||
if (loraData) {
|
||||
const strengthValue = isClipStrength ? loraData.clipStrength : loraData.strength;
|
||||
strengthInput.value = Number(strengthValue).toFixed(2);
|
||||
}
|
||||
if (strengthInput && typeof newStrength === 'number') {
|
||||
strengthInput.value = newStrength.toFixed(2);
|
||||
}
|
||||
|
||||
// Prevent showing the preview tooltip during drag
|
||||
@@ -226,23 +228,30 @@ export function initDrag(
|
||||
// Remove the class to restore normal cursor behavior
|
||||
document.body.classList.remove('lm-lora-strength-dragging');
|
||||
|
||||
// Only call onDragEnd and re-render if we actually dragged
|
||||
if (wasDragging) {
|
||||
if (typeof onDragEnd === 'function') {
|
||||
onDragEnd();
|
||||
}
|
||||
// Only call onDragEnd and re-render if we actually dragged.
|
||||
// try-finally guarantees __dragActive is always cleared, preventing a
|
||||
// permanent UI freeze if onDragEnd or setValue throws during cleanup.
|
||||
try {
|
||||
if (wasDragging) {
|
||||
if (typeof onDragEnd === 'function') {
|
||||
onDragEnd();
|
||||
}
|
||||
|
||||
// Commit final value through options.setValue so external observers are notified.
|
||||
// During drag, handleStrengthDrag mutates widgetValue in-place (updateWidget=false),
|
||||
// bypassing widget.value setter and options.setValue entirely. This assignment
|
||||
// flushes the in-place mutation through the setter so any setValue wrappers fire.
|
||||
widget.value = widget.value;
|
||||
if (typeof widget.callback === 'function') {
|
||||
widget.callback(widget.value);
|
||||
// Re-enable renderLoras in setValue and flush final value through setter.
|
||||
// The last handleStrengthDrag call already wrote the final strength to
|
||||
// widgetValue via setValue (with render suppressed). widget.value = widget.value
|
||||
// triggers setValue again, which now calls renderLoras since __dragActive is false.
|
||||
widget.__dragActive = false;
|
||||
widget.value = widget.value;
|
||||
if (typeof widget.callback === 'function') {
|
||||
widget.callback(widget.value);
|
||||
}
|
||||
}
|
||||
} finally {
|
||||
widget.__dragActive = false;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
dragEl.addEventListener('pointerup', endDrag);
|
||||
dragEl.addEventListener('pointercancel', endDrag);
|
||||
}
|
||||
@@ -285,6 +294,9 @@ export function initHeaderDrag(headerEl, widget, renderFunction) {
|
||||
activePointerId = e.pointerId;
|
||||
currentHeaderElement = e.currentTarget;
|
||||
|
||||
// Suppress renderLoras in setValue during drag (see initDrag for rationale)
|
||||
widget.__dragActive = true;
|
||||
|
||||
// Capture pointer to receive all subsequent events regardless of stopPropagation
|
||||
const target = e.currentTarget;
|
||||
target.setPointerCapture(e.pointerId);
|
||||
@@ -352,13 +364,20 @@ export function initHeaderDrag(headerEl, widget, renderFunction) {
|
||||
// Remove the class to restore normal cursor behavior
|
||||
document.body.classList.remove('lm-lora-strength-dragging');
|
||||
|
||||
// Only re-render if we actually dragged
|
||||
if (wasDragging) {
|
||||
// Commit final value through options.setValue so external observers are notified.
|
||||
widget.value = widget.value;
|
||||
if (typeof widget.callback === 'function') {
|
||||
widget.callback(widget.value);
|
||||
// Only re-render if we actually dragged.
|
||||
// try-finally guarantees __dragActive is always cleared, preventing a
|
||||
// permanent UI freeze if setValue throws during cleanup.
|
||||
try {
|
||||
if (wasDragging) {
|
||||
// Re-enable renderLoras in setValue and flush final value through setter
|
||||
widget.__dragActive = false;
|
||||
widget.value = widget.value;
|
||||
if (typeof widget.callback === 'function') {
|
||||
widget.callback(widget.value);
|
||||
}
|
||||
}
|
||||
} finally {
|
||||
widget.__dragActive = false;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -7,12 +7,16 @@ import { app } from "../../scripts/app.js";
|
||||
// Roles are stored in ``node.properties.lm_marker_role`` and automatically
|
||||
// persist with the workflow JSON.
|
||||
//
|
||||
// Two categories:
|
||||
// send_* – consumed by the standalone UI's "Send to Workflow" feature
|
||||
// meta_* – consumed by the metadata processor to override heuristic inference
|
||||
//
|
||||
// The workflow registry reads these markers and makes them available to the
|
||||
// standalone UI (e.g. ``sendEmbeddingToWorkflow`` also considers nodes marked
|
||||
// as ``send_prompt_target``).
|
||||
// =============================================================================
|
||||
|
||||
const ROLES = {
|
||||
const SEND_ROLES = {
|
||||
send_prompt_target: {
|
||||
label: "Send Prompt Target",
|
||||
emoji: "\uD83D\uDCDD",
|
||||
@@ -23,6 +27,28 @@ const ROLES = {
|
||||
},
|
||||
};
|
||||
|
||||
const META_ROLES = {
|
||||
meta_primary_model: {
|
||||
label: "Meta hints: Primary Model",
|
||||
emoji: "\uD83D\uDCA1",
|
||||
},
|
||||
meta_primary_sampler: {
|
||||
label: "Meta hints: Primary Sampler",
|
||||
emoji: "\uD83D\uDCA1",
|
||||
},
|
||||
meta_positive_prompt: {
|
||||
label: "Meta hints: Positive Prompt",
|
||||
emoji: "\uD83D\uDCA1",
|
||||
},
|
||||
meta_negative_prompt: {
|
||||
label: "Meta hints: Negative Prompt",
|
||||
emoji: "\uD83D\uDCA1",
|
||||
},
|
||||
};
|
||||
|
||||
// Flat lookup for setMarker / getMarker / clearMarker
|
||||
const ROLES = { ...SEND_ROLES, ...META_ROLES };
|
||||
|
||||
// ---- Helpers ----------------------------------------------------------------
|
||||
|
||||
function getMarker(node) {
|
||||
@@ -54,7 +80,7 @@ function clearMarker(node) {
|
||||
// Restore original title: prefer stripping emoji from current title
|
||||
// (captures user renames after marking), fall back to saved original.
|
||||
const cleaned = node.title?.replace(
|
||||
/^(\u2709\uFE0F?|\u2699\uFE0F?|\uD83D\uDCDD|\uD83C\uDF9B\uFE0F?|\uD83D\uDD27)\s*/,
|
||||
/^(\u2709\uFE0F?|\u2699\uFE0F?|\uD83D\uDCDD|\uD83C\uDF9B\uFE0F?|\uD83D\uDD27|\uD83D\uDCA1)\s*/,
|
||||
''
|
||||
);
|
||||
if (cleaned && cleaned !== node.title) {
|
||||
@@ -84,16 +110,23 @@ function buildSubmenuOptions(node) {
|
||||
const currentRole = getMarker(node);
|
||||
const options = [];
|
||||
|
||||
for (const [key, def] of Object.entries(ROLES)) {
|
||||
const isActive = currentRole === key;
|
||||
options.push({
|
||||
content: `${isActive ? "\u2713 " : ""}${def.label}`,
|
||||
disabled: isActive,
|
||||
callback: () => setMarker(node, key),
|
||||
});
|
||||
}
|
||||
const buildGroup = (roles) => {
|
||||
for (const [key, def] of Object.entries(roles)) {
|
||||
const isActive = currentRole === key;
|
||||
options.push({
|
||||
content: `${isActive ? "\u2713 " : ""}${def.label}`,
|
||||
disabled: isActive,
|
||||
callback: () => setMarker(node, key),
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
buildGroup(SEND_ROLES);
|
||||
options.push(null); // separator
|
||||
buildGroup(META_ROLES);
|
||||
|
||||
if (currentRole) {
|
||||
options.push(null); // separator
|
||||
options.push({
|
||||
content: "Clear marker",
|
||||
callback: () => clearMarker(node),
|
||||
|
||||
@@ -130,6 +130,35 @@ app.registerExtension({
|
||||
widget.serializeValue = () => {
|
||||
return applyTextReplacements(widget.value);
|
||||
};
|
||||
|
||||
// --- Conditional widget visibility for webp_method / jpeg_subsampling ---
|
||||
const formatWidget = getWidgetByName(this, "file_format");
|
||||
const webpMethodWidget = getWidgetByName(this, "webp_method");
|
||||
const jpegSubWidget = getWidgetByName(this, "jpeg_subsampling");
|
||||
|
||||
function updateFormatConditional() {
|
||||
const fmt = formatWidget?.value;
|
||||
if (webpMethodWidget) {
|
||||
webpMethodWidget.disabled = fmt !== "webp";
|
||||
webpMethodWidget.hidden = fmt !== "webp";
|
||||
}
|
||||
if (jpegSubWidget) {
|
||||
jpegSubWidget.disabled = fmt !== "jpeg";
|
||||
jpegSubWidget.hidden = fmt !== "jpeg";
|
||||
}
|
||||
}
|
||||
|
||||
// Set initial state
|
||||
updateFormatConditional();
|
||||
|
||||
// Watch for format changes
|
||||
if (formatWidget) {
|
||||
const origCallback = formatWidget.callback;
|
||||
formatWidget.callback = function (value) {
|
||||
origCallback?.call(this, value);
|
||||
updateFormatConditional();
|
||||
};
|
||||
}
|
||||
});
|
||||
},
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user