mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-09 07:20:15 -03:00
Compare commits
15 Commits
v1.1.8
...
e341e0b9d2
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e341e0b9d2 | ||
|
|
e6538c83bb | ||
|
|
92e1285ea5 | ||
|
|
2aabd1d90e | ||
|
|
7b8b778f83 | ||
|
|
7c8dc57d55 | ||
|
|
fe95fae5f2 | ||
|
|
ce8a95abf7 | ||
|
|
c8e7e543d6 | ||
|
|
a9dbb15ffa | ||
|
|
cf64043f7d | ||
|
|
ccaff92c18 | ||
|
|
585b5c922a | ||
|
|
ea80c2224c | ||
|
|
8b0f56c1a6 |
@@ -137,7 +137,13 @@ npm run test:coverage # Generate coverage report
|
||||
- Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json)
|
||||
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
||||
- Run `python scripts/sync_translation_keys.py` after adding UI strings to `locales/en.json`
|
||||
- Symlinks require normalized paths
|
||||
- Symlinks require normalized paths.
|
||||
**Business paths vs real paths**: All stored paths and operation routing use the
|
||||
original paths as they appear under configured model roots — symlinks are NOT
|
||||
resolved. `os.path.realpath` is only for scanner dedup and the symlink cache.
|
||||
Any path passed to `os.remove`/`os.rename`/`shutil.move` or validated by a
|
||||
containment check MUST use the business path (i.e. `os.path.abspath`, not
|
||||
`realpath`).
|
||||
|
||||
## Git / Commit Messages
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ try: # pragma: no cover - import fallback for pytest collection
|
||||
from .py.nodes.lora_cycler import LoraCyclerLM
|
||||
from .py.nodes.lora_info import LoraInfoLM
|
||||
from .py.nodes.lora_syntax_to_path import LoraSyntaxToPath
|
||||
from .py.nodes.create_hook_lora import CreateHookLoraLM
|
||||
from .py.metadata_collector import init as init_metadata_collector
|
||||
except (
|
||||
ImportError
|
||||
@@ -62,6 +63,9 @@ except (
|
||||
LoraSyntaxToPath = importlib.import_module(
|
||||
"py.nodes.lora_syntax_to_path"
|
||||
).LoraSyntaxToPath
|
||||
CreateHookLoraLM = importlib.import_module(
|
||||
"py.nodes.create_hook_lora"
|
||||
).CreateHookLoraLM
|
||||
init_metadata_collector = importlib.import_module("py.metadata_collector").init
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -83,6 +87,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
LoraCyclerLM.NAME: LoraCyclerLM,
|
||||
LoraInfoLM.NAME: LoraInfoLM,
|
||||
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
|
||||
CreateHookLoraLM.NAME: CreateHookLoraLM,
|
||||
}
|
||||
|
||||
WEB_DIRECTORY = "./web/comfyui"
|
||||
|
||||
117
py/nodes/create_hook_lora.py
Normal file
117
py/nodes/create_hook_lora.py
Normal file
@@ -0,0 +1,117 @@
|
||||
"""Create Hook LoRA (LoraManager) — multi-LoRA hook node compatible with ComfyUI's built-in hook pipeline.
|
||||
|
||||
Produces ``("HOOKS",)`` output that chains seamlessly with downstream hook consumers
|
||||
(ConditioningSetProperties, SetHookKeyframes, CombineHooks, SetClipHooks, etc.).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
|
||||
from ..utils.utils import get_lora_info_absolute
|
||||
from .utils import (
|
||||
FlexibleOptionalInputType,
|
||||
any_type,
|
||||
apply_lora_syntax_format,
|
||||
get_loras_list,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CreateHookLoraLM:
|
||||
NAME = "Create Hook LoRA (LoraManager)"
|
||||
CATEGORY = "Lora Manager/hooks"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"text": (
|
||||
"AUTOCOMPLETE_TEXT_LORAS",
|
||||
{
|
||||
"placeholder": "Search LoRAs to add...",
|
||||
"tooltip": (
|
||||
"Search and select LoRAs. Each LoRA gets its own "
|
||||
"model/clip strength. Hooks chain with prev_hooks."
|
||||
),
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": FlexibleOptionalInputType(any_type),
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("HOOKS", "STRING", "STRING")
|
||||
RETURN_NAMES = ("HOOKS", "trigger_words", "active_loras")
|
||||
FUNCTION = "create_hook"
|
||||
|
||||
def create_hook(self, text: str, **kwargs):
|
||||
"""Create a HookGroup from the selected LoRAs, chained with prev_hooks.
|
||||
|
||||
Each active LoRA from the widget is loaded and wrapped in a WeightHook
|
||||
via :func:`comfy.hooks.create_hook_lora`. All hooks are combined into a
|
||||
single group and returned alongside trigger words and a human-readable
|
||||
summary of the active LoRAs.
|
||||
"""
|
||||
del text # used by the frontend widget only
|
||||
|
||||
# Lazy imports: comfy is not available in CI/test environment at module level
|
||||
import comfy.hooks # type: ignore # noqa: C0415
|
||||
import comfy.utils # type: ignore # noqa: C0415
|
||||
|
||||
prev_hooks: comfy.hooks.HookGroup | None = kwargs.get("prev_hooks")
|
||||
|
||||
hook_group = prev_hooks.clone() if prev_hooks is not None else comfy.hooks.HookGroup()
|
||||
|
||||
all_trigger_words: list[str] = []
|
||||
active_loras: list[tuple[str, float, float]] = []
|
||||
|
||||
for lora in get_loras_list(kwargs):
|
||||
if not lora.get("active", False):
|
||||
continue
|
||||
|
||||
lora_name = apply_lora_syntax_format(lora["name"])
|
||||
model_strength = float(lora["strength"])
|
||||
clip_strength = float(lora.get("clipStrength", model_strength))
|
||||
|
||||
# Skip useless no-op entries (both strengths are zero)
|
||||
if model_strength == 0.0 and clip_strength == 0.0:
|
||||
continue
|
||||
|
||||
lora_path, trigger_words = get_lora_info_absolute(lora_name)
|
||||
if not lora_path or not os.path.isfile(lora_path):
|
||||
logger.warning("LoRA '%s' not found — skipping", lora_name)
|
||||
continue
|
||||
|
||||
try:
|
||||
lora_weights = comfy.utils.load_torch_file(lora_path, safe_load=True)
|
||||
|
||||
lora_hooks = comfy.hooks.create_hook_lora(
|
||||
lora=lora_weights,
|
||||
strength_model=model_strength,
|
||||
strength_clip=clip_strength,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to load LoRA '%s' — skipping", lora_name)
|
||||
continue
|
||||
hook_group = hook_group.clone_and_combine(lora_hooks)
|
||||
|
||||
active_loras.append((lora_name, model_strength, clip_strength))
|
||||
all_trigger_words.extend(trigger_words)
|
||||
|
||||
# Format trigger words (group mode separator)
|
||||
trigger_words_text = ",, ".join(all_trigger_words) if all_trigger_words else ""
|
||||
|
||||
# Format active LoRAs summary
|
||||
formatted_loras = []
|
||||
for name, model_s, clip_s in active_loras:
|
||||
if abs(model_s - clip_s) > 0.001:
|
||||
formatted_loras.append(
|
||||
f"<lora:{name}:{model_s}:{clip_s}>"
|
||||
)
|
||||
else:
|
||||
formatted_loras.append(f"<lora:{name}:{model_s}>")
|
||||
active_loras_text = " ".join(formatted_loras)
|
||||
|
||||
return (hook_group, trigger_words_text, active_loras_text)
|
||||
@@ -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__)
|
||||
|
||||
|
||||
@@ -142,148 +292,181 @@ 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
|
||||
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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -1389,7 +1389,17 @@ class DownloadManager:
|
||||
|
||||
# Update save directory with relative path if provided
|
||||
if relative_path:
|
||||
base_save_dir = save_dir
|
||||
save_dir = os.path.join(save_dir, relative_path)
|
||||
# Security: validate path containment after joining
|
||||
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
|
||||
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
|
||||
if not resolved_dir.startswith(base_dir + os.sep) and resolved_dir != base_dir:
|
||||
logger.warning(
|
||||
"Path traversal detected: %s escapes %s",
|
||||
resolved_dir, base_dir,
|
||||
)
|
||||
return {"success": False, "error": "Download path is outside allowed directory"}
|
||||
# Create directory if it doesn't exist
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
@@ -1827,6 +1837,9 @@ class DownloadManager:
|
||||
model_tags, model_type
|
||||
)
|
||||
|
||||
if not first_tag:
|
||||
first_tag = "no tags" # Default if no tags available
|
||||
|
||||
# Format the template with available data
|
||||
formatted_path = path_template
|
||||
formatted_path = formatted_path.replace("{base_model}", mapped_base_model)
|
||||
@@ -1842,6 +1855,15 @@ class DownloadManager:
|
||||
if model_type == "embedding":
|
||||
formatted_path = formatted_path.replace(" ", "_")
|
||||
|
||||
# Sanitize the resolved path to prevent path traversal:
|
||||
# - Strip leading slashes (prevents os.path.join from treating path as absolute)
|
||||
# - Collapse double slashes from empty placeholder substitutions
|
||||
# - Strip trailing slashes for cleanliness
|
||||
formatted_path = formatted_path.lstrip("/")
|
||||
while "//" in formatted_path:
|
||||
formatted_path = formatted_path.replace("//", "/")
|
||||
formatted_path = formatted_path.rstrip("/")
|
||||
|
||||
return formatted_path
|
||||
|
||||
async def _execute_download(
|
||||
|
||||
@@ -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]}"
|
||||
)
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ from abc import ABC, abstractmethod
|
||||
from ..utils.utils import calculate_relative_path_for_model, remove_empty_dirs
|
||||
from ..utils.constants import AUTO_ORGANIZE_BATCH_SIZE
|
||||
from ..services.settings_manager import get_settings_manager
|
||||
from ..services.model_lifecycle_service import _require_path_in_library_roots
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -493,6 +494,9 @@ class ModelMoveService:
|
||||
Dictionary with move result
|
||||
"""
|
||||
try:
|
||||
_require_path_in_library_roots(file_path, self.scanner, label="Source path")
|
||||
_require_path_in_library_roots(target_path, self.scanner, label="Target path")
|
||||
|
||||
if use_default_paths:
|
||||
# Find the model in cache to get metadata
|
||||
cache = await self.scanner.get_cached_data()
|
||||
|
||||
@@ -48,6 +48,36 @@ async def delete_model_artifacts(
|
||||
return deleted
|
||||
|
||||
|
||||
def _require_path_in_library_roots(file_path: str, scanner, *, label: str = "path") -> None:
|
||||
"""Raise ``ValueError`` if *file_path* is not inside a configured model root.
|
||||
|
||||
Uses ``os.path.abspath()`` (NOT ``realpath``) to resolve ``..`` and ``.``
|
||||
while preserving symlinks — this keeps the check in business-path space.
|
||||
Skips when the scanner does not expose ``get_model_roots`` or the list
|
||||
is empty.
|
||||
"""
|
||||
|
||||
roots = None
|
||||
if hasattr(scanner, "get_model_roots"):
|
||||
try:
|
||||
roots = scanner.get_model_roots()
|
||||
except NotImplementedError:
|
||||
roots = None
|
||||
if not roots:
|
||||
return
|
||||
|
||||
resolved = os.path.abspath(os.path.normpath(file_path))
|
||||
|
||||
for root in roots:
|
||||
root_resolved = os.path.abspath(os.path.normpath(root))
|
||||
if resolved == root_resolved or resolved.startswith(root_resolved + os.sep):
|
||||
return
|
||||
|
||||
raise ValueError(
|
||||
f"{label} '{file_path}' is outside configured library directories"
|
||||
)
|
||||
|
||||
|
||||
class ModelLifecycleService:
|
||||
"""Co-ordinate destructive and mutating model operations."""
|
||||
|
||||
@@ -74,6 +104,8 @@ class ModelLifecycleService:
|
||||
if not file_path:
|
||||
raise ValueError("Model path is required")
|
||||
|
||||
_require_path_in_library_roots(file_path, self._scanner, label="File path")
|
||||
|
||||
cache = await self._scanner.get_cached_data()
|
||||
|
||||
cached_entry = None
|
||||
@@ -182,6 +214,8 @@ class ModelLifecycleService:
|
||||
if not file_path:
|
||||
raise ValueError("Model path is required")
|
||||
|
||||
_require_path_in_library_roots(file_path, self._scanner, label="File path")
|
||||
|
||||
metadata_path = os.path.splitext(file_path)[0] + ".metadata.json"
|
||||
metadata = await self._metadata_loader(metadata_path)
|
||||
metadata["exclude"] = True
|
||||
@@ -229,6 +263,8 @@ class ModelLifecycleService:
|
||||
if not file_path:
|
||||
raise ValueError("Model path is required")
|
||||
|
||||
_require_path_in_library_roots(file_path, self._scanner, label="File path")
|
||||
|
||||
if not os.path.exists(file_path):
|
||||
raise ValueError("Model file does not exist")
|
||||
|
||||
@@ -270,6 +306,9 @@ class ModelLifecycleService:
|
||||
if not file_paths:
|
||||
raise ValueError("No file paths provided for deletion")
|
||||
|
||||
for path in file_paths:
|
||||
_require_path_in_library_roots(path, self._scanner, label="File path")
|
||||
|
||||
return await self._scanner.bulk_delete_models(file_paths)
|
||||
|
||||
async def rename_model(
|
||||
@@ -280,6 +319,8 @@ class ModelLifecycleService:
|
||||
if not file_path or not new_file_name:
|
||||
raise ValueError("File path and new file name are required")
|
||||
|
||||
_require_path_in_library_roots(file_path, self._scanner, label="File path")
|
||||
|
||||
invalid_chars = {"/", "\\", ":", "*", "?", '"', "<", ">", "|"}
|
||||
if any(char in new_file_name for char in invalid_chars):
|
||||
raise ValueError("Invalid characters in file name")
|
||||
|
||||
@@ -14,7 +14,7 @@ from ..utils.metadata_manager import MetadataManager
|
||||
from ..utils.civitai_utils import resolve_license_info
|
||||
from .model_cache import ModelCache
|
||||
from .model_hash_index import ModelHashIndex
|
||||
from .model_lifecycle_service import delete_model_artifacts
|
||||
from .model_lifecycle_service import delete_model_artifacts, _require_path_in_library_roots
|
||||
from .service_registry import ServiceRegistry
|
||||
from .websocket_manager import ws_manager
|
||||
from .persistent_model_cache import get_persistent_cache
|
||||
@@ -1394,6 +1394,9 @@ class ModelScanner:
|
||||
|
||||
base_name = os.path.splitext(os.path.basename(source_path))[0]
|
||||
source_dir = os.path.dirname(source_path)
|
||||
|
||||
_require_path_in_library_roots(source_path, self, label="Source path")
|
||||
_require_path_in_library_roots(target_path, self, label="Target path")
|
||||
|
||||
os.makedirs(target_path, exist_ok=True)
|
||||
|
||||
@@ -1971,6 +1974,8 @@ class ModelScanner:
|
||||
break
|
||||
|
||||
try:
|
||||
_require_path_in_library_roots(file_path, self, label="File path")
|
||||
|
||||
target_dir = os.path.dirname(file_path)
|
||||
base_name = os.path.basename(file_path)
|
||||
file_name, main_extension = os.path.splitext(base_name)
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -12,6 +12,7 @@ NODE_TYPES = {
|
||||
"Lora Loader (LoraManager)": 1,
|
||||
"Lora Stacker (LoraManager)": 2,
|
||||
"WanVideo Lora Select (LoraManager)": 3,
|
||||
"Create Hook LoRA (LoraManager)": 4,
|
||||
}
|
||||
|
||||
# Default ComfyUI node color when bgcolor is null
|
||||
|
||||
@@ -488,6 +488,12 @@ def calculate_relative_path_for_model(
|
||||
if model_type == "embedding":
|
||||
formatted_path = formatted_path.replace(" ", "_")
|
||||
|
||||
# Sanitize the resolved path to prevent path traversal
|
||||
formatted_path = formatted_path.lstrip("/")
|
||||
while "//" in formatted_path:
|
||||
formatted_path = formatted_path.replace("//", "/")
|
||||
formatted_path = formatted_path.rstrip("/")
|
||||
|
||||
return formatted_path
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-lora-manager"
|
||||
description = "Revolutionize your workflow with the ultimate LoRA companion for ComfyUI!"
|
||||
version = "1.1.8"
|
||||
version = "1.1.9"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = [
|
||||
"aiohttp",
|
||||
|
||||
@@ -49,10 +49,6 @@ export const MODEL_CONFIG = {
|
||||
* @returns {Object} Object containing all API endpoints for the model type
|
||||
*/
|
||||
export function getApiEndpoints(modelType) {
|
||||
if (!Object.values(MODEL_TYPES).includes(modelType)) {
|
||||
throw new Error(`Invalid model type: ${modelType}`);
|
||||
}
|
||||
|
||||
return {
|
||||
// Base CRUD operations
|
||||
list: `/api/lm/${modelType}/list`,
|
||||
|
||||
@@ -369,21 +369,24 @@ export function getMatureBlurThreshold(settings = {}) {
|
||||
export const NODE_TYPES = {
|
||||
LORA_LOADER: 1,
|
||||
LORA_STACKER: 2,
|
||||
WAN_VIDEO_LORA_SELECT: 3
|
||||
WAN_VIDEO_LORA_SELECT: 3,
|
||||
HOOK_LORA: 4
|
||||
};
|
||||
|
||||
// Node type names to IDs mapping
|
||||
export const NODE_TYPE_NAMES = {
|
||||
"Lora Loader (LoraManager)": NODE_TYPES.LORA_LOADER,
|
||||
"Lora Stacker (LoraManager)": NODE_TYPES.LORA_STACKER,
|
||||
"WanVideo Lora Select (LoraManager)": NODE_TYPES.WAN_VIDEO_LORA_SELECT
|
||||
"WanVideo Lora Select (LoraManager)": NODE_TYPES.WAN_VIDEO_LORA_SELECT,
|
||||
"Create Hook LoRA (LoraManager)": NODE_TYPES.HOOK_LORA
|
||||
};
|
||||
|
||||
// Node type icons
|
||||
export const NODE_TYPE_ICONS = {
|
||||
[NODE_TYPES.LORA_LOADER]: "fas fa-l",
|
||||
[NODE_TYPES.LORA_STACKER]: "fas fa-s",
|
||||
[NODE_TYPES.WAN_VIDEO_LORA_SELECT]: "fas fa-w"
|
||||
[NODE_TYPES.WAN_VIDEO_LORA_SELECT]: "fas fa-w",
|
||||
[NODE_TYPES.HOOK_LORA]: "fas fa-h"
|
||||
};
|
||||
|
||||
// Default ComfyUI node color when bgcolor is null
|
||||
|
||||
@@ -85,6 +85,7 @@ sys.modules['comfy.utils'] = comfy_mock.utils
|
||||
sys.modules['comfy.sd'] = comfy_mock.sd
|
||||
sys.modules['comfy.model_management'] = comfy_mock.model_management
|
||||
sys.modules['comfy.comfy_types'] = comfy_mock.comfy_types
|
||||
sys.modules['comfy.hooks'] = MockModule("comfy.hooks")
|
||||
|
||||
execution_mock = MockModule("execution")
|
||||
execution_mock.PromptExecutor = mock.MagicMock()
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -1189,6 +1189,109 @@ def test_relative_path_sanitizes_model_and_version_placeholders():
|
||||
assert relative_path == "Fancy_Model/Version_One"
|
||||
|
||||
|
||||
def test_relative_path_empty_first_tag_fallback():
|
||||
"""Test that empty first_tag falls back to 'no tags'."""
|
||||
manager = DownloadManager()
|
||||
settings_manager = get_settings_manager()
|
||||
settings_manager.settings["download_path_templates"]["lora"] = (
|
||||
"{base_model}/{first_tag}"
|
||||
)
|
||||
|
||||
version_info = {
|
||||
"baseModel": "SDXL",
|
||||
"model": {"name": "Test Model", "tags": []},
|
||||
"creator": {"username": "Author"},
|
||||
}
|
||||
|
||||
relative_path = manager._calculate_relative_path(version_info, "lora")
|
||||
|
||||
assert relative_path == "SDXL/no tags"
|
||||
|
||||
|
||||
def test_relative_path_empty_base_model_and_first_tag():
|
||||
"""Test that empty base_model + empty first_tag does NOT produce a leading slash."""
|
||||
manager = DownloadManager()
|
||||
settings_manager = get_settings_manager()
|
||||
settings_manager.settings["download_path_templates"]["lora"] = (
|
||||
"{base_model}/{first_tag}"
|
||||
)
|
||||
|
||||
version_info = {
|
||||
"baseModel": "",
|
||||
"model": {"name": "Test Model", "tags": []},
|
||||
"creator": {"username": "Author"},
|
||||
}
|
||||
|
||||
relative_path = manager._calculate_relative_path(version_info, "lora")
|
||||
|
||||
assert not relative_path.startswith("/")
|
||||
assert relative_path == "no tags"
|
||||
|
||||
|
||||
def test_relative_path_sanitizes_double_slashes():
|
||||
"""Test that empty placeholder substitutions don't produce double slashes."""
|
||||
manager = DownloadManager()
|
||||
settings_manager = get_settings_manager()
|
||||
settings_manager.settings["download_path_templates"]["lora"] = (
|
||||
"{base_model}/{first_tag}/{author}"
|
||||
)
|
||||
|
||||
version_info = {
|
||||
"baseModel": "SDXL",
|
||||
"model": {"name": "Test Model", "tags": []},
|
||||
"creator": {"username": "Author"},
|
||||
}
|
||||
|
||||
relative_path = manager._calculate_relative_path(version_info, "lora")
|
||||
|
||||
assert "//" not in relative_path
|
||||
assert relative_path == "SDXL/no tags/Author"
|
||||
|
||||
|
||||
def test_download_containment_accepts_symlink_save_dir(tmp_path):
|
||||
"""Verify the download path containment check (download_manager.py:1395-1397)
|
||||
accepts save directories reached through user-created symlinks inside the
|
||||
library root — reproducing the symlink scenario from issue #1028."""
|
||||
# Library root with a symlink subdirectory pointing to an external drive
|
||||
lora_root = tmp_path / "loras"
|
||||
lora_root.mkdir()
|
||||
|
||||
external_drive = tmp_path / "external" / "models"
|
||||
external_drive.mkdir(parents=True)
|
||||
|
||||
symlink = lora_root / "Krea 2"
|
||||
symlink.symlink_to(str(external_drive))
|
||||
|
||||
# Simulate a download: base_save_dir = library root,
|
||||
# relative_path = "Krea 2/concept/NewModel"
|
||||
base_save_dir = str(lora_root)
|
||||
save_dir = os.path.join(base_save_dir, "Krea 2", "concept", "NewModel")
|
||||
|
||||
# Replicate the exact containment check from download_manager.py
|
||||
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
|
||||
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
|
||||
|
||||
# Must NOT be rejected — symlinks are legitimate business paths
|
||||
assert resolved_dir.startswith(base_dir + os.sep)
|
||||
|
||||
|
||||
def test_download_containment_rejects_dot_dot_traversal(tmp_path):
|
||||
"""Verify the download path containment check still blocks ``..`` traversal
|
||||
after the realpath → abspath change."""
|
||||
lora_root = tmp_path / "loras"
|
||||
lora_root.mkdir()
|
||||
|
||||
base_save_dir = str(lora_root)
|
||||
save_dir = os.path.join(base_save_dir, "..", "..", "etc", "passwd")
|
||||
|
||||
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
|
||||
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
|
||||
|
||||
# Must be rejected — dot-dot escapes the library root
|
||||
assert not resolved_dir.startswith(base_dir + os.sep)
|
||||
assert resolved_dir != base_dir
|
||||
|
||||
|
||||
def test_distribute_preview_to_entries_moves_and_copies(tmp_path):
|
||||
"""Test that preview distribution moves file to first entry and copies to others."""
|
||||
manager = DownloadManager()
|
||||
|
||||
@@ -243,6 +243,56 @@ class TestLLMServiceChatCompletionJson:
|
||||
|
||||
assert result == {"key": "value"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_json_falls_back_on_response_format_rejection(
|
||||
self, llm_service,
|
||||
):
|
||||
"""Retry without response_format when provider rejects it (HTTP 400)."""
|
||||
error_response = MockResponse(
|
||||
400,
|
||||
text_data=(
|
||||
'{"error":"\'response_format.type\' must be '
|
||||
'\'json_schema\' or \'text\'"}'
|
||||
),
|
||||
)
|
||||
success_response = MockResponse(
|
||||
200,
|
||||
json_data={
|
||||
"choices": [{"message": {"content": '{"key": "value"}'}}],
|
||||
"usage": {},
|
||||
"model": "local-model",
|
||||
},
|
||||
)
|
||||
|
||||
call_index = 0
|
||||
|
||||
class FallbackMockSession:
|
||||
def __init__(self):
|
||||
self.last_url = None
|
||||
self.last_json = None
|
||||
|
||||
def post(self, url, json=None, headers=None):
|
||||
nonlocal call_index
|
||||
self.last_url = url
|
||||
self.last_json = json
|
||||
call_index += 1
|
||||
return error_response if call_index == 1 else success_response
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
pass
|
||||
|
||||
with mock.patch("aiohttp.ClientSession", return_value=FallbackMockSession()):
|
||||
result = await llm_service.chat_completion_json(
|
||||
system_prompt="You are helpful.",
|
||||
user_prompt="Return JSON.",
|
||||
)
|
||||
|
||||
assert result == {"key": "value"}
|
||||
assert call_index == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_json_raises_on_non_json(self, llm_service):
|
||||
# Non-JSON content raises LLMResponseError (salvage also fails)
|
||||
|
||||
@@ -1,13 +1,181 @@
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from py.services.model_lifecycle_service import ModelLifecycleService
|
||||
from py.services.model_lifecycle_service import ModelLifecycleService, _require_path_in_library_roots
|
||||
from py.utils.metadata_manager import MetadataManager
|
||||
from py.utils.models import LoraMetadata
|
||||
|
||||
|
||||
class ScannerWithRoots:
|
||||
def __init__(self, roots):
|
||||
self._roots = list(roots)
|
||||
|
||||
def get_model_roots(self):
|
||||
return self._roots
|
||||
|
||||
|
||||
class TestRequirePathInLibraryRoots:
|
||||
def test_accepts_path_within_root(self, tmp_path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
model = root / "model.safetensors"
|
||||
model.write_text("")
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
_require_path_in_library_roots(str(model), scanner)
|
||||
|
||||
def test_rejects_path_outside_roots(self, tmp_path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
outside = tmp_path / "outside" / "model.safetensors"
|
||||
outside.parent.mkdir(parents=True)
|
||||
outside.write_text("")
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
with pytest.raises(ValueError, match="outside configured library"):
|
||||
_require_path_in_library_roots(str(outside), scanner)
|
||||
|
||||
def test_passes_when_no_roots_configured(self, tmp_path):
|
||||
f = tmp_path / "model.safetensors"
|
||||
f.write_text("")
|
||||
|
||||
scanner = ScannerWithRoots([])
|
||||
_require_path_in_library_roots(str(f), scanner)
|
||||
|
||||
def test_accepts_path_matching_root_exactly(self, tmp_path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
_require_path_in_library_roots(str(root), scanner)
|
||||
|
||||
def test_accepts_symlink_within_root(self, tmp_path):
|
||||
"""Symlinks under a configured root are legitimate business paths
|
||||
and should be accepted — containment works on business-path space,
|
||||
not resolved physical paths."""
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
|
||||
outside_dir = tmp_path / "outside"
|
||||
outside_dir.mkdir()
|
||||
outside_file = outside_dir / "escaped.safetensors"
|
||||
outside_file.write_text("")
|
||||
|
||||
symlink = root / "link.safetensors"
|
||||
symlink.symlink_to(outside_file)
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
# Symlink path is under root in business-path space → accepted
|
||||
_require_path_in_library_roots(str(symlink), scanner)
|
||||
|
||||
def test_rejects_dot_dot_traversal(self, tmp_path):
|
||||
"""Verify that ``..`` components are still resolved and blocked —
|
||||
``abspath`` normalises dot-dot but does not resolve symlinks."""
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
|
||||
# A path that traverses up out of the root via ..
|
||||
escaped = os.path.join(str(root), "..", "..", "etc", "passwd")
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
with pytest.raises(ValueError, match="outside configured library"):
|
||||
_require_path_in_library_roots(escaped, scanner)
|
||||
|
||||
|
||||
class ScannerForDelete:
|
||||
def __init__(self, raw_data, roots, model_type="lora"):
|
||||
self.model_type = model_type
|
||||
self.cache = DummyCache(raw_data)
|
||||
self._hash_index = DummyHashIndex()
|
||||
self._roots = list(roots)
|
||||
self._persist_calls = []
|
||||
|
||||
def get_model_roots(self):
|
||||
return self._roots
|
||||
|
||||
async def get_cached_data(self):
|
||||
return self.cache
|
||||
|
||||
async def _persist_current_cache(self):
|
||||
self._persist_calls.append(True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_model_rejects_path_outside_roots(tmp_path: Path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
model = root / "model.safetensors"
|
||||
model.write_bytes(b"data")
|
||||
|
||||
scanner = ScannerForDelete(
|
||||
raw_data=[{"file_path": str(model)}],
|
||||
roots=[str(root)],
|
||||
)
|
||||
service = ModelLifecycleService(
|
||||
scanner=scanner,
|
||||
metadata_manager=DummyMetadataManager({"civitai": {"modelId": 1}}),
|
||||
metadata_loader=lambda x: {},
|
||||
)
|
||||
# Path within root should work (model file exists)
|
||||
result = await service.delete_model(str(model))
|
||||
assert result["success"] is True
|
||||
|
||||
# Path outside root should be rejected
|
||||
outside = tmp_path / "outside.safetensors"
|
||||
outside.write_bytes(b"data")
|
||||
scanner2 = ScannerForDelete(
|
||||
raw_data=[],
|
||||
roots=[str(root)],
|
||||
)
|
||||
service2 = ModelLifecycleService(
|
||||
scanner=scanner2,
|
||||
metadata_manager=DummyMetadataManager({}),
|
||||
metadata_loader=lambda x: {},
|
||||
)
|
||||
with pytest.raises(ValueError, match="outside configured library"):
|
||||
await service2.delete_model(str(outside))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rename_model_rejects_path_outside_roots(tmp_path: Path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
service = ModelLifecycleService(
|
||||
scanner=scanner,
|
||||
metadata_manager=DummyMetadataManager({}),
|
||||
metadata_loader=lambda x: {},
|
||||
)
|
||||
outside = tmp_path / "outside.safetensors"
|
||||
outside.write_bytes(b"data")
|
||||
|
||||
with pytest.raises(ValueError, match="outside configured library"):
|
||||
await service.rename_model(file_path=str(outside), new_file_name="new_name")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bulk_delete_rejects_any_path_outside_roots(tmp_path: Path):
|
||||
root = tmp_path / "loras"
|
||||
root.mkdir()
|
||||
model_ok = root / "model.safetensors"
|
||||
model_ok.write_bytes(b"data")
|
||||
outside = tmp_path / "outside.safetensors"
|
||||
outside.write_bytes(b"data")
|
||||
|
||||
scanner = ScannerWithRoots([str(root)])
|
||||
service = ModelLifecycleService(
|
||||
scanner=scanner,
|
||||
metadata_manager=DummyMetadataManager({}),
|
||||
metadata_loader=lambda x: {},
|
||||
)
|
||||
with pytest.raises(ValueError, match="outside configured library"):
|
||||
await service.bulk_delete_models([str(model_ok), str(outside)])
|
||||
|
||||
|
||||
class DummyCache:
|
||||
def __init__(self, raw_data):
|
||||
self.raw_data = raw_data
|
||||
|
||||
@@ -114,6 +114,38 @@ def test_calculate_relative_path_sanitizes_model_and_version_names(isolated_sett
|
||||
assert relative_path == "Fancy_Model/Version_One"
|
||||
|
||||
|
||||
def test_calculate_relative_path_sanitizes_leading_slash(isolated_settings):
|
||||
"""Test that empty base_model does NOT produce a leading slash in the path."""
|
||||
isolated_settings["download_path_templates"]["lora"] = "{base_model}/{first_tag}"
|
||||
|
||||
model_data = {
|
||||
"base_model": "",
|
||||
"tags": [],
|
||||
"civitai": {"id": 1, "creator": {"username": "Author"}},
|
||||
}
|
||||
|
||||
relative_path = calculate_relative_path_for_model(model_data, "lora")
|
||||
|
||||
assert not relative_path.startswith("/")
|
||||
assert relative_path == "no tags"
|
||||
|
||||
|
||||
def test_calculate_relative_path_sanitizes_double_slashes(isolated_settings):
|
||||
"""Test that empty substitutions don't produce double slashes."""
|
||||
isolated_settings["download_path_templates"]["lora"] = "{base_model}/{first_tag}/{author}"
|
||||
|
||||
model_data = {
|
||||
"base_model": "",
|
||||
"tags": [],
|
||||
"civitai": {"id": 1, "creator": {"username": "Author"}},
|
||||
}
|
||||
|
||||
relative_path = calculate_relative_path_for_model(model_data, "lora")
|
||||
|
||||
assert "//" not in relative_path
|
||||
assert relative_path == "no tags/Author"
|
||||
|
||||
|
||||
def test_calculate_recipe_fingerprint_filters_and_sorts():
|
||||
loras = [
|
||||
{"hash": "ABC", "strength": 0.1234},
|
||||
|
||||
@@ -16,6 +16,7 @@ export const LORA_PROVIDER_NODE_TYPES = [
|
||||
"Lora Stacker (LoraManager)",
|
||||
"Lora Randomizer (LoraManager)",
|
||||
"Lora Cycler (LoraManager)",
|
||||
"Create Hook LoRA (LoraManager)",
|
||||
] as const;
|
||||
|
||||
/**
|
||||
|
||||
143
web/comfyui/create_hook_lora.js
Normal file
143
web/comfyui/create_hook_lora.js
Normal file
@@ -0,0 +1,143 @@
|
||||
import { app } from "../../scripts/app.js";
|
||||
import {
|
||||
getActiveLorasFromNode,
|
||||
updateConnectedTriggerWords,
|
||||
chainCallback,
|
||||
mergeLoras,
|
||||
getWidgetByName,
|
||||
getWidgetSerializedValue,
|
||||
} from "./utils.js";
|
||||
import { addLorasWidget } from "./loras_widget.js";
|
||||
import { applyLoraValuesToText, debounce } from "./lora_syntax_utils.js";
|
||||
import { applySelectionHighlight } from "./trigger_word_highlight.js";
|
||||
import { updateConnectedLoraInfoNodes } from "./lora_info.js";
|
||||
|
||||
app.registerExtension({
|
||||
name: "LoraManager.CreateHookLora",
|
||||
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if (nodeType.comfyClass === "Create Hook LoRA (LoraManager)") {
|
||||
chainCallback(nodeType.prototype, "onNodeCreated", function () {
|
||||
// Enable widget serialization so loras widget state is persisted
|
||||
this.serialize_widgets = true;
|
||||
|
||||
this.addInput("prev_hooks", "HOOKS", {
|
||||
shape: 7,
|
||||
});
|
||||
|
||||
// Flags to prevent callback loops between text widget ↔ loras widget
|
||||
let isUpdating = false;
|
||||
let isSyncingInput = false;
|
||||
|
||||
// Get the text input widget (AUTOCOMPLETE_TEXT_LORAS type, created by Vue widgets)
|
||||
const inputWidget = getWidgetByName(this, "text");
|
||||
if (!inputWidget) {
|
||||
console.warn(
|
||||
"LoRA Manager: text widget not found for Create Hook LoRA"
|
||||
);
|
||||
return;
|
||||
}
|
||||
this.inputWidget = inputWidget;
|
||||
|
||||
const scheduleInputSync = debounce((lorasValue) => {
|
||||
if (isSyncingInput) {
|
||||
return;
|
||||
}
|
||||
|
||||
isSyncingInput = true;
|
||||
isUpdating = true;
|
||||
|
||||
try {
|
||||
const nextText = applyLoraValuesToText(
|
||||
inputWidget.value,
|
||||
lorasValue
|
||||
);
|
||||
|
||||
if (inputWidget.value !== nextText) {
|
||||
inputWidget.value = nextText;
|
||||
}
|
||||
} finally {
|
||||
isUpdating = false;
|
||||
isSyncingInput = false;
|
||||
}
|
||||
});
|
||||
|
||||
// Create the LoRA list widget
|
||||
const result = addLorasWidget(
|
||||
this,
|
||||
"loras",
|
||||
{
|
||||
onSelectionChange: (selection) => {
|
||||
applySelectionHighlight(this, selection);
|
||||
updateConnectedLoraInfoNodes(this, selection);
|
||||
},
|
||||
},
|
||||
(value) => {
|
||||
// Prevent recursive calls
|
||||
if (isUpdating) return;
|
||||
isUpdating = true;
|
||||
|
||||
try {
|
||||
// Update connected trigger word toggles with active LoRA names
|
||||
const activeLoraNames = new Set();
|
||||
value.forEach((lora) => {
|
||||
if (lora.active) {
|
||||
activeLoraNames.add(lora.name);
|
||||
}
|
||||
});
|
||||
updateConnectedTriggerWords(this, activeLoraNames);
|
||||
} finally {
|
||||
isUpdating = false;
|
||||
}
|
||||
|
||||
scheduleInputSync(value);
|
||||
}
|
||||
);
|
||||
|
||||
this.lorasWidget = result.widget;
|
||||
|
||||
// Set up callback for the text input widget to trigger merge logic
|
||||
inputWidget.callback = (value) => {
|
||||
if (isUpdating) return;
|
||||
isUpdating = true;
|
||||
|
||||
try {
|
||||
const currentLoras = this.lorasWidget?.value || [];
|
||||
const mergedLoras = mergeLoras(value, currentLoras);
|
||||
if (this.lorasWidget) {
|
||||
this.lorasWidget.value = mergedLoras;
|
||||
}
|
||||
|
||||
// Update connected trigger word toggles
|
||||
const activeLoraNames = getActiveLorasFromNode(this);
|
||||
updateConnectedTriggerWords(this, activeLoraNames);
|
||||
} finally {
|
||||
isUpdating = false;
|
||||
}
|
||||
};
|
||||
});
|
||||
}
|
||||
},
|
||||
|
||||
async loadedGraphNode(node) {
|
||||
if (node.comfyClass === "Create Hook LoRA (LoraManager)") {
|
||||
// Restore saved loras widget values on workflow load
|
||||
let existingLoras = [];
|
||||
if (node.widgets_values && node.widgets_values.length > 0) {
|
||||
const savedValue = getWidgetSerializedValue(node, "loras");
|
||||
existingLoras = savedValue || [];
|
||||
}
|
||||
// Merge the loras data from text widget with saved values
|
||||
const inputWidget =
|
||||
node.inputWidget || getWidgetByName(node, "text");
|
||||
if (!inputWidget) {
|
||||
console.warn(
|
||||
"LoRA Manager: text widget not found while restoring Create Hook LoRA"
|
||||
);
|
||||
return;
|
||||
}
|
||||
const mergedLoras = mergeLoras(inputWidget.value, existingLoras);
|
||||
node.lorasWidget.value = mergedLoras;
|
||||
}
|
||||
},
|
||||
});
|
||||
@@ -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;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ const LORA_NODE_CLASSES = new Set([
|
||||
"Lora Loader (LoraManager)",
|
||||
"Lora Stacker (LoraManager)",
|
||||
"WanVideo Lora Select (LoraManager)",
|
||||
"Create Hook LoRA (LoraManager)",
|
||||
]);
|
||||
|
||||
function normalizeTriggerWordList(triggerWords) {
|
||||
|
||||
@@ -8,6 +8,7 @@ export const LORA_PROVIDER_NODE_TYPES = [
|
||||
"Lora Stacker (LoraManager)",
|
||||
"Lora Randomizer (LoraManager)",
|
||||
"Lora Cycler (LoraManager)",
|
||||
"Create Hook LoRA (LoraManager)",
|
||||
];
|
||||
|
||||
export const LORA_STACK_AGGREGATOR_NODE_TYPES = [
|
||||
|
||||
@@ -15656,7 +15656,8 @@ function createVueWidgetCleanup(vueApp, onCleanup) {
|
||||
const LORA_PROVIDER_NODE_TYPES$1 = [
|
||||
"Lora Stacker (LoraManager)",
|
||||
"Lora Randomizer (LoraManager)",
|
||||
"Lora Cycler (LoraManager)"
|
||||
"Lora Cycler (LoraManager)",
|
||||
"Create Hook LoRA (LoraManager)"
|
||||
];
|
||||
const LORA_STACK_AGGREGATOR_NODE_TYPES$1 = [
|
||||
"Lora Stack Combiner (LoraManager)"
|
||||
@@ -15781,7 +15782,8 @@ const ROOT_GRAPH_ID = "root";
|
||||
const LORA_PROVIDER_NODE_TYPES = [
|
||||
"Lora Stacker (LoraManager)",
|
||||
"Lora Randomizer (LoraManager)",
|
||||
"Lora Cycler (LoraManager)"
|
||||
"Lora Cycler (LoraManager)",
|
||||
"Create Hook LoRA (LoraManager)"
|
||||
];
|
||||
const LORA_STACK_AGGREGATOR_NODE_TYPES = [
|
||||
"Lora Stack Combiner (LoraManager)"
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -9,6 +9,7 @@ const LORA_NODE_CLASSES = new Set([
|
||||
"Lora Loader (LoraManager)",
|
||||
"Lora Stacker (LoraManager)",
|
||||
"WanVideo Lora Select (LoraManager)",
|
||||
"Create Hook LoRA (LoraManager)",
|
||||
]);
|
||||
|
||||
const TARGET_WIDGET_NAMES = new Set(["ckpt_name", "unet_name"]);
|
||||
|
||||
Reference in New Issue
Block a user