From 8104e81a29461f8955df75da59a737afa32a8ad3 Mon Sep 17 00:00:00 2001 From: Artificial Sweetener Date: Sat, 8 Aug 2026 10:35:24 -0400 Subject: [PATCH] fix(runtime): centralize Comfy patcher lifecycle --- simple_syrup/runtime/checkpoint_loader.py | 62 ++--- .../runtime/contextual_diffusion_sampling.py | 13 +- .../runtime/differential_diffusion.py | 26 +- .../runtime/mixture_of_diffusers_sampling.py | 29 ++- .../runtime/multidiffusion_sampling.py | 29 ++- simple_syrup/runtime/patcher_lifecycle.py | 216 ++++++++++++++++ simple_syrup/runtime/regional_lora_hooks.py | 20 +- .../regional_multidiffusion_sampling.py | 29 ++- simple_syrup/runtime/vae_loader.py | 7 +- tests/test_checkpoint_loader.py | 11 +- tests/test_comfy_patcher_lifecycle.py | 239 ++++++++++++++++++ tests/test_contextual_diffusion_sampling.py | 9 +- tests/test_patcher_lifecycle_policy.py | 200 +++++++++++++++ tests/test_regional_lora_hooks.py | 5 +- 14 files changed, 808 insertions(+), 87 deletions(-) create mode 100644 simple_syrup/runtime/patcher_lifecycle.py create mode 100644 tests/test_comfy_patcher_lifecycle.py create mode 100644 tests/test_patcher_lifecycle_policy.py diff --git a/simple_syrup/runtime/checkpoint_loader.py b/simple_syrup/runtime/checkpoint_loader.py index 2507df2..e395fd0 100644 --- a/simple_syrup/runtime/checkpoint_loader.py +++ b/simple_syrup/runtime/checkpoint_loader.py @@ -9,8 +9,13 @@ from __future__ import annotations import importlib from collections.abc import Iterable from types import ModuleType -from typing import Any, Protocol, cast, runtime_checkable +from typing import Any, Protocol, cast +from .patcher_lifecycle import ( + PATCHER_LIFECYCLE, + ClipLayerMutation, + ComfyPatcherLifecycle, +) from .vae_loader import VaeLoaderService USE_CHECKPOINT_VAE_CHOICE = "Use Checkpoint VAE" @@ -25,11 +30,13 @@ class CheckpointLoaderService: self, folder_paths_module: ModuleType | None = None, vae_loader: VaeLoaderBoundary | None = None, + patcher_lifecycle: ComfyPatcherLifecycle | None = None, ) -> None: """Create a checkpoint loader with injectable runtime boundaries.""" self._folder_paths_module = folder_paths_module self._vae_loader = vae_loader or VaeLoaderService(folder_paths_module) + self._patcher_lifecycle = patcher_lifecycle or PATCHER_LIFECYCLE def load_checkpoint( self, @@ -57,12 +64,32 @@ class CheckpointLoaderService: model = loaded[0] clip = loaded[1] checkpoint_vae = loaded[2] - selected_clip = _selected_clip(clip, clip_skip) + selected_clip = self._selected_clip(clip, clip_skip) if vae_name == USE_CHECKPOINT_VAE_CHOICE: - return model, selected_clip, checkpoint_vae + selected_vae = checkpoint_vae + else: + selected_vae = self._vae_loader.load_vae(vae_name) - return model, selected_clip, self._vae_loader.load_vae(vae_name) + return ( + model, + selected_clip, + self._patcher_lifecycle.preserve_vae( + selected_vae, + operation="SimpleSyrup checkpoint loading", + ), + ) + + def _selected_clip(self, clip: object, clip_skip: bool) -> object: + """Return the loaded CLIP or a lifecycle-owned clip-skip derivation.""" + + if not clip_skip: + return clip + return self._patcher_lifecycle.derive_clip( + clip, + (ClipLayerMutation(CLIP_SKIP_LAYER),), + operation="SimpleSyrup checkpoint clip skip", + ) def _folder_paths(self) -> ModuleType: """Return the ComfyUI folder_paths module.""" @@ -83,17 +110,6 @@ class VaeLoaderBoundary(Protocol): """Load the named external VAE.""" -@runtime_checkable -class ClipLayerBoundary(Protocol): - """CLIP interface required to apply the ComfyUI clip-skip layer.""" - - def clone(self) -> ClipLayerBoundary: - """Return an independent CLIP object.""" - - def clip_layer(self, layer_idx: int) -> None: - """Set the CLIP layer index used during prompt encoding.""" - - def _validate_clip_skip(clip_skip: object) -> None: """Reject non-boolean clip-skip selections from runtime callers.""" @@ -101,22 +117,6 @@ def _validate_clip_skip(clip_skip: object) -> None: raise TypeError("clip_skip must be a boolean.") -def _selected_clip(clip: object, clip_skip: bool) -> object: - """Return the loaded CLIP or a cloned CLIP with clip skip applied.""" - - if not clip_skip: - return clip - - if not isinstance(clip, ClipLayerBoundary): - raise TypeError( - "clip_skip requires a CLIP object with clone() and clip_layer()." - ) - - selected_clip = clip.clone() - selected_clip.clip_layer(CLIP_SKIP_LAYER) - return selected_clip - - def _comfy_sd() -> Any: """Import ComfyUI's stable diffusion loading module lazily.""" diff --git a/simple_syrup/runtime/contextual_diffusion_sampling.py b/simple_syrup/runtime/contextual_diffusion_sampling.py index 53401ce..82017c6 100644 --- a/simple_syrup/runtime/contextual_diffusion_sampling.py +++ b/simple_syrup/runtime/contextual_diffusion_sampling.py @@ -18,6 +18,7 @@ from ..domain.contextual_diffusion import ( ) from ..shared.logging import get_logger from . import sampling_samplers, sampling_schedulers +from .patcher_lifecycle import PATCHER_LIFECYCLE, ModelUnetWrapperMutation from .tiled_sampling import ( ApplyModel, Latent, @@ -156,10 +157,9 @@ def clone_model_with_contextual_diffusion( sigmas: torch.Tensor, diffusion_mode: str, ) -> Any: - """Clone a model and install one pre-CFG contextual prediction wrapper.""" + """Derive a model with one pre-CFG contextual prediction wrapper.""" - cloned_model = model.clone() - old_wrapper = cloned_model.model_options.get("model_function_wrapper") + old_wrapper = model.model_options.get("model_function_wrapper") if old_wrapper is not None and not callable(old_wrapper): raise ValueError("Existing model_function_wrapper is not callable.") wrapper = ContextualDiffusionModelWrapper( @@ -169,8 +169,11 @@ def clone_model_with_contextual_diffusion( diffusion_mode=diffusion_mode, existing_wrapper=cast(ModelFunctionWrapper | None, old_wrapper), ) - cloned_model.set_model_unet_function_wrapper(wrapper) - return cloned_model + return PATCHER_LIFECYCLE.derive_model( + model, + (ModelUnetWrapperMutation(wrapper),), + operation="SimpleSyrup contextual diffusion", + ) class ContextualDiffusionModelWrapper: diff --git a/simple_syrup/runtime/differential_diffusion.py b/simple_syrup/runtime/differential_diffusion.py index 063a423..ba3ba0e 100644 --- a/simple_syrup/runtime/differential_diffusion.py +++ b/simple_syrup/runtime/differential_diffusion.py @@ -9,6 +9,8 @@ from __future__ import annotations from importlib import import_module from typing import Any +from .patcher_lifecycle import PATCHER_LIFECYCLE, ModelDenoiseMaskMutation + def has_denoise_mask_function(model: Any) -> bool: """Return whether a model patcher already has denoise-mask behavior.""" @@ -20,31 +22,29 @@ def has_denoise_mask_function(model: Any) -> bool: def clone_with_differential_diffusion(model: Any, strength: float = 1.0) -> Any: - """Return a clone patched with ComfyUI differential denoise masks.""" + """Return a lifecycle-owned MODEL with differential denoise masks.""" if has_denoise_mask_function(model): return model - cloned_model = model.clone() - install_differential_diffusion(cloned_model, strength=strength) - return cloned_model + return PATCHER_LIFECYCLE.derive_model( + model, + (differential_diffusion_mutation(strength=strength),), + operation="SimpleSyrup differential diffusion", + ) -def install_differential_diffusion(model: Any, strength: float = 1.0) -> Any: - """Install ComfyUI differential denoise-mask behavior on a model patcher.""" +def differential_diffusion_mutation( + strength: float = 1.0, +) -> ModelDenoiseMaskMutation: + """Build the ComfyUI differential denoise-mask mutation.""" - if has_denoise_mask_function(model): - return model - set_mask_function = getattr(model, "set_model_denoise_mask_function", None) - if not callable(set_mask_function): - raise ValueError("Model does not support differential diffusion denoise masks.") differential_diffusion = import_module( "comfy_extras.nodes_differential_diffusion" ).DifferentialDiffusion - set_mask_function( + return ModelDenoiseMaskMutation( lambda *args, **kwargs: differential_diffusion.forward( *args, **kwargs, strength=strength, ) ) - return model diff --git a/simple_syrup/runtime/mixture_of_diffusers_sampling.py b/simple_syrup/runtime/mixture_of_diffusers_sampling.py index 917b7ac..3d084d5 100644 --- a/simple_syrup/runtime/mixture_of_diffusers_sampling.py +++ b/simple_syrup/runtime/mixture_of_diffusers_sampling.py @@ -23,7 +23,15 @@ from ..domain.tiled_diffusion import ( from ..shared.logging import get_logger from . import sampling_samplers, sampling_schedulers from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback -from .differential_diffusion import install_differential_diffusion +from .differential_diffusion import ( + differential_diffusion_mutation, + has_denoise_mask_function, +) +from .patcher_lifecycle import ( + PATCHER_LIFECYCLE, + ModelMutation, + ModelUnetWrapperMutation, +) from .tiled_sampling import ( ApplyModel, Latent, @@ -176,7 +184,7 @@ def clone_model_with_mixture_of_diffusers( differential_diffusion: bool = False, tiled_plan: TiledDiffusionPlan | None = None, ) -> tuple[Any, TiledDiffusionPlan]: - """Return a model clone patched with a pre-CFG Mixture wrapper.""" + """Return a derived model patched with a pre-CFG Mixture wrapper.""" plan = tiled_plan or build_tiled_diffusion_plan( latent_width=latent_width, @@ -187,10 +195,7 @@ def clone_model_with_mixture_of_diffusers( tile_batch_size=tile_batch_size, ) _validate_supplied_plan(plan, latent_width, latent_height) - cloned_model = model.clone() - if differential_diffusion: - install_differential_diffusion(cloned_model) - old_wrapper = cloned_model.model_options.get("model_function_wrapper") + old_wrapper = model.model_options.get("model_function_wrapper") if old_wrapper is not None and not callable(old_wrapper): raise ValueError("Existing model_function_wrapper is not callable.") @@ -198,8 +203,16 @@ def clone_model_with_mixture_of_diffusers( plan=plan, existing_wrapper=cast(ModelFunctionWrapper | None, old_wrapper), ) - cloned_model.set_model_unet_function_wrapper(wrapper) - return cloned_model, plan + mutations: list[ModelMutation] = [] + if differential_diffusion and not has_denoise_mask_function(model): + mutations.append(differential_diffusion_mutation()) + mutations.append(ModelUnetWrapperMutation(wrapper)) + derived_model = PATCHER_LIFECYCLE.derive_model( + model, + mutations, + operation="SimpleSyrup Mixture of Diffusers", + ) + return derived_model, plan class MixtureOfDiffusersModelWrapper: diff --git a/simple_syrup/runtime/multidiffusion_sampling.py b/simple_syrup/runtime/multidiffusion_sampling.py index 5181b19..0ce1924 100644 --- a/simple_syrup/runtime/multidiffusion_sampling.py +++ b/simple_syrup/runtime/multidiffusion_sampling.py @@ -23,7 +23,15 @@ from ..domain.tiled_diffusion import ( from ..shared.logging import get_logger from . import sampling_samplers, sampling_schedulers from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback -from .differential_diffusion import install_differential_diffusion +from .differential_diffusion import ( + differential_diffusion_mutation, + has_denoise_mask_function, +) +from .patcher_lifecycle import ( + PATCHER_LIFECYCLE, + ModelMutation, + ModelUnetWrapperMutation, +) from .tiled_sampling import ( ApplyModel, Latent, @@ -179,7 +187,7 @@ def clone_model_with_multidiffusion( differential_diffusion: bool = False, tiled_plan: TiledDiffusionPlan | None = None, ) -> tuple[Any, TiledDiffusionPlan]: - """Return a model clone patched with a pre-CFG MultiDiffusion wrapper.""" + """Return a derived model patched with a pre-CFG MultiDiffusion wrapper.""" plan = tiled_plan or build_tiled_diffusion_plan( latent_width=latent_width, @@ -190,10 +198,7 @@ def clone_model_with_multidiffusion( tile_batch_size=tile_batch_size, ) _validate_supplied_plan(plan, latent_width, latent_height) - cloned_model = model.clone() - if differential_diffusion: - install_differential_diffusion(cloned_model) - old_wrapper = cloned_model.model_options.get("model_function_wrapper") + old_wrapper = model.model_options.get("model_function_wrapper") if old_wrapper is not None and not callable(old_wrapper): raise ValueError("Existing model_function_wrapper is not callable.") @@ -201,8 +206,16 @@ def clone_model_with_multidiffusion( plan=plan, existing_wrapper=cast(ModelFunctionWrapper | None, old_wrapper), ) - cloned_model.set_model_unet_function_wrapper(wrapper) - return cloned_model, plan + mutations: list[ModelMutation] = [] + if differential_diffusion and not has_denoise_mask_function(model): + mutations.append(differential_diffusion_mutation()) + mutations.append(ModelUnetWrapperMutation(wrapper)) + derived_model = PATCHER_LIFECYCLE.derive_model( + model, + mutations, + operation="SimpleSyrup MultiDiffusion", + ) + return derived_model, plan class MultiDiffusionModelWrapper: diff --git a/simple_syrup/runtime/patcher_lifecycle.py b/simple_syrup/runtime/patcher_lifecycle.py new file mode 100644 index 0000000..90b1628 --- /dev/null +++ b/simple_syrup/runtime/patcher_lifecycle.py @@ -0,0 +1,216 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Own lifecycle-safe derivation of ComfyUI MODEL and CLIP values.""" + +from __future__ import annotations + +from collections.abc import Callable, Iterable +from dataclasses import dataclass +from typing import Any, Protocol, TypeVar, cast + + +class ModelMutation(Protocol): + """Apply one supported mutation to an already-derived MODEL patcher.""" + + def apply(self, model: object) -> None: + """Apply the mutation through ComfyUI's public patcher API.""" + + +class ClipMutation(Protocol): + """Apply one supported mutation to an already-derived CLIP value.""" + + def apply(self, clip: object) -> None: + """Apply the mutation through ComfyUI's public CLIP API.""" + + +@dataclass(frozen=True) +class ModelDenoiseMaskMutation: + """Install a denoise-mask function through the MODEL patcher API.""" + + function: Callable[..., object] + + def apply(self, model: object) -> None: + """Install the configured denoise-mask function.""" + + setter = getattr(model, "set_model_denoise_mask_function", None) + if not callable(setter): + raise TypeError("MODEL does not support denoise-mask functions.") + setter(self.function) + + +@dataclass(frozen=True) +class ModelUnetWrapperMutation: + """Install a model-function wrapper through the MODEL patcher API.""" + + wrapper: Callable[..., object] + + def apply(self, model: object) -> None: + """Install the configured model-function wrapper.""" + + setter = getattr(model, "set_model_unet_function_wrapper", None) + if not callable(setter): + raise TypeError("MODEL does not support model-function wrappers.") + setter(self.wrapper) + + +@dataclass(frozen=True) +class ModelCalcCondBatchMutation: + """Install a calc-cond-batch function through the MODEL patcher API.""" + + function: Callable[..., object] + + def apply(self, model: object) -> None: + """Install the configured calc-cond-batch function.""" + + setter = getattr(model, "set_model_sampler_calc_cond_batch_function", None) + if not callable(setter): + raise TypeError("MODEL does not support calc-cond-batch functions.") + setter(self.function) + + +@dataclass(frozen=True) +class ClipLayerMutation: + """Select the text-encoder layer used by a derived CLIP value.""" + + layer_index: int + + def apply(self, clip: object) -> None: + """Apply the configured CLIP layer index.""" + + select_layer = getattr(clip, "clip_layer", None) + if not callable(select_layer): + raise TypeError("CLIP does not support layer selection.") + select_layer(self.layer_index) + + +@dataclass(frozen=True) +class ClipHookScheduleMutation: + """Install a native ComfyUI hook schedule on a derived CLIP value.""" + + hooks: object + target: object + + def apply(self, clip: object) -> None: + """Clone and register hooks through the derived CLIP patcher.""" + + patcher = _required_attribute(clip, "patcher", value_name="CLIP") + clone_hooks = getattr(self.hooks, "clone", None) + if not callable(clone_hooks): + raise TypeError("HOOKS does not support lifecycle-safe cloning.") + register_hooks = getattr(patcher, "register_all_hook_patches", None) + if not callable(register_hooks): + raise TypeError("CLIP patcher does not support hook registration.") + + patcher_boundary = cast(Any, patcher) + clip_boundary = cast(Any, clip) + patcher_boundary.forced_hooks = clone_hooks() + clip_boundary.use_clip_schedule = True + register_hooks(self.hooks, self.target) + + +PatcherValue = TypeVar("PatcherValue") + + +class ComfyPatcherLifecycle: + """Derive Comfy patchers while preserving their source lineage.""" + + def derive_model( + self, + source: PatcherValue, + mutations: Iterable[ModelMutation], + *, + operation: str, + ) -> PatcherValue: + """Clone one MODEL, verify its lineage, and apply all mutations.""" + + derived = self._clone(source, operation=operation) + self._require_direct_parent(source, derived, operation=operation) + for mutation in mutations: + mutation.apply(derived) + return derived + + def derive_clip( + self, + source: PatcherValue, + mutations: Iterable[ClipMutation], + *, + operation: str, + disable_dynamic: bool = False, + ) -> PatcherValue: + """Clone one CLIP, verify patcher lineage, and apply all mutations.""" + + clone = getattr(source, "clone", None) + if not callable(clone): + raise TypeError(f"{operation} requires a cloneable CLIP value.") + derived = cast( + PatcherValue, + clone(disable_dynamic=disable_dynamic) if disable_dynamic else clone(), + ) + if derived is source: + raise RuntimeError(f"{operation} returned the source CLIP from clone().") + + source_patcher = _required_attribute( + source, + "patcher", + value_name="source CLIP", + ) + derived_patcher = _required_attribute( + derived, + "patcher", + value_name="derived CLIP", + ) + self._require_direct_parent( + source_patcher, + derived_patcher, + operation=operation, + ) + for mutation in mutations: + mutation.apply(derived) + return derived + + def preserve_vae(self, vae: PatcherValue, *, operation: str) -> PatcherValue: + """Return an unmodified VAE and make the no-derivation contract explicit.""" + + if not operation.strip(): + raise ValueError("VAE lifecycle operations require a descriptive name.") + return vae + + @staticmethod + def _clone(source: PatcherValue, *, operation: str) -> PatcherValue: + """Clone one MODEL through its ComfyUI boundary.""" + + clone = getattr(source, "clone", None) + if not callable(clone): + raise TypeError(f"{operation} requires a cloneable MODEL value.") + derived = cast(PatcherValue, clone()) + if derived is source: + raise RuntimeError(f"{operation} returned the source MODEL from clone().") + return derived + + @staticmethod + def _require_direct_parent( + source: object, + derived: object, + *, + operation: str, + ) -> None: + """Require Comfy's parent link used by loaded-model cleanup.""" + + if getattr(derived, "parent", None) is not source: + raise RuntimeError( + f"{operation} produced a derived patcher without its source as parent." + ) + + +PATCHER_LIFECYCLE = ComfyPatcherLifecycle() + + +def _required_attribute(value: object, name: str, *, value_name: str) -> object: + """Return a required dynamic ComfyUI boundary attribute.""" + + attribute = getattr(value, name, None) + if attribute is None: + raise TypeError(f"{value_name} does not expose {name}.") + return attribute diff --git a/simple_syrup/runtime/regional_lora_hooks.py b/simple_syrup/runtime/regional_lora_hooks.py index 4929f4b..20d4d6d 100644 --- a/simple_syrup/runtime/regional_lora_hooks.py +++ b/simple_syrup/runtime/regional_lora_hooks.py @@ -10,6 +10,8 @@ import logging from importlib import import_module from typing import Any +from .patcher_lifecycle import PATCHER_LIFECYCLE, ClipHookScheduleMutation + LOGGER = logging.getLogger(__name__) @@ -32,12 +34,18 @@ def prepare_regional_lora_clip(clip: Any, hooks: object) -> tuple[Any, object]: "text-encoder hook group entries.", matching_hook_count, ) - prepared_clip = clip.clone(disable_dynamic=True) - prepared_clip.patcher.forced_hooks = hooks.clone() - prepared_clip.use_clip_schedule = True - prepared_clip.patcher.register_all_hook_patches( - hooks, - comfy_hooks.create_target_dict(comfy_hooks.EnumWeightTarget.Clip), + prepared_clip = PATCHER_LIFECYCLE.derive_clip( + clip, + ( + ClipHookScheduleMutation( + hooks=hooks, + target=comfy_hooks.create_target_dict( + comfy_hooks.EnumWeightTarget.Clip + ), + ), + ), + operation="SimpleSyrup regional LoRA CLIP preparation", + disable_dynamic=True, ) return prepared_clip, hooks diff --git a/simple_syrup/runtime/regional_multidiffusion_sampling.py b/simple_syrup/runtime/regional_multidiffusion_sampling.py index e06fa91..0b63ddd 100644 --- a/simple_syrup/runtime/regional_multidiffusion_sampling.py +++ b/simple_syrup/runtime/regional_multidiffusion_sampling.py @@ -22,7 +22,15 @@ from ..domain.regional_detailing import LatentRegion from ..shared.logging import get_logger from . import sampling_samplers, sampling_schedulers from .detail_previews import DetailPreviewContext, prepare_detail_preview_callback -from .differential_diffusion import install_differential_diffusion +from .differential_diffusion import ( + differential_diffusion_mutation, + has_denoise_mask_function, +) +from .patcher_lifecycle import ( + PATCHER_LIFECYCLE, + ModelCalcCondBatchMutation, + ModelMutation, +) from .tiled_sampling import ( Latent, reject_unsupported_conditioning, @@ -168,7 +176,7 @@ def clone_model_with_regional_multidiffusion( global_prompt_weight: float = 0.0, differential_diffusion: bool = False, ) -> tuple[Any, RegionalMultiDiffusionSummary]: - """Return a model clone patched with regional calc-cond-batch blending.""" + """Return a derived model patched with regional calc-cond-batch blending.""" _validate_sampling_controls( steps=1, @@ -180,10 +188,7 @@ def clone_model_with_regional_multidiffusion( latent_height=latent_height, regions=regions, ) - cloned_model = model.clone() - if differential_diffusion: - install_differential_diffusion(cloned_model) - old_wrapper = cloned_model.model_options.get("sampler_calc_cond_batch_function") + old_wrapper = model.model_options.get("sampler_calc_cond_batch_function") if old_wrapper is not None and not callable(old_wrapper): raise ValueError("Existing sampler_calc_cond_batch_function is not callable.") @@ -194,7 +199,15 @@ def clone_model_with_regional_multidiffusion( existing_calc_cond_batch=cast(CalcCondBatchFunction | None, old_wrapper), global_prompt_weight=global_prompt_weight, ) - cloned_model.set_model_sampler_calc_cond_batch_function(wrapper) + mutations: list[ModelMutation] = [] + if differential_diffusion and not has_denoise_mask_function(model): + mutations.append(differential_diffusion_mutation()) + mutations.append(ModelCalcCondBatchMutation(wrapper)) + derived_model = PATCHER_LIFECYCLE.derive_model( + model, + mutations, + operation="SimpleSyrup regional MultiDiffusion", + ) summary = RegionalMultiDiffusionSummary( latent_width=latent_width, latent_height=latent_height, @@ -208,7 +221,7 @@ def clone_model_with_regional_multidiffusion( default=0, ), ) - return cloned_model, summary + return derived_model, summary class RegionalMultiDiffusionCalcCondBatch: diff --git a/simple_syrup/runtime/vae_loader.py b/simple_syrup/runtime/vae_loader.py index 23fde5a..2b63b64 100644 --- a/simple_syrup/runtime/vae_loader.py +++ b/simple_syrup/runtime/vae_loader.py @@ -14,6 +14,8 @@ from typing import Any import torch +from .patcher_lifecycle import PATCHER_LIFECYCLE + VIDEO_TAES = ("taehv", "lighttaew2_2", "lighttaew2_1", "lighttaehy1_5", "taeltx_2") IMAGE_TAES = ("taesd", "taesdxl", "taesd3", "taef1") @@ -156,7 +158,10 @@ def _build_vae(sd: dict[str, object], metadata: object | None) -> object: comfy_sd = _comfy_sd() vae = comfy_sd.VAE(sd=sd, metadata=metadata) vae.throw_exception_if_invalid() - return vae + return PATCHER_LIFECYCLE.preserve_vae( + vae, + operation="SimpleSyrup VAE loading", + ) def _folder_paths() -> ModuleType: diff --git a/tests/test_checkpoint_loader.py b/tests/test_checkpoint_loader.py index db5380d..4c579aa 100644 --- a/tests/test_checkpoint_loader.py +++ b/tests/test_checkpoint_loader.py @@ -9,7 +9,7 @@ from __future__ import annotations import sys from dataclasses import dataclass, field from pathlib import Path -from types import ModuleType +from types import ModuleType, SimpleNamespace import pytest @@ -24,18 +24,23 @@ from simple_syrup.runtime.checkpoint_loader import ( class FakeClip: """CLIP double that records clone and layer selection behavior.""" - def __init__(self, name: str = "checkpoint_clip") -> None: + def __init__( + self, + name: str = "checkpoint_clip", + parent_patcher: object | None = None, + ) -> None: """Create a CLIP double with no selected layer.""" self.name = name self.layer: int | None = None self.clone_count = 0 + self.patcher = SimpleNamespace(parent=parent_patcher) def clone(self) -> FakeClip: """Return an independent CLIP double and record the clone call.""" self.clone_count += 1 - return FakeClip(f"{self.name}_clone") + return FakeClip(f"{self.name}_clone", parent_patcher=self.patcher) def clip_layer(self, layer: int) -> None: """Record the selected CLIP layer.""" diff --git a/tests/test_comfy_patcher_lifecycle.py b/tests/test_comfy_patcher_lifecycle.py new file mode 100644 index 0000000..54b8e1f --- /dev/null +++ b/tests/test_comfy_patcher_lifecycle.py @@ -0,0 +1,239 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Characterize ComfyUI's patcher lifecycle leak detector.""" + +from __future__ import annotations + +import gc +import logging +import weakref +from typing import Any + +import pytest +import torch + +from simple_syrup.runtime.patcher_lifecycle import ( + ClipLayerMutation, + ComfyPatcherLifecycle, + ModelCalcCondBatchMutation, + ModelDenoiseMaskMutation, + ModelUnetWrapperMutation, +) + + +class AnimaTEModel_(torch.nn.Module): + """Represent a tiny Anima text encoder without production weights.""" + + def __init__(self) -> None: + """Create one parameter so ComfyUI sees a normal torch module.""" + + super().__init__() + self.projection = torch.nn.Linear(1, 1) + + +class _AnimaClip: + """Expose an Anima-style encoder through the native CLIP clone shape.""" + + def __init__(self, patcher: Any, encoder: AnimaTEModel_) -> None: + """Create a CLIP value around one real Comfy ModelPatcher.""" + + self.patcher = patcher + self.cond_stage_model = encoder + self.layer_index: int | None = None + + def clone(self, disable_dynamic: bool = False) -> _AnimaClip: + """Clone the patcher with the same behavior as Comfy's CLIP wrapper.""" + + return _AnimaClip( + self.patcher.clone(disable_dynamic=disable_dynamic), + self.cond_stage_model, + ) + + def clip_layer(self, layer_index: int) -> None: + """Record the selected text-encoder layer.""" + + self.layer_index = layer_index + + +def test_real_comfy_anima_lifecycle_regression( + caplog: pytest.LogCaptureFixture, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Exercise the unsafe predicate and every safe owner path in one worker.""" + + _assert_comfy_marks_rootless_anima_patcher_dead() + _assert_comfy_returns_clone_to_live_source() + _assert_lifecycle_owned_anima_clip_is_safe(caplog, monkeypatch) + _assert_supported_model_mutations_share_one_clone() + + +def _assert_comfy_marks_rootless_anima_patcher_dead() -> None: + """Capture the exact condition behind ComfyUI's memory-leak warning.""" + + from comfy.model_management import LoadedModel + + encoder = AnimaTEModel_() + patcher = _patcher(encoder) + loaded = LoadedModel(patcher) + loaded.real_model = weakref.ref(encoder) + + del patcher + gc.collect() + + assert loaded.model is None + assert loaded.real_model() is encoder + assert loaded.is_dead() is True + + +def _assert_comfy_returns_clone_to_live_source() -> None: + """Prove valid clone lineage cannot satisfy the stale-patcher condition.""" + + from comfy.model_management import LoadedModel + + encoder = AnimaTEModel_() + source = _patcher(encoder) + derived = source.clone() + loaded = LoadedModel(derived) + loaded.real_model = weakref.ref(encoder) + + del derived + gc.collect() + + assert loaded.model is source + assert loaded.is_dead() is False + + +def _assert_lifecycle_owned_anima_clip_is_safe( + caplog: pytest.LogCaptureFixture, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A released first-party Anima CLIP derivation resolves to its live source.""" + + import comfy.model_management + from comfy.model_management import LoadedModel + + encoder = AnimaTEModel_() + source_clip = _AnimaClip(_patcher(encoder), encoder) + derived_clip = ComfyPatcherLifecycle().derive_clip( + source_clip, + (ClipLayerMutation(-2),), + operation="Anima CLIP regression", + ) + assert isinstance(derived_clip, _AnimaClip) + assert derived_clip.layer_index == -2 + loaded = LoadedModel(derived_clip.patcher) + loaded.real_model = weakref.ref(encoder) + monkeypatch.setattr(comfy.model_management, "current_loaded_models", [loaded]) + + del derived_clip + gc.collect() + with caplog.at_level(logging.INFO): + comfy.model_management.cleanup_models_gc() + + assert loaded.model is source_clip.patcher + assert loaded.is_dead() is False + assert "Potential memory leak detected" not in caplog.text + assert "WARNING, memory leak" not in caplog.text + + caplog.clear() + released_encoder = AnimaTEModel_() + released_source = _AnimaClip(_patcher(released_encoder), released_encoder) + released_clip = ComfyPatcherLifecycle().derive_clip( + released_source, + (ClipLayerMutation(-2),), + operation="released Anima CLIP regression", + ) + released = LoadedModel(released_clip.patcher) + released.real_model = weakref.ref(released_encoder) + monkeypatch.setattr(comfy.model_management, "current_loaded_models", [released]) + + del released_clip, released_source, released_encoder + gc.collect() + with caplog.at_level(logging.INFO): + comfy.model_management.cleanup_models_gc() + + assert released.model is None + assert released.real_model() is None + assert released.is_dead() is False + assert "Potential memory leak detected" not in caplog.text + assert "WARNING, memory leak" not in caplog.text + + +def _assert_supported_model_mutations_share_one_clone() -> None: + """All supported first-party MODEL changes share one verified derivation.""" + + source = _patcher(torch.nn.Linear(1, 1)) + + def denoise_mask(*args: object, **kwargs: object) -> object: + """Return a stable test sentinel.""" + + del args, kwargs + return object() + + def model_wrapper(args: object) -> object: + """Return the supplied model-wrapper arguments.""" + + return args + + def calc_cond_batch(args: object) -> object: + """Return the supplied calc-cond-batch arguments.""" + + return args + + derived = ComfyPatcherLifecycle().derive_model( + source, + ( + ModelDenoiseMaskMutation(denoise_mask), + ModelUnetWrapperMutation(model_wrapper), + ModelCalcCondBatchMutation(calc_cond_batch), + ), + operation="MODEL mutation regression", + ) + + assert derived.parent is source + assert derived.model_options["denoise_mask_function"] is denoise_mask + assert derived.model_options["model_function_wrapper"] is model_wrapper + assert derived.model_options["sampler_calc_cond_batch_function"] is calc_cond_batch + + +def test_lifecycle_rejects_a_clone_without_comfy_parent_lineage() -> None: + """A non-Comfy clone cannot silently enter a first-party lifecycle path.""" + + class BrokenModel: + """Return an unrelated object from clone().""" + + def clone(self) -> BrokenModel: + """Return a clone without a parent link.""" + + return BrokenModel() + + with pytest.raises(RuntimeError, match="without its source as parent"): + ComfyPatcherLifecycle().derive_model( + BrokenModel(), + (), + operation="broken regression", + ) + + +def test_lifecycle_preserves_vae_identity() -> None: + """The first-party VAE contract never clones or mutates the supplied value.""" + + vae = object() + + result = ComfyPatcherLifecycle().preserve_vae( + vae, + operation="VAE regression", + ) + + assert result is vae + + +def _patcher(model: torch.nn.Module) -> Any: + """Create a CPU patcher without loading model weights.""" + + from comfy.model_patcher import ModelPatcher + + device = torch.device("cpu") + return ModelPatcher(model, load_device=device, offload_device=device) diff --git a/tests/test_contextual_diffusion_sampling.py b/tests/test_contextual_diffusion_sampling.py index cf13be3..f366c6b 100644 --- a/tests/test_contextual_diffusion_sampling.py +++ b/tests/test_contextual_diffusion_sampling.py @@ -323,17 +323,22 @@ def _controls( class _FakeModel: """Provide the ModelPatcher surface used by the semantic runtime.""" - def __init__(self, model_options: dict[str, Any] | None = None) -> None: + def __init__( + self, + model_options: dict[str, Any] | None = None, + parent: _FakeModel | None = None, + ) -> None: """Create a CPU-backed fake model patcher.""" self.load_device = torch.device("cpu") self.model_options = {} if model_options is None else model_options self.wrapper: object | None = None + self.parent = parent def clone(self) -> _FakeModel: """Return a clone with copied model options.""" - return _FakeModel(self.model_options.copy()) + return _FakeModel(self.model_options.copy(), parent=self) def set_model_unet_function_wrapper(self, wrapper: object) -> None: """Capture the installed wrapper.""" diff --git a/tests/test_patcher_lifecycle_policy.py b/tests/test_patcher_lifecycle_policy.py new file mode 100644 index 0000000..d13922b --- /dev/null +++ b/tests/test_patcher_lifecycle_policy.py @@ -0,0 +1,200 @@ +# SimpleSyrup - workflow-focused ComfyUI extensions for image generation +# Copyright (C) 2026 Artificial Sweetener and contributors +# SPDX-License-Identifier: AGPL-3.0-or-later + +"""Enforce the single first-party Comfy patcher lifecycle boundary.""" + +from __future__ import annotations + +import ast +from collections import Counter +from pathlib import Path + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +SOURCE_ROOT = PROJECT_ROOT / "simple_syrup" +LIFECYCLE_MODULE = "simple_syrup/runtime/patcher_lifecycle.py" +FORBIDDEN_PATCHER_CALLS = frozenset( + { + "add_object_patch", + "add_patches", + "clip_layer", + "register_all_hook_patches", + "set_tokenizer_option", + } +) +FORBIDDEN_PATCHER_WRITES = frozenset({"forced_hooks", "use_clip_schedule"}) +APPROVED_VALUE_CLONES = Counter( + { + ("simple_syrup/domain/segs_tiled_diffusion.py", "mask"): 2, + ("simple_syrup/image/crop_composite.py", "image"): 1, + ( + "simple_syrup/image/resize_service.py", + "values.expand(cropped.shape[0], cropped.shape[1], " + "plan.output_height, plan.output_width)", + ): 1, + ( + "simple_syrup/masking/prompt_segs_with_sam_service.py", + "crop_image(image_tensor, crop_region).detach()", + ): 1, + ( + "simple_syrup/masking/prompt_segs_with_sam_service.py", + "crop_mask(final_mask, crop_region).detach()", + ): 1, + ("simple_syrup/runtime/sam_region_overlay_renderer.py", "image.detach()"): 2, + ( + "simple_syrup/services/detail_segs_as_regions_service.py", + "image_tensor", + ): 1, + ( + "simple_syrup/services/detail_segs_by_scale_factor_service.py", + "image_tensor", + ): 2, + ( + "simple_syrup/services/detail_segs_by_scale_factor_tiled_diffusion_service.py", + "image_tensor", + ): 2, + ( + "simple_syrup/services/mask_to_segs_service.py", + "crop_image(image_tensor, crop_region).detach()", + ): 1, + ( + "simple_syrup/services/mask_to_segs_service.py", + "cropped_segment_mask", + ): 1, + ( + "simple_syrup/services/segs_detection_service.py", + "crop_mask(mask, crop_region).detach()", + ): 1, + ( + "simple_syrup/services/segs_detection_service.py", + "crop_image(image_tensor, crop_region).detach()", + ): 1, + ( + "simple_syrup/services/segs_from_sam_output_service.py", + "image[:, crop_region.top:crop_region.bottom, " + "crop_region.left:crop_region.right, :].detach()", + ): 1, + ( + "simple_syrup/services/segs_from_sam_output_service.py", + "local_mask.unsqueeze(0).detach()", + ): 1, + ( + "simple_syrup/services/segs_output_service.py", + "crop_mask(combined_mask, crop_region).detach()", + ): 1, + ( + "simple_syrup/services/segs_output_service.py", + "crop_image(image_tensor, crop_region).detach()", + ): 1, + ( + "simple_syrup/services/simple_preview_segs_service.py", + "image.detach().cpu()", + ): 1, + } +) + + +def test_first_party_patcher_lifecycle_has_one_authoritative_owner() -> None: + """Reject every first-party clone or patcher mutation outside the owner.""" + + observed_clones: Counter[tuple[str, str]] = Counter() + violations: list[str] = [] + for path in sorted(SOURCE_ROOT.rglob("*.py")): + relative_path = path.relative_to(PROJECT_ROOT).as_posix() + if "third_party" in path.parts or relative_path == LIFECYCLE_MODULE: + continue + source = path.read_text(encoding="utf-8") + tree = ast.parse(source, filename=str(path)) + file_violations, clones = _lifecycle_bypasses(tree, relative_path) + violations.extend(file_violations) + observed_clones.update(clones) + + unexpected_clones = observed_clones - APPROVED_VALUE_CLONES + missing_clones = APPROVED_VALUE_CLONES - observed_clones + assert violations == [] + assert unexpected_clones == Counter() + assert missing_clones == Counter() + + +def test_policy_detects_a_new_direct_patcher_code_path() -> None: + """Prove the policy rejects clone, mutation, and object-patch bypasses.""" + + source = """ +def unsafe(model): + cloned = model.clone() + cloned.set_model_unet_function_wrapper(lambda args: args) + cloned.add_object_patch("encode", model.encode) + cloned.forced_hooks = object() + return cloned +""" + violations, clones = _lifecycle_bypasses( + ast.parse(source), + "simple_syrup/runtime/new_feature.py", + ) + + assert len(violations) == 3 + assert clones == Counter({("simple_syrup/runtime/new_feature.py", "model"): 1}) + + +def _lifecycle_bypasses( + tree: ast.AST, + relative_path: str, +) -> tuple[list[str], Counter[tuple[str, str]]]: + """Return direct patcher mutations and all clone callsites in one tree.""" + + violations: list[str] = [] + clones: Counter[tuple[str, str]] = Counter() + for node in ast.walk(tree): + if isinstance(node, ast.Attribute): + if node.attr == "clone": + clones[(relative_path, ast.unparse(node.value))] += 1 + elif _is_forbidden_patcher_call(node.attr): + violations.append(f"{relative_path}:{node.lineno}: {node.attr}") + elif isinstance(node, ast.Call) and _uses_forbidden_dynamic_attribute(node): + violations.append(f"{relative_path}:{node.lineno}: dynamic patcher access") + elif isinstance(node, (ast.Assign, ast.AnnAssign, ast.AugAssign)): + for target in _assignment_targets(node): + if ( + isinstance(target, ast.Attribute) + and target.attr in FORBIDDEN_PATCHER_WRITES + ): + violations.append( + f"{relative_path}:{node.lineno}: write {target.attr}" + ) + return violations, clones + + +def _uses_forbidden_dynamic_attribute(node: ast.Call) -> bool: + """Return whether getattr or setattr hides a protected patcher attribute.""" + + if not isinstance(node.func, ast.Name) or node.func.id not in { + "getattr", + "setattr", + }: + return False + if len(node.args) < 2 or not isinstance(node.args[1], ast.Constant): + return False + attribute_name = node.args[1].value + return isinstance(attribute_name, str) and ( + _is_forbidden_patcher_call(attribute_name) + or attribute_name in FORBIDDEN_PATCHER_WRITES + ) + + +def _is_forbidden_patcher_call(attribute_name: str) -> bool: + """Return whether an attribute mutates a managed MODEL or CLIP value.""" + + return ( + attribute_name.startswith("set_model_") + or attribute_name in FORBIDDEN_PATCHER_CALLS + ) + + +def _assignment_targets( + node: ast.Assign | ast.AnnAssign | ast.AugAssign, +) -> tuple[ast.expr, ...]: + """Normalize assignment node targets for lifecycle-policy inspection.""" + + if isinstance(node, ast.Assign): + return tuple(node.targets) + return (node.target,) diff --git a/tests/test_regional_lora_hooks.py b/tests/test_regional_lora_hooks.py index 4b64a10..87956c6 100644 --- a/tests/test_regional_lora_hooks.py +++ b/tests/test_regional_lora_hooks.py @@ -19,13 +19,14 @@ from simple_syrup.runtime.regional_lora_hooks import prepare_regional_lora_clip class _FakeClip: """Record whether regional preparation clones the text encoder.""" - def __init__(self) -> None: + def __init__(self, parent_patcher: object | None = None) -> None: """Create a clip and its minimal patcher collaboration.""" self.cond_stage_model = object() self.registrations: list[tuple[Any, Any]] = [] self.patcher = SimpleNamespace( forced_hooks=None, + parent=parent_patcher, register_all_hook_patches=self._register_hooks, ) self.use_clip_schedule = False @@ -35,7 +36,7 @@ class _FakeClip: """Return a distinct clip while recording clone policy.""" self.clone_calls.append(disable_dynamic) - clone = _FakeClip() + clone = _FakeClip(parent_patcher=self.patcher) clone.clone_calls = self.clone_calls clone.registrations = self.registrations clone.patcher.register_all_hook_patches = clone._register_hooks