diff --git a/__init__.py b/__init__.py index a25cbf2..84358d3 100644 --- a/__init__.py +++ b/__init__.py @@ -15,6 +15,9 @@ sys.modules.setdefault("simple_syrup", _simple_syrup_package) from .simple_syrup.runtime.external_llm_routes import ( # noqa: E402 register_external_llm_routes, ) +from .simple_syrup.runtime.mask_batch_preview_routes import ( # noqa: E402 + register_mask_batch_preview_routes, +) from .simple_syrup.runtime.settings_routes import register_settings_routes # noqa: E402 WEB_DIRECTORY = "./web/dist" @@ -40,6 +43,7 @@ async def comfy_entrypoint() -> object: register_settings_routes() register_external_llm_routes() +register_mask_batch_preview_routes() __all__ = [ "WEB_DIRECTORY", diff --git a/simple_syrup/domain/prompt_control_prompt.py b/simple_syrup/domain/prompt_control_prompt.py index a1bbe4a..eaf3311 100644 --- a/simple_syrup/domain/prompt_control_prompt.py +++ b/simple_syrup/domain/prompt_control_prompt.py @@ -25,9 +25,15 @@ class PreparedPromptChunk: @dataclass(frozen=True) class PreparedPromptSide: - """Store ordered prompt chunks and all scheduling tags for one prompt side.""" + """Store ordered prompt chunks for one positive or negative prompt side.""" chunks: tuple[PreparedPromptChunk, ...] + + +@dataclass(frozen=True) +class PromptSegmentHookPlan: + """Store the combined LoRA schedule for one aligned SEP position.""" + lora_tags: str @@ -44,7 +50,7 @@ def extract_lora_tags(text: str) -> str: def prepare_prompt_side(text: str, separator: str) -> PreparedPromptSide: - """Split a prompt side into cleaned chunks and aggregate LoRA tags.""" + """Split a prompt side into ordered cleaned chunks with local LoRA tags.""" chunks = tuple( PreparedPromptChunk( @@ -53,8 +59,7 @@ def prepare_prompt_side(text: str, separator: str) -> PreparedPromptSide: ) for chunk in split_prompt_batch(text, separator) ) - lora_tags = "\n".join(chunk.lora_tags for chunk in chunks if chunk.lora_tags) - return PreparedPromptSide(chunks=chunks, lora_tags=lora_tags) + return PreparedPromptSide(chunks=chunks) def apply_encode_style(encode_style: str, prompt_text: str) -> str: diff --git a/simple_syrup/domain/regional_prompting.py b/simple_syrup/domain/regional_prompting.py new file mode 100644 index 0000000..2c4f90f --- /dev/null +++ b/simple_syrup/domain/regional_prompting.py @@ -0,0 +1,74 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Pure policies for global-first regional prompt pairing.""" + +from __future__ import annotations + +import math +from dataclasses import dataclass + +MAX_REGIONAL_PROMPT_WEIGHT = 1.0 + + +@dataclass(frozen=True) +class RegionalConditioningPair: + """Map one regional conditioning entry to its authored mask index.""" + + conditioning_index: int + mask_index: int + + +@dataclass(frozen=True) +class RegionalConditioningPlan: + """Describe valid positional pairing after the global entry.""" + + region_count: int + pairs: tuple[RegionalConditioningPair, ...] + + +def validate_regional_prompt_weight(weight: float) -> None: + """Reject regional influence values outside the normalized blend range.""" + + if not math.isfinite(weight): + raise ValueError("regional_prompt_weight must be finite.") + if not 0.0 <= weight <= MAX_REGIONAL_PROMPT_WEIGHT: + raise ValueError( + "regional_prompt_weight must be between 0.0 and " + f"{MAX_REGIONAL_PROMPT_WEIGHT:.1f}." + ) + + +def build_regional_conditioning_plan( + *, + region_count: int, + conditioning_count: int, + input_name: str, +) -> RegionalConditioningPlan: + """Return the global-first positional plan or reject excess prompts.""" + + if region_count < 1: + raise ValueError("regional prompting requires at least one authored mask.") + if conditioning_count < 1: + raise ValueError( + f"{input_name} conditioning must contain a global entry at index 0." + ) + + regional_count = conditioning_count - 1 + if regional_count > region_count: + raise ValueError( + f"{input_name} conditioning contains {regional_count} regional " + f"entries but only {region_count} authored masks were provided." + ) + + return RegionalConditioningPlan( + region_count=region_count, + pairs=tuple( + RegionalConditioningPair( + conditioning_index=mask_index + 1, + mask_index=mask_index, + ) + for mask_index in range(regional_count) + ), + ) diff --git a/simple_syrup/masking/regional_prompt_masks.py b/simple_syrup/masking/regional_prompt_masks.py new file mode 100644 index 0000000..aee70eb --- /dev/null +++ b/simple_syrup/masking/regional_prompt_masks.py @@ -0,0 +1,61 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Mask preparation for full-context regional prompting.""" + +from __future__ import annotations + +from collections.abc import Sequence + +import torch + +from .detailer_masks import gaussian_feather_mask + + +def prepare_regional_mask_batch(mask: object, feather: int) -> torch.Tensor: + """Return a validated and optionally feathered BHW mask batch.""" + + if not isinstance(mask, torch.Tensor): + raise TypeError("regional prompting requires a torch MASK tensor.") + if feather < 0: + raise ValueError("region_mask_feather must be greater than or equal to 0.") + + working = mask.float() + if working.ndim == 2: + working = working.unsqueeze(0) + if working.ndim != 3: + raise ValueError("regional prompting requires an HW or BHW MASK tensor.") + if int(working.shape[0]) < 1: + raise ValueError("regional prompting requires at least one authored mask.") + if int(working.shape[1]) < 1 or int(working.shape[2]) < 1: + raise ValueError("regional masks must have non-empty height and width.") + + normalized = working.clamp(0.0, 1.0) + if feather == 0: + return normalized + return gaussian_feather_mask(normalized, feather) + + +def regional_mask(mask_batch: torch.Tensor, index: int) -> torch.Tensor: + """Return one positional region as a singleton BHW mask.""" + + if index < 0 or index >= int(mask_batch.shape[0]): + raise IndexError(f"regional mask index {index} is out of range.") + return mask_batch[index : index + 1] + + +def complementary_global_prompt_mask( + mask_batch: torch.Tensor, + mask_indices: Sequence[int], + regional_prompt_weight: float, +) -> torch.Tensor: + """Return global influence that recedes across accumulated region coverage.""" + + if not mask_indices: + raise ValueError("global prompt masking requires at least one regional mask.") + coverage = torch.zeros_like(mask_batch[0:1]) + for index in mask_indices: + coverage.add_(regional_mask(mask_batch, index)) + coverage.clamp_(0.0, 1.0) + return 1.0 - coverage * regional_prompt_weight diff --git a/simple_syrup/nodes/encode_prompt_batch.py b/simple_syrup/nodes/encode_prompt_batch.py index 6f54511..4702e3c 100644 --- a/simple_syrup/nodes/encode_prompt_batch.py +++ b/simple_syrup/nodes/encode_prompt_batch.py @@ -18,14 +18,12 @@ class EncodePromptBatch: RETURN_TYPES = ("CONDITIONING_BATCH", "CONDITIONING_BATCH") RETURN_NAMES = ("positive", "negative") OUTPUT_TOOLTIPS = ( - "Positive conditioning entries selected by SEGS order.", - "Negative conditioning entries selected by SEGS order.", + "Ordered positive conditioning entries for batch-aware consumers.", + "Ordered negative conditioning entries for batch-aware consumers.", ) FUNCTION = "encode" CATEGORY = "SimpleSyrup/Conditioning" - DESCRIPTION = ( - "Encodes [SEP]-separated prompts into per-segment conditioning batches." - ) + DESCRIPTION = "Encodes [SEP]-separated prompts into ordered conditioning batches." SEARCH_ALIASES = ["conditioning batch", "prompt batch", "segs prompts"] encoder_class: ClassVar[type[ComfyConditioningEncoder]] = ComfyConditioningEncoder @@ -51,7 +49,7 @@ class EncodePromptBatch: "default": "", "multiline": True, "tooltip": ( - "Positive prompts in SEGS order, separated by [SEP]." + "Ordered positive prompt entries separated by [SEP]." ), }, ), @@ -61,7 +59,7 @@ class EncodePromptBatch: "default": "", "multiline": True, "tooltip": ( - "Negative prompts in SEGS order, separated by [SEP]." + "Ordered negative prompt entries separated by [SEP]." ), }, ), diff --git a/simple_syrup/nodes/ksampler_extras.py b/simple_syrup/nodes/ksampler_extras.py index b44b00e..b78d8c5 100644 --- a/simple_syrup/nodes/ksampler_extras.py +++ b/simple_syrup/nodes/ksampler_extras.py @@ -6,13 +6,10 @@ from __future__ import annotations -from importlib import import_module -from typing import Any +from typing import Any, ClassVar -import torch - -from ..domain.conditioning_batch import ConditioningBatch, select_conditioning from ..runtime import sampling_samplers, sampling_schedulers +from ..services.ksampler_sampling_service import KSamplerSamplingService from . import tooltips Latent = dict[str, Any] @@ -30,6 +27,7 @@ class KSamplerExtras: "compatible workflows." ) SEARCH_ALIASES = ["ksampler", "sampler", "ays", "gits", "lcm"] + service_class: ClassVar[type[KSamplerSamplingService]] = KSamplerSamplingService @classmethod def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[Any, ...]]]: @@ -113,151 +111,16 @@ class KSamplerExtras: ) -> tuple[Latent]: """Sample a latent with ComfyUI samplers and extra scheduler sigmas.""" - sampler = sampling_samplers.resolve_sampler(sampler_name) - sigmas = sampling_schedulers.calculate_sigmas( + output = self.service_class().sample( model=model, - scheduler_name=scheduler, - sampler_name=sampler_name, + seed=seed, steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=positive, + negative=negative, + latent_image=latent_image, denoise=denoise, - ).to(model.load_device) - - latent_samples = latent_image["samples"] - comfy_sample = _comfy_sample() - comfy_utils = _comfy_utils() - latent_preview = _latent_preview() - - latent_samples = comfy_sample.fix_empty_latent_channels( - model, - latent_samples, - latent_image.get("downscale_ratio_spacial", None), ) - - batch_inds = ( - latent_image["batch_index"] if "batch_index" in latent_image else None - ) - noise = comfy_sample.prepare_noise(latent_samples, seed, batch_inds) - noise_mask = latent_image.get("noise_mask", None) - - callback = latent_preview.prepare_callback(model, steps) - disable_pbar = not comfy_utils.PROGRESS_BAR_ENABLED - if _uses_conditioning_batch(positive, negative): - samples = _sample_conditioning_batch( - comfy_sample=comfy_sample, - model=model, - noise=noise, - cfg=cfg, - sampler=sampler, - sigmas=sigmas, - positive=positive, - negative=negative, - latent_samples=latent_samples, - noise_mask=noise_mask, - callback=callback, - disable_pbar=disable_pbar, - seed=seed, - ) - else: - samples = comfy_sample.sample_custom( - model, - noise, - cfg, - sampler, - sigmas, - positive, - negative, - latent_samples, - noise_mask=noise_mask, - callback=callback, - disable_pbar=disable_pbar, - seed=seed, - ) - - output = latent_image.copy() - output.pop("downscale_ratio_spacial", None) - output["samples"] = samples return (output,) - - -def _comfy_sample() -> Any: - """Import ComfyUI sample helpers lazily.""" - - import comfy.sample - - return comfy.sample - - -def _uses_conditioning_batch(positive: Any, negative: Any) -> bool: - """Return whether either conditioning input needs per-item selection.""" - - return isinstance(positive, ConditioningBatch) or isinstance( - negative, - ConditioningBatch, - ) - - -def _sample_conditioning_batch( - *, - comfy_sample: Any, - model: Any, - noise: torch.Tensor, - cfg: float, - sampler: Any, - sigmas: torch.Tensor, - positive: Any, - negative: Any, - latent_samples: torch.Tensor, - noise_mask: Any, - callback: Any, - disable_pbar: bool, - seed: int, -) -> torch.Tensor: - """Sample each latent batch item with its selected conditioning.""" - - sampled: list[torch.Tensor] = [] - for index in range(int(latent_samples.shape[0])): - sampled.append( - comfy_sample.sample_custom( - model, - noise[index : index + 1], - cfg, - sampler, - sigmas, - select_conditioning(positive, index), - select_conditioning(negative, index), - latent_samples[index : index + 1], - noise_mask=_slice_noise_mask(noise_mask, index, latent_samples), - callback=callback, - disable_pbar=disable_pbar, - seed=seed, - ) - ) - return torch.cat(sampled, dim=0) - - -def _slice_noise_mask( - noise_mask: Any, - index: int, - latent_samples: torch.Tensor, -) -> Any: - """Return the noise mask slice matching one latent batch item.""" - - if isinstance(noise_mask, torch.Tensor) and noise_mask.shape[0] == int( - latent_samples.shape[0], - ): - return noise_mask[index : index + 1] - return noise_mask - - -def _comfy_utils() -> Any: - """Import ComfyUI utility state lazily.""" - - import comfy.utils - - return comfy.utils - - -def _latent_preview() -> Any: - """Import ComfyUI preview helpers lazily.""" - - return import_module("latent_preview") diff --git a/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py b/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py index 22cd362..c9c18f4 100644 --- a/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py +++ b/simple_syrup/nodes/schedule_and_encode_prompts_with_prompt_control.py @@ -23,15 +23,15 @@ class ScheduleAndEncodePromptsWithPromptControl: ) RETURN_NAMES = ("model", "positive", "negative") OUTPUT_TOOLTIPS = ( - "Model after LoRA tags from positive and negative prompts are scheduled.", + "Model with single-prompt LoRAs applied; SEP-local LoRAs stay on conditioning.", "Positive conditioning or SimpleSyrup conditioning batch.", "Negative conditioning or SimpleSyrup conditioning batch.", ) FUNCTION = "execute" CATEGORY = "SimpleSyrup/Conditioning" DESCRIPTION = ( - "Schedules Prompt-Control LoRAs and encodes prompts, using [SEP] to " - "create SimpleSyrup conditioning batches." + "Schedules Prompt-Control LoRAs and encodes prompts. [SEP] creates " + "conditioning batches with segment-local LoRA hooks." ) SEARCH_ALIASES = ["prompt control", "schedule prompts", "encode prompts"] @@ -67,8 +67,8 @@ class ScheduleAndEncodePromptsWithPromptControl: "default": "", "multiline": False, "tooltip": ( - "Positive Prompt-Control text. [SEP] creates a " - "conditioning batch for SimpleSyrup batch-aware nodes." + "Positive Prompt-Control text; [SEP] creates ordered " + "entries with segment-local LoRA hooks." ), }, ), @@ -78,8 +78,8 @@ class ScheduleAndEncodePromptsWithPromptControl: "default": "", "multiline": False, "tooltip": ( - "Negative Prompt-Control text. [SEP] creates a " - "conditioning batch for SimpleSyrup batch-aware nodes." + "Negative Prompt-Control text; [SEP] creates ordered " + "entries sharing each index's LoRA hooks." ), }, ), diff --git a/simple_syrup/nodes_v3/__init__.py b/simple_syrup/nodes_v3/__init__.py index 940afed..aa2de4a 100644 --- a/simple_syrup/nodes_v3/__init__.py +++ b/simple_syrup/nodes_v3/__init__.py @@ -14,7 +14,10 @@ def get_nodes() -> list[type[object]]: from .batch_region_conditioning import BatchRegionConditioningV3 from .batch_segs import BatchSEGSV3 + from .compose_regional_conditioning import ComposeRegionalConditioningV3 from .external_llm_prompt import ExternalLLMPromptV3 + from .ksampler_prompt_by_region import KSamplerPromptByRegionV3 + from .ksampler_prompt_by_tiled_region import KSamplerPromptByTiledRegionV3 from .legacy_node_wrappers import ( ConditioningBatchAppendV3, ConditioningBatchStartV3, @@ -41,6 +44,7 @@ def get_nodes() -> list[type[object]]: UpscaleLatentFromImageV3, ViTMatteModelLoaderV3, ) + from .load_mask_batch import LoadMaskBatchV3 from .mask_to_segs import MaskToSEGSV3 from .scale_factor import ScaleFactorV3 from .simple_load_checkpoint import SimpleLoadCheckpointV3 @@ -56,6 +60,7 @@ def get_nodes() -> list[type[object]]: BatchSEGSV3, ConditioningBatchAppendV3, ConditioningBatchStartV3, + ComposeRegionalConditioningV3, DetailSEGSAsRegionsV3, DetailSEGSByScaleFactorTiledDiffusionV3, DetailSEGSByScaleFactorV3, @@ -65,10 +70,13 @@ def get_nodes() -> list[type[object]]: GroundedSAMModelInfoV3, GroundingDINOModelLoaderV3, KSamplerExtrasV3, + KSamplerPromptByRegionV3, + KSamplerPromptByTiledRegionV3, KSamplerTiledDiffusionV3, LatentDiagnosticsV3, LayerStyleSAMModelsAdapterV3, LoadUltralyticsModelV3, + LoadMaskBatchV3, MaskToSEGSV3, PromptEncodeStyleAndNormalizationV3, PromptEncodeStyleV3, diff --git a/simple_syrup/nodes_v3/compose_regional_conditioning.py b/simple_syrup/nodes_v3/compose_regional_conditioning.py new file mode 100644 index 0000000..8858208 --- /dev/null +++ b/simple_syrup/nodes_v3/compose_regional_conditioning.py @@ -0,0 +1,83 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Comfy v3 node for composing standard masked regional conditioning.""" + +from __future__ import annotations + +from importlib import import_module +from typing import TYPE_CHECKING, Any, ClassVar + +from ..services.regional_conditioning_service import RegionalConditioningService +from .regional_ksampler_schema import regional_conditioning_inputs + +if TYPE_CHECKING: + + class _ComfyNodeBase: + """Type-checking base for Comfy v3 nodes.""" + + RETURN_TYPES: ClassVar[list[str]] + RETURN_NAMES: ClassVar[list[str]] + +else: + _ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode + +_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io + + +class ComposeRegionalConditioningV3(_ComfyNodeBase): + """Compose global-first batches and ordered masks for native samplers.""" + + conditioning_service_class: ClassVar[type[RegionalConditioningService]] = ( + RegionalConditioningService + ) + + @classmethod + def define_schema(cls) -> Any: + """Declare the standard regional-conditioning composition schema.""" + + return _comfy_io.Schema( + node_id="SimpleSyrup.ComposeRegionalConditioning", + display_name="Compose Regional Conditioning", + category="SimpleSyrup/Conditioning", + description=( + "Pairs global-first prompt batches with ordered masks and returns " + "standard masked conditioning for native Comfy samplers." + ), + search_aliases=["regional prompt", "masked conditioning"], + inputs=regional_conditioning_inputs(_comfy_io), + outputs=[ + _comfy_io.Conditioning.Output( + "positive", + tooltip=( + "Standard masked positive conditioning with hooks preserved." + ), + ), + _comfy_io.Conditioning.Output( + "negative", + tooltip=( + "Standard masked negative conditioning with hooks preserved." + ), + ), + ], + ) + + @classmethod + def execute( + cls, + positive: object, + negative: object, + region_masks: object, + regional_prompt_weight: float, + region_mask_feather: int, + ) -> tuple[object, object]: + """Return hook-preserving standard Comfy conditioning values.""" + + return cls.conditioning_service_class().assemble( + positive=positive, + negative=negative, + masks=region_masks, + regional_prompt_weight=regional_prompt_weight, + region_mask_feather=region_mask_feather, + ) diff --git a/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py b/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py index cc6a68e..b009edf 100644 --- a/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py +++ b/simple_syrup/nodes_v3/encode_prompt_batch_with_prompt_control.py @@ -48,7 +48,7 @@ class EncodePromptBatchWithPromptControl(_ComfyNodeBase): category="SimpleSyrup/Conditioning", description=( "Encodes [SEP]-separated prompts into per-segment Prompt Control " - "conditioning batches." + "conditioning batches with segment-local LoRA hooks." ), inputs=[ _comfy_io.Clip.Input( @@ -64,8 +64,8 @@ class EncodePromptBatchWithPromptControl(_ComfyNodeBase): multiline=True, default="", tooltip=( - "Positive Prompt Control prompts in SEGS order, separated " - "by the separator text." + "Positive Prompt Control prompts in positional order; each " + "segment keeps its aligned LoRA hooks." ), ), _comfy_io.String.Input( @@ -73,8 +73,8 @@ class EncodePromptBatchWithPromptControl(_ComfyNodeBase): multiline=True, default="", tooltip=( - "Negative Prompt Control prompts in SEGS order, separated " - "by the separator text." + "Negative Prompt Control prompts in positional order; each " + "segment shares hooks with the matching positive index." ), ), _comfy_io.String.Input( diff --git a/simple_syrup/nodes_v3/ksampler_prompt_by_region.py b/simple_syrup/nodes_v3/ksampler_prompt_by_region.py new file mode 100644 index 0000000..7a8f8e4 --- /dev/null +++ b/simple_syrup/nodes_v3/ksampler_prompt_by_region.py @@ -0,0 +1,103 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Comfy v3 KSampler for full-context regional prompting.""" + +from __future__ import annotations + +from importlib import import_module +from typing import TYPE_CHECKING, Any, ClassVar + +from ..nodes import tooltips +from ..services.ksampler_sampling_service import KSamplerSamplingService +from ..services.regional_conditioning_service import RegionalConditioningService +from .regional_ksampler_schema import regional_ksampler_inputs + +if TYPE_CHECKING: + + class _ComfyNodeBase: + """Type-checking base for Comfy v3 nodes.""" + + RETURN_TYPES: ClassVar[list[str]] + RETURN_NAMES: ClassVar[list[str]] + +else: + _ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode + +_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io + + +class KSamplerPromptByRegionV3(_ComfyNodeBase): + """Sample a full latent with global-first ordered regional prompts.""" + + conditioning_service_class: ClassVar[type[RegionalConditioningService]] = ( + RegionalConditioningService + ) + sampling_service_class: ClassVar[type[KSamplerSamplingService]] = ( + KSamplerSamplingService + ) + + @classmethod + def define_schema(cls) -> Any: + """Declare the non-tiled regional KSampler schema.""" + + return _comfy_io.Schema( + node_id="SimpleSyrup.KSamplerPromptByRegion", + display_name="KSampler (Prompt by Region)", + category="SimpleSyrup/Sampling", + description=( + "Denoises the full latent with one global prompt and ordered " + "mask-bound regional prompts." + ), + search_aliases=["ksampler", "regional prompt", "masked prompt"], + inputs=regional_ksampler_inputs(_comfy_io), + outputs=[ + _comfy_io.Latent.Output( + "latent", + tooltip=tooltips.DENOISED_LATENT_OUTPUT, + ) + ], + ) + + @classmethod + def execute( + cls, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: object, + negative: object, + region_masks: object, + regional_prompt_weight: float, + region_mask_feather: int, + latent_image: dict[str, Any], + denoise: float, + ) -> tuple[dict[str, Any]]: + """Assemble regional conditioning and sample the full latent.""" + + assembled_positive, assembled_negative = ( + cls.conditioning_service_class().assemble( + positive=positive, + negative=negative, + masks=region_masks, + regional_prompt_weight=regional_prompt_weight, + region_mask_feather=region_mask_feather, + ) + ) + output = cls.sampling_service_class().sample( + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=assembled_positive, + negative=assembled_negative, + latent_image=latent_image, + denoise=denoise, + ) + return (output,) diff --git a/simple_syrup/nodes_v3/ksampler_prompt_by_tiled_region.py b/simple_syrup/nodes_v3/ksampler_prompt_by_tiled_region.py new file mode 100644 index 0000000..cc31294 --- /dev/null +++ b/simple_syrup/nodes_v3/ksampler_prompt_by_tiled_region.py @@ -0,0 +1,123 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Comfy v3 tiled KSampler for full-context regional prompting.""" + +from __future__ import annotations + +from importlib import import_module +from typing import TYPE_CHECKING, Any, ClassVar + +from ..nodes import tooltips +from ..services.regional_conditioning_service import RegionalConditioningService +from ..services.tiled_diffusion_sampling_service import TiledDiffusionSamplingService +from .regional_ksampler_schema import regional_ksampler_inputs, tiled_regional_inputs + +if TYPE_CHECKING: + + class _ComfyNodeBase: + """Type-checking base for Comfy v3 nodes.""" + + RETURN_TYPES: ClassVar[list[str]] + RETURN_NAMES: ClassVar[list[str]] + +else: + _ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode + +_comfy_io: Any = None if TYPE_CHECKING else import_module("comfy_api.latest").io + + +class KSamplerPromptByTiledRegionV3(_ComfyNodeBase): + """Sample tiled latents with global-first ordered regional prompts.""" + + conditioning_service_class: ClassVar[type[RegionalConditioningService]] = ( + RegionalConditioningService + ) + sampling_service_class: ClassVar[type[TiledDiffusionSamplingService]] = ( + TiledDiffusionSamplingService + ) + + @classmethod + def define_schema(cls) -> Any: + """Declare the tiled regional KSampler schema.""" + + return _comfy_io.Schema( + node_id="SimpleSyrup.KSamplerPromptByTiledRegion", + display_name="KSampler (Prompt by Tiled Region)", + category="SimpleSyrup/Sampling", + description=( + "Denoises large latents in overlapping tiles while preserving " + "global and ordered mask-bound regional prompts." + ), + search_aliases=[ + "ksampler", + "regional prompt", + "tiled regional prompt", + "regional hires fix", + ], + inputs=[ + *regional_ksampler_inputs(_comfy_io), + *tiled_regional_inputs(_comfy_io), + ], + outputs=[ + _comfy_io.Latent.Output( + "latent", + tooltip=tooltips.DENOISED_LATENT_OUTPUT, + ) + ], + ) + + @classmethod + def execute( + cls, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: object, + negative: object, + region_masks: object, + regional_prompt_weight: float, + region_mask_feather: int, + latent_image: dict[str, Any], + denoise: float, + diffusion_mode: str, + latent_tile_width: int, + latent_tile_height: int, + latent_tile_overlap: int, + latent_tile_batch_size: int, + ) -> tuple[dict[str, Any]]: + """Assemble regional conditioning and sample overlapping latent tiles.""" + + assembled_positive, assembled_negative = ( + cls.conditioning_service_class().assemble( + positive=positive, + negative=negative, + masks=region_masks, + regional_prompt_weight=regional_prompt_weight, + region_mask_feather=region_mask_feather, + ) + ) + output = cls.sampling_service_class().sample( + diffusion_mode=diffusion_mode, + model=model, + seed=seed, + steps=steps, + cfg=cfg, + sampler_name=sampler_name, + scheduler=scheduler, + positive=assembled_positive, + negative=assembled_negative, + latent_image=latent_image, + denoise=denoise, + latent_tile_width=latent_tile_width, + latent_tile_height=latent_tile_height, + latent_tile_overlap=latent_tile_overlap, + latent_tile_batch_size=latent_tile_batch_size, + preview_context=None, + allow_full_context_masks=True, + ) + return (output,) diff --git a/simple_syrup/nodes_v3/load_mask_batch.py b/simple_syrup/nodes_v3/load_mask_batch.py new file mode 100644 index 0000000..d63acfb --- /dev/null +++ b/simple_syrup/nodes_v3/load_mask_batch.py @@ -0,0 +1,110 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Comfy v3 node for loading one or many authored masks with native widgets.""" + +from __future__ import annotations + +from importlib import import_module +from typing import TYPE_CHECKING, Any, ClassVar + +from ..runtime.mask_file_loader import MASK_CHANNELS +from ..services.load_mask_batch_service import LoadMaskBatchService + +if TYPE_CHECKING: + + class _ComfyNodeBase: + """Type-checking base for Comfy v3 nodes.""" + + RETURN_TYPES: ClassVar[list[str]] + RETURN_NAMES: ClassVar[list[str]] + +else: + _ComfyNodeBase = import_module("comfy_api.latest").io.ComfyNode + +_comfy_api: Any = None if TYPE_CHECKING else import_module("comfy_api.latest") +_comfy_io: Any = None if TYPE_CHECKING else _comfy_api.io +_comfy_ui: Any = None if TYPE_CHECKING else _comfy_api.UI + + +class LoadMaskBatchV3(_ComfyNodeBase): + """Load an ordered set of authored files as one Comfy MASK batch.""" + + service_class: ClassVar[type[LoadMaskBatchService]] = LoadMaskBatchService + + @classmethod + def define_schema(cls) -> Any: + """Declare a native image-upload combo that accepts one or many files.""" + + choices = list(cls.service_class().available_files()) + return _comfy_io.Schema( + node_id="SimpleSyrup.LoadMaskBatch", + display_name="Load Mask Batch", + category="SimpleSyrup/Loaders", + description=( + "Loads one or many authored mask files in selection order as a " + "single mask batch." + ), + search_aliases=["load masks", "mask batch", "regional masks"], + has_intermediate_output=True, + inputs=[ + _comfy_io.MultiCombo.Input( + "image", + options=choices, + default=[], + placeholder="Select one or more masks", + chip=True, + tooltip=( + "Select one or more ordered mask files; each selected " + "file becomes one regional mask." + ), + extra_dict={ + "image_upload": True, + "image_folder": "input", + "allow_batch": True, + }, + ), + _comfy_io.Combo.Input( + "channel", + options=list(MASK_CHANNELS), + default="alpha", + tooltip=( + "Image channel read from every file using ComfyUI mask " + "loading semantics." + ), + ), + ], + outputs=[ + _comfy_io.Mask.Output( + "mask", + tooltip="Ordered BHW mask batch containing one mask per file.", + ), + ], + ) + + @classmethod + def execute(cls, image: str | list[str], channel: str) -> Any: + """Load the selected files and return a native mask-batch preview.""" + + mask_batch = cls.service_class().load(image, channel) + return _comfy_io.NodeOutput( + mask_batch, + ui=_comfy_ui.PreviewMask(mask_batch, cls=cls), + ) + + @classmethod + def validate_inputs(cls, image: str | list[str], channel: str) -> bool | str: + """Validate native widget values before execution or filesystem reads.""" + + try: + cls.service_class().validate(image, channel) + except (TypeError, ValueError) as error: + return str(error) + return True + + @classmethod + def fingerprint_inputs(cls, image: str | list[str], channel: str) -> str: + """Fingerprint the ordered selected files and shared channel.""" + + return cls.service_class().fingerprint(image, channel) diff --git a/simple_syrup/nodes_v3/regional_ksampler_schema.py b/simple_syrup/nodes_v3/regional_ksampler_schema.py new file mode 100644 index 0000000..cde5f9a --- /dev/null +++ b/simple_syrup/nodes_v3/regional_ksampler_schema.py @@ -0,0 +1,169 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Shared Comfy v3 schema declarations for regional KSamplers.""" + +from __future__ import annotations + +from typing import Any + +from ..domain.regional_prompting import MAX_REGIONAL_PROMPT_WEIGHT +from ..domain.tiled_diffusion import TILED_DIFFUSION_MODES +from ..nodes import tooltips +from ..nodes.ksampler_tiled_diffusion import MAX_LATENT_TILE_SIZE +from ..runtime import sampling_samplers, sampling_schedulers + + +def regional_ksampler_inputs(comfy_io: Any) -> list[Any]: + """Return common regional KSampler inputs in workflow order.""" + + return [ + comfy_io.Model.Input("model", tooltip=tooltips.SAMPLING_MODEL), + comfy_io.Int.Input( + "seed", + default=0, + min=0, + max=0xFFFFFFFFFFFFFFFF, + control_after_generate=True, + tooltip=tooltips.SAMPLING_SEED, + ), + comfy_io.Int.Input( + "steps", + default=20, + min=1, + max=10000, + tooltip=tooltips.SAMPLING_STEPS, + ), + comfy_io.Float.Input( + "cfg", + default=8.0, + min=0.0, + max=100.0, + step=0.1, + round=0.01, + tooltip=tooltips.SAMPLING_CFG, + ), + comfy_io.Combo.Input( + "sampler_name", + options=list(sampling_samplers.available_samplers()), + tooltip=tooltips.SAMPLER_NAME, + ), + comfy_io.Combo.Input( + "scheduler", + options=list(sampling_schedulers.available_schedulers()), + tooltip=tooltips.SCHEDULER, + ), + *regional_conditioning_inputs(comfy_io), + comfy_io.Latent.Input("latent_image", tooltip=tooltips.LATENT_IMAGE), + comfy_io.Float.Input( + "denoise", + default=1.0, + min=0.0, + max=1.0, + step=0.01, + tooltip=tooltips.DENOISE_STRENGTH, + ), + ] + + +def regional_conditioning_inputs(comfy_io: Any) -> list[Any]: + """Return the authoritative ordered regional-composition inputs.""" + + conditioning_batch = comfy_io.Custom("CONDITIONING_BATCH") + return [ + comfy_io.MultiType.Input( + "positive", + [comfy_io.Conditioning, conditioning_batch], + tooltip=( + "Positive conditioning whose first batch entry is global and " + "later entries pair with masks in order." + ), + ), + comfy_io.MultiType.Input( + "negative", + [comfy_io.Conditioning, conditioning_batch], + tooltip=( + "Negative conditioning whose first batch entry is global and " + "later entries pair with masks in order." + ), + ), + comfy_io.Mask.Input( + "region_masks", + tooltip=( + "Ordered authored masks; mask 0 pairs with conditioning batch entry 1." + ), + ), + comfy_io.Float.Input( + "regional_prompt_weight", + default=0.5, + min=0.0, + max=MAX_REGIONAL_PROMPT_WEIGHT, + step=0.01, + round=0.01, + tooltip=( + "Balances regional prompts against the global prompt; 0 uses " + "only global prompting, 1 uses only regional prompting inside " + "solid masks, and overlaps reduce the global share further." + ), + ), + comfy_io.Int.Input( + "region_mask_feather", + default=0, + min=0, + max=512, + step=1, + tooltip=( + "Softens regional mask edges by this many image pixels; 0 " + "preserves authored mask values." + ), + ), + ] + + +def tiled_regional_inputs(comfy_io: Any) -> list[Any]: + """Return tile controls matching KSampler Tiled Diffusion.""" + + return [ + comfy_io.Combo.Input( + "diffusion_mode", + options=list(TILED_DIFFUSION_MODES), + default="multidiffusion", + tooltip=( + "Tiled blend method; MultiDiffusion is steady while Mixture of " + "Diffusers weights tile centers more strongly." + ), + ), + comfy_io.Int.Input( + "latent_tile_width", + default=128, + min=16, + max=MAX_LATENT_TILE_SIZE, + step=16, + tooltip=tooltips.LATENT_TILE_WIDTH, + ), + comfy_io.Int.Input( + "latent_tile_height", + default=128, + min=16, + max=MAX_LATENT_TILE_SIZE, + step=16, + tooltip=tooltips.LATENT_TILE_HEIGHT, + ), + comfy_io.Int.Input( + "latent_tile_overlap", + default=16, + min=0, + max=256, + step=4, + tooltip=tooltips.LATENT_TILE_OVERLAP, + ), + comfy_io.Int.Input( + "latent_tile_batch_size", + default=4, + min=1, + max=8, + step=1, + tooltip=tooltips.LATENT_TILE_BATCH_SIZE, + ), + ] diff --git a/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py b/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py index 5b33d56..994c2b0 100644 --- a/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py +++ b/simple_syrup/nodes_v3/schedule_and_encode_prompts_with_prompt_control.py @@ -49,8 +49,8 @@ class ScheduleAndEncodePromptsWithPromptControl(_ComfyNodeBase): enable_expand=True, category="SimpleSyrup/Conditioning", description=( - "Schedules Prompt-Control LoRAs and encodes prompts, using [SEP] " - "to create SimpleSyrup conditioning batches." + "Schedules Prompt-Control LoRAs and encodes prompts. With [SEP], " + "each conditioning entry keeps only its segment's LoRA hooks." ), inputs=[ _comfy_io.Model.Input( @@ -84,8 +84,8 @@ class ScheduleAndEncodePromptsWithPromptControl(_ComfyNodeBase): multiline=False, default="", tooltip=( - "Positive Prompt-Control text. [SEP] creates a " - "conditioning batch for SimpleSyrup batch-aware nodes." + "Positive Prompt-Control text; [SEP] creates ordered " + "conditioning entries with segment-local LoRA hooks." ), ), _comfy_io.String.Input( @@ -93,8 +93,8 @@ class ScheduleAndEncodePromptsWithPromptControl(_ComfyNodeBase): multiline=False, default="", tooltip=( - "Negative Prompt-Control text. [SEP] creates a " - "conditioning batch for SimpleSyrup batch-aware nodes." + "Negative Prompt-Control text; [SEP] creates ordered " + "conditioning entries sharing each index's LoRA hooks." ), ), ], @@ -102,8 +102,8 @@ class ScheduleAndEncodePromptsWithPromptControl(_ComfyNodeBase): _comfy_io.Model.Output( "model", tooltip=( - "Model after LoRA tags from positive and negative prompts " - "are scheduled." + "Model with single-prompt LoRAs applied globally; SEP-local " + "LoRAs travel on their conditioning entries instead." ), ), MixedConditioningIO.Output( diff --git a/simple_syrup/runtime/mask_batch_preview_routes.py b/simple_syrup/runtime/mask_batch_preview_routes.py new file mode 100644 index 0000000..7348d80 --- /dev/null +++ b/simple_syrup/runtime/mask_batch_preview_routes.py @@ -0,0 +1,185 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""HTTP adapter for channel-accurate authored-mask previews.""" + +from __future__ import annotations + +import sys +from collections.abc import Callable, Coroutine, Sequence +from importlib import import_module +from typing import Any, Protocol, cast + +from aiohttp import web + +from ..shared.logging import get_logger + +LOGGER = get_logger(__name__) +MASK_BATCH_PREVIEW_ROUTE = "/simple-syrup/mask-batch/preview" + +Handler = Callable[[Any], Coroutine[Any, Any, web.Response]] +_REGISTERED_PROMPT_SERVERS: set[int] = set() + + +class RoutesProtocol(Protocol): + """Subset of Comfy's route table needed for preview registration.""" + + def post(self, path: str) -> Callable[[Handler], Handler]: + """Return a POST route decorator.""" + + +class PromptServerProtocol(Protocol): + """Subset of Comfy's PromptServer needed for route registration.""" + + routes: RoutesProtocol + + +class MaskBatchLoaderProtocol(Protocol): + """Load ordered masks with the same channel semantics as node execution.""" + + def load_each(self, files: Sequence[str], channel: str) -> Sequence[object]: + """Load ordered files independently through the selected mask channel.""" + + +class MaskBatchPreviewRendererProtocol(Protocol): + """Render loaded masks into Comfy's native execution-output shape.""" + + def render(self, masks: Sequence[object]) -> dict[str, object]: + """Return a JSON-compatible native preview payload.""" + + +class NativeMaskBatchPreviewRenderer: + """Render mask tensors through Comfy's authoritative PreviewMask helper.""" + + def render(self, masks: Sequence[object]) -> dict[str, object]: + """Save one native temporary preview per independently sized mask.""" + + comfy_api: Any = import_module("comfy_api.latest") + images: list[object] = [] + for mask in masks: + preview: Any = comfy_api.UI.PreviewMask(mask) + payload: object = preview.as_dict() + if not isinstance(payload, dict): + raise TypeError( + "Comfy PreviewMask returned an invalid preview payload." + ) + preview_images = payload.get("images") + if not isinstance(preview_images, (list, tuple)): + raise TypeError("Comfy PreviewMask returned invalid preview images.") + images.extend(preview_images) + return {"images": images, "animated": (False,)} + + +class MaskBatchPreviewHandlers: + """Handle previews using native mask semantics without batching constraints.""" + + def __init__( + self, + loader: MaskBatchLoaderProtocol, + renderer: MaskBatchPreviewRendererProtocol, + ) -> None: + """Create handlers with explicit loading and rendering collaborators.""" + + self._loader = loader + self._renderer = renderer + + async def post_preview(self, request: Any) -> web.Response: + """Return native previews for ordered files and one selected channel.""" + + try: + payload: object = await request.json() + except Exception as error: + LOGGER.warning( + "invalid mask batch preview request body", + extra={"route": MASK_BATCH_PREVIEW_ROUTE, "reason": str(error)}, + ) + return web.json_response( + {"error": ("Load Mask Batch preview request body must be valid JSON.")}, + status=400, + ) + + try: + files, channel = self._request_values(payload) + masks = self._loader.load_each(files, channel) + except (TypeError, ValueError) as error: + return web.json_response({"error": str(error)}, status=400) + + try: + preview = self._renderer.render(masks) + except Exception as error: + LOGGER.exception( + "mask batch preview rendering failed", + extra={"route": MASK_BATCH_PREVIEW_ROUTE, "reason": str(error)}, + ) + return web.json_response( + {"error": "ComfyUI could not render the mask batch preview."}, + status=500, + ) + return web.json_response(preview) + + def _request_values(self, payload: object) -> tuple[tuple[str, ...], str]: + """Validate and narrow a JSON preview request.""" + + if not isinstance(payload, dict): + raise TypeError("Load Mask Batch preview payload must be an object.") + + files = payload.get("files") + if not isinstance(files, list): + raise TypeError("Load Mask Batch preview files must be a list.") + if not files: + raise ValueError("Load Mask Batch preview requires at least one mask file.") + if any(not isinstance(path, str) or not path for path in files): + raise TypeError("Load Mask Batch preview files must be non-empty strings.") + + channel = payload.get("channel") + if not isinstance(channel, str) or not channel: + raise TypeError( + "Load Mask Batch preview channel must be a non-empty string." + ) + return tuple(files), channel + + +def register_mask_batch_preview_routes( + loader: MaskBatchLoaderProtocol | None = None, + renderer: MaskBatchPreviewRendererProtocol | None = None, + prompt_server: PromptServerProtocol | None = None, +) -> bool: + """Register mask preview routes with Comfy's PromptServer when available.""" + + server_instance = prompt_server or _prompt_server_instance() + if server_instance is None: + return False + + server_key = id(server_instance) + if prompt_server is None and server_key in _REGISTERED_PROMPT_SERVERS: + return True + + if loader is None: + from ..services.load_mask_batch_service import LoadMaskBatchService + + loader = LoadMaskBatchService() + handlers = MaskBatchPreviewHandlers( + loader, + renderer or NativeMaskBatchPreviewRenderer(), + ) + server_instance.routes.post(MASK_BATCH_PREVIEW_ROUTE)(handlers.post_preview) + if prompt_server is None: + _REGISTERED_PROMPT_SERVERS.add(server_key) + return True + + +def _prompt_server_instance() -> PromptServerProtocol | None: + """Return Comfy's PromptServer instance without importing it eagerly.""" + + try: + server_module = sys.modules["server"] + prompt_server = server_module.PromptServer + instance = prompt_server.instance + except (KeyError, AttributeError) as error: + LOGGER.debug( + "PromptServer unavailable for mask batch preview routes", + extra={"reason": str(error)}, + ) + return None + return cast(PromptServerProtocol, instance) diff --git a/simple_syrup/runtime/mask_file_loader.py b/simple_syrup/runtime/mask_file_loader.py new file mode 100644 index 0000000..cc934fc --- /dev/null +++ b/simple_syrup/runtime/mask_file_loader.py @@ -0,0 +1,91 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""ComfyUI filesystem adapter for authored mask files.""" + +from __future__ import annotations + +import hashlib +from importlib import import_module +from pathlib import Path + +import torch +from PIL import Image + +MASK_CHANNELS: tuple[str, ...] = ("alpha", "red", "green", "blue") + + +class MaskFileLoader: + """Load one static mask through ComfyUI's native mask implementation.""" + + def available_files(self) -> tuple[str, ...]: + """Return choices declared by ComfyUI's native mask loader.""" + + declaration = import_module("nodes").LoadImageMask.INPUT_TYPES() + image_input = declaration["required"]["image"] + choices = image_input[0] + if not isinstance(choices, (list, tuple)): + raise TypeError("ComfyUI LoadImageMask returned invalid image choices.") + return tuple(str(value) for value in choices) + + def validate(self, annotated_path: str, channel: str) -> None: + """Validate one path with ComfyUI's native mask-loader contract.""" + + self._validate_channel(channel) + result = import_module("nodes").LoadImageMask.VALIDATE_INPUTS(annotated_path) + if result is not True: + raise ValueError(str(result)) + + def load(self, annotated_path: str, channel: str) -> torch.Tensor: + """Return exactly one BHW mask from a validated annotated path.""" + + self.validate(annotated_path, channel) + path = self._resolve_path(annotated_path) + with Image.open(path) as image: + frame_count = int(getattr(image, "n_frames", 1)) + if frame_count != 1: + raise ValueError( + f"Load Mask Batch requires one mask per file; {annotated_path!r} " + f"contains {frame_count} frames." + ) + + load_image_mask = import_module("nodes").LoadImageMask() + result = load_image_mask.load_image_mask(annotated_path, channel) + mask = result[0] + if not isinstance(mask, torch.Tensor): + raise TypeError(f"ComfyUI did not return a MASK for {annotated_path!r}.") + if mask.ndim != 3 or int(mask.shape[0]) != 1: + raise ValueError( + f"Load Mask Batch requires one mask per file; {annotated_path!r} " + f"returned shape {tuple(mask.shape)}." + ) + return mask.float().clamp(0.0, 1.0) + + def fingerprint(self, annotated_path: str) -> str: + """Return the content fingerprint for one validated mask file.""" + + path = self._resolve_path(annotated_path) + digest = hashlib.sha256() + with path.open("rb") as handle: + for chunk in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + def _resolve_path(self, annotated_path: str) -> Path: + """Resolve an existing Comfy annotated path without accepting arbitrary IO.""" + + if not isinstance(annotated_path, str) or not annotated_path: + raise ValueError("Load Mask Batch requires a non-empty mask file path.") + folder_paths = import_module("folder_paths") + if not folder_paths.exists_annotated_filepath(annotated_path): + raise ValueError(f"Mask file does not exist: {annotated_path!r}.") + return Path(folder_paths.get_annotated_filepath(annotated_path)) + + def _validate_channel(self, channel: str) -> None: + """Reject channels outside ComfyUI's native mask choices.""" + + if channel not in MASK_CHANNELS: + raise ValueError( + f"mask channel must be one of {MASK_CHANNELS}; received {channel!r}." + ) diff --git a/simple_syrup/runtime/mixture_of_diffusers_sampling.py b/simple_syrup/runtime/mixture_of_diffusers_sampling.py index dc35eb5..fda3501 100644 --- a/simple_syrup/runtime/mixture_of_diffusers_sampling.py +++ b/simple_syrup/runtime/mixture_of_diffusers_sampling.py @@ -62,6 +62,7 @@ def sample_mixture_of_diffusers( latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, differential_diffusion: bool = False, + allow_full_context_masks: bool = False, ) -> Latent: """Sample a latent with a cloned model patched for Mixture of Diffusers.""" @@ -72,8 +73,16 @@ def sample_mixture_of_diffusers( latent_tile_height=latent_tile_height, latent_tile_batch_size=latent_tile_batch_size, ) - reject_unsupported_conditioning(positive, sampler_label=SAMPLER_LABEL) - reject_unsupported_conditioning(negative, sampler_label=SAMPLER_LABEL) + reject_unsupported_conditioning( + positive, + sampler_label=SAMPLER_LABEL, + allow_full_context_masks=allow_full_context_masks, + ) + reject_unsupported_conditioning( + negative, + sampler_label=SAMPLER_LABEL, + allow_full_context_masks=allow_full_context_masks, + ) sampler = sampling_samplers.resolve_sampler(sampler_name) sigmas = sampling_schedulers.calculate_sigmas( diff --git a/simple_syrup/runtime/multidiffusion_sampling.py b/simple_syrup/runtime/multidiffusion_sampling.py index 472b213..c9a732a 100644 --- a/simple_syrup/runtime/multidiffusion_sampling.py +++ b/simple_syrup/runtime/multidiffusion_sampling.py @@ -62,6 +62,7 @@ def sample_multidiffusion( latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, differential_diffusion: bool = False, + allow_full_context_masks: bool = False, ) -> Latent: """Sample a latent with a cloned model patched for MultiDiffusion.""" @@ -73,8 +74,16 @@ def sample_multidiffusion( latent_tile_batch_size=latent_tile_batch_size, ) _reject_unipc_sampler(sampler_name) - reject_unsupported_conditioning(positive, sampler_label=SAMPLER_LABEL) - reject_unsupported_conditioning(negative, sampler_label=SAMPLER_LABEL) + reject_unsupported_conditioning( + positive, + sampler_label=SAMPLER_LABEL, + allow_full_context_masks=allow_full_context_masks, + ) + reject_unsupported_conditioning( + negative, + sampler_label=SAMPLER_LABEL, + allow_full_context_masks=allow_full_context_masks, + ) sampler = sampling_samplers.resolve_sampler(sampler_name) sigmas = sampling_schedulers.calculate_sigmas( diff --git a/simple_syrup/runtime/prompt_control_batch_graph.py b/simple_syrup/runtime/prompt_control_batch_graph.py index 5179a05..43f806a 100644 --- a/simple_syrup/runtime/prompt_control_batch_graph.py +++ b/simple_syrup/runtime/prompt_control_batch_graph.py @@ -2,16 +2,17 @@ # Copyright (C) 2026 Artificial Sweetener and contributors # SPDX-License-Identifier: AGPL-3.0-or-later -"""Runtime graph expansion for Prompt Control conditioning batches.""" +"""Orchestrate Prompt Control SEP conditioning-batch graph expansion.""" from __future__ import annotations -import sys -from importlib import import_module -from typing import Any, cast +from typing import Any -from ..domain.conditioning_batch import split_prompt_batch -from .prompt_control_availability import find_prompt_control_install +from ..domain.prompt_control_prompt import PreparedPromptSide +from ..services.prompt_control_segment_planning_service import ( + PromptControlSegmentPlanningService, +) +from .prompt_control_graph_adapter import PromptControlGraphAdapter PROMPT_CONTROL_MISSING_MESSAGE = ( "Encode Prompt Batch w/ Prompt Control requires comfyui-prompt-control. " @@ -20,7 +21,10 @@ PROMPT_CONTROL_MISSING_MESSAGE = ( class PromptControlBatchGraphBuilder: - """Build lazy Prompt Control graphs that return conditioning batches.""" + """Build segment-local hook-aware Prompt Control conditioning batches.""" + + planning_service_class = PromptControlSegmentPlanningService + graph_adapter_class = PromptControlGraphAdapter def build( self, @@ -29,105 +33,63 @@ class PromptControlBatchGraphBuilder: negative_prompt: str, separator: str, ) -> Any: - """Return an io.NodeOutput with positive and negative batch links.""" + """Return positive and negative conditioning-batch graph links.""" - io, graph_utils, lazy_node = self._prompt_control_dependencies() - positive_chunks = split_prompt_batch(positive_prompt, separator) - negative_chunks = split_prompt_batch(negative_prompt, separator) - - expand: dict[str, dict[str, Any]] = {} - positive_output, positive_expand = self._encode_chunks( - chunks=positive_chunks, - clip=clip, - graph_utils=graph_utils, - lazy_node=lazy_node, + plan = self.planning_service_class().prepare( + positive_prompt=positive_prompt, + negative_prompt=negative_prompt, + separator=separator, ) - negative_output, negative_expand = self._encode_chunks( - chunks=negative_chunks, - clip=clip, - graph_utils=graph_utils, - lazy_node=lazy_node, - ) - expand.update(positive_expand) - expand.update(negative_expand) - return io.NodeOutput(positive_output, negative_output, expand=expand) - - def _encode_chunks( - self, - chunks: tuple[str, ...], - clip: Any, - graph_utils: Any, - lazy_node: Any, - ) -> tuple[list[Any], dict[str, dict[str, Any]]]: - """Encode chunks with Prompt Control and pack the resulting links.""" - + adapter = self.graph_adapter_class.load(PROMPT_CONTROL_MISSING_MESSAGE) expand: dict[str, dict[str, Any]] = {} - conditioning_outputs: list[Any] = [] - for chunk in chunks: - node_output = lazy_node.execute( + segment_clips = tuple( + adapter.clip_with_hooks( clip=clip, - text=chunk, - tags="", - start=0.0, - end=1.0, - num_steps=0, + lora_tags=hook.lora_tags, + expand=expand, + label=f"segment {index}", ) - node_expand = cast(dict[str, dict[str, Any]], node_output.expand or {}) - overlap = set(expand).intersection(node_expand) - if overlap: - overlapping_ids = ", ".join(sorted(overlap)) - raise ValueError( - "Prompt Control generated duplicate graph node ids: " - f"{overlapping_ids}." - ) - expand.update(node_expand) - conditioning_outputs.append(node_output.args[0]) - - pack_graph = graph_utils.GraphBuilder() - current = pack_graph.node( - "SimpleSyrup.ConditioningBatchStart", - conditioning=conditioning_outputs[0], + for index, hook in enumerate(plan.hooks) ) - for conditioning in conditioning_outputs[1:]: - current = pack_graph.node( - "SimpleSyrup.ConditioningBatchAppend", - batch=current.out(0), - conditioning=conditioning, + positive = self._encode_side( + plan.positive, + segment_clips=segment_clips, + adapter=adapter, + expand=expand, + label="positive", + ) + negative = self._encode_side( + plan.negative, + segment_clips=segment_clips, + adapter=adapter, + expand=expand, + label="negative", + ) + return adapter.io.NodeOutput(positive, negative, expand=expand) + + def _encode_side( + self, + side: PreparedPromptSide, + *, + segment_clips: tuple[Any, ...], + adapter: PromptControlGraphAdapter, + expand: dict[str, dict[str, Any]], + label: str, + ) -> Any: + """Encode and pack all existing chunks from one prompt side.""" + + outputs = [ + adapter.encode_segment( + clip=segment_clips[index], + text=chunk.text, + expand=expand, + label=f"{label} segment {index}", ) - pack_expand = cast(dict[str, dict[str, Any]], pack_graph.finalize()) - overlap = set(expand).intersection(pack_expand) - if overlap: - overlapping_ids = ", ".join(sorted(overlap)) - raise ValueError( - f"SimpleSyrup generated duplicate graph node ids: {overlapping_ids}." - ) - expand.update(pack_expand) - return current.out(0), expand - - def _prompt_control_dependencies(self) -> tuple[Any, Any, Any]: - """Import Prompt Control and Comfy v3 dependencies on demand.""" - - try: - io = import_module("comfy_api.latest.io") - except ModuleNotFoundError: - comfy_api = import_module("comfy_api.latest") - io = comfy_api.io - try: - graph_utils = import_module("comfy_execution.graph_utils") - lazy_nodes = self._import_prompt_control_lazy_nodes() - except ModuleNotFoundError as exc: - raise RuntimeError(PROMPT_CONTROL_MISSING_MESSAGE) from exc - return io, graph_utils, lazy_nodes.PCLazyTextEncodeAdvanced - - def _import_prompt_control_lazy_nodes(self) -> Any: - """Import Prompt Control lazy nodes from normal or sibling extension paths.""" - - try: - return import_module("prompt_control.nodes_lazy") - except ModuleNotFoundError: - availability = find_prompt_control_install() - if availability.root_path is not None: - root_path = str(availability.root_path) - if root_path not in sys.path: - sys.path.insert(0, root_path) - return import_module("prompt_control.nodes_lazy") + for index, chunk in enumerate(side.chunks) + ] + return adapter.pack_conditionings( + outputs, + expand=expand, + label=label, + always_batch=True, + ) diff --git a/simple_syrup/runtime/prompt_control_graph_adapter.py b/simple_syrup/runtime/prompt_control_graph_adapter.py new file mode 100644 index 0000000..fca4347 --- /dev/null +++ b/simple_syrup/runtime/prompt_control_graph_adapter.py @@ -0,0 +1,187 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Adapt Prompt Control and Comfy graph nodes for SEP prompt orchestration.""" + +from __future__ import annotations + +import sys +from importlib import import_module +from typing import Any, cast + +from .prompt_control_availability import find_prompt_control_install + + +class PromptControlGraphAdapter: + """Build hook-aware lazy graph fragments behind one runtime boundary.""" + + def __init__(self, io: Any, graph_utils: Any, lazy_nodes: Any) -> None: + """Store imported host graph APIs for one expansion request.""" + + self.io = io + self._graph_utils = graph_utils + self._lazy_nodes = lazy_nodes + + @classmethod + def load(cls, missing_message: str) -> PromptControlGraphAdapter: + """Load Comfy and Prompt Control graph APIs or raise an actionable error.""" + + try: + io = import_module("comfy_api.latest.io") + except ModuleNotFoundError: + io = import_module("comfy_api.latest").io + try: + graph_utils = import_module("comfy_execution.graph_utils") + lazy_nodes = cls._import_prompt_control_lazy_nodes() + except ModuleNotFoundError as exc: + raise RuntimeError(missing_message) from exc + return cls(io, graph_utils, lazy_nodes) + + def schedule_global_loras( + self, + *, + model: Any, + clip: Any, + positive_tags: str, + negative_tags: str, + expand: dict[str, dict[str, Any]], + ) -> tuple[Any, Any]: + """Preserve single-segment model/CLIP scheduling behavior.""" + + positive = self._lazy_nodes.PCLazyLoraLoaderAdvanced.execute( + model=model, + clip=clip, + text=positive_tags, + apply_hooks=True, + tags="", + start=0.0, + end=1.0, + num_steps=0, + ) + self.merge_expand(expand, positive.expand, "positive LoRA scheduling") + negative = self._lazy_nodes.PCLazyLoraLoaderAdvanced.execute( + model=positive.args[0], + clip=positive.args[1], + text=negative_tags, + apply_hooks=True, + tags="", + start=0.0, + end=1.0, + num_steps=0, + ) + self.merge_expand(expand, negative.expand, "negative LoRA scheduling") + return negative.args[0], negative.args[1] + + def encode_segment( + self, + *, + clip: Any, + text: str, + expand: dict[str, dict[str, Any]], + label: str, + ) -> Any: + """Encode one segment with a previously prepared CLIP link.""" + + output = self._lazy_nodes.PCLazyTextEncodeAdvanced.execute( + clip=clip, + text=text, + tags="", + start=0.0, + end=1.0, + num_steps=0, + ) + self.merge_expand(expand, output.expand, f"{label} text encoding") + return output.args[0] + + def clip_with_hooks( + self, + *, + clip: Any, + lora_tags: str, + expand: dict[str, dict[str, Any]], + label: str, + ) -> Any: + """Return a CLIP link sharing one segment's hooks across prompt sides.""" + + if not lora_tags: + return clip + graph = self._graph_utils.GraphBuilder() + hooks = graph.node("PCLoraHooksFromText", text=lora_tags) + hooked_clip = graph.node( + "SetClipHooks", + clip=clip, + hooks=hooks.out(0), + apply_to_conds=True, + schedule_clip=True, + ) + self.merge_expand( + expand, + cast(dict[str, dict[str, Any]], graph.finalize()), + f"{label} LoRA hooks", + ) + return hooked_clip.out(0) + + def pack_conditionings( + self, + conditionings: list[Any], + *, + expand: dict[str, dict[str, Any]], + label: str, + always_batch: bool, + ) -> Any: + """Return one conditioning or a SimpleSyrup conditioning batch link.""" + + if len(conditionings) == 1 and not always_batch: + return conditionings[0] + graph = self._graph_utils.GraphBuilder() + current = graph.node( + "SimpleSyrup.ConditioningBatchStart", + conditioning=conditionings[0], + ) + for conditioning in conditionings[1:]: + current = graph.node( + "SimpleSyrup.ConditioningBatchAppend", + batch=current.out(0), + conditioning=conditioning, + ) + self.merge_expand( + expand, + cast(dict[str, dict[str, Any]], graph.finalize()), + f"{label} batch packing", + ) + return current.out(0) + + def merge_expand( + self, + target: dict[str, dict[str, Any]], + source: object, + operation: str, + ) -> None: + """Merge a graph fragment while rejecting duplicate generated ids.""" + + if not source: + return + fragment = cast(dict[str, dict[str, Any]], source) + overlap = set(target).intersection(fragment) + if overlap: + overlapping_ids = ", ".join(sorted(overlap)) + raise ValueError( + "Prompt-Control graph expansion generated duplicate node ids " + f"during {operation}: {overlapping_ids}." + ) + target.update(fragment) + + @staticmethod + def _import_prompt_control_lazy_nodes() -> Any: + """Import Prompt Control from its normal or sibling extension path.""" + + try: + return import_module("prompt_control.nodes_lazy") + except ModuleNotFoundError: + availability = find_prompt_control_install() + if availability.root_path is not None: + root_path = str(availability.root_path) + if root_path not in sys.path: + sys.path.insert(0, root_path) + return import_module("prompt_control.nodes_lazy") diff --git a/simple_syrup/runtime/prompt_control_schedule_encode_graph.py b/simple_syrup/runtime/prompt_control_schedule_encode_graph.py index 4a34052..b2d7ced 100644 --- a/simple_syrup/runtime/prompt_control_schedule_encode_graph.py +++ b/simple_syrup/runtime/prompt_control_schedule_encode_graph.py @@ -2,20 +2,18 @@ # Copyright (C) 2026 Artificial Sweetener and contributors # SPDX-License-Identifier: AGPL-3.0-or-later -"""Runtime graph expansion for Prompt-Control scheduling and prompt encoding.""" +"""Orchestrate Prompt-Control model scheduling and SEP prompt encoding.""" from __future__ import annotations -import sys -from importlib import import_module -from typing import Any, cast +from typing import Any -from ..domain.prompt_control_prompt import ( - PreparedPromptSide, - apply_encode_style, - prepare_prompt_side, +from ..domain.prompt_control_prompt import PreparedPromptSide, apply_encode_style +from ..services.prompt_control_segment_planning_service import ( + PromptControlSegmentPlan, + PromptControlSegmentPlanningService, ) -from .prompt_control_availability import find_prompt_control_install +from .prompt_control_graph_adapter import PromptControlGraphAdapter PROMPT_CONTROL_MISSING_MESSAGE = ( "Schedule & Encode Prompts requires comfyui-prompt-control. " @@ -25,7 +23,10 @@ PROMPT_BATCH_SEPARATOR = "[SEP]" class PromptControlScheduleEncodeGraphBuilder: - """Build lazy Prompt-Control graphs for LoRA scheduling and prompt encoding.""" + """Build scheduled model and hook-aware conditioning graph outputs.""" + + planning_service_class = PromptControlSegmentPlanningService + graph_adapter_class = PromptControlGraphAdapter def build( self, @@ -35,161 +36,118 @@ class PromptControlScheduleEncodeGraphBuilder: negative_prompt: str, encode_style: str = "", ) -> Any: - """Return an io.NodeOutput for scheduled model and encoded prompts.""" - - io, graph_utils, lazy_nodes = self._prompt_control_dependencies() - positive_side = prepare_prompt_side(positive_prompt, PROMPT_BATCH_SEPARATOR) - negative_side = prepare_prompt_side(negative_prompt, PROMPT_BATCH_SEPARATOR) + """Return model plus single or SEP-batched conditioning outputs.""" + plan = self.planning_service_class().prepare( + positive_prompt=positive_prompt, + negative_prompt=negative_prompt, + separator=PROMPT_BATCH_SEPARATOR, + ) + adapter = self.graph_adapter_class.load(PROMPT_CONTROL_MISSING_MESSAGE) expand: dict[str, dict[str, Any]] = {} - positive_lora = lazy_nodes.PCLazyLoraLoaderAdvanced.execute( + scheduled_model, encoding_clip = self._sampling_inputs( model=model, clip=clip, - text=positive_side.lora_tags, - apply_hooks=True, - tags="", - start=0.0, - end=1.0, - num_steps=0, - ) - self._merge_expand(expand, positive_lora.expand, "positive LoRA scheduling") - - negative_lora = lazy_nodes.PCLazyLoraLoaderAdvanced.execute( - model=positive_lora.args[0], - clip=positive_lora.args[1], - text=negative_side.lora_tags, - apply_hooks=True, - tags="", - start=0.0, - end=1.0, - num_steps=0, - ) - self._merge_expand(expand, negative_lora.expand, "negative LoRA scheduling") - - scheduled_model = negative_lora.args[0] - scheduled_clip = negative_lora.args[1] - positive_conditioning = self._encode_side( - side=positive_side, - clip=scheduled_clip, - encode_style=encode_style, - graph_utils=graph_utils, - lazy_text_encoder=lazy_nodes.PCLazyTextEncodeAdvanced, + plan=plan, + adapter=adapter, expand=expand, - label="positive prompt encoding", ) - negative_conditioning = self._encode_side( - side=negative_side, - clip=scheduled_clip, - encode_style=encode_style, - graph_utils=graph_utils, - lazy_text_encoder=lazy_nodes.PCLazyTextEncodeAdvanced, + segment_clips = self._segment_clips( + plan=plan, + clip=encoding_clip, + adapter=adapter, expand=expand, - label="negative prompt encoding", ) - - return io.NodeOutput( + positive = self._encode_side( + plan.positive, + segment_clips=segment_clips, + encode_style=encode_style, + adapter=adapter, + expand=expand, + label="positive", + ) + negative = self._encode_side( + plan.negative, + segment_clips=segment_clips, + encode_style=encode_style, + adapter=adapter, + expand=expand, + label="negative", + ) + return adapter.io.NodeOutput( scheduled_model, - positive_conditioning, - negative_conditioning, + positive, + negative, expand=expand, ) + def _sampling_inputs( + self, + *, + model: Any, + clip: Any, + plan: PromptControlSegmentPlan, + adapter: PromptControlGraphAdapter, + expand: dict[str, dict[str, Any]], + ) -> tuple[Any, Any]: + """Keep single prompts global and batched prompts segment-local.""" + + if plan.is_batched: + return model, clip + return adapter.schedule_global_loras( + model=model, + clip=clip, + positive_tags=plan.positive.chunks[0].lora_tags, + negative_tags=plan.negative.chunks[0].lora_tags, + expand=expand, + ) + + def _segment_clips( + self, + *, + plan: PromptControlSegmentPlan, + clip: Any, + adapter: PromptControlGraphAdapter, + expand: dict[str, dict[str, Any]], + ) -> tuple[Any, ...]: + """Create one shared hooked CLIP link per batched segment index.""" + + if not plan.is_batched: + return (clip,) + return tuple( + adapter.clip_with_hooks( + clip=clip, + lora_tags=hook.lora_tags, + expand=expand, + label=f"segment {index}", + ) + for index, hook in enumerate(plan.hooks) + ) + def _encode_side( self, - *, side: PreparedPromptSide, - clip: Any, + *, + segment_clips: tuple[Any, ...], encode_style: str, - graph_utils: Any, - lazy_text_encoder: Any, + adapter: PromptControlGraphAdapter, expand: dict[str, dict[str, Any]], label: str, ) -> Any: - """Encode one prompt side and return conditioning or conditioning batch.""" + """Encode one side with the shared hook plan at each existing index.""" - conditioning_outputs: list[Any] = [] - for index, chunk in enumerate(side.chunks): - text = apply_encode_style(encode_style, chunk.text) - node_output = lazy_text_encoder.execute( - clip=clip, - text=text, - tags="", - start=0.0, - end=1.0, - num_steps=0, + outputs = [ + adapter.encode_segment( + clip=segment_clips[index], + text=apply_encode_style(encode_style, chunk.text), + expand=expand, + label=f"{label} segment {index}", ) - self._merge_expand( - expand, - node_output.expand, - f"{label} chunk {index}", - ) - conditioning_outputs.append(node_output.args[0]) - - if len(conditioning_outputs) == 1: - return conditioning_outputs[0] - - pack_graph = graph_utils.GraphBuilder() - current = pack_graph.node( - "SimpleSyrup.ConditioningBatchStart", - conditioning=conditioning_outputs[0], + for index, chunk in enumerate(side.chunks) + ] + return adapter.pack_conditionings( + outputs, + expand=expand, + label=label, + always_batch=False, ) - for conditioning in conditioning_outputs[1:]: - current = pack_graph.node( - "SimpleSyrup.ConditioningBatchAppend", - batch=current.out(0), - conditioning=conditioning, - ) - self._merge_expand( - expand, - cast(dict[str, dict[str, Any]], pack_graph.finalize()), - f"{label} batch packing", - ) - return current.out(0) - - def _prompt_control_dependencies(self) -> tuple[Any, Any, Any]: - """Import Prompt-Control and Comfy graph helpers on demand.""" - - try: - io = import_module("comfy_api.latest.io") - except ModuleNotFoundError: - comfy_api = import_module("comfy_api.latest") - io = comfy_api.io - try: - graph_utils = import_module("comfy_execution.graph_utils") - lazy_nodes = self._import_prompt_control_lazy_nodes() - except ModuleNotFoundError as exc: - raise RuntimeError(PROMPT_CONTROL_MISSING_MESSAGE) from exc - return io, graph_utils, lazy_nodes - - def _import_prompt_control_lazy_nodes(self) -> Any: - """Import Prompt-Control lazy nodes from installed or sibling paths.""" - - try: - return import_module("prompt_control.nodes_lazy") - except ModuleNotFoundError: - availability = find_prompt_control_install() - if availability.root_path is not None: - root_path = str(availability.root_path) - if root_path not in sys.path: - sys.path.insert(0, root_path) - return import_module("prompt_control.nodes_lazy") - - def _merge_expand( - self, - target: dict[str, dict[str, Any]], - source: object, - operation: str, - ) -> None: - """Merge a lazy expand graph and reject duplicate generated node ids.""" - - if not source: - return - expand = cast(dict[str, dict[str, Any]], source) - overlap = set(target).intersection(expand) - if overlap: - overlapping_ids = ", ".join(sorted(overlap)) - raise ValueError( - f"Prompt-Control graph expansion generated duplicate node ids " - f"during {operation}: {overlapping_ids}." - ) - target.update(expand) diff --git a/simple_syrup/runtime/tiled_sampling.py b/simple_syrup/runtime/tiled_sampling.py index e0f17f7..1c1503e 100644 --- a/simple_syrup/runtime/tiled_sampling.py +++ b/simple_syrup/runtime/tiled_sampling.py @@ -21,7 +21,7 @@ Latent: TypeAlias = dict[str, Any] ApplyModel: TypeAlias = Callable[..., torch.Tensor] ModelFunctionWrapper: TypeAlias = Callable[[ApplyModel, dict[str, Any]], torch.Tensor] -UNSUPPORTED_CONDITIONING_KEYS = frozenset({"area", "mask", "control", "gligen"}) +UNSUPPORTED_CONDITIONING_KEYS = frozenset({"area", "control", "gligen"}) def validate_sampling_controls( @@ -89,27 +89,49 @@ def reject_unsupported_conditioning( conditioning: object, *, sampler_label: str, + allow_full_context_masks: bool = False, ) -> None: - """Reject regional and external-control conditioning for basic tiled samplers.""" + """Reject conditioning that the selected tiled path cannot preserve.""" - if contains_unsupported_conditioning_key(conditioning): + if contains_unsupported_conditioning_key( + conditioning, + allow_full_context_masks=allow_full_context_masks, + ): raise ValueError( f"{sampler_label} does not support regional conditioning or " "ControlNet in the first implementation." ) -def contains_unsupported_conditioning_key(value: object) -> bool: - """Return whether a nested conditioning object contains unsupported keys.""" +def contains_unsupported_conditioning_key( + value: object, + *, + allow_full_context_masks: bool = False, +) -> bool: + """Return whether nested conditioning exceeds the tiled support policy.""" if isinstance(value, dict): if any(key in UNSUPPORTED_CONDITIONING_KEYS for key in value): return True + if "mask" in value and ( + not allow_full_context_masks or value.get("set_area_to_bounds") is not False + ): + return True return any( - contains_unsupported_conditioning_key(item) for item in value.values() + contains_unsupported_conditioning_key( + item, + allow_full_context_masks=allow_full_context_masks, + ) + for item in value.values() ) if isinstance(value, list | tuple): - return any(contains_unsupported_conditioning_key(item) for item in value) + return any( + contains_unsupported_conditioning_key( + item, + allow_full_context_masks=allow_full_context_masks, + ) + for item in value + ) return False diff --git a/simple_syrup/services/ksampler_sampling_service.py b/simple_syrup/services/ksampler_sampling_service.py new file mode 100644 index 0000000..f3fb921 --- /dev/null +++ b/simple_syrup/services/ksampler_sampling_service.py @@ -0,0 +1,186 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Application service for ordinary KSampler-style latent sampling.""" + +from __future__ import annotations + +from importlib import import_module +from typing import Any, TypeAlias + +import torch + +from ..domain.conditioning_batch import ConditioningBatch, select_conditioning +from ..runtime import sampling_samplers, sampling_schedulers +from ..shared.logging import get_logger + +Latent: TypeAlias = dict[str, Any] +LOGGER = get_logger(__name__) + + +class KSamplerSamplingService: + """Own reusable full-latent sampling while preserving batch semantics.""" + + def sample( + self, + *, + model: Any, + seed: int, + steps: int, + cfg: float, + sampler_name: str, + scheduler: str, + positive: Any, + negative: Any, + latent_image: Latent, + denoise: float, + ) -> Latent: + """Sample a latent with configured SimpleSyrup sampler extensions.""" + + sampler = sampling_samplers.resolve_sampler(sampler_name) + sigmas = sampling_schedulers.calculate_sigmas( + model=model, + scheduler_name=scheduler, + sampler_name=sampler_name, + steps=steps, + denoise=denoise, + ).to(model.load_device) + latent_samples = latent_image["samples"] + if not isinstance(latent_samples, torch.Tensor): + raise TypeError("KSampler latent samples must be a torch.Tensor.") + + comfy_sample = import_module("comfy.sample") + comfy_utils = import_module("comfy.utils") + latent_samples = comfy_sample.fix_empty_latent_channels( + model, + latent_samples, + latent_image.get("downscale_ratio_spacial", None), + ) + if not isinstance(latent_samples, torch.Tensor): + raise TypeError( + "KSampler normalized latent samples must be a torch.Tensor." + ) + noise = comfy_sample.prepare_noise( + latent_samples, + seed, + latent_image.get("batch_index"), + ) + noise_mask = latent_image.get("noise_mask") + callback = import_module("latent_preview").prepare_callback(model, steps) + disable_pbar = not comfy_utils.PROGRESS_BAR_ENABLED + if self._uses_conditioning_batch(positive, negative): + samples = self._sample_conditioning_batch( + comfy_sample=comfy_sample, + model=model, + noise=noise, + cfg=cfg, + sampler=sampler, + sigmas=sigmas, + positive=positive, + negative=negative, + latent_samples=latent_samples, + noise_mask=noise_mask, + callback=callback, + disable_pbar=disable_pbar, + seed=seed, + ) + else: + samples = comfy_sample.sample_custom( + model, + noise, + cfg, + sampler, + sigmas, + positive, + negative, + latent_samples, + noise_mask=noise_mask, + callback=callback, + disable_pbar=disable_pbar, + seed=seed, + ) + if not isinstance(samples, torch.Tensor): + raise TypeError("KSampler output samples must be a torch.Tensor.") + + output = latent_image.copy() + output.pop("downscale_ratio_spacial", None) + output["samples"] = samples + LOGGER.info( + "KSampler pass completed", + extra={ + "operation": "ksampler_sample", + "sampler": sampler_name, + "scheduler": scheduler, + "steps": steps, + "denoise": denoise, + "latent_batch_size": int(samples.shape[0]), + "latent_height": int(samples.shape[-2]), + "latent_width": int(samples.shape[-1]), + }, + ) + return output + + def _sample_conditioning_batch( + self, + *, + comfy_sample: Any, + model: Any, + noise: torch.Tensor, + cfg: float, + sampler: Any, + sigmas: torch.Tensor, + positive: Any, + negative: Any, + latent_samples: torch.Tensor, + noise_mask: Any, + callback: Any, + disable_pbar: bool, + seed: int, + ) -> torch.Tensor: + """Sample latent items with existing per-item batch selection.""" + + sampled: list[torch.Tensor] = [] + for index in range(int(latent_samples.shape[0])): + sampled.append( + comfy_sample.sample_custom( + model, + noise[index : index + 1], + cfg, + sampler, + sigmas, + select_conditioning(positive, index), + select_conditioning(negative, index), + latent_samples[index : index + 1], + noise_mask=self._slice_noise_mask( + noise_mask, + index, + latent_samples, + ), + callback=callback, + disable_pbar=disable_pbar, + seed=seed, + ) + ) + return torch.cat(sampled, dim=0) + + def _slice_noise_mask( + self, + noise_mask: Any, + index: int, + latent_samples: torch.Tensor, + ) -> Any: + """Return a per-item noise mask when its batch matches the latent.""" + + if isinstance(noise_mask, torch.Tensor) and noise_mask.shape[0] == int( + latent_samples.shape[0] + ): + return noise_mask[index : index + 1] + return noise_mask + + def _uses_conditioning_batch(self, positive: Any, negative: Any) -> bool: + """Return whether existing per-latent conditioning selection applies.""" + + return isinstance(positive, ConditioningBatch) or isinstance( + negative, ConditioningBatch + ) diff --git a/simple_syrup/services/load_mask_batch_service.py b/simple_syrup/services/load_mask_batch_service.py new file mode 100644 index 0000000..3fd4106 --- /dev/null +++ b/simple_syrup/services/load_mask_batch_service.py @@ -0,0 +1,97 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Application service for ordered authored-mask loading.""" + +from __future__ import annotations + +import hashlib +from collections.abc import Sequence +from typing import ClassVar + +import torch + +from ..runtime.mask_file_loader import MaskFileLoader +from ..shared.logging import get_logger + +LOGGER = get_logger(__name__) + + +class LoadMaskBatchService: + """Load ordered files into one dimensionally consistent MASK batch.""" + + loader_class: ClassVar[type[MaskFileLoader]] = MaskFileLoader + + def validate(self, files: Sequence[str], channel: str) -> None: + """Validate every selected file through the native Comfy mask loader.""" + + ordered_files = self._validate_files(files) + loader = self.loader_class() + for path in ordered_files: + loader.validate(path, channel) + + def load(self, files: Sequence[str], channel: str) -> torch.Tensor: + """Load one or many authored masks without changing their order.""" + + masks = self.load_each(files, channel) + expected_shape = tuple(masks[0].shape[1:]) + for index, mask in enumerate(masks[1:], start=1): + if tuple(mask.shape[1:]) != expected_shape: + raise ValueError( + "Load Mask Batch requires every mask to have identical " + f"dimensions; mask 0 is {expected_shape[0]}x{expected_shape[1]} " + f"but mask {index} is {mask.shape[1]}x{mask.shape[2]}." + ) + batch = torch.cat(masks, dim=0) + LOGGER.info( + "Authored mask batch loaded", + extra={ + "operation": "load_mask_batch", + "mask_count": int(batch.shape[0]), + "mask_height": int(batch.shape[1]), + "mask_width": int(batch.shape[2]), + "channel": channel, + }, + ) + return batch + + def load_each(self, files: Sequence[str], channel: str) -> tuple[torch.Tensor, ...]: + """Load ordered masks without requiring batch-compatible dimensions.""" + + ordered_files = self._validate_files(files) + loader = self.loader_class() + return tuple(loader.load(path, channel) for path in ordered_files) + + def fingerprint(self, files: Sequence[str], channel: str) -> str: + """Return an order-sensitive fingerprint for files and channel.""" + + ordered_files = self._validate_files(files) + loader = self.loader_class() + digest = hashlib.sha256() + digest.update(channel.encode("utf-8")) + for path in ordered_files: + encoded_path = path.encode("utf-8") + digest.update(len(encoded_path).to_bytes(8, "big")) + digest.update(encoded_path) + digest.update(loader.fingerprint(path).encode("ascii")) + return digest.hexdigest() + + def available_files(self) -> tuple[str, ...]: + """Return Comfy input images eligible for selection.""" + + return self.loader_class().available_files() + + def _validate_files(self, files: Sequence[str]) -> tuple[str, ...]: + """Return a non-empty ordered immutable file sequence.""" + + ordered: tuple[str, ...] + if isinstance(files, str): + ordered = (files,) + else: + ordered = tuple(files) + if not ordered: + raise ValueError("Load Mask Batch requires at least one mask file.") + if any(not isinstance(path, str) or not path for path in ordered): + raise TypeError("Load Mask Batch mask files must be non-empty strings.") + return ordered diff --git a/simple_syrup/services/prompt_control_segment_planning_service.py b/simple_syrup/services/prompt_control_segment_planning_service.py new file mode 100644 index 0000000..ae325c3 --- /dev/null +++ b/simple_syrup/services/prompt_control_segment_planning_service.py @@ -0,0 +1,73 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Plan aligned Prompt-Control SEP segments without runtime dependencies.""" + +from __future__ import annotations + +from dataclasses import dataclass + +from ..domain.prompt_control_prompt import ( + PreparedPromptSide, + PromptSegmentHookPlan, + prepare_prompt_side, +) + + +@dataclass(frozen=True) +class PromptControlSegmentPlan: + """Store two prompt sides and the LoRA hooks shared at each position.""" + + positive: PreparedPromptSide + negative: PreparedPromptSide + hooks: tuple[PromptSegmentHookPlan, ...] + + @property + def is_batched(self) -> bool: + """Return whether either prompt side contains multiple SEP segments.""" + + return len(self.positive.chunks) > 1 or len(self.negative.chunks) > 1 + + +class PromptControlSegmentPlanningService: + """Prepare prompt sides and align their segment-local LoRA schedules.""" + + def prepare( + self, + *, + positive_prompt: str, + negative_prompt: str, + separator: str, + ) -> PromptControlSegmentPlan: + """Return a deterministic plan without padding either prompt side.""" + + positive = prepare_prompt_side(positive_prompt, separator) + negative = prepare_prompt_side(negative_prompt, separator) + hook_count = max(len(positive.chunks), len(negative.chunks)) + hooks = tuple( + PromptSegmentHookPlan( + lora_tags=self._combined_lora_tags(positive, negative, index) + ) + for index in range(hook_count) + ) + return PromptControlSegmentPlan( + positive=positive, + negative=negative, + hooks=hooks, + ) + + def _combined_lora_tags( + self, + positive: PreparedPromptSide, + negative: PreparedPromptSide, + index: int, + ) -> str: + """Join existing positive then negative tags for one segment index.""" + + tags: list[str] = [] + if index < len(positive.chunks) and positive.chunks[index].lora_tags: + tags.append(positive.chunks[index].lora_tags) + if index < len(negative.chunks) and negative.chunks[index].lora_tags: + tags.append(negative.chunks[index].lora_tags) + return "\n".join(tags) diff --git a/simple_syrup/services/regional_conditioning_service.py b/simple_syrup/services/regional_conditioning_service.py new file mode 100644 index 0000000..6bdf33c --- /dev/null +++ b/simple_syrup/services/regional_conditioning_service.py @@ -0,0 +1,163 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Application service for assembling full-context regional conditioning.""" + +from __future__ import annotations + +from typing import Any, TypeAlias + +import torch + +from ..domain.conditioning_batch import ConditioningBatch +from ..domain.regional_prompting import ( + build_regional_conditioning_plan, + validate_regional_prompt_weight, +) +from ..masking.regional_prompt_masks import ( + complementary_global_prompt_mask, + prepare_regional_mask_batch, + regional_mask, +) +from ..shared.logging import get_logger + +Conditioning: TypeAlias = list[list[Any]] +LOGGER = get_logger(__name__) + + +class RegionalConditioningService: + """Combine global and ordered regional conditioning without sampling.""" + + def assemble( + self, + *, + positive: object, + negative: object, + masks: object, + regional_prompt_weight: float, + region_mask_feather: int, + ) -> tuple[Conditioning, Conditioning]: + """Return standard Comfy positive and negative masked conditioning.""" + + validate_regional_prompt_weight(regional_prompt_weight) + mask_batch = prepare_regional_mask_batch(masks, region_mask_feather) + assembled_positive = self._assemble_input( + positive, + mask_batch, + regional_prompt_weight=regional_prompt_weight, + input_name="positive", + ) + assembled_negative = self._assemble_input( + negative, + mask_batch, + regional_prompt_weight=regional_prompt_weight, + input_name="negative", + ) + LOGGER.info( + "Regional conditioning assembled", + extra={ + "operation": "assemble_regional_conditioning", + "region_count": int(mask_batch.shape[0]), + "positive_entry_count": len(assembled_positive), + "negative_entry_count": len(assembled_negative), + "regional_prompt_weight": regional_prompt_weight, + "region_mask_feather": region_mask_feather, + }, + ) + return assembled_positive, assembled_negative + + def _assemble_input( + self, + value: object, + mask_batch: torch.Tensor, + *, + regional_prompt_weight: float, + input_name: str, + ) -> Conditioning: + """Assemble one global-first conditioning input.""" + + if isinstance(value, ConditioningBatch): + entries = value.entries + else: + entries = (value,) + plan = build_regional_conditioning_plan( + region_count=int(mask_batch.shape[0]), + conditioning_count=len(entries), + input_name=input_name, + ) + global_conditioning = self._validate_conditioning( + entries[0], + input_name=f"{input_name} global", + ) + if not plan.pairs or regional_prompt_weight == 0.0: + return self._copy_conditioning(global_conditioning) + + global_mask = complementary_global_prompt_mask( + mask_batch, + tuple(pair.mask_index for pair in plan.pairs), + regional_prompt_weight, + ) + assembled = self._with_mask( + global_conditioning, + global_mask, + mask_strength=1.0, + ) + for pair in plan.pairs: + conditioning = self._validate_conditioning( + entries[pair.conditioning_index], + input_name=(f"{input_name} regional entry {pair.conditioning_index}"), + ) + mask = regional_mask(mask_batch, pair.mask_index) + assembled.extend( + self._with_mask( + conditioning, + mask, + mask_strength=regional_prompt_weight, + ) + ) + return assembled + + def _validate_conditioning( + self, + value: object, + *, + input_name: str, + ) -> Conditioning: + """Return a structurally valid standard Comfy conditioning list.""" + + if not isinstance(value, list): + raise TypeError(f"{input_name} must be a standard CONDITIONING value.") + for index, item in enumerate(value): + if ( + not isinstance(item, list | tuple) + or len(item) != 2 + or not isinstance(item[1], dict) + ): + raise ValueError( + f"{input_name} item {index} must contain a tensor and metadata." + ) + return value + + def _copy_conditioning(self, conditioning: Conditioning) -> Conditioning: + """Copy conditioning containers and metadata without cloning tensors.""" + + return [[item[0], dict(item[1])] for item in conditioning] + + def _with_mask( + self, + conditioning: Conditioning, + mask: torch.Tensor, + *, + mask_strength: float, + ) -> Conditioning: + """Copy conditioning entries with one full-context regional mask.""" + + masked: Conditioning = [] + for item in conditioning: + metadata = dict(item[1]) + metadata["mask"] = mask + metadata["mask_strength"] = mask_strength + metadata["set_area_to_bounds"] = False + masked.append([item[0], metadata]) + return masked diff --git a/simple_syrup/services/tiled_diffusion_sampling_service.py b/simple_syrup/services/tiled_diffusion_sampling_service.py index 33d9738..fc3b177 100644 --- a/simple_syrup/services/tiled_diffusion_sampling_service.py +++ b/simple_syrup/services/tiled_diffusion_sampling_service.py @@ -41,6 +41,7 @@ class TiledDiffusionSamplingService: latent_tile_batch_size: int, preview_context: DetailPreviewContext | None = None, differential_diffusion: bool = False, + allow_full_context_masks: bool = False, ) -> Latent: """Sample a latent with the selected tiled diffusion method.""" @@ -64,6 +65,7 @@ class TiledDiffusionSamplingService: latent_tile_batch_size=latent_tile_batch_size, preview_context=preview_context, differential_diffusion=differential_diffusion, + allow_full_context_masks=allow_full_context_masks, ) if diffusion_mode == "multidiffusion": return multidiffusion_sampling.sample_multidiffusion( @@ -83,6 +85,7 @@ class TiledDiffusionSamplingService: latent_tile_batch_size=latent_tile_batch_size, preview_context=preview_context, differential_diffusion=differential_diffusion, + allow_full_context_masks=allow_full_context_masks, ) return mixture_of_diffusers_sampling.sample_mixture_of_diffusers( model=model, @@ -101,6 +104,7 @@ class TiledDiffusionSamplingService: latent_tile_batch_size=latent_tile_batch_size, preview_context=preview_context, differential_diffusion=differential_diffusion, + allow_full_context_masks=allow_full_context_masks, ) def _sample_conditioning_batch( @@ -123,6 +127,7 @@ class TiledDiffusionSamplingService: latent_tile_batch_size: int, preview_context: DetailPreviewContext | None, differential_diffusion: bool, + allow_full_context_masks: bool, ) -> Latent: """Sample latent batch items one at a time with selected conditioning.""" @@ -151,6 +156,7 @@ class TiledDiffusionSamplingService: latent_tile_batch_size=latent_tile_batch_size, preview_context=preview_context, differential_diffusion=differential_diffusion, + allow_full_context_masks=allow_full_context_masks, ) output_samples = output["samples"] if not isinstance(output_samples, torch.Tensor): diff --git a/tests/test_compose_regional_conditioning_v3_node.py b/tests/test_compose_regional_conditioning_v3_node.py new file mode 100644 index 0000000..3665b46 --- /dev/null +++ b/tests/test_compose_regional_conditioning_v3_node.py @@ -0,0 +1,79 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the standard regional-conditioning composer node.""" + +from __future__ import annotations + +from typing import ClassVar + +import torch + +from simple_syrup.nodes_v3.compose_regional_conditioning import ( + ComposeRegionalConditioningV3, +) + + +class FakeConditioningService: + """Record one composition request and return recognizable outputs.""" + + calls: ClassVar[list[dict[str, object]]] = [] + + def assemble(self, **kwargs: object) -> tuple[object, object]: + """Record arguments and return standard-conditioning stand-ins.""" + + type(self).calls.append(kwargs) + return "positive-conditioning", "negative-conditioning" + + +def test_composer_schema_outputs_native_conditioning() -> None: + """The composer exposes the same regional inputs and standard outputs.""" + + schema = ComposeRegionalConditioningV3.define_schema() + + assert schema.node_id == "SimpleSyrup.ComposeRegionalConditioning" + assert schema.display_name == "Compose Regional Conditioning" + assert schema.category == "SimpleSyrup/Conditioning" + assert [input_item.id for input_item in schema.inputs] == [ + "positive", + "negative", + "region_masks", + "regional_prompt_weight", + "region_mask_feather", + ] + assert schema.inputs[3].default == 0.5 + assert [output.io_type for output in schema.outputs] == [ + "CONDITIONING", + "CONDITIONING", + ] + + +def test_composer_delegates_to_authoritative_regional_service() -> None: + """The node adds no policy beyond its application service.""" + + original = ComposeRegionalConditioningV3.conditioning_service_class + ComposeRegionalConditioningV3.conditioning_service_class = FakeConditioningService # type: ignore[assignment] + FakeConditioningService.calls = [] + masks = torch.ones((2, 4, 4)) + try: + output = ComposeRegionalConditioningV3.execute( + positive="positive-batch", + negative="negative-batch", + region_masks=masks, + regional_prompt_weight=0.75, + region_mask_feather=3, + ) + finally: + ComposeRegionalConditioningV3.conditioning_service_class = original + + assert output == ("positive-conditioning", "negative-conditioning") + assert FakeConditioningService.calls == [ + { + "positive": "positive-batch", + "negative": "negative-batch", + "masks": masks, + "regional_prompt_weight": 0.75, + "region_mask_feather": 3, + } + ] diff --git a/tests/test_load_mask_batch_service.py b/tests/test_load_mask_batch_service.py new file mode 100644 index 0000000..bafe8e4 --- /dev/null +++ b/tests/test_load_mask_batch_service.py @@ -0,0 +1,148 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for ordered authored-mask loading orchestration.""" + +from __future__ import annotations + +from collections.abc import Sequence +from typing import ClassVar + +import pytest +import torch + +from simple_syrup.runtime.mask_file_loader import MaskFileLoader +from simple_syrup.services.load_mask_batch_service import LoadMaskBatchService + + +class FakeMaskFileLoader(MaskFileLoader): + """Provide deterministic masks and fingerprints for service tests.""" + + masks: ClassVar[dict[str, torch.Tensor]] = {} + loaded: ClassVar[list[tuple[str, str]]] = [] + validated: ClassVar[list[tuple[str, str]]] = [] + + def available_files(self) -> tuple[str, ...]: + """Return deterministic selectable paths.""" + + return ("first.png", "second.png") + + def load(self, path: str, channel: str) -> torch.Tensor: + """Record and return the configured singleton mask.""" + + self.loaded.append((path, channel)) + return self.masks[path] + + def validate(self, path: str, channel: str) -> None: + """Record validation through the runtime boundary.""" + + self.validated.append((path, channel)) + + def fingerprint(self, path: str) -> str: + """Return a recognizable per-path fingerprint.""" + + return f"digest:{path}" + + +class FakeLoadMaskBatchService(LoadMaskBatchService): + """Use the fake runtime adapter in application-service tests.""" + + loader_class = FakeMaskFileLoader + + +def _service_with_masks(masks: dict[str, torch.Tensor]) -> FakeLoadMaskBatchService: + """Return a service configured with fake ordered masks.""" + + FakeMaskFileLoader.masks = masks + FakeMaskFileLoader.loaded = [] + FakeMaskFileLoader.validated = [] + return FakeLoadMaskBatchService() + + +def test_loads_one_or_many_files_in_supplied_order() -> None: + """The file sequence directly defines MASK batch order.""" + + service = _service_with_masks( + { + "b.png": torch.full((1, 2, 3), 0.8), + "a.png": torch.full((1, 2, 3), 0.2), + } + ) + + one = service.load(["a.png"], "red") + many = service.load(["b.png", "a.png"], "blue") + + assert one.shape == (1, 2, 3) + assert torch.all(one == 0.2) + assert torch.all(many[0] == 0.8) + assert torch.all(many[1] == 0.2) + assert FakeMaskFileLoader.loaded == [ + ("a.png", "red"), + ("b.png", "blue"), + ("a.png", "blue"), + ] + + +def test_rejects_empty_selection_and_dimension_mismatch() -> None: + """A batch always contains compatible one-file regions.""" + + service = _service_with_masks( + { + "small.png": torch.zeros((1, 2, 2)), + "wide.png": torch.zeros((1, 2, 3)), + } + ) + + with pytest.raises(ValueError, match="at least one mask file"): + service.load([], "alpha") + with pytest.raises(ValueError, match="identical dimensions"): + service.load(["small.png", "wide.png"], "alpha") + + +def test_fingerprint_is_sensitive_to_order_channel_and_contents() -> None: + """Caching distinguishes every workflow-relevant loader input.""" + + service = _service_with_masks({}) + + first = service.fingerprint(["a.png", "b.png"], "red") + reordered = service.fingerprint(["b.png", "a.png"], "red") + other_channel = service.fingerprint(["a.png", "b.png"], "blue") + + assert first != reordered + assert first != other_channel + + +def test_available_files_delegates_to_runtime_adapter() -> None: + """Schema choices come from the Comfy filesystem boundary.""" + + assert FakeLoadMaskBatchService().available_files() == ( + "first.png", + "second.png", + ) + + +def test_validation_preserves_selection_order() -> None: + """Every native widget value is validated in its authored order.""" + + service = _service_with_masks({}) + + service.validate(["second.png", "first.png"], "green") + + assert FakeMaskFileLoader.validated == [ + ("second.png", "green"), + ("first.png", "green"), + ] + + +@pytest.mark.parametrize("files", ["single.png", ("single.png",)]) +def test_string_and_sequence_inputs_both_mean_one_mask( + files: str | Sequence[str], +) -> None: + """API callers may serialize one selected path as a scalar or list.""" + + service = _service_with_masks({"single.png": torch.ones((1, 2, 2))}) + + result = service.load(files, "alpha") + + assert result.shape == (1, 2, 2) diff --git a/tests/test_load_mask_batch_v3_node.py b/tests/test_load_mask_batch_v3_node.py new file mode 100644 index 0000000..7014307 --- /dev/null +++ b/tests/test_load_mask_batch_v3_node.py @@ -0,0 +1,110 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the Load Mask Batch Comfy v3 node.""" + +from __future__ import annotations + +from typing import Any, ClassVar + +import torch + +from simple_syrup.nodes_v3 import load_mask_batch as node_module +from simple_syrup.nodes_v3.load_mask_batch import LoadMaskBatchV3 + + +class FakeService: + """Provide deterministic loader behavior at the node boundary.""" + + loaded: ClassVar[tuple[list[str], str] | None] = None + + def available_files(self) -> tuple[str, ...]: + """Return node combo choices.""" + + return ("b.png", "a.png") + + def validate(self, files: str | list[str], channel: str) -> None: + """Accept deterministic widget values.""" + + del files, channel + + def load(self, files: list[str], channel: str) -> torch.Tensor: + """Record input order and return a two-mask batch.""" + + type(self).loaded = (files, channel) + return torch.stack([torch.zeros((2, 2)), torch.ones((2, 2))]) + + def fingerprint(self, files: list[str], channel: str) -> str: + """Return an order-visible node fingerprint.""" + + return f"{channel}:{'|'.join(files)}" + + +def test_schema_exposes_native_ordered_mask_multiselect() -> None: + """The node uses one native multi-value file list and disk uploader.""" + + original = LoadMaskBatchV3.service_class + LoadMaskBatchV3.service_class = FakeService # type: ignore[assignment] + try: + schema = LoadMaskBatchV3.define_schema() + finally: + LoadMaskBatchV3.service_class = original + + inputs = {value.id: value for value in schema.inputs} + image = inputs["image"].as_dict() + assert schema.node_id == "SimpleSyrup.LoadMaskBatch" + assert schema.display_name == "Load Mask Batch" + assert schema.has_intermediate_output is True + assert image["image_upload"] is True + assert image["allow_batch"] is True + assert image["image_folder"] == "input" + assert image["multiselect"] is True + assert image["multi_select"] == { + "placeholder": "Select one or more masks", + "chip": True, + } + assert image["default"] == [] + assert image["options"] == ["b.png", "a.png"] + assert [output.io_type for output in schema.outputs] == ["MASK"] + + +def test_execute_preserves_order_and_returns_native_preview(monkeypatch: Any) -> None: + """Execution returns one MASK batch and previews that exact tensor.""" + + previews: list[tuple[torch.Tensor, object]] = [] + + def fake_preview(mask: torch.Tensor, cls: object) -> str: + """Record the native preview payload.""" + + previews.append((mask, cls)) + return "preview" + + original = LoadMaskBatchV3.service_class + LoadMaskBatchV3.service_class = FakeService # type: ignore[assignment] + monkeypatch.setattr(node_module._comfy_ui, "PreviewMask", fake_preview) + try: + output = LoadMaskBatchV3.execute(["b.png", "a.png"], "green") + fingerprint = LoadMaskBatchV3.fingerprint_inputs(["b.png", "a.png"], "green") + finally: + LoadMaskBatchV3.service_class = original + + assert FakeService.loaded == (["b.png", "a.png"], "green") + assert output.result is not None + mask_batch = output.result[0] + assert mask_batch.shape == (2, 2, 2) + assert previews == [(mask_batch, LoadMaskBatchV3)] + assert output.ui == "preview" + assert fingerprint == "green:b.png|a.png" + + +def test_validate_inputs_accepts_native_scalar_or_batch_values() -> None: + """The regular native combo may serialize one path or an ordered path list.""" + + original = LoadMaskBatchV3.service_class + LoadMaskBatchV3.service_class = FakeService # type: ignore[assignment] + try: + assert LoadMaskBatchV3.validate_inputs("a.png", "alpha") is True + assert LoadMaskBatchV3.validate_inputs(["b.png", "a.png"], "red") is True + finally: + LoadMaskBatchV3.service_class = original diff --git a/tests/test_mask_batch_preview_routes.py b/tests/test_mask_batch_preview_routes.py new file mode 100644 index 0000000..f1db509 --- /dev/null +++ b/tests/test_mask_batch_preview_routes.py @@ -0,0 +1,353 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for channel-accurate authored-mask preview routes.""" + +from __future__ import annotations + +import asyncio +import json +from collections.abc import Callable, Sequence +from importlib import import_module +from pathlib import Path +from typing import cast + +import pytest +from aiohttp import web +from PIL import Image + +import simple_syrup.runtime.mask_batch_preview_routes as preview_routes +from simple_syrup.runtime.mask_batch_preview_routes import ( + MASK_BATCH_PREVIEW_ROUTE, + Handler, + MaskBatchLoaderProtocol, + MaskBatchPreviewHandlers, + MaskBatchPreviewRendererProtocol, + NativeMaskBatchPreviewRenderer, + PromptServerProtocol, + register_mask_batch_preview_routes, +) +from simple_syrup.services.load_mask_batch_service import LoadMaskBatchService + + +class FakeRoutes: + """Route table double that records decorated POST handlers.""" + + def __init__(self) -> None: + """Create an empty fake route table.""" + + self.post_handlers: dict[str, Handler] = {} + + def post(self, path: str) -> Callable[[Handler], Handler]: + """Record a POST route handler.""" + + def decorator(handler: Handler) -> Handler: + self.post_handlers[path] = handler + return handler + + return decorator + + +class FakePromptServer: + """PromptServer double exposing a route table.""" + + def __init__(self) -> None: + """Create a fake PromptServer.""" + + self.routes = FakeRoutes() + + +class FakeRequest: + """Request double with injectable JSON body behavior.""" + + def __init__(self, payload: object | BaseException) -> None: + """Create a request that returns or raises from `json()`.""" + + self._payload = payload + + async def json(self) -> object: + """Return the configured JSON payload.""" + + if isinstance(self._payload, BaseException): + raise self._payload + return self._payload + + +class FakeMaskBatchLoader: + """Record preview loads and return recognizable mask objects.""" + + def __init__(self, failure: ValueError | None = None) -> None: + """Create a loader with an optional validation failure.""" + + self.failure = failure + self.loaded: list[tuple[tuple[str, ...], str]] = [] + self.masks = (object(), object()) + + def load_each(self, files: Sequence[str], channel: str) -> Sequence[object]: + """Record the ordered request and return configured masks.""" + + if self.failure is not None: + raise self.failure + self.loaded.append((tuple(files), channel)) + return self.masks + + +class FakeMaskBatchPreviewRenderer: + """Return a recognizable native-preview response.""" + + def __init__(self) -> None: + """Create a renderer that records the loaded batch.""" + + self.rendered: list[tuple[object, ...]] = [] + self.payload: dict[str, object] = { + "images": [ + { + "filename": "ComfyUI_temp_mask.png", + "subfolder": "", + "type": "temp", + } + ], + "animated": [False], + } + + def render(self, masks: Sequence[object]) -> dict[str, object]: + """Record and return the native preview payload.""" + + self.rendered.append(tuple(masks)) + return self.payload + + +def test_preview_route_loads_selected_channel_and_returns_native_output() -> None: + """The preview uses the same ordered files and channel as node execution.""" + + prompt_server = FakePromptServer() + loader = FakeMaskBatchLoader() + renderer = FakeMaskBatchPreviewRenderer() + register_fake_routes(prompt_server, loader, renderer) + + response = asyncio.run( + prompt_server.routes.post_handlers[MASK_BATCH_PREVIEW_ROUTE]( + FakeRequest( + { + "files": ["characters/left.png", "characters/right.png"], + "channel": "green", + } + ) + ) + ) + + assert response.status == 200 + assert json.loads(response_text(response)) == renderer.payload + assert loader.loaded == [(("characters/left.png", "characters/right.png"), "green")] + assert renderer.rendered == [loader.masks] + + +def test_preview_route_uses_real_native_loader_and_renderer_on_synthetic_masks( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Synthetic masks preview through the same native path as node execution.""" + + input_directory = tmp_path / "input" + preview_directory = tmp_path / "preview" + input_directory.mkdir() + preview_directory.mkdir() + paths = { + "left.png": input_directory / "left.png", + "right.png": input_directory / "right.png", + } + Image.new("RGBA", (8, 6), color=(255, 255, 255, 0)).save(paths["left.png"]) + Image.new("RGBA", (8, 6), color=(0, 0, 0, 255)).save(paths["right.png"]) + + folder_paths = import_module("folder_paths") + monkeypatch.setattr( + folder_paths, + "exists_annotated_filepath", + lambda value: value in paths, + ) + monkeypatch.setattr( + folder_paths, + "get_annotated_filepath", + lambda value: str(paths[value]), + ) + monkeypatch.setattr( + folder_paths, + "get_temp_directory", + lambda: str(preview_directory), + ) + + handler = MaskBatchPreviewHandlers( + LoadMaskBatchService(), + NativeMaskBatchPreviewRenderer(), + ) + response = asyncio.run( + handler.post_preview( + FakeRequest({"files": ["left.png", "right.png"], "channel": "alpha"}) + ) + ) + + payload = json.loads(response_text(response)) + assert response.status == 200 + assert len(payload["images"]) == 2 + assert payload["animated"] == [False] + for image in payload["images"]: + assert image["type"] == "temp" + preview_path = preview_directory / image["subfolder"] / image["filename"] + with Image.open(preview_path) as preview: + assert preview.size == (8, 6) + + +def test_preview_route_renders_mixed_dimensions_independently( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Preview every selected mask even when execution cannot form one batch.""" + + input_directory = tmp_path / "input" + preview_directory = tmp_path / "preview" + input_directory.mkdir() + preview_directory.mkdir() + paths = { + "small.png": input_directory / "small.png", + "large.png": input_directory / "large.png", + } + Image.new("RGBA", (8, 6), color=(255, 255, 255, 0)).save(paths["small.png"]) + Image.new("RGBA", (12, 10), color=(0, 0, 0, 255)).save(paths["large.png"]) + + folder_paths = import_module("folder_paths") + monkeypatch.setattr( + folder_paths, + "exists_annotated_filepath", + lambda value: value in paths, + ) + monkeypatch.setattr( + folder_paths, + "get_annotated_filepath", + lambda value: str(paths[value]), + ) + monkeypatch.setattr( + folder_paths, + "get_temp_directory", + lambda: str(preview_directory), + ) + + service = LoadMaskBatchService() + handler = MaskBatchPreviewHandlers( + service, + NativeMaskBatchPreviewRenderer(), + ) + response = asyncio.run( + handler.post_preview( + FakeRequest({"files": ["small.png", "large.png"], "channel": "red"}) + ) + ) + + payload = json.loads(response_text(response)) + assert response.status == 200 + assert len(payload["images"]) == 2 + preview_sizes = [] + for image in payload["images"]: + preview_path = preview_directory / image["subfolder"] / image["filename"] + with Image.open(preview_path) as preview: + preview_sizes.append(preview.size) + assert preview_sizes == [(8, 6), (12, 10)] + with pytest.raises(ValueError, match="identical dimensions"): + service.load(["small.png", "large.png"], "red") + + +@pytest.mark.parametrize( + "payload, message", + [ + ({"channel": "alpha"}, "files"), + ({"files": [], "channel": "alpha"}, "at least one"), + ({"files": ["mask.png"], "channel": 1}, "channel"), + ], +) +def test_preview_route_rejects_malformed_payloads( + payload: object, + message: str, +) -> None: + """Malformed preview requests fail before filesystem loading.""" + + prompt_server = FakePromptServer() + loader = FakeMaskBatchLoader() + register_fake_routes(prompt_server, loader, FakeMaskBatchPreviewRenderer()) + + response = asyncio.run( + prompt_server.routes.post_handlers[MASK_BATCH_PREVIEW_ROUTE]( + FakeRequest(payload) + ) + ) + + assert response.status == 400 + assert message in response_text(response) + assert loader.loaded == [] + + +def test_preview_route_surfaces_loader_validation_errors() -> None: + """Native path, channel, and shape errors remain actionable to the client.""" + + prompt_server = FakePromptServer() + loader = FakeMaskBatchLoader(ValueError("mask dimensions do not match")) + register_fake_routes(prompt_server, loader, FakeMaskBatchPreviewRenderer()) + + response = asyncio.run( + prompt_server.routes.post_handlers[MASK_BATCH_PREVIEW_ROUTE]( + FakeRequest({"files": ["mask.png"], "channel": "alpha"}) + ) + ) + + assert response.status == 400 + assert json.loads(response_text(response)) == { + "error": "mask dimensions do not match" + } + + +def test_preview_route_rejects_invalid_json() -> None: + """Bodies that cannot be decoded never reach the mask loader.""" + + prompt_server = FakePromptServer() + loader = FakeMaskBatchLoader() + register_fake_routes(prompt_server, loader, FakeMaskBatchPreviewRenderer()) + + response = asyncio.run( + prompt_server.routes.post_handlers[MASK_BATCH_PREVIEW_ROUTE]( + FakeRequest(ValueError("bad json")) + ) + ) + + assert response.status == 400 + assert "valid JSON" in response_text(response) + assert loader.loaded == [] + + +def test_route_registration_is_import_safe_without_prompt_server( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Missing Comfy PromptServer does not import or initialize mask services.""" + + monkeypatch.setattr(preview_routes, "_prompt_server_instance", lambda: None) + + assert register_mask_batch_preview_routes(prompt_server=None) is False + + +def register_fake_routes( + prompt_server: FakePromptServer, + loader: FakeMaskBatchLoader, + renderer: FakeMaskBatchPreviewRenderer, +) -> bool: + """Register preview routes against structurally typed test doubles.""" + + return register_mask_batch_preview_routes( + loader=cast(MaskBatchLoaderProtocol, loader), + renderer=cast(MaskBatchPreviewRendererProtocol, renderer), + prompt_server=cast(PromptServerProtocol, prompt_server), + ) + + +def response_text(response: web.Response) -> str: + """Return response text after asserting aiohttp populated it.""" + + assert response.text is not None + return response.text diff --git a/tests/test_mask_file_loader.py b/tests/test_mask_file_loader.py new file mode 100644 index 0000000..256e418 --- /dev/null +++ b/tests/test_mask_file_loader.py @@ -0,0 +1,127 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for the Comfy authored-mask file adapter.""" + +from __future__ import annotations + +from importlib import import_module +from pathlib import Path + +import pytest +import torch +from PIL import Image + +from simple_syrup.runtime.mask_file_loader import MaskFileLoader + + +def _patch_path( + monkeypatch: pytest.MonkeyPatch, + path: Path, + *, + exists: bool = True, +) -> None: + """Route Comfy annotated-path helpers to one temporary file.""" + + folder_paths = import_module("folder_paths") + monkeypatch.setattr(import_module("nodes"), "folder_paths", folder_paths) + + monkeypatch.setattr(folder_paths, "exists_annotated_filepath", lambda value: exists) + monkeypatch.setattr(folder_paths, "get_annotated_filepath", lambda value: str(path)) + + +def test_load_uses_native_mask_loader_and_validated_path( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Static files are decoded through Comfy's LoadImageMask behavior.""" + + path = tmp_path / "mask.png" + Image.new("RGB", (3, 2), color=(255, 0, 0)).save(path) + _patch_path(monkeypatch, path) + calls: list[tuple[str, str]] = [] + + class FakeNativeLoader: + """Return one recognizable native mask.""" + + @classmethod + def VALIDATE_INPUTS(cls, value: str) -> bool: + """Accept the annotated path through the native contract.""" + + del cls, value + return True + + def load_image_mask(self, value: str, channel: str) -> tuple[torch.Tensor]: + """Record native loader arguments.""" + + calls.append((value, channel)) + return (torch.full((1, 2, 3), 0.75),) + + nodes = import_module("nodes") + + monkeypatch.setattr(nodes, "LoadImageMask", FakeNativeLoader) + + result = MaskFileLoader().load("mask.png", "red") + + assert calls == [("mask.png", "red")] + assert torch.all(result == 0.75) + + +def test_available_files_uses_native_mask_loader_choices( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The node's file combo is sourced from native LoadImageMask inputs.""" + + class FakeNativeLoader: + """Expose recognizable native widget choices.""" + + @classmethod + def INPUT_TYPES(cls) -> dict[str, dict[str, tuple[list[str], dict[str, bool]]]]: + """Return the native image-upload declaration shape.""" + + del cls + return {"required": {"image": (["z.png", "a.png"], {"image_upload": True})}} + + monkeypatch.setattr(import_module("nodes"), "LoadImageMask", FakeNativeLoader) + + assert MaskFileLoader().available_files() == ("z.png", "a.png") + + +def test_rejects_missing_multiframe_and_invalid_channel( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Files cannot silently add regions or escape native channel semantics.""" + + path = tmp_path / "animated.gif" + first = Image.new("L", (2, 2), color=0) + second = Image.new("L", (2, 2), color=255) + first.save(path, save_all=True, append_images=[second]) + _patch_path(monkeypatch, path) + loader = MaskFileLoader() + + with pytest.raises(ValueError, match="contains 2 frames"): + loader.load("animated.gif", "red") + with pytest.raises(ValueError, match="mask channel must be one of"): + loader.load("animated.gif", "luminance") + + _patch_path(monkeypatch, path, exists=False) + with pytest.raises(ValueError, match="does not exist"): + loader.fingerprint("missing.png") + + +def test_fingerprint_reads_validated_file_contents( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Content changes invalidate the loader node cache.""" + + path = tmp_path / "mask.png" + path.write_bytes(b"first") + _patch_path(monkeypatch, path) + loader = MaskFileLoader() + first = loader.fingerprint("mask.png") + path.write_bytes(b"second") + + assert loader.fingerprint("mask.png") != first diff --git a/tests/test_native_hooked_regional_conditioning.py b/tests/test_native_hooked_regional_conditioning.py new file mode 100644 index 0000000..5c68c2b --- /dev/null +++ b/tests/test_native_hooked_regional_conditioning.py @@ -0,0 +1,195 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Integration tests for native Comfy hook-aware regional evaluation.""" + +from __future__ import annotations + +from collections.abc import Callable +from typing import Any + +import comfy.conds +import comfy.hooks +import comfy.sampler_helpers +import comfy.samplers +import pytest +import torch + +from simple_syrup.domain.conditioning_batch import ConditioningBatch +from simple_syrup.domain.tiled_diffusion import build_tiled_diffusion_plan +from simple_syrup.runtime.mixture_of_diffusers_sampling import ( + MixtureOfDiffusersModelWrapper, +) +from simple_syrup.runtime.multidiffusion_sampling import MultiDiffusionModelWrapper +from simple_syrup.services.regional_conditioning_service import ( + RegionalConditioningService, +) + +ModelWrapperFactory = Callable[[], object] + + +class FakePatcher: + """Record hook preparation and activation performed by Comfy's sampler.""" + + def __init__(self) -> None: + """Initialize hook observations for one model evaluation.""" + + self.prepared: list[comfy.hooks.HookGroup] = [] + self.applied: list[comfy.hooks.HookGroup | None] = [] + self.active: comfy.hooks.HookGroup | None = None + + def prepare_hook_patches_current_keyframe( + self, + timestep: torch.Tensor, + hooks: comfy.hooks.HookGroup, + model_options: dict[str, Any], + ) -> None: + """Record the hook group prepared for the current sigma.""" + + del timestep, model_options + self.prepared.append(hooks) + + def prepare_state( + self, + timestep: torch.Tensor, + model_options: dict[str, Any], + ) -> None: + """Accept Comfy's per-step patcher state preparation.""" + + del timestep, model_options + + def get_free_memory(self, device: torch.device) -> float: + """Return ample deterministic capacity for conditioning batches.""" + + del device + return 1.0e12 + + def apply_hooks( + self, + *, + hooks: comfy.hooks.HookGroup | None, + ) -> dict[str, Any]: + """Record the hook group activated before the model call.""" + + self.applied.append(hooks) + self.active = hooks + return {} + + +class FakeModel: + """Provide the minimal model surface used by native conditioning evaluation.""" + + def __init__(self) -> None: + """Create a model with an observable hook-aware patcher.""" + + self.current_patcher = FakePatcher() + self.model_calls: list[comfy.hooks.HookGroup | None] = [] + + def memory_required( + self, + input_shape: list[int], + *, + cond_shapes: dict[str, list[list[int]]], + ) -> float: + """Return a fixed estimate so Comfy can select a batch size.""" + + del input_shape, cond_shapes + return 1.0 + + def apply_model( + self, + input_x: torch.Tensor, + timestep: torch.Tensor, + **conditioning: Any, + ) -> torch.Tensor: + """Record the active hook group and return a deterministic prediction.""" + + del timestep, conditioning + self.model_calls.append(self.current_patcher.active) + return torch.ones_like(input_x) + + +@pytest.mark.parametrize( + "wrapper_factory", + [ + None, + lambda: MultiDiffusionModelWrapper( + plan=build_tiled_diffusion_plan( + latent_width=8, + latent_height=8, + tile_width=4, + tile_height=4, + overlap=1, + tile_batch_size=2, + ), + existing_wrapper=None, + ), + lambda: MixtureOfDiffusersModelWrapper( + plan=build_tiled_diffusion_plan( + latent_width=8, + latent_height=8, + tile_width=4, + tile_height=4, + overlap=1, + tile_batch_size=2, + ), + existing_wrapper=None, + ), + ], + ids=["native", "multidiffusion", "mixture-of-diffusers"], +) +def test_native_sampler_activates_hooks_from_masked_regional_conditioning( + wrapper_factory: ModelWrapperFactory | None, +) -> None: + """Comfy activates each preserved hook group through direct and tiled paths.""" + + global_hooks = comfy.hooks.HookGroup() + regional_hooks = comfy.hooks.HookGroup() + assembled, _ = RegionalConditioningService().assemble( + positive=ConditioningBatch( + ( + _conditioning(global_hooks), + _conditioning(regional_hooks), + ) + ), + negative=_conditioning(None), + masks=torch.ones((1, 8, 8)), + regional_prompt_weight=0.5, + region_mask_feather=0, + ) + converted = comfy.sampler_helpers.convert_cond(assembled) + for conditioning in converted: + cross_attn = conditioning.pop("cross_attn") + conditioning["model_conds"] = { + "c_crossattn": comfy.conds.CONDCrossAttn(cross_attn) + } + + model = FakeModel() + model_options: dict[str, Any] = {} + if wrapper_factory is not None: + model_options["model_function_wrapper"] = wrapper_factory() + + outputs = comfy.samplers.calc_cond_batch( + model, + [converted], + torch.zeros((1, 1, 8, 8)), + torch.ones((1,)), + model_options, + ) + + assert len(outputs) == 1 + assert set(model.current_patcher.prepared) == {global_hooks, regional_hooks} + assert set(model.current_patcher.applied) == {global_hooks, regional_hooks} + assert set(model.model_calls) == {global_hooks, regional_hooks} + + +def _conditioning( + hooks: comfy.hooks.HookGroup | None, +) -> list[list[object]]: + """Return one raw Comfy conditioning entry with optional hooks.""" + + metadata: dict[str, object] = {} + if hooks is not None: + metadata["hooks"] = hooks + return [[torch.ones((1, 2, 3)), metadata]] diff --git a/tests/test_node_tooltips.py b/tests/test_node_tooltips.py index 8150e48..2592db2 100644 --- a/tests/test_node_tooltips.py +++ b/tests/test_node_tooltips.py @@ -7,6 +7,7 @@ from __future__ import annotations import sys +from pathlib import Path from types import ModuleType from typing import Any, Protocol @@ -44,6 +45,21 @@ class _FakeFolderPaths(ModuleType): return self._files.get(folder_name, []) + def get_input_directory(self) -> str: + """Return an existing directory for native image-loader discovery.""" + + return str(Path(__file__).parent) + + def filter_files_content_types( + self, + files: list[str], + content_types: list[str], + ) -> list[str]: + """Return no images from the test directory.""" + + del files, content_types + return [] + def test_v3_nodes_provide_tooltip_metadata(monkeypatch: pytest.MonkeyPatch) -> None: """All supported v3 schemas expose descriptions and field-level help.""" @@ -92,6 +108,12 @@ def test_high_impact_tooltips_explain_direction_and_units( assert "seams" in tiled_inputs["latent_tile_overlap"].tooltip.lower() assert "memory" in tiled_inputs["latent_tile_batch_size"].tooltip.lower() + regional_inputs = _inputs_by_id(schemas["SimpleSyrup.KSamplerPromptByRegion"]) + regional_weight_tooltip = regional_inputs["regional_prompt_weight"].tooltip.lower() + assert "0" in regional_weight_tooltip + assert "1" in regional_weight_tooltip + assert "overlaps" in regional_weight_tooltip + def _inputs_by_id(schema: Any) -> dict[str, Any]: """Return schema inputs keyed by id.""" diff --git a/tests/test_prompt_control_batch_graph.py b/tests/test_prompt_control_batch_graph.py index d631134..f9ace8d 100644 --- a/tests/test_prompt_control_batch_graph.py +++ b/tests/test_prompt_control_batch_graph.py @@ -86,6 +86,38 @@ def test_prompt_control_batch_graph_builds_pack_chain_for_multiple_chunks( assert output.args[1] == ["BATCH.0.5.2", 0] +def test_prompt_control_batch_graph_attaches_segment_local_lora_hooks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The clip-only encoder shares each aligned hook plan across both sides.""" + + calls = _install_fake_prompt_control(monkeypatch) + graph_utils = import_module("comfy_execution.graph_utils") + graph_utils.GraphBuilder.set_default_prefix("HOOKS", 0, 0) + + output = PromptControlBatchGraphBuilder().build( + clip=[0, 0], + positive_prompt="face [SEP] hair ", + negative_prompt="blur [SEP] noise", + separator="[SEP]", + ) + + assert output.expand is not None + hook_nodes = [ + node + for node in output.expand.values() + if node["class_type"] == "PCLoraHooksFromText" + ] + assert [node["inputs"]["text"] for node in hook_nodes] == [ + "\n", + "", + ] + assert [call["text"] for call in calls] == ["face ", "hair ", "blur ", "noise"] + assert calls[0]["clip"] == calls[2]["clip"] + assert calls[1]["clip"] == calls[3]["clip"] + assert calls[0]["clip"] != calls[1]["clip"] + + def test_prompt_control_batch_graph_reports_missing_prompt_control( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -97,7 +129,7 @@ def test_prompt_control_batch_graph_reports_missing_prompt_control( return import_module(name) monkeypatch.setattr( - "simple_syrup.runtime.prompt_control_batch_graph.import_module", + "simple_syrup.runtime.prompt_control_graph_adapter.import_module", fake_import_module, ) @@ -111,11 +143,14 @@ def test_prompt_control_batch_graph_reports_missing_prompt_control( assert PROMPT_CONTROL_MISSING_MESSAGE.startswith("Encode Prompt Batch") -def _install_fake_prompt_control(monkeypatch: pytest.MonkeyPatch) -> None: +def _install_fake_prompt_control( + monkeypatch: pytest.MonkeyPatch, +) -> list[dict[str, Any]]: """Install a small Prompt Control lazy-node double for graph tests.""" prompt_control = ModuleType("prompt_control") nodes_lazy = ModuleType("prompt_control.nodes_lazy") + calls: list[dict[str, Any]] = [] class FakePCLazyTextEncodeAdvanced: """Graph-expanding stand-in for Prompt Control's lazy text encoder.""" @@ -132,6 +167,7 @@ def _install_fake_prompt_control(monkeypatch: pytest.MonkeyPatch) -> None: """Return one lazy text encode node output.""" del tags, start, end, num_steps + calls.append({"clip": clip, "text": text}) graph_utils = import_module("comfy_execution.graph_utils") io = import_module("comfy_api.latest").io graph = graph_utils.GraphBuilder() @@ -146,3 +182,4 @@ def _install_fake_prompt_control(monkeypatch: pytest.MonkeyPatch) -> None: cast(Any, prompt_control).nodes_lazy = nodes_lazy monkeypatch.setitem(sys.modules, "prompt_control", prompt_control) monkeypatch.setitem(sys.modules, "prompt_control.nodes_lazy", nodes_lazy) + return calls diff --git a/tests/test_prompt_control_prompt.py b/tests/test_prompt_control_prompt.py index af1ccaf..d72a4a8 100644 --- a/tests/test_prompt_control_prompt.py +++ b/tests/test_prompt_control_prompt.py @@ -38,8 +38,8 @@ def test_extract_lora_tags_returns_angle_tags_joined_with_newlines() -> None: assert extract_lora_tags(text) == "\n" -def test_prepare_prompt_side_splits_ordered_chunks_and_aggregates_loras() -> None: - """Separator-delimited chunks preserve order and collect all tags.""" +def test_prepare_prompt_side_preserves_ordered_segment_local_loras() -> None: + """Separator-delimited chunks retain their own tags without aggregation.""" side = prepare_prompt_side( "face [SEP] hair ", @@ -51,7 +51,6 @@ def test_prepare_prompt_side_splits_ordered_chunks_and_aggregates_loras() -> Non "", "", ] - assert side.lora_tags == "\n" def test_prepare_prompt_side_preserves_empty_chunks() -> None: @@ -61,7 +60,6 @@ def test_prepare_prompt_side_preserves_empty_chunks() -> None: assert [chunk.text for chunk in side.chunks] == ["face", ""] assert [chunk.lora_tags for chunk in side.chunks] == ["", ""] - assert side.lora_tags == "" def test_prepare_prompt_side_rejects_empty_separator() -> None: diff --git a/tests/test_prompt_control_schedule_encode_graph.py b/tests/test_prompt_control_schedule_encode_graph.py index 91f5f8c..883e630 100644 --- a/tests/test_prompt_control_schedule_encode_graph.py +++ b/tests/test_prompt_control_schedule_encode_graph.py @@ -88,22 +88,41 @@ def test_schedule_encode_graph_packs_only_multichunk_sides( ] -def test_schedule_encode_graph_collects_loras_from_all_chunks( +def test_schedule_encode_graph_keeps_loras_local_to_aligned_segments( monkeypatch: pytest.MonkeyPatch, ) -> None: - """LoRA scheduling sees all tags from every separator-delimited chunk.""" + """Each SEP index receives one hook plan shared by both prompt sides.""" calls = _install_fake_prompt_control(monkeypatch) - PromptControlScheduleEncodeGraphBuilder().build( + output = PromptControlScheduleEncodeGraphBuilder().build( model=["model", 0], clip=["clip", 0], positive_prompt="face [SEP] hair ", - negative_prompt="blur [SEP] noise ", + negative_prompt="blur [SEP] noise ", ) - assert calls["lora"][0]["text"] == "\n" - assert calls["lora"][1]["text"] == "\n" + assert output.args[0] == ["model", 0] + assert calls["lora"] == [] + assert output.expand is not None + hook_nodes = [ + node + for node in output.expand.values() + if node["class_type"] == "PCLoraHooksFromText" + ] + assert [node["inputs"]["text"] for node in hook_nodes] == [ + "\n", + "\n", + ] + clip_nodes = [ + node for node in output.expand.values() if node["class_type"] == "SetClipHooks" + ] + assert len(clip_nodes) == 2 + assert all(node["inputs"]["apply_to_conds"] is True for node in clip_nodes) + assert all(node["inputs"]["schedule_clip"] is True for node in clip_nodes) + assert calls["encode"][0]["clip"] == calls["encode"][2]["clip"] + assert calls["encode"][1]["clip"] == calls["encode"][3]["clip"] + assert calls["encode"][0]["clip"] != calls["encode"][1]["clip"] def test_schedule_encode_graph_reports_duplicate_expand_ids( @@ -133,7 +152,7 @@ def test_schedule_encode_graph_reports_missing_prompt_control( return import_module(name) monkeypatch.setattr( - "simple_syrup.runtime.prompt_control_schedule_encode_graph.import_module", + "simple_syrup.runtime.prompt_control_graph_adapter.import_module", fake_import_module, ) diff --git a/tests/test_prompt_control_segment_planning_service.py b/tests/test_prompt_control_segment_planning_service.py new file mode 100644 index 0000000..69903fd --- /dev/null +++ b/tests/test_prompt_control_segment_planning_service.py @@ -0,0 +1,69 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for Prompt-Control SEP segment and hook planning.""" + +from __future__ import annotations + +from simple_syrup.services.prompt_control_segment_planning_service import ( + PromptControlSegmentPlanningService, +) + + +def test_planner_combines_lora_tags_only_at_aligned_indexes() -> None: + """Positive and negative LoRAs share a plan only within one SEP position.""" + + plan = PromptControlSegmentPlanningService().prepare( + positive_prompt=( + "global [SEP] " + "left [SEP] right" + ), + negative_prompt=( + "bad [SEP] worse " + ), + separator="[SEP]", + ) + + assert plan.is_batched is True + assert [chunk.text for chunk in plan.positive.chunks] == [ + "global ", + "left ", + "right", + ] + assert [hook.lora_tags for hook in plan.hooks] == [ + "\n", + "\n", + "", + ] + + +def test_planner_preserves_explicit_empties_without_padding_shorter_side() -> None: + """Empty chunks remain positional while unequal side lengths remain unequal.""" + + plan = PromptControlSegmentPlanningService().prepare( + positive_prompt="global [SEP] [SEP] right ", + negative_prompt="negative", + separator="[SEP]", + ) + + assert [chunk.text for chunk in plan.positive.chunks] == [ + "global", + "", + "right ", + ] + assert [chunk.text for chunk in plan.negative.chunks] == ["negative"] + assert [hook.lora_tags for hook in plan.hooks] == ["", "", ""] + + +def test_planner_marks_single_segment_prompts_as_unbatched() -> None: + """No-SEP prompts retain the compatibility path used by the scheduler.""" + + plan = PromptControlSegmentPlanningService().prepare( + positive_prompt="portrait ", + negative_prompt="blur", + separator="[SEP]", + ) + + assert plan.is_batched is False + assert plan.hooks[0].lora_tags == "" diff --git a/tests/test_regional_conditioning_service.py b/tests/test_regional_conditioning_service.py new file mode 100644 index 0000000..08b8295 --- /dev/null +++ b/tests/test_regional_conditioning_service.py @@ -0,0 +1,222 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for shared regional conditioning assembly.""" + +from __future__ import annotations + +import pytest +import torch + +from simple_syrup.domain.conditioning_batch import ConditioningBatch +from simple_syrup.services.regional_conditioning_service import ( + RegionalConditioningService, +) + + +def _conditioning(name: str) -> list[list[object]]: + """Return recognizable standard Comfy conditioning.""" + + return [[name, {"source": name}]] + + +def test_normal_conditioning_remains_global_only() -> None: + """Normal conditioning is copied and broadcast globally.""" + + positive = _conditioning("global positive") + negative = _conditioning("global negative") + + assembled_positive, assembled_negative = RegionalConditioningService().assemble( + positive=positive, + negative=negative, + masks=torch.ones((2, 4, 4)), + regional_prompt_weight=0.5, + region_mask_feather=0, + ) + + assert assembled_positive == positive + assert assembled_negative == negative + assert assembled_positive is not positive + assert assembled_positive[0][1] is not positive[0][1] + + +def test_batches_pair_global_and_regions_independently() -> None: + """Positive and negative regional counts may differ without fallback reuse.""" + + masks = torch.stack( + [torch.zeros((3, 3)), torch.ones((3, 3)), torch.full((3, 3), 0.5)] + ) + positive = ConditioningBatch( + ( + _conditioning("positive global"), + _conditioning("positive region 0"), + _conditioning("positive region 1"), + ) + ) + negative = ConditioningBatch( + (_conditioning("negative global"), _conditioning("negative region 0")) + ) + + assembled_positive, assembled_negative = RegionalConditioningService().assemble( + positive=positive, + negative=negative, + masks=masks, + regional_prompt_weight=0.75, + region_mask_feather=0, + ) + + assert [item[0] for item in assembled_positive] == [ + "positive global", + "positive region 0", + "positive region 1", + ] + assert [item[0] for item in assembled_negative] == [ + "negative global", + "negative region 0", + ] + assert torch.equal( + assembled_positive[0][1]["mask"], + torch.full((1, 3, 3), 0.25), + ) + assert torch.equal( + assembled_negative[0][1]["mask"], + torch.ones((1, 3, 3)), + ) + assert assembled_positive[0][1]["mask_strength"] == 1.0 + assert assembled_negative[0][1]["mask_strength"] == 1.0 + assert torch.equal(assembled_positive[1][1]["mask"], masks[0:1]) + assert torch.equal(assembled_positive[2][1]["mask"], masks[1:2]) + assert torch.equal(assembled_negative[1][1]["mask"], masks[0:1]) + assert assembled_positive[1][1]["mask_strength"] == 0.75 + assert assembled_positive[2][1]["mask_strength"] == 0.75 + assert assembled_negative[1][1]["mask_strength"] == 0.75 + assert assembled_positive[1][1]["set_area_to_bounds"] is False + + +def test_excess_positive_or_negative_regions_fail() -> None: + """Either conditioning side rejects regional entries without masks.""" + + service = RegionalConditioningService() + masks = torch.ones((1, 2, 2)) + excessive = ConditioningBatch( + (_conditioning("global"), _conditioning("one"), _conditioning("two")) + ) + + with pytest.raises(ValueError, match="positive conditioning contains 2"): + service.assemble( + positive=excessive, + negative=_conditioning("negative"), + masks=masks, + regional_prompt_weight=0.5, + region_mask_feather=0, + ) + with pytest.raises(ValueError, match="negative conditioning contains 2"): + service.assemble( + positive=_conditioning("positive"), + negative=excessive, + masks=masks, + regional_prompt_weight=0.5, + region_mask_feather=0, + ) + + +def test_feathering_preserves_inputs_and_softens_regional_copy() -> None: + """Optional feathering changes attached masks without mutating the input.""" + + masks = torch.zeros((1, 9, 9)) + masks[:, 3:6, 3:6] = 1.0 + original = masks.clone() + + positive, _ = RegionalConditioningService().assemble( + positive=ConditioningBatch((_conditioning("global"), _conditioning("region"))), + negative=_conditioning("negative"), + masks=masks, + regional_prompt_weight=0.5, + region_mask_feather=2, + ) + + attached = positive[1][1]["mask"] + assert isinstance(attached, torch.Tensor) + assert torch.equal(masks, original) + assert not torch.equal(attached, original) + assert bool(torch.any((attached > 0.0) & (attached < 1.0))) + + +def test_mask_composition_preserves_segment_lora_hook_metadata() -> None: + """Global and regional hook groups survive standard mask composition.""" + + global_hooks = object() + regional_hooks = object() + positive = ConditioningBatch( + ( + [["global", {"hooks": global_hooks, "other": "global metadata"}]], + [["region", {"hooks": regional_hooks, "other": "region metadata"}]], + ) + ) + + assembled, _ = RegionalConditioningService().assemble( + positive=positive, + negative=_conditioning("negative"), + masks=torch.ones((1, 3, 3)), + regional_prompt_weight=0.5, + region_mask_feather=0, + ) + + assert assembled[0][1]["hooks"] is global_hooks + assert assembled[0][1]["other"] == "global metadata" + assert assembled[1][1]["hooks"] is regional_hooks + assert assembled[1][1]["other"] == "region metadata" + + +def test_zero_regional_prompt_weight_returns_only_unchanged_global_entries() -> None: + """The zero endpoint disables regional conditioning completely.""" + + positive, negative = RegionalConditioningService().assemble( + positive=ConditioningBatch((_conditioning("global"), _conditioning("region"))), + negative=ConditioningBatch( + (_conditioning("global negative"), _conditioning("region negative")) + ), + masks=torch.ones((1, 2, 2)), + regional_prompt_weight=0.0, + region_mask_feather=0, + ) + + assert positive == _conditioning("global") + assert negative == _conditioning("global negative") + assert "mask" not in positive[0][1] + assert "mask" not in negative[0][1] + + +def test_full_regional_prompt_weight_complements_global_inside_mask() -> None: + """The one endpoint removes global influence only inside solid coverage.""" + + mask = torch.tensor([[[1.0, 0.0], [0.5, 0.0]]]) + positive, _ = RegionalConditioningService().assemble( + positive=ConditioningBatch((_conditioning("global"), _conditioning("region"))), + negative=_conditioning("negative"), + masks=mask, + regional_prompt_weight=1.0, + region_mask_feather=0, + ) + + assert torch.equal(positive[0][1]["mask"], 1.0 - mask) + assert positive[0][1]["mask_strength"] == 1.0 + assert torch.equal(positive[1][1]["mask"], mask) + assert positive[1][1]["mask_strength"] == 1.0 + + +@pytest.mark.parametrize("weight", [-0.01, 1.01, float("nan")]) +def test_invalid_regional_prompt_weight_fails_before_mask_processing( + weight: float, +) -> None: + """The service rejects invalid influence before touching mask inputs.""" + + with pytest.raises(ValueError, match="regional_prompt_weight"): + RegionalConditioningService().assemble( + positive=_conditioning("positive"), + negative=_conditioning("negative"), + masks=object(), + regional_prompt_weight=weight, + region_mask_feather=0, + ) diff --git a/tests/test_regional_ksampler_v3_nodes.py b/tests/test_regional_ksampler_v3_nodes.py new file mode 100644 index 0000000..a1bf4c5 --- /dev/null +++ b/tests/test_regional_ksampler_v3_nodes.py @@ -0,0 +1,194 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for full-latent and tiled regional KSampler v3 nodes.""" + +from __future__ import annotations + +from typing import Any, ClassVar + +import torch + +from simple_syrup.nodes_v3.ksampler_prompt_by_region import ( + KSamplerPromptByRegionV3, +) +from simple_syrup.nodes_v3.ksampler_prompt_by_tiled_region import ( + KSamplerPromptByTiledRegionV3, +) + + +class FakeConditioningService: + """Record regional assembly and return recognizable conditioning.""" + + calls: ClassVar[list[dict[str, object]]] = [] + + def assemble(self, **kwargs: object) -> tuple[object, object]: + """Record one shared assembly request.""" + + type(self).calls.append(kwargs) + return "assembled-positive", "assembled-negative" + + +class FakeSamplingService: + """Record normal or tiled sampling arguments.""" + + calls: ClassVar[list[dict[str, Any]]] = [] + output: ClassVar[dict[str, Any]] = {"samples": torch.ones((1, 4, 2, 2))} + + def sample(self, **kwargs: Any) -> dict[str, Any]: + """Record and return one latent.""" + + type(self).calls.append(kwargs) + return self.output + + +def test_schemas_expose_exact_names_and_shared_regional_contract() -> None: + """Both nodes expose global-first regional inputs with only tiled extras.""" + + normal = KSamplerPromptByRegionV3.define_schema() + tiled = KSamplerPromptByTiledRegionV3.define_schema() + normal_ids = [value.id for value in normal.inputs] + tiled_ids = [value.id for value in tiled.inputs] + + assert normal.node_id == "SimpleSyrup.KSamplerPromptByRegion" + assert normal.display_name == "KSampler (Prompt by Region)" + assert tiled.node_id == "SimpleSyrup.KSamplerPromptByTiledRegion" + assert tiled.display_name == "KSampler (Prompt by Tiled Region)" + assert normal_ids == [ + "model", + "seed", + "steps", + "cfg", + "sampler_name", + "scheduler", + "positive", + "negative", + "region_masks", + "regional_prompt_weight", + "region_mask_feather", + "latent_image", + "denoise", + ] + assert tiled_ids[: len(normal_ids)] == normal_ids + assert tiled_ids[len(normal_ids) :] == [ + "diffusion_mode", + "latent_tile_width", + "latent_tile_height", + "latent_tile_overlap", + "latent_tile_batch_size", + ] + regional_weight = normal.inputs[9] + assert regional_weight.default == 0.5 + assert regional_weight.min == 0.0 + assert regional_weight.max == 1.0 + assert [output.io_type for output in normal.outputs] == ["LATENT"] + assert [output.io_type for output in tiled.outputs] == ["LATENT"] + + +def test_non_tiled_node_assembles_then_samples_full_latent() -> None: + """The non-tiled node delegates regional and sampling concerns once.""" + + _reset_fakes() + original_conditioning = KSamplerPromptByRegionV3.conditioning_service_class + original_sampling = KSamplerPromptByRegionV3.sampling_service_class + KSamplerPromptByRegionV3.conditioning_service_class = FakeConditioningService # type: ignore[assignment] + KSamplerPromptByRegionV3.sampling_service_class = FakeSamplingService # type: ignore[assignment] + masks = torch.ones((2, 8, 8)) + latent = {"samples": torch.zeros((3, 4, 2, 2))} + try: + (output,) = KSamplerPromptByRegionV3.execute( + model="model", + seed=3, + steps=10, + cfg=5.0, + sampler_name="euler", + scheduler="normal", + positive="positive-batch", + negative="negative-batch", + region_masks=masks, + regional_prompt_weight=0.75, + region_mask_feather=4, + latent_image=latent, + denoise=0.7, + ) + finally: + KSamplerPromptByRegionV3.conditioning_service_class = original_conditioning + KSamplerPromptByRegionV3.sampling_service_class = original_sampling + + assert output is FakeSamplingService.output + assert FakeConditioningService.calls == [ + { + "positive": "positive-batch", + "negative": "negative-batch", + "masks": masks, + "regional_prompt_weight": 0.75, + "region_mask_feather": 4, + } + ] + call = FakeSamplingService.calls[0] + assert call["positive"] == "assembled-positive" + assert call["negative"] == "assembled-negative" + assert call["latent_image"] is latent + assert "diffusion_mode" not in call + + +def test_tiled_node_enables_only_full_context_regional_masks() -> None: + """The tiled node shares assembly and explicitly opts into mask support.""" + + _reset_fakes() + original_conditioning = KSamplerPromptByTiledRegionV3.conditioning_service_class + original_sampling = KSamplerPromptByTiledRegionV3.sampling_service_class + KSamplerPromptByTiledRegionV3.conditioning_service_class = FakeConditioningService # type: ignore[assignment] + KSamplerPromptByTiledRegionV3.sampling_service_class = FakeSamplingService # type: ignore[assignment] + masks = torch.ones((1, 8, 8)) + latent = {"samples": torch.zeros((1, 4, 16, 16))} + try: + (output,) = KSamplerPromptByTiledRegionV3.execute( + model="model", + seed=3, + steps=10, + cfg=5.0, + sampler_name="euler", + scheduler="normal", + positive="positive-batch", + negative="negative-batch", + region_masks=masks, + regional_prompt_weight=0.6, + region_mask_feather=0, + latent_image=latent, + denoise=0.7, + diffusion_mode="mixture_of_diffusers", + latent_tile_width=8, + latent_tile_height=8, + latent_tile_overlap=2, + latent_tile_batch_size=3, + ) + finally: + KSamplerPromptByTiledRegionV3.conditioning_service_class = original_conditioning + KSamplerPromptByTiledRegionV3.sampling_service_class = original_sampling + + assert output is FakeSamplingService.output + assert FakeConditioningService.calls == [ + { + "positive": "positive-batch", + "negative": "negative-batch", + "masks": masks, + "regional_prompt_weight": 0.6, + "region_mask_feather": 0, + } + ] + call = FakeSamplingService.calls[0] + assert call["diffusion_mode"] == "mixture_of_diffusers" + assert call["allow_full_context_masks"] is True + assert call["latent_tile_width"] == 8 + assert call["latent_tile_height"] == 8 + assert call["latent_tile_overlap"] == 2 + assert call["latent_tile_batch_size"] == 3 + + +def _reset_fakes() -> None: + """Clear shared fake records between node tests.""" + + FakeConditioningService.calls = [] + FakeSamplingService.calls = [] diff --git a/tests/test_regional_prompt_masks.py b/tests/test_regional_prompt_masks.py new file mode 100644 index 0000000..31b1b99 --- /dev/null +++ b/tests/test_regional_prompt_masks.py @@ -0,0 +1,49 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for full-context regional prompt mask preparation.""" + +from __future__ import annotations + +import pytest +import torch + +from simple_syrup.masking.regional_prompt_masks import ( + complementary_global_prompt_mask, +) + + +def test_complementary_global_mask_uses_clamped_accumulated_coverage() -> None: + """Overlapping regions accumulate while global coverage remains normalized.""" + + masks = torch.tensor( + [ + [[1.0, 1.0, 0.0]], + [[0.0, 1.0, 1.0]], + ] + ) + + global_mask = complementary_global_prompt_mask(masks, (0, 1), 0.5) + regional_total = masks.sum(dim=0, keepdim=True) * 0.5 + global_share = global_mask / (global_mask + regional_total) + + assert torch.equal(global_mask, torch.tensor([[[0.5, 0.5, 0.5]]])) + assert torch.allclose(global_share, torch.tensor([[[0.5, 1.0 / 3.0, 0.5]]])) + + +def test_complementary_global_mask_uses_only_paired_indices() -> None: + """Extra authored masks do not reduce global prompt influence.""" + + masks = torch.stack([torch.zeros((2, 2)), torch.ones((2, 2))]) + + global_mask = complementary_global_prompt_mask(masks, (0,), 1.0) + + assert torch.equal(global_mask, torch.ones((1, 2, 2))) + + +def test_complementary_global_mask_requires_a_regional_pair() -> None: + """Coverage cannot be calculated without a paired regional prompt.""" + + with pytest.raises(ValueError, match="at least one regional mask"): + complementary_global_prompt_mask(torch.ones((1, 2, 2)), (), 0.5) diff --git a/tests/test_regional_prompt_workflow.py b/tests/test_regional_prompt_workflow.py new file mode 100644 index 0000000..0875f3d --- /dev/null +++ b/tests/test_regional_prompt_workflow.py @@ -0,0 +1,142 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Integration test for the authored-mask regional prompt workflow.""" + +from __future__ import annotations + +from typing import Any, ClassVar, cast + +import pytest +import torch + +from simple_syrup.domain.conditioning_batch import ConditioningBatch +from simple_syrup.nodes.encode_prompt_batch import EncodePromptBatch +from simple_syrup.nodes_v3 import load_mask_batch as load_node_module +from simple_syrup.nodes_v3.ksampler_prompt_by_region import ( + KSamplerPromptByRegionV3, +) +from simple_syrup.nodes_v3.load_mask_batch import LoadMaskBatchV3 + + +class WorkflowEncoder: + """Encode prompt chunks into standard recognizable conditioning.""" + + def encode_batch( + self, + clip: Any, + chunks: tuple[str, ...], + ) -> ConditioningBatch: + """Return ordered standard conditioning entries.""" + + del clip + return ConditioningBatch(tuple([[chunk, {}]] for chunk in chunks)) + + +class WorkflowMaskLoader: + """Return three authored masks in selected-file order.""" + + def load(self, files: list[str], channel: str) -> torch.Tensor: + """Return values that expose positional ordering.""" + + assert files == ["left.png", "right.png", "extra.png"] + assert channel == "red" + return torch.stack( + [ + torch.zeros((4, 4)), + torch.ones((4, 4)), + torch.full((4, 4), 0.5), + ] + ) + + def validate(self, files: str | list[str], channel: str) -> None: + """Accept integration-test widget values.""" + + del files, channel + + def fingerprint(self, files: list[str], channel: str) -> str: + """Provide the loader protocol for the node class.""" + + return f"{channel}:{files!r}" + + def available_files(self) -> tuple[str, ...]: + """Provide the schema protocol for the node class.""" + + return () + + +class WorkflowSampler: + """Capture assembled conditioning at the final sampling boundary.""" + + calls: ClassVar[list[dict[str, Any]]] = [] + + def sample(self, **kwargs: Any) -> dict[str, Any]: + """Return the input latent after recording the complete request.""" + + type(self).calls.append(kwargs) + return cast(dict[str, Any], kwargs["latent_image"]) + + +def test_sep_prompts_and_disk_masks_flow_into_regional_sampler( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The public nodes compose with global-first and missing-region behavior.""" + + monkeypatch.setattr(EncodePromptBatch, "encoder_class", WorkflowEncoder) + monkeypatch.setattr(LoadMaskBatchV3, "service_class", WorkflowMaskLoader) + monkeypatch.setattr( + load_node_module._comfy_ui, "PreviewMask", lambda mask, cls: None + ) + monkeypatch.setattr( + KSamplerPromptByRegionV3, + "sampling_service_class", + WorkflowSampler, + ) + WorkflowSampler.calls = [] + + positive, negative = EncodePromptBatch().encode( + clip=object(), + positive_prompt="global [SEP] left [SEP] right", + negative_prompt="global negative [SEP] left negative", + separator="[SEP]", + ) + mask_output = LoadMaskBatchV3.execute( + image=["left.png", "right.png", "extra.png"], + channel="red", + ) + assert mask_output.result is not None + masks = mask_output.result[0] + latent = {"samples": torch.zeros((2, 4, 1, 1))} + + (output,) = KSamplerPromptByRegionV3.execute( + model=object(), + seed=1, + steps=2, + cfg=3.0, + sampler_name="euler", + scheduler="normal", + positive=positive, + negative=negative, + region_masks=masks, + regional_prompt_weight=0.8, + region_mask_feather=0, + latent_image=latent, + denoise=1.0, + ) + + assert output is latent + call = WorkflowSampler.calls[0] + assembled_positive = call["positive"] + assembled_negative = call["negative"] + assert [item[0] for item in assembled_positive] == ["global", "left", "right"] + assert [item[0] for item in assembled_negative] == [ + "global negative", + "left negative", + ] + assert torch.equal(assembled_positive[1][1]["mask"], masks[0:1]) + assert torch.equal(assembled_positive[2][1]["mask"], masks[1:2]) + assert assembled_positive[1][1]["mask_strength"] == 0.8 + assert assembled_positive[2][1]["mask_strength"] == 0.8 + assert assembled_negative[1][1]["mask_strength"] == 0.8 + assert call["latent_image"]["samples"].shape[0] == 2 diff --git a/tests/test_regional_prompting_domain.py b/tests/test_regional_prompting_domain.py new file mode 100644 index 0000000..2727ba6 --- /dev/null +++ b/tests/test_regional_prompting_domain.py @@ -0,0 +1,94 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Tests for global-first regional prompt pairing policy.""" + +from __future__ import annotations + +import pytest + +from simple_syrup.domain.regional_prompting import ( + RegionalConditioningPair, + build_regional_conditioning_plan, + validate_regional_prompt_weight, +) + + +def test_pairs_every_available_regional_entry_in_order() -> None: + """Entry zero stays global and later entries map directly to masks.""" + + plan = build_regional_conditioning_plan( + region_count=3, + conditioning_count=3, + input_name="positive", + ) + + assert plan.region_count == 3 + assert plan.pairs == ( + RegionalConditioningPair(conditioning_index=1, mask_index=0), + RegionalConditioningPair(conditioning_index=2, mask_index=1), + ) + + +def test_fewer_regional_entries_than_masks_is_valid() -> None: + """Unpaired authored masks intentionally retain global conditioning only.""" + + plan = build_regional_conditioning_plan( + region_count=4, + conditioning_count=1, + input_name="negative", + ) + + assert plan.pairs == () + + +def test_more_regional_entries_than_masks_fails_actionably() -> None: + """Excess prompts cannot silently lose their positional meaning.""" + + with pytest.raises( + ValueError, + match="positive conditioning contains 2 regional entries but only 1", + ): + build_regional_conditioning_plan( + region_count=1, + conditioning_count=3, + input_name="positive", + ) + + +@pytest.mark.parametrize( + ("region_count", "conditioning_count", "message"), + [ + (0, 1, "at least one authored mask"), + (1, 0, "global entry at index 0"), + ], +) +def test_invalid_plan_counts_fail( + region_count: int, + conditioning_count: int, + message: str, +) -> None: + """A plan always contains masks and a structural global entry.""" + + with pytest.raises(ValueError, match=message): + build_regional_conditioning_plan( + region_count=region_count, + conditioning_count=conditioning_count, + input_name="positive", + ) + + +@pytest.mark.parametrize("weight", [0.0, 0.5, 1.0]) +def test_regional_prompt_weight_accepts_normalized_range(weight: float) -> None: + """Regional influence accepts both endpoints and the default midpoint.""" + + validate_regional_prompt_weight(weight) + + +@pytest.mark.parametrize("weight", [-0.01, 1.01, float("inf"), float("nan")]) +def test_regional_prompt_weight_rejects_invalid_values(weight: float) -> None: + """Invalid regional influence fails before conditioning assembly.""" + + with pytest.raises(ValueError, match="regional_prompt_weight"): + validate_regional_prompt_weight(weight) diff --git a/tests/test_registration.py b/tests/test_registration.py index e33d4a1..5c2f4c6 100644 --- a/tests/test_registration.py +++ b/tests/test_registration.py @@ -21,6 +21,7 @@ BASE_NODE_IDS = [ "SimpleSyrup.BatchSEGS", "SimpleSyrup.ConditioningBatchAppend", "SimpleSyrup.ConditioningBatchStart", + "SimpleSyrup.ComposeRegionalConditioning", "SimpleSyrup.DetailSEGSAsRegions", "SimpleSyrup.DetailSEGSByScaleFactorTiledDiffusion", "SimpleSyrup.DetailSEGSByScaleFactor", @@ -30,10 +31,13 @@ BASE_NODE_IDS = [ "SimpleSyrup.GroundedSAMModelInfo", "SimpleSyrup.GroundingDINOModelLoader", "SimpleSyrup.KSamplerExtras", + "SimpleSyrup.KSamplerPromptByRegion", + "SimpleSyrup.KSamplerPromptByTiledRegion", "SimpleSyrup.KSamplerTiledDiffusion", "SimpleSyrup.LatentDiagnostics", "SimpleSyrup.LayerStyleSAMModelsAdapter", "SimpleSyrup.LoadUltralyticsModel", + "SimpleSyrup.LoadMaskBatch", "SimpleSyrup.MaskToSEGS", "SimpleSyrup.PromptEncodeStyleAndNormalization", "SimpleSyrup.PromptEncodeStyle", @@ -120,7 +124,7 @@ def test_package_imports_from_custom_nodes_parent_path() -> None: def test_comfy_import_exposes_stable_internal_package_alias() -> None: - """ComfyUI-style import exposes `simple_syrup` for vendored runtime imports.""" + """ComfyUI-style import exposes `simple_syrup` to nested vendored packages.""" project_root = Path(__file__).resolve().parents[1] custom_nodes_root = project_root.parent @@ -132,9 +136,9 @@ def test_comfy_import_exposes_stable_internal_package_alias() -> None: f"sys.path.insert(0, {str(custom_nodes_root)!r}); " "importlib.import_module('SimpleSyrup'); " "runtime = importlib.import_module(" - "'simple_syrup.third_party.groundingdino_runtime.models'" + "'simple_syrup.third_party.groundingdino_runtime.util'" "); " - "assert runtime.__name__.endswith('groundingdino_runtime.models')" + "assert runtime.__name__.endswith('groundingdino_runtime.util')" ) result = subprocess.run( diff --git a/tests/test_tiled_diffusion_sampling_service.py b/tests/test_tiled_diffusion_sampling_service.py index 66f8886..fcc002b 100644 --- a/tests/test_tiled_diffusion_sampling_service.py +++ b/tests/test_tiled_diffusion_sampling_service.py @@ -269,4 +269,5 @@ def _sample_kwargs( "latent_tile_batch_size": 3, "preview_context": preview_context, "differential_diffusion": False, + "allow_full_context_masks": False, } diff --git a/tests/test_tiled_sampling_runtime.py b/tests/test_tiled_sampling_runtime.py index c319319..d78a398 100644 --- a/tests/test_tiled_sampling_runtime.py +++ b/tests/test_tiled_sampling_runtime.py @@ -191,3 +191,22 @@ def test_contains_unsupported_conditioning_key_finds_nested_values() -> None: conditioning = [{"model_conds": {"nested": [{"mask": torch.ones((1, 1))}]}}] assert tiled_sampling.contains_unsupported_conditioning_key(conditioning) + + +def test_full_context_masks_require_explicit_tiled_support() -> None: + """Only explicit non-cropped masks pass the regional tiled policy.""" + + supported = [ + ["tensor", {"mask": torch.ones((1, 2, 2)), "set_area_to_bounds": False}] + ] + cropped = [["tensor", {"mask": torch.ones((1, 2, 2)), "set_area_to_bounds": True}]] + + assert tiled_sampling.contains_unsupported_conditioning_key(supported) + assert not tiled_sampling.contains_unsupported_conditioning_key( + supported, + allow_full_context_masks=True, + ) + assert tiled_sampling.contains_unsupported_conditioning_key( + cropped, + allow_full_context_masks=True, + ) diff --git a/web/dist/simple-syrup.js b/web/dist/simple-syrup.js index 9579688..43db19a 100644 --- a/web/dist/simple-syrup.js +++ b/web/dist/simple-syrup.js @@ -6,6 +6,34 @@ var SETTINGS_ROUTE = "/simple-syrup/settings"; var EXTERNAL_LLM_SETTINGS_ROUTE = "/simple-syrup/external-llm/settings"; var EXTERNAL_LLM_API_KEY_ROUTE = "/simple-syrup/external-llm/api-key"; var EXTERNAL_LLM_MODELS_REFRESH_ROUTE = "/simple-syrup/external-llm/models/refresh"; +var MASK_BATCH_PREVIEW_ROUTE = "/simple-syrup/mask-batch/preview"; +async function getMaskBatchPreview(files, channel, fetchImpl = fetch) { + const response = await fetchImpl(MASK_BATCH_PREVIEW_ROUTE, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ files, channel }) + }); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `Could not render Load Mask Batch preview. Backend returned ${String(response.status)}.` + ) + ); + } + return parseMaskBatchPreview(await response.json()); +} +function parseMaskBatchPreview(payload) { + if (!isMaskBatchPreviewPayload(payload)) { + throw new Error( + "SimpleSyrup mask batch preview payload is invalid. Expected native images and animation flags." + ); + } + return { + images: payload.images.map((image) => ({ ...image })), + animated: [...payload.animated] + }; +} async function getSettings(fetchImpl = fetch) { const response = await fetchImpl(SETTINGS_ROUTE); if (!response.ok) { @@ -123,6 +151,16 @@ function isExternalLLMSettingsPayload(payload) { (model) => typeof model === "string" ) === true && typeof payload.default_model === "string" && typeof payload.has_api_key === "boolean"; } +function isMaskBatchPreviewPayload(payload) { + if (typeof payload !== "object" || payload === null) return false; + const candidate = payload; + return Array.isArray(candidate.images) && candidate.images.every(isComfyImageResult) && Array.isArray(candidate.animated) && candidate.animated.every((value) => typeof value === "boolean"); +} +function isComfyImageResult(value) { + if (typeof value !== "object" || value === null) return false; + const candidate = value; + return typeof candidate.filename === "string" && typeof candidate.subfolder === "string" && (candidate.type === "input" || candidate.type === "output" || candidate.type === "temp"); +} async function backendErrorMessage(response, fallback) { try { const payload = await response.json(); @@ -510,8 +548,356 @@ function registerExternalLLMRefreshHook(app2, api = { refreshExternalLLMModels } }; } +// web/src/maskBatchPreview.ts +var EMPTY_PREVIEW = { + images: [], + animated: [] +}; +var PREVIEW_PROMPT_ID = "simple-syrup-mask-batch-preview"; +var MaskBatchPreviewController = class { + constructor(executionEvents, node, loadPreview = getMaskBatchPreview, logger = console, clearNativePreview2 = () => void 0) { + this.executionEvents = executionEvents; + this.node = node; + this.loadPreview = loadPreview; + this.logger = logger; + this.clearNativePreview = clearNativePreview2; + } + executionEvents; + node; + loadPreview; + logger; + clearNativePreview; + requestVersion = 0; + /** Replace the visible preview with the exact selected-channel mask output. */ + refresh(files, channel) { + if (this.node.id === void 0 || channel === void 0) return; + const requestVersion = ++this.requestVersion; + this.clearNativePreview(); + this.publish(EMPTY_PREVIEW); + if (files.length === 0) return; + void this.loadPreview([...files], channel).then((output) => { + if (requestVersion === this.requestVersion) this.publish(output); + }).catch((error) => { + if (requestVersion !== this.requestVersion) return; + this.logger.warn("Could not refresh Load Mask Batch preview.", error); + }); + } + /** Publish an execution-shaped result through Comfy's native preview. */ + publish(output) { + if (this.node.id === void 0) return; + this.executionEvents.dispatchEvent( + new CustomEvent("executed", { + detail: { + node: String(this.node.id), + output, + merge: false, + prompt_id: PREVIEW_PROMPT_ID + } + }) + ); + this.node.graph?.setDirtyCanvas?.(true, true); + } +}; + +// web/src/maskBatchUpload.ts +var LOAD_MASK_BATCH_NODE_ID = "SimpleSyrup.LoadMaskBatch"; +var REPLACE_MASKS_LABEL = "Replace masks..."; +var ADD_MASKS_LABEL = "Add masks..."; +var REMOVE_MASK_LABEL = "Remove selected mask"; +var EMPTY_MASK_SELECTION_LABEL = "No masks loaded"; +function registerMaskBatchUpload(app2, executionEvents, loadPreview, logger = console) { + const extension = { + name: "SimpleSyrup.LoadMaskBatchUpload", + nodeCreated(node) { + try { + configureMaskBatchNode(node, executionEvents, loadPreview, logger, app2); + } catch (error) { + logger.warn( + `Could not configure Load Mask Batch native controls: ${errorMessage2(error)}`, + error + ); + throw error; + } + } + }; + app2.registerExtension(extension); +} +function errorMessage2(error) { + return error instanceof Error && error.message ? error.message : String(error); +} +function configureMaskBatchNode(candidate, executionEvents, loadPreview, logger = console, app2) { + if (!isMaskBatchNode(candidate)) return; + const imageWidget = findWidget(candidate, "image"); + const channelWidget = findWidget(candidate, "channel"); + const uploadWidget = findNativeUploadWidget(candidate); + if (!imageWidget || !channelWidget || !uploadWidget?.callback) return; + hideInternalWidget(imageWidget); + hideInternalWidget(uploadWidget); + const preview = new MaskBatchPreviewController( + executionEvents, + candidate, + loadPreview, + logger, + () => { + clearNativePreview(candidate, app2); + } + ); + let selectionIntent = "replace"; + let appendBase = []; + let uploadPending = false; + let programmaticSelectionUpdate = false; + let ignoredUploadCallbackFiles; + const nativeUploadCallback = uploadWidget.callback; + const replaceWidget = candidate.addWidget( + "button", + "simple_syrup_replace_masks", + "image", + () => { + selectionIntent = "replace"; + appendBase = []; + uploadPending = true; + nativeUploadCallback.call(uploadWidget); + }, + nativeButtonOptions( + "Choose one or more masks and replace the current ordered list." + ) + ); + replaceWidget.label = REPLACE_MASKS_LABEL; + const setSelectedMasks = (files) => { + programmaticSelectionUpdate = true; + try { + imageWidget.value = [...files]; + } finally { + programmaticSelectionUpdate = false; + } + }; + uploadWidget.callback = (value) => { + selectionIntent = "replace"; + appendBase = []; + uploadPending = true; + nativeUploadCallback.call(uploadWidget, value); + }; + const addWidget = candidate.addWidget( + "button", + "simple_syrup_add_masks", + "image", + () => { + selectionIntent = "append"; + appendBase = selectedMaskFiles(imageWidget); + uploadPending = true; + nativeUploadCallback.call(uploadWidget); + }, + nativeButtonOptions( + "Upload one or more masks and append them after the current ordered list." + ) + ); + addWidget.label = ADD_MASKS_LABEL; + const selectedMaskWidget = candidate.addWidget( + "combo", + "simple_syrup_selected_mask", + EMPTY_MASK_SELECTION_LABEL, + (value) => { + const selectedIndex = selectedMaskIndex(selectedMaskWidget, value); + candidate.imageIndex = selectedIndex; + candidate.graph?.setDirtyCanvas?.(true, true); + }, + { + serialize: false, + tooltip: "Select one loaded mask by its ordered position for preview or removal.", + values: [] + } + ); + selectedMaskWidget.label = "selected mask"; + const removeWidget = candidate.addWidget( + "button", + "simple_syrup_remove_mask", + "image", + () => { + const previous = selectedMaskFiles(imageWidget); + const index = activeMaskIndex(candidate, selectedMaskWidget, previous); + if (index === null) return; + const current = previous.toSpliced(index, 1); + setSelectedMasks(current); + updateMaskSelection(selectedMaskWidget, current, index); + candidate.imageIndex = current.length === 0 ? null : Math.min(index, current.length - 1); + notifySelectionChanged(candidate, imageWidget, { current, previous }); + updateRemoveAvailability(removeWidget, current.length); + preview.refresh(current, selectedChannel(channelWidget)); + }, + nativeButtonOptions( + "Open a mask in the preview gallery, then remove that position from the loaded list." + ) + ); + removeWidget.label = REMOVE_MASK_LABEL; + imageWidget.callback = (value) => { + if (programmaticSelectionUpdate) return; + const callbackValue = value ?? imageWidget.value; + if (uploadPending && !Array.isArray(callbackValue)) return; + const incoming = normalizeMaskFiles(callbackValue); + if (ignoredUploadCallbackFiles && sameMaskFiles(incoming, ignoredUploadCallbackFiles)) { + ignoredUploadCallbackFiles = void 0; + return; + } + ignoredUploadCallbackFiles = uploadPending ? [...incoming] : void 0; + uploadPending = false; + const current = selectionIntent === "append" ? [...appendBase, ...incoming] : incoming; + const appended = selectionIntent === "append"; + selectionIntent = "replace"; + appendBase = []; + if (appended) setSelectedMasks(current); + candidate.imageIndex = null; + updateMaskSelection(selectedMaskWidget, current); + updateRemoveAvailability(removeWidget, current.length); + preview.refresh(current, selectedChannel(channelWidget)); + }; + const originalChannelCallback = channelWidget.callback; + channelWidget.callback = (value) => { + originalChannelCallback?.call(channelWidget, value); + const files = selectedMaskFiles(imageWidget); + if (files.length > 0) { + preview.refresh(files, selectedChannel(channelWidget, value)); + } + }; + resetIntentForExternalUploads(candidate, () => { + selectionIntent = "replace"; + appendBase = []; + uploadPending = false; + }); + const originalOnGraphConfigured = candidate.onGraphConfigured; + candidate.onGraphConfigured = function(...args) { + const result = originalOnGraphConfigured?.apply(this, args); + const restored = selectedMaskFiles(imageWidget); + if (!Array.isArray(imageWidget.value)) setSelectedMasks(restored); + updateMaskSelection(selectedMaskWidget, restored); + updateRemoveAvailability(removeWidget, restored.length); + if (restored.length > 0) { + preview.refresh(restored, selectedChannel(channelWidget)); + } + return result; + }; + const initial = selectedMaskFiles(imageWidget); + if (!Array.isArray(imageWidget.value)) setSelectedMasks(initial); + updateMaskSelection(selectedMaskWidget, initial); + updateRemoveAvailability(removeWidget, initial.length); + if (initial.length > 0) { + preview.refresh(initial, selectedChannel(channelWidget)); + } +} +function hideInternalWidget(widget) { + widget.options ??= {}; + widget.options.hidden = true; + widget.hidden = true; + widget.computeSize = () => [0, -4]; +} +function updateMaskSelection(widget, files, preferredIndex = 0) { + const values = files.map( + (file, index) => `${String(index + 1)}. ${file}` + ); + const hasMasks = values.length > 0; + widget.options ??= {}; + widget.options.values = hasMasks ? values : [EMPTY_MASK_SELECTION_LABEL]; + widget.disabled = !hasMasks; + widget.options.disabled = !hasMasks; + widget.value = hasMasks ? values[Math.min(Math.max(preferredIndex, 0), values.length - 1)] : EMPTY_MASK_SELECTION_LABEL; +} +function selectedMaskIndex(widget, callbackValue) { + if (widget.disabled) return null; + const value = callbackValue ?? widget.value; + const values = widget.options?.values ?? []; + const index = typeof value === "string" ? values.indexOf(value) : -1; + return index >= 0 ? index : null; +} +function clearNativePreview(node, app2) { + node.imgs = void 0; + node.images = []; + if (node.id !== void 0) delete app2?.nodeOutputs?.[String(node.id)]; + node.graph?.setDirtyCanvas?.(true, true); +} +function activeMaskIndex(node, selectedMaskWidget, files) { + const galleryIndex = node.imageIndex; + if (galleryIndex !== null && galleryIndex !== void 0 && galleryIndex >= 0 && galleryIndex < files.length) { + return galleryIndex; + } + const selectedIndex = selectedMaskIndex(selectedMaskWidget); + return selectedIndex !== null && selectedIndex < files.length ? selectedIndex : null; +} +function findWidget(node, name) { + return node.widgets?.find((widget) => widget.name === name); +} +function findNativeUploadWidget(node) { + return node.widgets?.find( + (widget) => widget.type === "button" && widget.value === "image" && widget.options?.serialize === false && widget.options.canvasOnly === true + ); +} +function nativeButtonOptions(tooltip) { + return { serialize: false, tooltip }; +} +function selectedChannel(channelWidget, callbackValue) { + const value = callbackValue ?? channelWidget.value; + return typeof value === "string" && value.length > 0 ? value : void 0; +} +function selectedMaskFiles(widget) { + return normalizeMaskFiles(widget.value); +} +function normalizeMaskFiles(value) { + const values = Array.isArray(value) ? value : [value]; + return values.filter( + (item) => typeof item === "string" && item.length > 0 + ); +} +function sameMaskFiles(left, right) { + return left.length === right.length && left.every((file, index) => file === right[index]); +} +function updateRemoveAvailability(widget, count) { + const disabled = count === 0; + widget.disabled = disabled; + widget.options ??= {}; + widget.options.disabled = disabled; +} +function notifySelectionChanged(node, imageWidget, change) { + node.onWidgetChanged?.( + imageWidget.name, + [...change.current], + [...change.previous], + imageWidget + ); + node.graph?.setDirtyCanvas?.(true, true); +} +function resetIntentForExternalUploads(node, reset) { + const originalPasteFiles = node.pasteFiles; + const originalOnDragDrop = node.onDragDrop; + const originalOnRemoved = node.onRemoved; + const wrappedPasteFiles = originalPasteFiles ? (...args) => { + reset(); + return originalPasteFiles.apply(node, args); + } : void 0; + const wrappedOnDragDrop = originalOnDragDrop ? (...args) => { + reset(); + return originalOnDragDrop.apply(node, args); + } : void 0; + if (wrappedPasteFiles) node.pasteFiles = wrappedPasteFiles; + if (wrappedOnDragDrop) node.onDragDrop = wrappedOnDragDrop; + node.onRemoved = function(...args) { + if (node.pasteFiles === wrappedPasteFiles) { + if (originalPasteFiles) node.pasteFiles = originalPasteFiles; + else delete node.pasteFiles; + } + if (node.onDragDrop === wrappedOnDragDrop) { + if (originalOnDragDrop) node.onDragDrop = originalOnDragDrop; + else delete node.onDragDrop; + } + return originalOnRemoved?.apply(this, args); + }; +} +function isMaskBatchNode(candidate) { + if (typeof candidate !== "object" || candidate === null) return false; + const node = candidate; + return node.constructor?.comfyClass === LOAD_MASK_BATCH_NODE_ID && typeof node.addWidget === "function"; +} + // web/src/main.ts var comfyApp = app; +var comfyExecutionEvents = window.comfyAPI.api.api; comfyApp.registerExtension({ name: "SimpleSyrup.Settings", async setup(appInstance) { @@ -519,3 +905,4 @@ comfyApp.registerExtension({ registerExternalLLMRefreshHook(appInstance); } }); +registerMaskBatchUpload(comfyApp, comfyExecutionEvents); diff --git a/web/src/api.ts b/web/src/api.ts index 7dc44c8..22c5094 100644 --- a/web/src/api.ts +++ b/web/src/api.ts @@ -2,6 +2,8 @@ // Copyright (C) 2026 Artificial Sweetener and contributors // SPDX-License-Identifier: AGPL-3.0-or-later +import type { ComfyImageResult, ComfyNodeExecutionOutput } from "./types"; + export interface SimpleSyrupSettings { show_downloadable_models: boolean; } @@ -32,6 +34,42 @@ const EXTERNAL_LLM_SETTINGS_ROUTE = "/simple-syrup/external-llm/settings"; const EXTERNAL_LLM_API_KEY_ROUTE = "/simple-syrup/external-llm/api-key"; const EXTERNAL_LLM_MODELS_REFRESH_ROUTE = "/simple-syrup/external-llm/models/refresh"; +const MASK_BATCH_PREVIEW_ROUTE = "/simple-syrup/mask-batch/preview"; + +export async function getMaskBatchPreview( + files: string[], + channel: string, + fetchImpl: FetchLike = fetch +): Promise { + const response = await fetchImpl(MASK_BATCH_PREVIEW_ROUTE, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ files, channel }) + }); + if (!response.ok) { + throw new Error( + await backendErrorMessage( + response, + `Could not render Load Mask Batch preview. Backend returned ${String(response.status)}.` + ) + ); + } + return parseMaskBatchPreview(await response.json()); +} + +export function parseMaskBatchPreview( + payload: unknown +): ComfyNodeExecutionOutput { + if (!isMaskBatchPreviewPayload(payload)) { + throw new Error( + "SimpleSyrup mask batch preview payload is invalid. Expected native images and animation flags." + ); + } + return { + images: payload.images.map((image) => ({ ...image })), + animated: [...payload.animated] + }; +} export async function getSettings( fetchImpl: FetchLike = fetch @@ -210,6 +248,31 @@ function isExternalLLMSettingsPayload( ); } +function isMaskBatchPreviewPayload( + payload: unknown +): payload is Required { + if (typeof payload !== "object" || payload === null) return false; + const candidate = payload as Partial; + return ( + Array.isArray(candidate.images) && + candidate.images.every(isComfyImageResult) && + Array.isArray(candidate.animated) && + candidate.animated.every((value) => typeof value === "boolean") + ); +} + +function isComfyImageResult(value: unknown): value is ComfyImageResult { + if (typeof value !== "object" || value === null) return false; + const candidate = value as Partial; + return ( + typeof candidate.filename === "string" && + typeof candidate.subfolder === "string" && + (candidate.type === "input" || + candidate.type === "output" || + candidate.type === "temp") + ); +} + async function backendErrorMessage( response: Response, fallback: string diff --git a/web/src/main.ts b/web/src/main.ts index 30c66d1..8f62569 100644 --- a/web/src/main.ts +++ b/web/src/main.ts @@ -7,9 +7,16 @@ import { app } from "../../../scripts/app.js"; import { registerSimpleSyrupSettings } from "./settings"; import { registerExternalLLMRefreshHook } from "./refresh"; -import type { ComfyApp } from "./types"; +import { registerMaskBatchUpload } from "./maskBatchUpload"; +import type { ComfyApp, ComfyExecutionEvents } from "./types"; + +interface ComfyRuntimeWindow extends Window { + comfyAPI: { api: { api: ComfyExecutionEvents } }; +} const comfyApp = app as unknown as ComfyApp; +const comfyExecutionEvents = (window as unknown as ComfyRuntimeWindow).comfyAPI + .api.api; comfyApp.registerExtension({ name: "SimpleSyrup.Settings", @@ -18,3 +25,5 @@ comfyApp.registerExtension({ registerExternalLLMRefreshHook(appInstance); } }); + +registerMaskBatchUpload(comfyApp, comfyExecutionEvents); diff --git a/web/src/maskBatchPreview.ts b/web/src/maskBatchPreview.ts new file mode 100644 index 0000000..1b0383c --- /dev/null +++ b/web/src/maskBatchPreview.ts @@ -0,0 +1,73 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import { getMaskBatchPreview } from "./api"; +import type { + ComfyExecutionEvents, + ComfyNodeExecutionOutput, + Logger +} from "./types"; + +const EMPTY_PREVIEW: ComfyNodeExecutionOutput = { + images: [], + animated: [] +}; +const PREVIEW_PROMPT_ID = "simple-syrup-mask-batch-preview"; + +export type MaskBatchPreviewClient = ( + files: string[], + channel: string +) => Promise; + +export interface MaskBatchPreviewNode { + id?: string | number; + graph?: { setDirtyCanvas?: (foreground: boolean, background: boolean) => void }; +} + +/** Coordinate asynchronous native previews without publishing stale results. */ +export class MaskBatchPreviewController { + private requestVersion = 0; + + constructor( + private readonly executionEvents: ComfyExecutionEvents, + private readonly node: MaskBatchPreviewNode, + private readonly loadPreview: MaskBatchPreviewClient = getMaskBatchPreview, + private readonly logger: Logger = console, + private readonly clearNativePreview: () => void = () => undefined + ) {} + + /** Replace the visible preview with the exact selected-channel mask output. */ + refresh(files: string[], channel: string | undefined): void { + if (this.node.id === undefined || channel === undefined) return; + + const requestVersion = ++this.requestVersion; + this.clearNativePreview(); + this.publish(EMPTY_PREVIEW); + if (files.length === 0) return; + void this.loadPreview([...files], channel) + .then((output) => { + if (requestVersion === this.requestVersion) this.publish(output); + }) + .catch((error: unknown) => { + if (requestVersion !== this.requestVersion) return; + this.logger.warn("Could not refresh Load Mask Batch preview.", error); + }); + } + + /** Publish an execution-shaped result through Comfy's native preview. */ + private publish(output: ComfyNodeExecutionOutput): void { + if (this.node.id === undefined) return; + this.executionEvents.dispatchEvent( + new CustomEvent("executed", { + detail: { + node: String(this.node.id), + output, + merge: false, + prompt_id: PREVIEW_PROMPT_ID + } + }) + ); + this.node.graph?.setDirtyCanvas?.(true, true); + } +} diff --git a/web/src/maskBatchUpload.ts b/web/src/maskBatchUpload.ts new file mode 100644 index 0000000..9fe4ec0 --- /dev/null +++ b/web/src/maskBatchUpload.ts @@ -0,0 +1,486 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import { + MaskBatchPreviewController, + type MaskBatchPreviewClient +} from "./maskBatchPreview"; +import type { + ComfyApp, + ComfyExecutionEvents, + ComfyExtension, + ComfyImageResult, + Logger +} from "./types"; + +const LOAD_MASK_BATCH_NODE_ID = "SimpleSyrup.LoadMaskBatch"; +const REPLACE_MASKS_LABEL = "Replace masks..."; +const ADD_MASKS_LABEL = "Add masks..."; +const REMOVE_MASK_LABEL = "Remove selected mask"; +const EMPTY_MASK_SELECTION_LABEL = "No masks loaded"; + +type SelectionIntent = "append" | "replace"; +type NativeCallback = (...args: unknown[]) => unknown; + +interface MaskBatchWidget { + name: string; + value: unknown; + type?: string; + label?: string; + callback?: (value?: unknown) => void; + computeSize?: (width?: number) => [number, number]; + hidden?: boolean; + disabled?: boolean; + options?: { + canvasOnly?: boolean; + disabled?: boolean; + serialize?: boolean; + tooltip?: string; + hidden?: boolean; + values?: string[]; + }; +} + +interface MaskBatchNode { + constructor: { comfyClass?: string }; + id?: string | number; + widgets?: MaskBatchWidget[]; + imageIndex?: number | null; + images?: ComfyImageResult[]; + imgs?: unknown[] | undefined; + pasteFiles?: NativeCallback; + onDragDrop?: NativeCallback; + onRemoved?: NativeCallback; + onGraphConfigured?: NativeCallback; + onExecuted?: NativeCallback; + onWidgetChanged?: ( + name: string, + value: unknown, + previousValue: unknown, + widget: MaskBatchWidget + ) => void; + graph?: { setDirtyCanvas?: (foreground: boolean, background: boolean) => void }; + addWidget( + type: "button" | "combo", + name: string, + value: string | undefined, + callback: (value?: unknown) => void, + options: MaskBatchWidget["options"] + ): MaskBatchWidget; +} + +/** Register native upload-list controls for the mask batch loader. */ +export function registerMaskBatchUpload( + app: ComfyApp, + executionEvents: ComfyExecutionEvents, + loadPreview?: MaskBatchPreviewClient, + logger: Logger = console +): void { + const extension: ComfyExtension = { + name: "SimpleSyrup.LoadMaskBatchUpload", + nodeCreated(node: unknown) { + try { + configureMaskBatchNode(node, executionEvents, loadPreview, logger, app); + } catch (error: unknown) { + logger.warn( + `Could not configure Load Mask Batch native controls: ${errorMessage(error)}`, + error + ); + throw error; + } + } + }; + app.registerExtension(extension); +} + +/** Return a stable diagnostic message for an unknown frontend failure. */ +function errorMessage(error: unknown): string { + return error instanceof Error && error.message + ? error.message + : String(error); +} + +/** Wire native Comfy upload, append, remove, persistence, and preview behavior. */ +export function configureMaskBatchNode( + candidate: unknown, + executionEvents: ComfyExecutionEvents, + loadPreview?: MaskBatchPreviewClient, + logger: Logger = console, + app?: ComfyApp +): void { + if (!isMaskBatchNode(candidate)) return; + + const imageWidget = findWidget(candidate, "image"); + const channelWidget = findWidget(candidate, "channel"); + const uploadWidget = findNativeUploadWidget(candidate); + if (!imageWidget || !channelWidget || !uploadWidget?.callback) return; + + hideInternalWidget(imageWidget); + hideInternalWidget(uploadWidget); + + const preview = new MaskBatchPreviewController( + executionEvents, + candidate, + loadPreview, + logger, + () => { + clearNativePreview(candidate, app); + } + ); + let selectionIntent: SelectionIntent = "replace"; + let appendBase: string[] = []; + let uploadPending = false; + let programmaticSelectionUpdate = false; + let ignoredUploadCallbackFiles: string[] | undefined; + const nativeUploadCallback = uploadWidget.callback; + + const replaceWidget = candidate.addWidget( + "button", + "simple_syrup_replace_masks", + "image", + () => { + selectionIntent = "replace"; + appendBase = []; + uploadPending = true; + nativeUploadCallback.call(uploadWidget); + }, + nativeButtonOptions( + "Choose one or more masks and replace the current ordered list." + ) + ); + replaceWidget.label = REPLACE_MASKS_LABEL; + + const setSelectedMasks = (files: string[]): void => { + programmaticSelectionUpdate = true; + try { + imageWidget.value = [...files]; + } finally { + programmaticSelectionUpdate = false; + } + }; + + uploadWidget.callback = (value?: unknown) => { + selectionIntent = "replace"; + appendBase = []; + uploadPending = true; + nativeUploadCallback.call(uploadWidget, value); + }; + + const addWidget = candidate.addWidget( + "button", + "simple_syrup_add_masks", + "image", + () => { + selectionIntent = "append"; + appendBase = selectedMaskFiles(imageWidget); + uploadPending = true; + nativeUploadCallback.call(uploadWidget); + }, + nativeButtonOptions( + "Upload one or more masks and append them after the current ordered list." + ) + ); + addWidget.label = ADD_MASKS_LABEL; + + const selectedMaskWidget = candidate.addWidget( + "combo", + "simple_syrup_selected_mask", + EMPTY_MASK_SELECTION_LABEL, + (value?: unknown) => { + const selectedIndex = selectedMaskIndex(selectedMaskWidget, value); + candidate.imageIndex = selectedIndex; + candidate.graph?.setDirtyCanvas?.(true, true); + }, + { + serialize: false, + tooltip: + "Select one loaded mask by its ordered position for preview or removal.", + values: [] + } + ); + selectedMaskWidget.label = "selected mask"; + + const removeWidget = candidate.addWidget( + "button", + "simple_syrup_remove_mask", + "image", + () => { + const previous = selectedMaskFiles(imageWidget); + const index = activeMaskIndex(candidate, selectedMaskWidget, previous); + if (index === null) return; + const current = previous.toSpliced(index, 1); + + setSelectedMasks(current); + updateMaskSelection(selectedMaskWidget, current, index); + candidate.imageIndex = current.length === 0 ? null : Math.min(index, current.length - 1); + notifySelectionChanged(candidate, imageWidget, { current, previous }); + updateRemoveAvailability(removeWidget, current.length); + preview.refresh(current, selectedChannel(channelWidget)); + }, + nativeButtonOptions( + "Open a mask in the preview gallery, then remove that position from the loaded list." + ) + ); + removeWidget.label = REMOVE_MASK_LABEL; + + imageWidget.callback = (value?: unknown) => { + if (programmaticSelectionUpdate) return; + const callbackValue = value ?? imageWidget.value; + if (uploadPending && !Array.isArray(callbackValue)) return; + const incoming = normalizeMaskFiles(callbackValue); + if ( + ignoredUploadCallbackFiles && + sameMaskFiles(incoming, ignoredUploadCallbackFiles) + ) { + ignoredUploadCallbackFiles = undefined; + return; + } + ignoredUploadCallbackFiles = uploadPending ? [...incoming] : undefined; + uploadPending = false; + + const current = + selectionIntent === "append" ? [...appendBase, ...incoming] : incoming; + const appended = selectionIntent === "append"; + selectionIntent = "replace"; + appendBase = []; + + if (appended) setSelectedMasks(current); + candidate.imageIndex = null; + updateMaskSelection(selectedMaskWidget, current); + updateRemoveAvailability(removeWidget, current.length); + preview.refresh(current, selectedChannel(channelWidget)); + }; + + const originalChannelCallback = channelWidget.callback; + channelWidget.callback = (value?: unknown) => { + originalChannelCallback?.call(channelWidget, value); + const files = selectedMaskFiles(imageWidget); + if (files.length > 0) { + preview.refresh(files, selectedChannel(channelWidget, value)); + } + }; + + resetIntentForExternalUploads(candidate, () => { + selectionIntent = "replace"; + appendBase = []; + uploadPending = false; + }); + + const originalOnGraphConfigured = candidate.onGraphConfigured; + candidate.onGraphConfigured = function (...args: unknown[]): unknown { + const result = originalOnGraphConfigured?.apply(this, args); + const restored = selectedMaskFiles(imageWidget); + if (!Array.isArray(imageWidget.value)) setSelectedMasks(restored); + updateMaskSelection(selectedMaskWidget, restored); + updateRemoveAvailability(removeWidget, restored.length); + if (restored.length > 0) { + preview.refresh(restored, selectedChannel(channelWidget)); + } + return result; + }; + + const initial = selectedMaskFiles(imageWidget); + if (!Array.isArray(imageWidget.value)) setSelectedMasks(initial); + updateMaskSelection(selectedMaskWidget, initial); + updateRemoveAvailability(removeWidget, initial.length); + if (initial.length > 0) { + preview.refresh(initial, selectedChannel(channelWidget)); + } +} + +/** Keep serialized helper widgets out of both native node renderers. */ +function hideInternalWidget(widget: MaskBatchWidget): void { + widget.options ??= {}; + widget.options.hidden = true; + widget.hidden = true; + widget.computeSize = () => [0, -4]; +} + +/** Replace the native selector choices with ordered, duplicate-safe labels. */ +function updateMaskSelection( + widget: MaskBatchWidget, + files: string[], + preferredIndex = 0 +): void { + const values = files.map( + (file, index) => `${String(index + 1)}. ${file}` + ); + const hasMasks = values.length > 0; + widget.options ??= {}; + widget.options.values = hasMasks ? values : [EMPTY_MASK_SELECTION_LABEL]; + widget.disabled = !hasMasks; + widget.options.disabled = !hasMasks; + widget.value = hasMasks + ? values[Math.min(Math.max(preferredIndex, 0), values.length - 1)] + : EMPTY_MASK_SELECTION_LABEL; +} + +/** Resolve the selected ordered position from the native list widget. */ +function selectedMaskIndex( + widget: MaskBatchWidget, + callbackValue?: unknown +): number | null { + if (widget.disabled) return null; + const value = callbackValue ?? widget.value; + const values = widget.options?.values ?? []; + const index = typeof value === "string" ? values.indexOf(value) : -1; + return index >= 0 ? index : null; +} + +/** Clear classic and Nodes 2.0 preview state before publishing new output. */ +function clearNativePreview(node: MaskBatchNode, app?: ComfyApp): void { + node.imgs = undefined; + node.images = []; + if (node.id !== undefined) delete app?.nodeOutputs?.[String(node.id)]; + node.graph?.setDirtyCanvas?.(true, true); +} + +/** Prefer Comfy's active gallery image, then the native list selection. */ +function activeMaskIndex( + node: MaskBatchNode, + selectedMaskWidget: MaskBatchWidget, + files: string[] +): number | null { + const galleryIndex = node.imageIndex; + if ( + galleryIndex !== null && + galleryIndex !== undefined && + galleryIndex >= 0 && + galleryIndex < files.length + ) { + return galleryIndex; + } + const selectedIndex = selectedMaskIndex(selectedMaskWidget); + return selectedIndex !== null && selectedIndex < files.length + ? selectedIndex + : null; +} + +/** Return one native widget by its stable input name. */ +function findWidget( + node: MaskBatchNode, + name: string +): MaskBatchWidget | undefined { + return node.widgets?.find((widget) => widget.name === name); +} + +/** Find Comfy's image upload button without depending on localized labels. */ +function findNativeUploadWidget( + node: MaskBatchNode +): MaskBatchWidget | undefined { + return node.widgets?.find( + (widget) => + widget.type === "button" && + widget.value === "image" && + widget.options?.serialize === false && + widget.options.canvasOnly === true + ); +} + +/** Return options for a native button rendered by classic and Nodes 2.0. */ +function nativeButtonOptions(tooltip: string): MaskBatchWidget["options"] { + return { serialize: false, tooltip }; +} + +/** Return the newly selected native channel when it is usable. */ +function selectedChannel( + channelWidget: MaskBatchWidget, + callbackValue?: unknown +): string | undefined { + const value = callbackValue ?? channelWidget.value; + return typeof value === "string" && value.length > 0 ? value : undefined; +} + +/** Return the native multi-value widget as an ordered file list. */ +function selectedMaskFiles(widget: MaskBatchWidget): string[] { + return normalizeMaskFiles(widget.value); +} + +/** Narrow native scalar or multi-value widget state to valid file paths. */ +function normalizeMaskFiles(value: unknown): string[] { + const values = Array.isArray(value) ? value : [value]; + return values.filter( + (item): item is string => typeof item === "string" && item.length > 0 + ); +} + +/** Compare ordered upload results without relying on reactive proxy identity. */ +function sameMaskFiles(left: string[], right: string[]): boolean { + return ( + left.length === right.length && + left.every((file, index) => file === right[index]) + ); +} + +/** Disable removal only when the ordered mask list is empty. */ +function updateRemoveAvailability(widget: MaskBatchWidget, count: number): void { + const disabled = count === 0; + widget.disabled = disabled; + widget.options ??= {}; + widget.options.disabled = disabled; +} + +/** Notify Comfy when a non-upload control mutates persisted widget state. */ +function notifySelectionChanged( + node: MaskBatchNode, + imageWidget: MaskBatchWidget, + change: { current: string[]; previous: string[] } +): void { + node.onWidgetChanged?.( + imageWidget.name, + [...change.current], + [...change.previous], + imageWidget + ); + node.graph?.setDirtyCanvas?.(true, true); +} + +/** Prevent a cancelled Add action from affecting later paste or drop uploads. */ +function resetIntentForExternalUploads( + node: MaskBatchNode, + reset: () => void +): void { + const originalPasteFiles = node.pasteFiles; + const originalOnDragDrop = node.onDragDrop; + const originalOnRemoved = node.onRemoved; + + const wrappedPasteFiles: NativeCallback | undefined = originalPasteFiles + ? (...args: unknown[]): unknown => { + reset(); + return originalPasteFiles.apply(node, args); + } + : undefined; + const wrappedOnDragDrop: NativeCallback | undefined = originalOnDragDrop + ? (...args: unknown[]): unknown => { + reset(); + return originalOnDragDrop.apply(node, args); + } + : undefined; + + if (wrappedPasteFiles) node.pasteFiles = wrappedPasteFiles; + if (wrappedOnDragDrop) node.onDragDrop = wrappedOnDragDrop; + node.onRemoved = function (...args: unknown[]): unknown { + if (node.pasteFiles === wrappedPasteFiles) { + if (originalPasteFiles) node.pasteFiles = originalPasteFiles; + else delete node.pasteFiles; + } + if (node.onDragDrop === wrappedOnDragDrop) { + if (originalOnDragDrop) node.onDragDrop = originalOnDragDrop; + else delete node.onDragDrop; + } + return originalOnRemoved?.apply(this, args); + }; +} + +/** Return true when the candidate is the native-widget batch mask loader. */ +function isMaskBatchNode(candidate: unknown): candidate is MaskBatchNode { + if (typeof candidate !== "object" || candidate === null) return false; + const node = candidate as Partial; + return ( + node.constructor?.comfyClass === LOAD_MASK_BATCH_NODE_ID && + typeof node.addWidget === "function" + ); +} + +export type { MaskBatchPreviewClient } from "./maskBatchPreview"; diff --git a/web/src/types.ts b/web/src/types.ts index 4eb8286..628c659 100644 --- a/web/src/types.ts +++ b/web/src/types.ts @@ -27,6 +27,7 @@ export interface ComfySettingsApi { } export interface ComfyApp { + nodeOutputs?: Record; ui: { settings: ComfySettingsApi; }; @@ -34,7 +35,22 @@ export interface ComfyApp { registerExtension(extension: ComfyExtension): void; } +/** Native Comfy event target used to publish execution-shaped node output. */ +export type ComfyExecutionEvents = Pick; + +export interface ComfyImageResult { + filename: string; + subfolder: string; + type: "input" | "output" | "temp"; +} + +export interface ComfyNodeExecutionOutput { + images?: ComfyImageResult[]; + animated?: boolean[]; +} + export interface ComfyExtension { name: string; - setup(app: ComfyApp): void | Promise; + setup?(app: ComfyApp): void | Promise; + nodeCreated?(node: unknown): void | Promise; } diff --git a/web/tests/api.test.ts b/web/tests/api.test.ts index 12d6d4a..c5053d7 100644 --- a/web/tests/api.test.ts +++ b/web/tests/api.test.ts @@ -7,8 +7,10 @@ import { describe, expect, it, vi } from "vitest"; import { deleteExternalLLMApiKey, getExternalLLMSettings, + getMaskBatchPreview, getSettings, parseExternalLLMSettings, + parseMaskBatchPreview, parseSettings, refreshExternalLLMModels, saveExternalLLMApiKey, @@ -73,6 +75,62 @@ describe("settings API", () => { }); }); +describe("mask batch preview API", () => { + const preview = { + images: [ + { + filename: "ComfyUI_temp_mask.png", + subfolder: "", + type: "temp" as const + } + ], + animated: [false] + }; + + it("requests previews for the ordered files and selected channel", async () => { + const fetchImpl = vi + .fn() + .mockResolvedValue(createJsonResponse(preview)); + + await expect( + getMaskBatchPreview(["right.png", "left.png"], "blue", fetchImpl) + ).resolves.toEqual(preview); + expect(fetchImpl).toHaveBeenCalledWith( + "/simple-syrup/mask-batch/preview", + expect.objectContaining({ + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + files: ["right.png", "left.png"], + channel: "blue" + }) + }) + ); + }); + + it("surfaces preview validation errors", async () => { + const fetchImpl = vi.fn().mockResolvedValue( + createJsonResponse( + { error: "mask dimensions do not match" }, + { status: 400 } + ) + ); + + await expect( + getMaskBatchPreview(["right.png", "left.png"], "alpha", fetchImpl) + ).rejects.toThrow("mask dimensions do not match"); + }); + + it("rejects malformed native preview responses", () => { + expect(() => + parseMaskBatchPreview({ + images: [{ filename: "mask.png", subfolder: "", type: "other" }], + animated: [false] + }) + ).toThrow("mask batch preview payload is invalid"); + }); +}); + describe("external LLM settings API", () => { const payload = { base_url: "https://provider.example/v1", diff --git a/web/tests/maskBatchPreview.test.ts b/web/tests/maskBatchPreview.test.ts new file mode 100644 index 0000000..f03d879 --- /dev/null +++ b/web/tests/maskBatchPreview.test.ts @@ -0,0 +1,157 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import { describe, expect, it, vi } from "vitest"; + +import { + MaskBatchPreviewController, + type MaskBatchPreviewClient +} from "../src/maskBatchPreview"; +import type { + ComfyNodeExecutionOutput, + Logger +} from "../src/types"; + +const OUTPUT: ComfyNodeExecutionOutput = { + images: [{ filename: "mask.png", subfolder: "", type: "temp" }], + animated: [false] +}; + +function deferred(): { + promise: Promise; + resolve: (value: T) => void; +} { + let resolvePromise: ((value: T) => void) | undefined; + const promise = new Promise((resolve) => { + resolvePromise = resolve; + }); + return { + promise, + resolve(value: T) { + if (!resolvePromise) throw new Error("Deferred promise was not initialized."); + resolvePromise(value); + } + }; +} + +describe("MaskBatchPreviewController", () => { + it("clears without requesting a preview for an empty native selection", () => { + const executionEvents = new EventTarget(); + const published: ComfyNodeExecutionOutput[] = []; + executionEvents.addEventListener("executed", (event) => { + published.push( + (event as CustomEvent<{ output: ComfyNodeExecutionOutput }>).detail + .output + ); + }); + const client = vi.fn(); + const clearNativePreview = vi.fn(); + const controller = new MaskBatchPreviewController( + executionEvents, + { id: 3 }, + client, + console, + clearNativePreview + ); + + controller.refresh([], "alpha"); + + expect(published).toEqual([{ images: [], animated: [] }]); + expect(clearNativePreview).toHaveBeenCalledOnce(); + expect(client).not.toHaveBeenCalled(); + }); + + it("publishes through Comfy's native execution pipeline and clears first", async () => { + const node = { id: 4 }; + const executionEvents = new EventTarget(); + const published: ComfyNodeExecutionOutput[] = []; + executionEvents.addEventListener("executed", (event) => { + published.push( + (event as CustomEvent<{ output: ComfyNodeExecutionOutput }>).detail + .output + ); + }); + const client = vi.fn().mockResolvedValue(OUTPUT); + const controller = new MaskBatchPreviewController( + executionEvents, + node, + client + ); + + controller.refresh(["one.png"], "red"); + + expect(published).toEqual([{ images: [], animated: [] }]); + await vi.waitFor(() => { + expect(published.at(-1)).toEqual(OUTPUT); + }); + expect(client).toHaveBeenCalledWith(["one.png"], "red"); + }); + + it("discards responses made stale by a newer channel request", async () => { + const first = deferred(); + const second = deferred(); + const client = vi + .fn() + .mockReturnValueOnce(first.promise) + .mockReturnValueOnce(second.promise); + const executionEvents = new EventTarget(); + const published: ComfyNodeExecutionOutput[] = []; + executionEvents.addEventListener("executed", (event) => { + published.push( + (event as CustomEvent<{ output: ComfyNodeExecutionOutput }>).detail + .output + ); + }); + const controller = new MaskBatchPreviewController( + executionEvents, + { id: 5 }, + client + ); + const latest: ComfyNodeExecutionOutput = { + images: [{ filename: "blue.png", subfolder: "", type: "temp" }], + animated: [false] + }; + + controller.refresh(["mask.png"], "red"); + controller.refresh(["mask.png"], "blue"); + second.resolve(latest); + await vi.waitFor(() => { + expect(published.at(-1)).toEqual(latest); + }); + first.resolve(OUTPUT); + await Promise.resolve(); + + expect(published.at(-1)).toEqual(latest); + }); + + it("keeps the cleared preview and reports current request failures", async () => { + const error = new Error("preview failed"); + const client = vi.fn().mockRejectedValue(error); + const logger: Logger = { warn: vi.fn() }; + const executionEvents = new EventTarget(); + const published: ComfyNodeExecutionOutput[] = []; + executionEvents.addEventListener("executed", (event) => { + published.push( + (event as CustomEvent<{ output: ComfyNodeExecutionOutput }>).detail + .output + ); + }); + const controller = new MaskBatchPreviewController( + executionEvents, + { id: 6 }, + client, + logger + ); + + controller.refresh(["mask.png"], "alpha"); + + await vi.waitFor(() => { + expect(logger.warn).toHaveBeenCalledWith( + "Could not refresh Load Mask Batch preview.", + error + ); + }); + expect(published).toEqual([{ images: [], animated: [] }]); + }); +}); diff --git a/web/tests/maskBatchUpload.test.ts b/web/tests/maskBatchUpload.test.ts new file mode 100644 index 0000000..b49c121 --- /dev/null +++ b/web/tests/maskBatchUpload.test.ts @@ -0,0 +1,500 @@ +// SimpleSyrup - workflow-focused ComfyUI extensions for image generation +// Copyright (C) 2026 Artificial Sweetener and contributors +// SPDX-License-Identifier: AGPL-3.0-or-later + +import { describe, expect, it, vi } from "vitest"; + +import { + configureMaskBatchNode, + registerMaskBatchUpload +} from "../src/maskBatchUpload"; +import type { MaskBatchPreviewClient } from "../src/maskBatchUpload"; +import type { + ComfyApp, + ComfyExtension, + ComfyNodeExecutionOutput +} from "../src/types"; + +interface TestWidget { + name: string; + value: unknown; + type?: string; + label?: string; + callback?: (value?: unknown) => unknown; + computeSize?: (width?: number) => [number, number]; + hidden?: boolean; + serializeValue?: () => unknown; + disabled?: boolean; + options?: { + canvasOnly?: boolean; + disabled?: boolean; + serialize?: boolean; + tooltip?: string; + hidden?: boolean; + values?: string[]; + }; +} + +const PREVIEW_OUTPUT: ComfyNodeExecutionOutput = { + images: [ + { + filename: "ComfyUI_temp_mask.png", + subfolder: "", + type: "temp" + } + ], + animated: [false] +}; + +function createNode(value: unknown = []) { + const nativeUploadCallback = vi.fn(); + const channelCallback = vi.fn(); + const onWidgetChanged = vi.fn(); + const setDirtyCanvas = vi.fn(); + const pasteFiles = vi.fn(() => true); + const onDragDrop = vi.fn(() => Promise.resolve(true)); + const originalOnRemoved = vi.fn(); + const imageWidget: TestWidget = { + name: "image", + value, + type: "combo", + options: { canvasOnly: true } + }; + const channelWidget: TestWidget = { + name: "channel", + value: "alpha", + type: "combo", + callback: channelCallback + }; + const uploadWidget: TestWidget = { + name: "upload", + value: "image", + type: "button", + label: "choose file to upload", + callback: nativeUploadCallback, + options: { serialize: false, canvasOnly: true } + }; + const widgets = [imageWidget, channelWidget, uploadWidget]; + const node = { + constructor: { comfyClass: "SimpleSyrup.LoadMaskBatch" }, + id: 7, + widgets, + properties: {} as Record, + imageIndex: null as number | null, + images: undefined as ComfyNodeExecutionOutput["images"], + imgs: undefined as unknown[] | undefined, + pasteFiles, + onDragDrop, + onRemoved: originalOnRemoved as (...args: unknown[]) => unknown, + onGraphConfigured: undefined as ((...args: unknown[]) => unknown) | undefined, + onWidgetChanged, + graph: { setDirtyCanvas }, + addWidget( + type: "button", + name: string, + widgetValue: string | undefined, + callback: (value?: unknown) => void, + options: TestWidget["options"] + ): TestWidget { + const widget: TestWidget = { + type, + name, + value: widgetValue, + callback + }; + if (options) widget.options = options; + widgets.push(widget); + return widget; + } + }; + return { + channelCallback, + channelWidget, + imageWidget, + nativeUploadCallback, + node, + onDragDrop, + onWidgetChanged, + originalOnRemoved, + pasteFiles, + setDirtyCanvas, + uploadWidget + }; +} + +function createExecutionEvents(): EventTarget { + return new EventTarget(); +} + +function previewClient( + output: ComfyNodeExecutionOutput = PREVIEW_OUTPUT +): ReturnType> { + return vi.fn().mockResolvedValue(output); +} + +function button(node: ReturnType["node"], name: string) { + const widget = node.widgets.find((candidate) => candidate.name === name); + if (!widget) throw new Error(`Missing test button ${name}.`); + return widget; +} + +function selectFiles(imageWidget: TestWidget, files: string[]): void { + imageWidget.value = files; + imageWidget.callback?.(files); +} + +describe("Load Mask Batch native upload integration", () => { + it("registers a node-created hook", () => { + let extension: ComfyExtension | undefined; + const app = { + registerExtension(value: ComfyExtension) { + extension = value; + } + } as ComfyApp; + + registerMaskBatchUpload(app, createExecutionEvents()); + + expect(extension?.name).toBe("SimpleSyrup.LoadMaskBatchUpload"); + expect(typeof extension?.nodeCreated).toBe("function"); + }); + + it("reports native node-configuration failures with context", () => { + let extension: ComfyExtension | undefined; + const logger = { warn: vi.fn() }; + const app = { + registerExtension(value: ComfyExtension) { + extension = value; + } + } as ComfyApp; + registerMaskBatchUpload(app, createExecutionEvents(), undefined, logger); + const { node } = createNode(); + node.addWidget = () => { + throw new Error("native widget failed"); + }; + + expect(() => extension?.nodeCreated?.(node)).toThrow( + "native widget failed" + ); + expect(logger.warn).toHaveBeenCalledWith( + "Could not configure Load Mask Batch native controls: native widget failed", + expect.any(Error) + ); + }); + + it("uses native cross-renderer controls and hides internal upload state", () => { + const { imageWidget, node, uploadWidget } = createNode(["only.png"]); + configureMaskBatchNode(node, createExecutionEvents()); + + const replace = button(node, "simple_syrup_replace_masks"); + const add = button(node, "simple_syrup_add_masks"); + const selected = button(node, "simple_syrup_selected_mask"); + const remove = button(node, "simple_syrup_remove_mask"); + expect(uploadWidget.label).toBe("choose file to upload"); + expect(imageWidget.options).toMatchObject({ + canvasOnly: true, + hidden: true + }); + expect(uploadWidget.options).toMatchObject({ + canvasOnly: true, + hidden: true + }); + expect(imageWidget.computeSize?.()).toEqual([0, -4]); + expect(imageWidget.hidden).toBe(true); + expect(uploadWidget.computeSize?.()).toEqual([0, -4]); + expect(replace.label).toBe("Replace masks..."); + expect(add.label).toBe("Add masks..."); + expect(selected.label).toBe("selected mask"); + expect(selected.value).toBe("1. only.png"); + expect(selected.options?.values).toEqual(["1. only.png"]); + expect(remove.label).toBe("Remove selected mask"); + expect(add.type).toBe("button"); + expect(add.options).toMatchObject({ serialize: false }); + expect(add.options?.canvasOnly).toBeUndefined(); + expect(remove.disabled).toBe(false); + expect(remove.options?.disabled).toBe(false); + }); + + it("shows an explicit disabled empty state instead of undefined", () => { + const { node } = createNode([]); + + configureMaskBatchNode(node, createExecutionEvents()); + + const selected = button(node, "simple_syrup_selected_mask"); + expect(selected.value).toBe("No masks loaded"); + expect(selected.options?.values).toEqual(["No masks loaded"]); + expect(selected.disabled).toBe(true); + expect(selected.options?.disabled).toBe(true); + expect(button(node, "simple_syrup_remove_mask").disabled).toBe(true); + }); + + it("uses the native multi-value widget as the only selection owner", async () => { + const { imageWidget, node } = createNode([]); + const loadPreview = previewClient(); + configureMaskBatchNode(node, createExecutionEvents(), loadPreview); + loadPreview.mockClear(); + + button(node, "simple_syrup_add_masks").callback?.(); + const uploaded = ["one.png", "two.png", "three.png"]; + selectFiles(imageWidget, uploaded); + + expect(imageWidget.value).toEqual(uploaded); + expect(imageWidget.serializeValue).toBeUndefined(); + expect(node.properties).toEqual({}); + await vi.waitFor(() => { + expect(loadPreview).toHaveBeenCalledWith(uploaded, "alpha"); + }); + }); + + it("accepts direct native list edits and clears without a backend request", async () => { + const { imageWidget, node } = createNode(["one.png", "two.png"]); + const loadPreview = previewClient(); + configureMaskBatchNode(node, createExecutionEvents(), loadPreview); + loadPreview.mockClear(); + + selectFiles(imageWidget, ["two.png"]); + await vi.waitFor(() => { + expect(loadPreview).toHaveBeenCalledWith(["two.png"], "alpha"); + }); + + loadPreview.mockClear(); + selectFiles(imageWidget, []); + + expect(imageWidget.value).toEqual([]); + expect(loadPreview).not.toHaveBeenCalled(); + expect(button(node, "simple_syrup_remove_mask").disabled).toBe(true); + }); + + it("normalizes a serialized scalar from an older workflow", async () => { + const { imageWidget, node } = createNode("saved-mask.png"); + const loadPreview = previewClient(); + configureMaskBatchNode(node, createExecutionEvents(), loadPreview); + loadPreview.mockClear(); + + node.onGraphConfigured?.(); + + expect(imageWidget.value).toEqual(["saved-mask.png"]); + await vi.waitFor(() => { + expect(loadPreview).toHaveBeenCalledWith(["saved-mask.png"], "alpha"); + }); + }); + + it("replaces the ordered list through the explicit native control", async () => { + const { imageWidget, nativeUploadCallback, node } = createNode([ + "old.png" + ]); + const loadPreview = previewClient(); + configureMaskBatchNode(node, createExecutionEvents(), loadPreview); + loadPreview.mockClear(); + + button(node, "simple_syrup_replace_masks").callback?.(); + const replacement = ["right.png", "left.png"]; + selectFiles(imageWidget, replacement); + + expect(nativeUploadCallback).toHaveBeenCalledOnce(); + expect(imageWidget.value).toEqual(replacement); + expect(imageWidget.serializeValue).toBeUndefined(); + expect(node.properties).toEqual({}); + await vi.waitFor(() => { + expect(loadPreview).toHaveBeenCalledWith(replacement, "alpha"); + }); + }); + + it("appends native upload results in order and preserves duplicate paths", async () => { + const { imageWidget, nativeUploadCallback, node } = createNode(); + const loadPreview = previewClient(); + configureMaskBatchNode(node, createExecutionEvents(), loadPreview); + selectFiles(imageWidget, ["one.png", "duplicate.png"]); + loadPreview.mockClear(); + + button(node, "simple_syrup_add_masks").callback?.(); + imageWidget.value = "two.png"; + imageWidget.callback?.("two.png"); + const uploaded = ["two.png", "duplicate.png"]; + selectFiles(imageWidget, uploaded); + + const expected = [ + "one.png", + "duplicate.png", + "two.png", + "duplicate.png" + ]; + expect(nativeUploadCallback).toHaveBeenCalledOnce(); + expect(uploaded).toEqual(["two.png", "duplicate.png"]); + expect(imageWidget.value).toEqual(expected); + expect(imageWidget.serializeValue).toBeUndefined(); + expect(node.properties).toEqual({}); + await vi.waitFor(() => { + expect(loadPreview).toHaveBeenCalledWith(expected, "alpha"); + }); + }); + + it("avoids recursive updates from the Nodes 2.0 reactive multiselect", async () => { + const { imageWidget, node } = createNode(["one.png"]); + let storedValue = imageWidget.value; + Object.defineProperty(imageWidget, "value", { + configurable: true, + get: () => storedValue, + set: (value: unknown) => { + storedValue = value; + imageWidget.callback?.(value); + } + }); + const loadPreview = previewClient(); + configureMaskBatchNode(node, createExecutionEvents(), loadPreview); + loadPreview.mockClear(); + + button(node, "simple_syrup_add_masks").callback?.(); + const uploaded = ["two.png", "three.png"]; + imageWidget.value = uploaded; + imageWidget.callback?.([...uploaded]); + + expect(imageWidget.value).toEqual([ + "one.png", + "two.png", + "three.png" + ]); + await vi.waitFor(() => { + expect(loadPreview).toHaveBeenCalledTimes(1); + }); + }); + + it("cancelled Add cannot affect later replace, paste, or drop selections", () => { + const { imageWidget, node, onDragDrop, pasteFiles, uploadWidget } = + createNode(); + configureMaskBatchNode(node, createExecutionEvents()); + selectFiles(imageWidget, ["existing.png"]); + + button(node, "simple_syrup_add_masks").callback?.(); + uploadWidget.callback?.(); + selectFiles(imageWidget, ["replace.png"]); + expect(imageWidget.value).toEqual(["replace.png"]); + + button(node, "simple_syrup_add_masks").callback?.(); + node.pasteFiles(); + selectFiles(imageWidget, ["paste.png"]); + expect(imageWidget.value).toEqual(["paste.png"]); + expect(pasteFiles).toHaveBeenCalledOnce(); + + button(node, "simple_syrup_add_masks").callback?.(); + void node.onDragDrop(); + selectFiles(imageWidget, ["drop.png"]); + expect(imageWidget.value).toEqual(["drop.png"]); + expect(onDragDrop).toHaveBeenCalledOnce(); + }); + + it("removes the active gallery position and never deletes by filename", async () => { + const { imageWidget, node, onWidgetChanged } = createNode(); + const loadPreview = previewClient(); + configureMaskBatchNode(node, createExecutionEvents(), loadPreview); + selectFiles(imageWidget, ["duplicate.png", "middle.png", "duplicate.png"]); + loadPreview.mockClear(); + onWidgetChanged.mockClear(); + node.imageIndex = 2; + + button(node, "simple_syrup_remove_mask").callback?.(); + + const remaining = ["duplicate.png", "middle.png"]; + expect(imageWidget.value).toEqual(remaining); + expect(node.imageIndex).toBe(1); + expect(onWidgetChanged).toHaveBeenCalledWith( + "image", + remaining, + ["duplicate.png", "middle.png", "duplicate.png"], + imageWidget + ); + await vi.waitFor(() => { + expect(loadPreview).toHaveBeenCalledWith(remaining, "alpha"); + }); + }); + + it("removes the native list selection and permits clearing the final mask", () => { + const { imageWidget, node, onWidgetChanged } = createNode(); + configureMaskBatchNode(node, createExecutionEvents()); + selectFiles(imageWidget, ["one.png", "two.png"]); + const selected = button(node, "simple_syrup_selected_mask"); + const remove = button(node, "simple_syrup_remove_mask"); + + selected.callback?.("2. two.png"); + remove.callback?.(); + expect(imageWidget.value).toEqual(["one.png"]); + expect(selected.value).toBe("1. one.png"); + expect(remove.disabled).toBe(false); + + remove.callback?.(); + expect(imageWidget.value).toEqual([]); + expect(selected.value).toBe("No masks loaded"); + expect(selected.options?.values).toEqual(["No masks loaded"]); + expect(selected.disabled).toBe(true); + expect(remove.disabled).toBe(true); + expect(onWidgetChanged).toHaveBeenCalledTimes(2); + }); + + it("clears stale native preview state when the final mask is removed", () => { + const { imageWidget, node } = createNode(["one.png"]); + const loadPreview = previewClient(); + const app = { + nodeOutputs: { "7": PREVIEW_OUTPUT }, + registerExtension: vi.fn(), + ui: { settings: { addSetting: vi.fn() } } + } as unknown as ComfyApp; + configureMaskBatchNode( + node, + createExecutionEvents(), + loadPreview, + console, + app + ); + node.imgs = [{}]; + node.images = PREVIEW_OUTPUT.images; + app.nodeOutputs = { "7": PREVIEW_OUTPUT }; + + selectFiles(imageWidget, []); + + expect(node.imgs).toBeUndefined(); + expect(node.images).toEqual([]); + expect(node.imageIndex).toBeNull(); + expect(app.nodeOutputs).toEqual({}); + }); + + it("restores a persisted batch and refreshes the selected channel", async () => { + const { channelWidget, imageWidget, node } = createNode([ + "one.png", + "two.png" + ]); + const loadPreview = previewClient(); + configureMaskBatchNode(node, createExecutionEvents(), loadPreview); + loadPreview.mockClear(); + channelWidget.value = "blue"; + + node.onGraphConfigured?.(); + + expect(imageWidget.value).toEqual(["one.png", "two.png"]); + await vi.waitFor(() => { + expect(loadPreview).toHaveBeenCalledWith( + ["one.png", "two.png"], + "blue" + ); + }); + }); + + it("restores native paste and drop handlers before native node cleanup", () => { + const { node, onDragDrop, originalOnRemoved, pasteFiles } = createNode(); + configureMaskBatchNode(node, createExecutionEvents()); + + node.onRemoved(); + + expect(originalOnRemoved).toHaveBeenCalledOnce(); + expect(node.pasteFiles).toBe(pasteFiles); + expect(node.onDragDrop).toBe(onDragDrop); + }); + + it("does not modify unrelated nodes", () => { + const { node, uploadWidget } = createNode(); + node.constructor.comfyClass = "LoadImageMask"; + + configureMaskBatchNode(node, createExecutionEvents()); + + expect(uploadWidget.label).toBe("choose file to upload"); + expect(node.widgets).toHaveLength(3); + }); +});