mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-08 23:10:15 -03:00
Compare commits
9 Commits
v1.1.9
...
a8283a0d00
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
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)
|
- Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json)
|
||||||
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
||||||
- Run `python scripts/sync_translation_keys.py` after adding UI strings to `locales/en.json`
|
- 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
|
## Git / Commit Messages
|
||||||
|
|
||||||
|
|||||||
@@ -16,6 +16,156 @@ from PIL import Image, PngImagePlugin
|
|||||||
import piexif
|
import piexif
|
||||||
import logging
|
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__)
|
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.",
|
"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": (
|
"embed_workflow": (
|
||||||
"BOOLEAN",
|
"BOOLEAN",
|
||||||
{
|
{
|
||||||
"default": False,
|
"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": (
|
"save_with_metadata": (
|
||||||
@@ -142,148 +310,181 @@ class SaveImageLM:
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def format_metadata(self, metadata_dict):
|
def _resolve_model_cache_entry(self, scanner_type: str, name: str):
|
||||||
"""Format metadata in the requested format similar to userComment example"""
|
"""Resolve model hash, civitai metadata, and base_model from scanner cache.
|
||||||
if not metadata_dict:
|
Returns (hash_str, civitai_dict, base_model_str). All values are empty defaults when not found."""
|
||||||
return ""
|
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
|
entry = self._get_cached_model_by_name(scanner, name)
|
||||||
def add_param_if_not_none(param_list, label, value):
|
if entry is None:
|
||||||
if value is not None:
|
basename = os.path.splitext(os.path.basename(name))[0]
|
||||||
param_list.append(f"{label}: {value}")
|
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", "")
|
prompt = metadata_dict.get("prompt", "")
|
||||||
negative_prompt = metadata_dict.get("negative_prompt", "")
|
negative_prompt = metadata_dict.get("negative_prompt", "")
|
||||||
|
steps = metadata_dict.get("steps")
|
||||||
# Extract loras from the prompt if present
|
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", "")
|
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:
|
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>
|
# Resolve checkpoint hash and Civitai data from local cache
|
||||||
lora_matches = re.findall(r"<lora:([^:]+):([^>]+)>", loras_text)
|
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
|
# Resolve LoRA hash and Civitai data from local cache
|
||||||
for lora_name, strength in lora_matches:
|
loras_data: list[dict] = []
|
||||||
hash_value = self.get_lora_hash(lora_name)
|
for lora_name, strength in lora_entries:
|
||||||
if hash_value:
|
lora_hash, lora_civitai, lora_base_model = self._resolve_model_cache_entry(
|
||||||
lora_hashes[lora_name] = hash_value
|
"lora_scanner", lora_name
|
||||||
else:
|
)
|
||||||
prompt_with_loras = prompt
|
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)
|
# Build Hashes JSON (A1111 / Civitai standard format)
|
||||||
metadata_parts = [prompt_with_loras]
|
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:
|
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: list[str] = []
|
||||||
params = []
|
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
|
lines.append(", ".join(params))
|
||||||
if "steps" in metadata_dict:
|
return "\n".join(lines)
|
||||||
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)
|
|
||||||
|
|
||||||
# credit to nkchocoai
|
# credit to nkchocoai
|
||||||
# Add format_filename method to handle pattern substitution
|
# Add format_filename method to handle pattern substitution
|
||||||
@@ -573,6 +774,8 @@ class SaveImageLM:
|
|||||||
extra_pnginfo=None,
|
extra_pnginfo=None,
|
||||||
lossless_webp=True,
|
lossless_webp=True,
|
||||||
quality=100,
|
quality=100,
|
||||||
|
webp_method=6,
|
||||||
|
jpeg_subsampling=0,
|
||||||
embed_workflow=False,
|
embed_workflow=False,
|
||||||
save_with_metadata=True,
|
save_with_metadata=True,
|
||||||
add_counter_to_filename=True,
|
add_counter_to_filename=True,
|
||||||
@@ -627,15 +830,14 @@ class SaveImageLM:
|
|||||||
elif file_format == "jpeg":
|
elif file_format == "jpeg":
|
||||||
file = base_filename + ".jpg"
|
file = base_filename + ".jpg"
|
||||||
file_extension = ".jpg"
|
file_extension = ".jpg"
|
||||||
save_kwargs = {"quality": quality, "optimize": True}
|
save_kwargs = {"quality": quality, "optimize": True, "subsampling": jpeg_subsampling}
|
||||||
elif file_format == "webp":
|
elif file_format == "webp":
|
||||||
file = base_filename + ".webp"
|
file = base_filename + ".webp"
|
||||||
file_extension = ".webp"
|
file_extension = ".webp"
|
||||||
# Add optimization param to control performance
|
|
||||||
save_kwargs = {
|
save_kwargs = {
|
||||||
"quality": quality,
|
"quality": quality,
|
||||||
"lossless": lossless_webp,
|
"lossless": lossless_webp,
|
||||||
"method": 0,
|
"method": webp_method,
|
||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported file format: {file_format}")
|
raise ValueError(f"Unsupported file format: {file_format}")
|
||||||
@@ -722,6 +924,8 @@ class SaveImageLM:
|
|||||||
extra_pnginfo=None,
|
extra_pnginfo=None,
|
||||||
lossless_webp=True,
|
lossless_webp=True,
|
||||||
quality=100,
|
quality=100,
|
||||||
|
webp_method=6,
|
||||||
|
jpeg_subsampling=0,
|
||||||
embed_workflow=False,
|
embed_workflow=False,
|
||||||
save_with_metadata=True,
|
save_with_metadata=True,
|
||||||
add_counter_to_filename=True,
|
add_counter_to_filename=True,
|
||||||
@@ -751,6 +955,8 @@ class SaveImageLM:
|
|||||||
extra_pnginfo,
|
extra_pnginfo,
|
||||||
lossless_webp,
|
lossless_webp,
|
||||||
quality,
|
quality,
|
||||||
|
webp_method,
|
||||||
|
jpeg_subsampling,
|
||||||
embed_workflow,
|
embed_workflow,
|
||||||
save_with_metadata,
|
save_with_metadata,
|
||||||
add_counter_to_filename,
|
add_counter_to_filename,
|
||||||
|
|||||||
@@ -537,6 +537,7 @@ class ModelManagementHandler:
|
|||||||
# Update model_data with new hash
|
# Update model_data with new hash
|
||||||
model_data["sha256"] = sha256
|
model_data["sha256"] = sha256
|
||||||
model_data["hash_status"] = "completed"
|
model_data["hash_status"] = "completed"
|
||||||
|
hash_status = "completed"
|
||||||
else:
|
else:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "No SHA256 hash found"}, status=400
|
{"success": False, "error": "No SHA256 hash found"}, status=400
|
||||||
@@ -544,6 +545,32 @@ class ModelManagementHandler:
|
|||||||
|
|
||||||
await MetadataManager.hydrate_model_data(model_data)
|
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(
|
success, error = await self._metadata_sync.fetch_and_update_model(
|
||||||
sha256=model_data["sha256"],
|
sha256=model_data["sha256"],
|
||||||
file_path=file_path,
|
file_path=file_path,
|
||||||
@@ -566,7 +593,12 @@ class ModelManagementHandler:
|
|||||||
{"success": False, "error": OFFLINE_FRIENDLY_MESSAGE},
|
{"success": False, "error": OFFLINE_FRIENDLY_MESSAGE},
|
||||||
status=503,
|
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)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
async def relink_civitai(self, request: web.Request) -> web.Response:
|
async def relink_civitai(self, request: web.Request) -> web.Response:
|
||||||
|
|||||||
@@ -1392,8 +1392,8 @@ class DownloadManager:
|
|||||||
base_save_dir = save_dir
|
base_save_dir = save_dir
|
||||||
save_dir = os.path.join(save_dir, relative_path)
|
save_dir = os.path.join(save_dir, relative_path)
|
||||||
# Security: validate path containment after joining
|
# Security: validate path containment after joining
|
||||||
resolved_dir = os.path.realpath(os.path.normpath(save_dir))
|
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
|
||||||
base_dir = os.path.realpath(os.path.normpath(base_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:
|
if not resolved_dir.startswith(base_dir + os.sep) and resolved_dir != base_dir:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Path traversal detected: %s escapes %s",
|
"Path traversal detected: %s escapes %s",
|
||||||
|
|||||||
@@ -566,18 +566,52 @@ class LLMService:
|
|||||||
if effective_max is None:
|
if effective_max is None:
|
||||||
effective_max = 4096
|
effective_max = 4096
|
||||||
|
|
||||||
result = await self.chat_completion(
|
# Use json_schema (not json_object) for broader provider compatibility:
|
||||||
messages=messages,
|
# LM Studio and some other OpenAI-compatible servers reject
|
||||||
model=model,
|
# json_object but accept json_schema. {"type": "object"} is
|
||||||
temperature=temperature,
|
# functionally equivalent — it accepts any JSON object without
|
||||||
response_format={"type": "json_object"},
|
# constraining specific fields.
|
||||||
max_tokens=effective_max,
|
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 ""
|
content = result.get("content", "") or ""
|
||||||
if not content:
|
if not content:
|
||||||
raise LLMResponseError(
|
raise LLMResponseError(
|
||||||
"LLM returned empty content in json_object mode. "
|
"LLM returned empty content. "
|
||||||
f"Raw response: {json.dumps(result)[:500]}"
|
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:
|
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.
|
"""Raise ``ValueError`` if *file_path* is not inside a configured model root.
|
||||||
|
|
||||||
Uses ``os.path.realpath()`` to resolve symlinks before comparing,
|
Uses ``os.path.abspath()`` (NOT ``realpath``) to resolve ``..`` and ``.``
|
||||||
so symlink-based escapes are also caught. Skips when the scanner
|
while preserving symlinks — this keeps the check in business-path space.
|
||||||
does not expose ``get_model_roots`` or the list is empty.
|
Skips when the scanner does not expose ``get_model_roots`` or the list
|
||||||
|
is empty.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
roots = None
|
roots = None
|
||||||
@@ -65,10 +66,10 @@ def _require_path_in_library_roots(file_path: str, scanner, *, label: str = "pat
|
|||||||
if not roots:
|
if not roots:
|
||||||
return
|
return
|
||||||
|
|
||||||
resolved = os.path.realpath(os.path.normpath(file_path))
|
resolved = os.path.abspath(os.path.normpath(file_path))
|
||||||
|
|
||||||
for root in roots:
|
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):
|
if resolved == root_resolved or resolved.startswith(root_resolved + os.sep):
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|||||||
@@ -126,6 +126,7 @@ class BulkMetadataRefreshUseCase:
|
|||||||
if sha256:
|
if sha256:
|
||||||
model["sha256"] = sha256
|
model["sha256"] = sha256
|
||||||
model["hash_status"] = "completed"
|
model["hash_status"] = "completed"
|
||||||
|
hash_status = "completed"
|
||||||
else:
|
else:
|
||||||
self._logger.error(f"Failed to calculate hash for {file_path}")
|
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"})
|
failures.append({"name": model.get("model_name", file_path or "Unknown"), "error": "Failed to calculate hash"})
|
||||||
@@ -148,6 +149,16 @@ class BulkMetadataRefreshUseCase:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
await MetadataManager.hydrate_model_data(model)
|
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(
|
result, error_msg = await self._metadata_sync.fetch_and_update_model(
|
||||||
sha256=model["sha256"],
|
sha256=model["sha256"],
|
||||||
file_path=model["file_path"],
|
file_path=model["file_path"],
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ def test_save_image_defaults_to_writing_png_metadata(monkeypatch, tmp_path):
|
|||||||
|
|
||||||
image_path = tmp_path / "sample_00001_.png"
|
image_path = tmp_path / "sample_00001_.png"
|
||||||
with Image.open(image_path) as img:
|
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(
|
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 recipe["gen_params"] == {"prompt": "prompt text", "seed": 123}
|
||||||
assert scanner._json_path_map[recipe["id"]] == os.path.normpath(str(recipe_files[0]))
|
assert scanner._json_path_map[recipe["id"]] == os.path.normpath(str(recipe_files[0]))
|
||||||
assert scanner.fts_updates == [(recipe["id"], "add")]
|
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"
|
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):
|
def test_distribute_preview_to_entries_moves_and_copies(tmp_path):
|
||||||
"""Test that preview distribution moves file to first entry and copies to others."""
|
"""Test that preview distribution moves file to first entry and copies to others."""
|
||||||
manager = DownloadManager()
|
manager = DownloadManager()
|
||||||
|
|||||||
@@ -243,6 +243,56 @@ class TestLLMServiceChatCompletionJson:
|
|||||||
|
|
||||||
assert result == {"key": "value"}
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_chat_completion_json_raises_on_non_json(self, llm_service):
|
async def test_chat_completion_json_raises_on_non_json(self, llm_service):
|
||||||
# Non-JSON content raises LLMResponseError (salvage also fails)
|
# Non-JSON content raises LLMResponseError (salvage also fails)
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import json
|
import json
|
||||||
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -51,11 +52,12 @@ class TestRequirePathInLibraryRoots:
|
|||||||
scanner = ScannerWithRoots([str(root)])
|
scanner = ScannerWithRoots([str(root)])
|
||||||
_require_path_in_library_roots(str(root), scanner)
|
_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 = tmp_path / "loras"
|
||||||
root.mkdir()
|
root.mkdir()
|
||||||
model = root / "model.safetensors"
|
|
||||||
model.write_text("")
|
|
||||||
|
|
||||||
outside_dir = tmp_path / "outside"
|
outside_dir = tmp_path / "outside"
|
||||||
outside_dir.mkdir()
|
outside_dir.mkdir()
|
||||||
@@ -65,9 +67,22 @@ class TestRequirePathInLibraryRoots:
|
|||||||
symlink = root / "link.safetensors"
|
symlink = root / "link.safetensors"
|
||||||
symlink.symlink_to(outside_file)
|
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)])
|
scanner = ScannerWithRoots([str(root)])
|
||||||
with pytest.raises(ValueError, match="outside configured library"):
|
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:
|
class ScannerForDelete:
|
||||||
|
|||||||
@@ -37,22 +37,28 @@ app.registerExtension({
|
|||||||
|
|
||||||
// Handle broadcast mode (for Desktop/non-browser support)
|
// Handle broadcast mode (for Desktop/non-browser support)
|
||||||
if (numericNodeId === -1) {
|
if (numericNodeId === -1) {
|
||||||
// Find all Lora Loader nodes in the current graph
|
// Find all compatible nodes in the current graph
|
||||||
const loraLoaderNodes = getAllGraphNodes(app.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)
|
.map(({ node }) => node)
|
||||||
.filter((node) => node?.comfyClass === "Lora Loader (LoraManager)");
|
.filter((node) => compatibleClasses.has(node?.comfyClass));
|
||||||
|
|
||||||
// Update each Lora Loader node found
|
// Update each node found
|
||||||
if (loraLoaderNodes.length > 0) {
|
if (targetNodes.length > 0) {
|
||||||
loraLoaderNodes.forEach((node) => {
|
targetNodes.forEach((node) => {
|
||||||
this.updateNodeLoraCode(node, loraCode, mode);
|
this.updateNodeLoraCode(node, loraCode, mode);
|
||||||
});
|
});
|
||||||
console.log(
|
console.log(
|
||||||
`Updated ${loraLoaderNodes.length} Lora Loader nodes in broadcast mode`
|
`Updated ${targetNodes.length} nodes in broadcast mode`
|
||||||
);
|
);
|
||||||
} else {
|
} else {
|
||||||
console.warn(
|
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 ||
|
||||||
(node.comfyClass !== "Lora Loader (LoraManager)" &&
|
(node.comfyClass !== "Lora Loader (LoraManager)" &&
|
||||||
node.comfyClass !== "Lora Stacker (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(
|
console.warn(
|
||||||
"Node not found or not a LoraLoader:",
|
"Node not found or not a compatible LoRA node:",
|
||||||
graphId ?? "root",
|
graphId ?? "root",
|
||||||
nodeId
|
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,
|
hideOnZoom: true,
|
||||||
selectOn: ['click', 'focus']
|
selectOn: ['click', 'focus']
|
||||||
|
|||||||
@@ -37,17 +37,18 @@ export function handleStrengthDrag(name, initialStrength, initialX, event, widge
|
|||||||
syncClipStrengthIfCollapsed(lorasData[loraIndex]);
|
syncClipStrengthIfCollapsed(lorasData[loraIndex]);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Update the widget value only if updateWidget flag is true
|
// Always write back to widget.value to persist the mutation.
|
||||||
// This allows us to update inputs directly during drag without triggering re-render
|
// During drag (updateWidget=false), setValue skips renderLoras via __dragActive flag,
|
||||||
if (updateWidget) {
|
// so the DOM survives and pointer capture is preserved.
|
||||||
widget.value = formatLoraValue(lorasData);
|
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) {
|
if (updateWidget && widget.callback) {
|
||||||
widget.callback(widget.value);
|
widget.callback(widget.value);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return newStrength;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Function to handle proportional strength adjustment for all LoRAs via header dragging
|
// 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);
|
lorasData[index].clipStrength = Number(newClipStrength);
|
||||||
});
|
});
|
||||||
|
|
||||||
// Update widget value only if updateWidget flag is true
|
// Always write back to widget.value to persist mutations.
|
||||||
if (updateWidget) {
|
// During drag (updateWidget=false), setValue skips renderLoras via __dragActive flag.
|
||||||
widget.value = formatLoraValue(lorasData);
|
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) {
|
if (updateWidget && widget.callback) {
|
||||||
widget.callback(widget.value);
|
widget.callback(widget.value);
|
||||||
}
|
}
|
||||||
@@ -149,6 +149,13 @@ export function initDrag(
|
|||||||
activePointerId = e.pointerId;
|
activePointerId = e.pointerId;
|
||||||
currentDragElement = e.currentTarget;
|
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
|
// Capture pointer to receive all subsequent events regardless of stopPropagation
|
||||||
const target = e.currentTarget;
|
const target = e.currentTarget;
|
||||||
target.setPointerCapture(e.pointerId);
|
target.setPointerCapture(e.pointerId);
|
||||||
@@ -181,17 +188,12 @@ export function initDrag(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Call the strength adjustment function without updating widget.value during drag
|
// 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
|
// Update strength input directly instead of re-rendering to avoid losing event listeners
|
||||||
const strengthInput = currentDragElement.querySelector('.lm-lora-strength-input');
|
const strengthInput = currentDragElement.querySelector('.lm-lora-strength-input');
|
||||||
if (strengthInput) {
|
if (strengthInput && typeof newStrength === 'number') {
|
||||||
const lorasData = parseLoraValue(widget.value);
|
strengthInput.value = newStrength.toFixed(2);
|
||||||
const loraData = lorasData.find(l => l.name === name);
|
|
||||||
if (loraData) {
|
|
||||||
const strengthValue = isClipStrength ? loraData.clipStrength : loraData.strength;
|
|
||||||
strengthInput.value = Number(strengthValue).toFixed(2);
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Prevent showing the preview tooltip during drag
|
// Prevent showing the preview tooltip during drag
|
||||||
@@ -226,23 +228,30 @@ export function initDrag(
|
|||||||
// Remove the class to restore normal cursor behavior
|
// Remove the class to restore normal cursor behavior
|
||||||
document.body.classList.remove('lm-lora-strength-dragging');
|
document.body.classList.remove('lm-lora-strength-dragging');
|
||||||
|
|
||||||
// Only call onDragEnd and re-render if we actually dragged
|
// Only call onDragEnd and re-render if we actually dragged.
|
||||||
if (wasDragging) {
|
// try-finally guarantees __dragActive is always cleared, preventing a
|
||||||
if (typeof onDragEnd === 'function') {
|
// permanent UI freeze if onDragEnd or setValue throws during cleanup.
|
||||||
onDragEnd();
|
try {
|
||||||
}
|
if (wasDragging) {
|
||||||
|
if (typeof onDragEnd === 'function') {
|
||||||
|
onDragEnd();
|
||||||
|
}
|
||||||
|
|
||||||
// Commit final value through options.setValue so external observers are notified.
|
// Re-enable renderLoras in setValue and flush final value through setter.
|
||||||
// During drag, handleStrengthDrag mutates widgetValue in-place (updateWidget=false),
|
// The last handleStrengthDrag call already wrote the final strength to
|
||||||
// bypassing widget.value setter and options.setValue entirely. This assignment
|
// widgetValue via setValue (with render suppressed). widget.value = widget.value
|
||||||
// flushes the in-place mutation through the setter so any setValue wrappers fire.
|
// triggers setValue again, which now calls renderLoras since __dragActive is false.
|
||||||
widget.value = widget.value;
|
widget.__dragActive = false;
|
||||||
if (typeof widget.callback === 'function') {
|
widget.value = widget.value;
|
||||||
widget.callback(widget.value);
|
if (typeof widget.callback === 'function') {
|
||||||
|
widget.callback(widget.value);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
} finally {
|
||||||
|
widget.__dragActive = false;
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
dragEl.addEventListener('pointerup', endDrag);
|
dragEl.addEventListener('pointerup', endDrag);
|
||||||
dragEl.addEventListener('pointercancel', endDrag);
|
dragEl.addEventListener('pointercancel', endDrag);
|
||||||
}
|
}
|
||||||
@@ -285,6 +294,9 @@ export function initHeaderDrag(headerEl, widget, renderFunction) {
|
|||||||
activePointerId = e.pointerId;
|
activePointerId = e.pointerId;
|
||||||
currentHeaderElement = e.currentTarget;
|
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
|
// Capture pointer to receive all subsequent events regardless of stopPropagation
|
||||||
const target = e.currentTarget;
|
const target = e.currentTarget;
|
||||||
target.setPointerCapture(e.pointerId);
|
target.setPointerCapture(e.pointerId);
|
||||||
@@ -352,13 +364,20 @@ export function initHeaderDrag(headerEl, widget, renderFunction) {
|
|||||||
// Remove the class to restore normal cursor behavior
|
// Remove the class to restore normal cursor behavior
|
||||||
document.body.classList.remove('lm-lora-strength-dragging');
|
document.body.classList.remove('lm-lora-strength-dragging');
|
||||||
|
|
||||||
// Only re-render if we actually dragged
|
// Only re-render if we actually dragged.
|
||||||
if (wasDragging) {
|
// try-finally guarantees __dragActive is always cleared, preventing a
|
||||||
// Commit final value through options.setValue so external observers are notified.
|
// permanent UI freeze if setValue throws during cleanup.
|
||||||
widget.value = widget.value;
|
try {
|
||||||
if (typeof widget.callback === 'function') {
|
if (wasDragging) {
|
||||||
widget.callback(widget.value);
|
// 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;
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -130,6 +130,35 @@ app.registerExtension({
|
|||||||
widget.serializeValue = () => {
|
widget.serializeValue = () => {
|
||||||
return applyTextReplacements(widget.value);
|
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