fix(runtime): centralize Comfy patcher lifecycle

This commit is contained in:
Artificial Sweetener
2026-08-08 10:35:24 -04:00
parent 4a6b8eb15d
commit 8104e81a29
14 changed files with 808 additions and 87 deletions
+31 -31
View File
@@ -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."""
@@ -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:
+13 -13
View File
@@ -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
@@ -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:
@@ -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:
+216
View File
@@ -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
+14 -6
View File
@@ -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
@@ -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:
+6 -1
View File
@@ -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:
+8 -3
View File
@@ -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."""
+239
View File
@@ -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)
+7 -2
View File
@@ -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."""
+200
View File
@@ -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,)
+3 -2
View File
@@ -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