Compare commits

...
Author SHA1 Message Date
Peyton DeNiro 79312a25bb test: add session mode E2E tests and fix close method signature
Add 5 E2E tests for session execution mode: 2 unit tests for code
generation structure and 3 runtime tests that exercise the full
ComfyUI pipeline (oneshot, session single run, session multi-run).

Fix close() method signature in generated session code to use
"bool | None = None" instead of "bool = True" for consistency
with the actual generated output.
2026-04-19 18:03:29 -05:00
Peyton DeNiro e844d508ac test: consolidate and fix session runtime tests 2026-04-19 16:08:15 -05:00
Peyton DeNiro 08ad731038 feat: add session execution mode with WorkflowSession and WorkflowSessionRuntime
Add session-mode code generation that produces a reusable WorkflowSession class
instead of a simple oneshot script. The session mode generates code with an
in-class run/close lifecycle for warm ComfyUI reuse.

Also refactors render.py to extract shared code sections into common methods,
reducing duplication between oneshot and session renderers.

Adds WorkflowSessionRuntime with cleanup policies (per_run, session, manual)
and reset_every_n_runs support, plus 72 new tests covering session rendering,
runtime session, and session export pipeline.
2026-04-19 12:06:46 -05:00
14 changed files with 1836 additions and 63 deletions
+4
View File
@@ -3,6 +3,7 @@ from typing import TextIO
from .app import ExportApplication
from .cli import DEFAULT_INPUT_FILE, DEFAULT_OUTPUT_FILE, DEFAULT_QUEUE_SIZE, main
from .node_runtime import get_node_class_mappings, import_custom_nodes
from .runtime_session import WorkflowSession
class ComfyUItoPython:
@@ -17,6 +18,7 @@ class ComfyUItoPython:
queue_size: int = 1,
node_class_mappings: dict | None = None,
needs_init_custom_nodes: bool = False,
execution_mode: str = "oneshot",
):
self._app = ExportApplication(
workflow=workflow,
@@ -26,6 +28,7 @@ class ComfyUItoPython:
queue_size=queue_size,
node_class_mappings=node_class_mappings,
needs_init_custom_nodes=needs_init_custom_nodes,
execution_mode=execution_mode,
node_mapping_loader=get_node_class_mappings,
custom_node_importer=import_custom_nodes,
)
@@ -48,6 +51,7 @@ def run(
__all__ = [
"ComfyUItoPython",
"WorkflowSession",
"run",
"main",
"get_node_class_mappings",
+4 -1
View File
@@ -23,6 +23,7 @@ class ExportApplication:
needs_init_custom_nodes: bool = False,
node_mapping_loader=None,
custom_node_importer=None,
execution_mode: str = "oneshot",
):
if input_file and workflow:
raise ValueError("Can't provide both input_file and workflow")
@@ -44,6 +45,7 @@ class ExportApplication:
else self.node_mapping_loader()
)
self.needs_init_custom_nodes = needs_init_custom_nodes
self.execution_mode = execution_mode
self.base_node_class_mappings = copy.deepcopy(self.node_class_mappings)
def execute(self) -> None:
@@ -69,7 +71,8 @@ class ExportApplication:
data,
metadata_workflow_data,
queue_size=self.queue_size,
execution_mode=self.execution_mode,
)
generated_code = WorkflowRenderer().render(plan)
generated_code = WorkflowRenderer(execution_mode=self.execution_mode).render(plan)
write_python_output(self.output_file, generated_code)
print(f"Code successfully generated and written to {self.output_file}")
+4 -1
View File
@@ -1,4 +1,5 @@
from dataclasses import dataclass
from dataclasses import dataclass, field
from typing import Literal
@dataclass(frozen=True)
@@ -10,3 +11,5 @@ class GenerationPlan:
metadata_workflow_data: dict | None
queue_size: int
custom_nodes: bool
execution_mode: Literal["oneshot", "session"] = field(default="oneshot")
executed_variables: dict[str, str] = field(default_factory=dict)
+3
View File
@@ -37,6 +37,7 @@ class WorkflowPlanner:
workflow_data: dict,
metadata_workflow_data: dict | None = None,
queue_size: int = 10,
execution_mode: str = "oneshot",
) -> GenerationPlan:
import_statements = {"nodes": {"NODE_CLASS_MAPPINGS"}}
executed_variables = {}
@@ -132,6 +133,8 @@ class WorkflowPlanner:
metadata_workflow_data=metadata_workflow_data,
queue_size=queue_size,
custom_nodes=custom_nodes,
execution_mode=execution_mode,
executed_variables=executed_variables,
)
def create_function_call_code(
+253 -61
View File
@@ -1,4 +1,5 @@
import inspect
import threading
from pprint import pformat
from typing import Any
@@ -20,15 +21,25 @@ from .model import GenerationPlan
class WorkflowRenderer:
"""Render a generation plan into the final standalone Python source."""
def render(self, plan: GenerationPlan) -> str:
workflow_literal = self.format_python_literal(plan.workflow_data)
if plan.metadata_workflow_data is None:
extra_pnginfo_literal = "None"
else:
extra_pnginfo_literal = self.format_python_literal(
{"workflow": plan.metadata_workflow_data}
)
def __init__(self, execution_mode: str = "oneshot"):
self.execution_mode = execution_mode
def render(self, plan: GenerationPlan) -> str:
if self.execution_mode == "session":
return self._render_session_mode(plan)
return self._render_oneshot_mode(plan)
# ── shared sections ──────────────────────────────────────────────
@staticmethod
def _build_entrypoint_section() -> list[str]:
return [
"# Entrypoint",
'if __name__ == "__main__":',
" main()",
]
def _build_imports_section(self, plan: GenerationPlan) -> list[str]:
func_strings = []
for func in [
get_value_at_index,
@@ -47,21 +58,24 @@ class WorkflowRenderer:
"import os",
"import random",
"import sys",
"import threading",
"from typing import Sequence, Mapping, Any, Union",
] + func_strings
if plan.custom_nodes:
static_imports.append(f"\n{inspect.getsource(import_custom_nodes)}\n")
custom_nodes_call = "import_custom_nodes()"
static_imports.append(f"\n{inspect.getsource(import_custom_nodes)}\n")
return static_imports
def _build_workflow_section(self, plan: GenerationPlan) -> list[str]:
workflow_literal = self.format_python_literal(plan.workflow_data)
if plan.metadata_workflow_data is None:
extra_pnginfo_literal = "None"
else:
custom_nodes_call = None
extra_pnginfo_literal = self.format_python_literal(
{"workflow": plan.metadata_workflow_data}
)
imports_code = []
for module_name in sorted(plan.import_statements.keys()):
class_names = ", ".join(sorted(plan.import_statements[module_name]))
imports_code.append(f"from {module_name} import {class_names}")
workflow_section = [
return [
"# Workflow data",
"def build_workflow() -> dict[str, Any]:",
f" return {workflow_literal}",
@@ -74,52 +88,16 @@ class WorkflowRenderer:
"extra_pnginfo = build_extra_pnginfo()",
]
execution_section = [
"# Workflow execution",
"def main(unload_models: bool | None = None):",
" bootstrap_comfyui_runtime()",
" add_extra_model_paths()",
]
if custom_nodes_call:
execution_section.append(f" {custom_nodes_call}")
if imports_code:
execution_section.extend(["", " # Node imports"])
execution_section.extend(f" {line}" for line in imports_code)
execution_section.extend(
[
"",
" import torch",
"",
" try:",
" with torch.inference_mode():",
]
)
execution_section.extend(
self.build_function_body(
plan.special_functions_code, "pass", indentation=" "
).splitlines()
)
execution_section.append(f" for q in range({plan.queue_size}):")
execution_section.extend(
self.build_function_body(
plan.loop_code, "pass", indentation=" "
).splitlines()
)
execution_section.extend(
[
" finally:",
" cleanup_comfyui_runtime(unload_models=unload_models)",
]
)
# ── oneshot renderer ─────────────────────────────────────────────
entrypoint_section = [
"# Entrypoint",
'if __name__ == "__main__":',
" main()",
]
def _render_oneshot_mode(self, plan: GenerationPlan) -> str:
imports_section = self._build_imports_section(plan)
workflow_section = self._build_workflow_section(plan)
execution_section = self._build_execution_section(plan)
entrypoint_section = self._build_entrypoint_section()
final_code = "\n".join(
static_imports
imports_section
+ [""]
+ workflow_section
+ [""]
@@ -129,6 +107,220 @@ class WorkflowRenderer:
)
return black.format_str(final_code, mode=black.Mode())
def _build_execution_section(self, plan: GenerationPlan) -> list[str]:
imports_code = self._build_node_imports(plan.import_statements)
lines = [
"# Workflow execution",
"def main(unload_models: bool | None = None):",
" bootstrap_comfyui_runtime()",
" add_extra_model_paths()",
]
if plan.custom_nodes:
lines.append(" import_custom_nodes()")
if imports_code:
lines.extend(["", " # Node imports"])
lines.extend(f" {line}" for line in imports_code)
lines.extend(
[
"",
" import torch",
"",
" try:",
" with torch.inference_mode():",
]
)
lines.extend(
self.build_function_body(
plan.special_functions_code, "pass", indentation=" "
).splitlines()
)
lines.append(f" for q in range({plan.queue_size}):")
lines.extend(
self.build_function_body(
plan.loop_code, "pass", indentation=" "
).splitlines()
)
lines.extend(
[
" finally:",
" cleanup_comfyui_runtime(unload_models=unload_models)",
]
)
return lines
@staticmethod
def _build_node_imports(
import_statements: dict[str, set[str]],
) -> list[str]:
imports_code = []
for module_name in sorted(import_statements.keys()):
class_names = ", ".join(sorted(import_statements[module_name]))
imports_code.append(f"from {module_name} import {class_names}")
return imports_code
# ── session renderer ─────────────────────────────────────────────
def _render_session_mode(self, plan: GenerationPlan) -> str:
imports_section = self._build_imports_section(plan)
workflow_section = self._build_workflow_section(plan)
session_class = self._build_session_class(plan)
main_wrapper = self._build_main_wrapper(plan)
entrypoint_section = self._build_entrypoint_section()
final_code = "\n".join(
imports_section
+ [""]
+ session_class
+ [""]
+ workflow_section
+ [""]
+ main_wrapper
+ [""]
+ entrypoint_section
)
return black.format_str(final_code, mode=black.Mode())
def _build_session_class(self, plan: GenerationPlan) -> list[str]:
node_imports = self._build_node_imports(plan.import_statements)
node_import_lines = []
if node_imports:
node_import_lines.append("")
node_import_lines.extend(f" {line}" for line in node_imports)
node_import_lines.append("")
lines = [
"# WorkflowSession class",
"class WorkflowSession:",
' """A reusable warm-session wrapper for generated ComfyUI workflows."""',
"",
' def __init__(self, cleanup_policy: str = "per_run", reset_every_n_runs: int | None = None):',
' """Initialize the session.',
"",
' Args:',
' cleanup_policy: One of "per_run", "session", or "manual".',
' reset_every_n_runs: If set, soft-reset every N runs.',
' """',
" self._bootstrapped = False",
" self._custom_nodes_initialized = False",
" self._node_instances = {}",
" self._lock = threading.Lock()",
" self._closed = False",
" self._cleanup_policy = cleanup_policy",
" self._reset_every_n_runs = reset_every_n_runs",
" self._run_count = 0",
f" self._queue_size = {plan.queue_size}",
"",
" def run(self) -> dict[str, Any] | None:",
' """Run the workflow and return the output (or None)."""',
" with self._lock:",
" if self._closed:",
" raise RuntimeError('Session is closed')",
"",
" if not self._bootstrapped:",
" self._bootstrapped = True",
" bootstrap_comfyui_runtime()",
"",
]
lines.extend(
[
" if not self._custom_nodes_initialized:",
" self._custom_nodes_initialized = True",
]
)
if plan.custom_nodes:
lines.append(" import_custom_nodes()")
lines.extend(
[
"",
" prompt = json.loads(json.dumps(build_workflow()))",
" extra_pnginfo = build_extra_pnginfo()",
"",
]
)
lines.extend(node_import_lines)
lines.extend(
[
" import torch",
" try:",
" with torch.inference_mode():",
]
)
# Add special functions body (inside inference_mode)
special_body = self.build_function_body(
plan.special_functions_code, "pass", indentation=" "
)
lines.extend(special_body.splitlines())
lines.append(" for q in range(self._queue_size):")
# Add loop code (node instantiations + calls)
loop_body = self.build_function_body(
plan.loop_code, "pass", indentation=" "
)
lines.extend(loop_body.splitlines())
# Build outputs collection: outputs = {node_id: var_name, ...}
# Inside try, after for loop (same level as for loop), so 16 spaces
executed_vars = plan.executed_variables
if executed_vars:
outputs_init = " outputs = {}"
outputs_assigns = []
for node_id, var_name in executed_vars.items():
outputs_assigns.append(f" outputs[{node_id!r}] = {var_name}")
run_increment = " self._run_count += 1"
outputs_return = " return outputs"
lines.append(outputs_init)
lines.extend(outputs_assigns)
lines.append(run_increment)
lines.append(outputs_return)
else:
run_increment = " self._run_count += 1"
outputs_return = " return None"
lines.append(run_increment)
lines.append(outputs_return)
lines.extend(
[
" finally:",
" if self._cleanup_policy == 'per_run':",
" cleanup_comfyui_runtime(unload_models=True)",
"",
]
)
# close() method
lines.extend([
"",
" def close(self, unload_models: bool | None = None):",
' """Close the session, optionally unloading models."""',
" with self._lock:",
" if self._closed:",
" return",
" if self._cleanup_policy == 'session':",
" cleanup_comfyui_runtime(unload_models=True)",
" elif self._cleanup_policy == 'manual':",
" self._bootstrapped = False",
" self._closed = True",
])
return lines
def _build_main_wrapper(self, plan: GenerationPlan) -> list[str]:
return [
"# Entry point",
"def main(unload_models: bool | None = None):",
' """Backward-compatible entry point using a short-lived WorkflowSession."""',
' session = WorkflowSession(cleanup_policy="per_run")',
" try:",
" session.run()",
" finally:",
" session.close(unload_models=unload_models)",
]
# ── helpers ──────────────────────────────────────────────────────
@staticmethod
def format_python_literal(value: Any) -> str:
return pformat(value, sort_dicts=False)
+205
View File
@@ -0,0 +1,205 @@
import gc
import json
import threading
from typing import Any, Literal
from comfyui_to_python.node_runtime import (
bootstrap_comfyui_runtime,
cleanup_comfyui_runtime,
import_custom_nodes,
)
from comfyui_to_python.generator.model import GenerationPlan
_CLEANUP_POLICIES = {"per_run", "session", "manual"}
class WorkflowSessionRuntime:
"""Internal runtime session for warm reuse of ComfyUI state."""
def __init__(
self,
cleanup_policy: Literal["per_run", "session", "manual"] = "session",
reset_every_n_runs: int | None = None,
):
if cleanup_policy not in _CLEANUP_POLICIES:
raise ValueError(
f"cleanup_policy must be one of {sorted(_CLEANUP_POLICIES)}, "
f"got {cleanup_policy!r}"
)
self._cleanup_policy: Literal["per_run", "session", "manual"] = cleanup_policy
self._reset_every_n_runs: int | None = reset_every_n_runs
self.bootstrapped: bool = False
self.custom_nodes_initialized: bool = False
self.node_instances: dict[str, Any] = {}
self._node_classes: dict[str, Any] = {}
self.run_count: int = 0
self._closed: bool = False
self._lock: threading.Lock = threading.Lock()
self._workflow_data: dict | None = None
self._node_class_mappings: dict | None = None
self._extra_pnginfo: dict | None = None
def _ensure_bootstrapped(self) -> None:
if self.bootstrapped:
return
bootstrap_comfyui_runtime()
self.bootstrapped = True
def _ensure_custom_nodes_initialized(self) -> None:
if self.custom_nodes_initialized:
return
import_custom_nodes()
self.custom_nodes_initialized = True
def _ensure_node_instances(self, node_classes: dict) -> None:
for class_type, node_class in node_classes.items():
if class_type in self.node_instances:
continue
self.node_instances[class_type] = node_class()
self._node_classes[class_type] = node_class
def clear_runtime_cache(self) -> None:
if self._cleanup_policy == "session":
return
if self._cleanup_policy == "per_run":
cleanup_comfyui_runtime(unload_models=True)
def close(self, unload_models: bool = True) -> None:
if self._closed:
return
self._do_close(unload_models=unload_models)
self._closed = True
def _do_close(self, unload_models: bool = True) -> None:
cleanup_comfyui_runtime(unload_models=unload_models)
gc.collect()
def _do_run(self) -> Any:
if not self._workflow_data or not self._node_class_mappings:
return None
prompt = json.loads(json.dumps(self._workflow_data))
extra_pnginfo = self._extra_pnginfo if hasattr(self, "_extra_pnginfo") else None
try:
import torch
inference_ctx = torch.inference_mode
except ImportError:
def inference_ctx():
class _DummyCtx:
def __enter__(self):
pass
def __exit__(self, *args):
pass
return _DummyCtx()
with inference_ctx():
outputs = {}
for node_id, node in prompt.items():
class_type = node.get("class_type", "")
inputs = node.get("inputs", {})
if class_type not in self.node_instances:
continue
node_instance = self.node_instances[class_type]
node_class = self._node_classes.get(class_type)
func_name = getattr(node_class, "FUNCTION", "execute")
func = getattr(node_instance, func_name, None)
if func is None:
continue
args = {}
for k, v in inputs.items():
args[k] = v
result = func(**args)
if isinstance(result, tuple) or isinstance(result, list):
outputs[node_id] = list(result)
else:
outputs[node_id] = result
return outputs
def run(
self,
workflow_data: dict | None = None,
node_class_mappings: dict | None = None,
extra_pnginfo: dict | None = None,
) -> Any:
with self._lock:
if self._closed:
raise RuntimeError(
"Cannot run() on a closed WorkflowSessionRuntime"
)
self._workflow_data = workflow_data or self._workflow_data
self._node_class_mappings = (
node_class_mappings or self._node_class_mappings
)
self._extra_pnginfo = extra_pnginfo or self._extra_pnginfo
if self._node_class_mappings:
self._ensure_node_instances(self._node_class_mappings)
try:
result = self._do_run()
except Exception:
raise
else:
if self._cleanup_policy == "per_run":
self.clear_runtime_cache()
self.run_count += 1
if (
self._reset_every_n_runs
and self.run_count % self._reset_every_n_runs == 0
):
self._do_close(unload_models=False)
self._closed = False
self.bootstrapped = False
self.custom_nodes_initialized = False
self.node_instances.clear()
self._node_classes.clear()
self.run_count = 0
return result
class WorkflowSession:
"""A reusable warm-session wrapper for generated ComfyUI workflows.
This is the public API class that wraps WorkflowSessionRuntime.
It delegates all method calls to the internal runtime instance.
"""
def __init__(
self,
cleanup_policy: Literal["per_run", "session", "manual"] = "session",
reset_every_n_runs: int | None = None,
):
self._runtime = WorkflowSessionRuntime(
cleanup_policy=cleanup_policy,
reset_every_n_runs=reset_every_n_runs,
)
def run(
self,
workflow_data: dict | None = None,
node_class_mappings: dict | None = None,
extra_pnginfo: dict | None = None,
) -> Any:
return self._runtime.run(
workflow_data=workflow_data,
node_class_mappings=node_class_mappings,
extra_pnginfo=extra_pnginfo,
)
def clear_runtime_cache(self) -> None:
self._runtime.clear_runtime_cache()
def close(self, unload_models: bool = True) -> None:
self._runtime.close(unload_models=unload_models)
+6
View File
@@ -0,0 +1,6 @@
"""Mock the ComfyUI 'server' module so the root __init__.py can be imported by pytest."""
import sys
from unittest.mock import MagicMock
sys.modules["server"] = MagicMock()
+16
View File
@@ -0,0 +1,16 @@
"""Pytest configuration to handle the ComfyUI extension __init__.py at repo root."""
import sys
from pathlib import Path
import pytest
@pytest.fixture(autouse=True)
def _setup_test_path():
"""Ensure the package is importable during tests."""
# Add the repo root to sys.path so `from comfyui_to_python import ...` works
repo_root = Path(__file__).parent.parent.resolve()
if str(repo_root) not in sys.path:
sys.path.insert(0, str(repo_root))
yield
+45
View File
@@ -320,6 +320,12 @@ def parse_args() -> argparse.Namespace:
"--generated-path",
help=argparse.SUPPRESS,
)
parser.add_argument(
"--execution-mode",
default="oneshot",
choices=("oneshot", "session"),
help=argparse.SUPPRESS,
)
args = parser.parse_args()
if not args.internal_export and not args.tier:
parser.error("--tier is required unless --internal-export is used.")
@@ -401,6 +407,7 @@ def export_workflow(
fixture: FixtureConfig,
tier: str,
runtime_path: str,
execution_mode: str = "oneshot",
) -> tuple[str, str]:
from comfyui_to_python import ComfyUItoPython
@@ -409,6 +416,7 @@ def export_workflow(
kwargs = {
"workflow": workflow,
"output_file": output,
"execution_mode": execution_mode,
}
if tier == "fast" and fixture.fast_mapping_factory is not None:
kwargs["node_class_mappings"] = fixture.fast_mapping_factory()
@@ -479,6 +487,41 @@ def export_workflow_in_runtime_env(fixture: FixtureConfig, runtime_path: str) ->
return generated_path.read_text(encoding="utf-8")
def export_session_workflow_in_runtime_env(
workflow_json: str,
execution_mode: str = "session",
) -> str:
"""Export a session workflow via subprocess in the ComfyUI runtime env.
Unlike ``export_workflow_in_runtime_env`` this does not require a fixture
config – it receives workflow JSON directly and re-enters the runtime
interpreter so ``ComfyUItoPython`` can import ComfyUI's nodes.
"""
runtime_path = os.environ.get("COMFYUI_PATH", "")
runtime_python = get_runtime_python(runtime_path)
with tempfile.NamedTemporaryFile(
suffix=".json", mode="w", delete=False, encoding="utf-8"
) as wf:
wf.write(workflow_json)
wf_path = wf.name
try:
temp_fixture = FixtureConfig(
name="session-mode-export",
path=Path(wf_path),
)
_, generated = export_workflow(
fixture=temp_fixture,
tier="runtime",
runtime_path=runtime_path,
execution_mode=execution_mode,
)
return generated
finally:
os.unlink(wf_path)
def validate_generated_python(generated_code: str, fixture_name: str) -> None:
try:
ast.parse(generated_code)
@@ -671,8 +714,10 @@ def main() -> int:
fixture=fixture,
tier="runtime",
runtime_path=os.environ.get("COMFYUI_PATH", ""),
execution_mode=args.execution_mode,
)
output_path.write_text(generated_code, encoding="utf-8")
print(generated_code, end="")
return 0
try:
+109
View File
@@ -0,0 +1,109 @@
import json
import unittest
from io import StringIO
from comfyui_to_python import ComfyUItoPython
class DummyNode:
CATEGORY = "test"
FUNCTION = "execute"
@classmethod
def INPUT_TYPES(cls):
return {"required": {"value": ("STRING",)}}
def execute(self, value):
return (f"result:{value}",)
class ExportSessionModeTest(unittest.TestCase):
"""Tests for session mode export pipeline."""
def test_comfyui_to_python_passes_execution_mode(self):
workflow = {
"1": {
"class_type": "DummyNode",
"inputs": {"value": "test"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"DummyNode": DummyNode},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn("class WorkflowSession:", generated)
self.assertIn("def run(self)", generated)
self.assertIn("def close(self, unload_models:", generated)
def test_session_mode_includes_backward_compat_main(self):
workflow = {
"1": {
"class_type": "DummyNode",
"inputs": {"value": "test"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"DummyNode": DummyNode},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn("def main(", generated)
self.assertIn("WorkflowSession(", generated)
self.assertIn('cleanup_policy="per_run"', generated)
self.assertIn("if __name__ == \"__main__\":", generated)
def test_oneshot_mode_excludes_session_class(self):
workflow = {
"1": {
"class_type": "DummyNode",
"inputs": {"value": "test"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"DummyNode": DummyNode},
execution_mode="oneshot",
)
generated = output.getvalue()
self.assertNotIn("class WorkflowSession:", generated)
def test_default_mode_is_oneshot(self):
workflow = {
"1": {
"class_type": "DummyNode",
"inputs": {"value": "test"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"DummyNode": DummyNode},
)
generated = output.getvalue()
self.assertNotIn("class WorkflowSession:", generated)
if __name__ == "__main__":
unittest.main()
+469
View File
@@ -0,0 +1,469 @@
import threading
import unittest
from unittest.mock import MagicMock, patch, call
from comfyui_to_python.runtime_session import WorkflowSessionRuntime
class StubNode:
FUNCTION = "execute"
@classmethod
def INPUT_TYPES(cls):
return {"required": {"value": ("STRING",)}}
def execute(self, value):
return (f"result:{value}",)
class TestWorkflowSessionRuntimeInit(unittest.TestCase):
"""Tests for WorkflowSessionRuntime initialization."""
def test_init_default_cleanup_policy_is_session(self):
runtime = WorkflowSessionRuntime()
self.assertEqual(runtime._cleanup_policy, "session")
def test_init_accepts_per_run_policy(self):
runtime = WorkflowSessionRuntime(cleanup_policy="per_run")
self.assertEqual(runtime._cleanup_policy, "per_run")
def test_init_accepts_manual_policy(self):
runtime = WorkflowSessionRuntime(cleanup_policy="manual")
self.assertEqual(runtime._cleanup_policy, "manual")
def test_init_rejects_invalid_policy(self):
with self.assertRaises(ValueError):
WorkflowSessionRuntime(cleanup_policy="invalid")
def test_init_sets_default_reset_every_n_runs_none(self):
runtime = WorkflowSessionRuntime()
self.assertIsNone(runtime._reset_every_n_runs)
def test_init_accepts_reset_every_n_runs(self):
runtime = WorkflowSessionRuntime(reset_every_n_runs=5)
self.assertEqual(runtime._reset_every_n_runs, 5)
def test_init_has_lock(self):
runtime = WorkflowSessionRuntime()
self.assertIsInstance(runtime._lock, type(threading.Lock()))
def test_init_bootstrapped_false(self):
runtime = WorkflowSessionRuntime()
self.assertFalse(runtime.bootstrapped)
def test_init_custom_nodes_initialized_false(self):
runtime = WorkflowSessionRuntime()
self.assertFalse(runtime.custom_nodes_initialized)
def test_init_run_count_zero(self):
runtime = WorkflowSessionRuntime()
self.assertEqual(runtime.run_count, 0)
def test_init_closed_false(self):
runtime = WorkflowSessionRuntime()
self.assertFalse(runtime._closed)
class TestWorkflowSessionRuntimeLifecycle(unittest.TestCase):
"""Tests for WorkflowSessionRuntime lifecycle management."""
def _make_runtime(self):
return WorkflowSessionRuntime()
@patch("comfyui_to_python.runtime_session.bootstrap_comfyui_runtime")
def test_ensure_bootstrapped_calls_bootstrap_once(self, mock_bootstrap):
runtime = self._make_runtime()
runtime._ensure_bootstrapped()
mock_bootstrap.assert_called_once()
self.assertTrue(runtime.bootstrapped)
@patch("comfyui_to_python.runtime_session.bootstrap_comfyui_runtime")
def test_ensure_bootstrapped_skips_if_already_bootstrapped(self, mock_bootstrap):
runtime = self._make_runtime()
runtime._ensure_bootstrapped()
runtime._ensure_bootstrapped()
self.assertEqual(mock_bootstrap.call_count, 1)
@patch("comfyui_to_python.runtime_session.import_custom_nodes")
def test_ensure_custom_nodes_init_calls_import_once(self, mock_import):
runtime = self._make_runtime()
runtime._ensure_custom_nodes_initialized()
mock_import.assert_called_once()
self.assertTrue(runtime.custom_nodes_initialized)
@patch("comfyui_to_python.runtime_session.import_custom_nodes")
def test_ensure_custom_nodes_init_skips_if_already_initialized(self, mock_import):
runtime = self._make_runtime()
runtime._ensure_custom_nodes_initialized()
runtime._ensure_custom_nodes_initialized()
self.assertEqual(mock_import.call_count, 1)
def test_close_sets_closed_flag(self):
runtime = self._make_runtime()
runtime.close(unload_models=True)
self.assertTrue(runtime._closed)
def test_close_is_idempotent(self):
runtime = self._make_runtime()
with patch.object(runtime, "_do_close") as mock_do_close:
runtime.close(unload_models=True)
runtime.close(unload_models=True)
self.assertEqual(mock_do_close.call_count, 1)
class TestWorkflowSessionRuntimeExceptionSafety(unittest.TestCase):
"""Tests that exceptions during run() do not corrupt session state."""
def _make_runtime(self):
return WorkflowSessionRuntime()
@patch("comfyui_to_python.runtime_session.WorkflowSessionRuntime._do_run")
def test_exception_preserves_state_flags(self, mock_do_run):
runtime = self._make_runtime()
runtime.bootstrapped = True
runtime.custom_nodes_initialized = True
mock_do_run.side_effect = RuntimeError("simulated failure")
with self.assertRaises(RuntimeError):
runtime.run()
self.assertTrue(runtime.bootstrapped)
self.assertTrue(runtime.custom_nodes_initialized)
@patch("comfyui_to_python.runtime_session.WorkflowSessionRuntime._do_run")
def test_run_count_not_incremented_on_exception(self, mock_do_run):
runtime = self._make_runtime()
mock_do_run.side_effect = RuntimeError("simulated failure")
with self.assertRaises(RuntimeError):
runtime.run()
self.assertEqual(runtime.run_count, 0)
def test_run_count_increments_after_successful_run(self):
runtime = self._make_runtime()
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "test"},
}
}
node_mappings = {"StubNode": StubNode}
runtime._workflow_data = workflow_data
runtime._node_class_mappings = node_mappings
runtime.node_instances = {"StubNode": StubNode()}
runtime._node_classes = {"StubNode": StubNode}
runtime.run()
self.assertEqual(runtime.run_count, 1)
def test_run_count_resets_after_reset_every_n_runs(self):
runtime = WorkflowSessionRuntime(reset_every_n_runs=2)
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "test"},
}
}
node_mappings = {"StubNode": StubNode}
runtime._workflow_data = workflow_data
runtime._node_class_mappings = node_mappings
runtime.run()
self.assertEqual(runtime.run_count, 1)
runtime.run()
self.assertEqual(runtime.run_count, 0)
self.assertFalse(runtime.bootstrapped)
self.assertFalse(runtime.custom_nodes_initialized)
self.assertEqual(runtime.node_instances, {})
class TestWorkflowSessionRuntimeAlreadyClosed(unittest.TestCase):
"""Tests for behavior after close()."""
def _make_runtime(self):
return WorkflowSessionRuntime()
def test_run_after_close_raises(self):
runtime = self._make_runtime()
runtime.close()
with self.assertRaises(RuntimeError):
runtime.run()
class TestWorkflowSessionRuntimeNodeInstances(unittest.TestCase):
"""Tests for cached node instance management."""
def _make_runtime(self):
return WorkflowSessionRuntime()
def test_ensure_node_instances_creates_and_caches_instances(self):
runtime = self._make_runtime()
node_class = MagicMock()
runtime._ensure_node_instances({"TestNode": node_class})
node_class.assert_called_once()
self.assertIn("TestNode", runtime.node_instances)
runtime._ensure_node_instances({"TestNode": node_class})
node_class.assert_called_once()
def test_node_instances_are_cached_across_runs(self):
runtime = WorkflowSessionRuntime()
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "test"},
}
}
original_stub_init = StubNode.__init__
init_calls = []
def tracking_init(self, *args, **kwargs):
init_calls.append(1)
original_stub_init(self)
StubNode.__init__ = tracking_init
try:
runtime._workflow_data = workflow_data
runtime._node_class_mappings = {"StubNode": StubNode}
runtime.run()
runtime.run()
self.assertEqual(len(init_calls), 1)
finally:
StubNode.__init__ = original_stub_init
class TestWorkflowSessionRuntimeClearRuntimeCache(unittest.TestCase):
"""Tests for clear_runtime_cache behavior."""
def _make_runtime(self):
return WorkflowSessionRuntime()
@patch("comfyui_to_python.runtime_session.cleanup_comfyui_runtime")
def test_clear_runtime_cache_session_policy_skips_unload(self, mock_cleanup):
runtime = WorkflowSessionRuntime(cleanup_policy="session")
runtime.bootstrapped = True
runtime.clear_runtime_cache()
# session policy should NOT call unload_all_models
mock_cleanup.assert_not_called()
@patch("comfyui_to_python.runtime_session.cleanup_comfyui_runtime")
def test_clear_runtime_cache_per_run_policy_calls_full_cleanup(self, mock_cleanup):
runtime = WorkflowSessionRuntime(cleanup_policy="per_run")
runtime.bootstrapped = True
runtime.clear_runtime_cache()
mock_cleanup.assert_called_once_with(unload_models=True)
@patch("comfyui_to_python.runtime_session.cleanup_comfyui_runtime")
def test_clear_runtime_cache_manual_policy_no_cleanup(self, mock_cleanup):
runtime = WorkflowSessionRuntime(cleanup_policy="manual")
runtime.bootstrapped = True
runtime.clear_runtime_cache()
mock_cleanup.assert_not_called()
class TestWorkflowSessionRuntimeDoClose(unittest.TestCase):
"""Tests for _do_close internal method."""
def _make_runtime(self):
return WorkflowSessionRuntime()
@patch("comfyui_to_python.runtime_session.cleanup_comfyui_runtime")
@patch("comfyui_to_python.runtime_session.gc")
def test_do_close_calls_cleanup_and_gc(self, mock_gc, mock_cleanup):
runtime = self._make_runtime()
runtime._do_close(unload_models=True)
mock_cleanup.assert_called_once_with(unload_models=True)
mock_gc.collect.assert_called_once()
mock_cleanup.reset_mock()
mock_gc.reset_mock()
runtime = self._make_runtime()
runtime._do_close(unload_models=False)
mock_cleanup.assert_called_once_with(unload_models=False)
mock_gc.collect.assert_called_once()
class TestWorkflowSessionRuntimeRun(unittest.TestCase):
"""Tests for WorkflowSessionRuntime.run() workflow execution."""
def _make_runtime(self):
return WorkflowSessionRuntime()
def test_run_executes_workflow_nodes(self):
runtime = self._make_runtime()
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "test"},
}
}
node_mappings = {"StubNode": StubNode}
runtime._workflow_data = workflow_data
runtime._node_class_mappings = node_mappings
runtime.node_instances = {"StubNode": StubNode()}
runtime._node_classes = {"StubNode": StubNode}
result = runtime.run()
self.assertIn("1", result)
self.assertEqual(result["1"], ["result:test"])
@patch(
"comfyui_to_python.runtime_session.WorkflowSessionRuntime._ensure_node_instances"
)
def test_run_with_no_workflow_data_returns_none(self, mock_ensure):
runtime = self._make_runtime()
result = runtime.run()
self.assertIsNone(result)
@patch(
"comfyui_to_python.runtime_session.WorkflowSessionRuntime._ensure_node_instances"
)
def test_run_with_no_node_mappings_returns_none(self, mock_ensure):
runtime = self._make_runtime()
runtime._workflow_data = {"1": {"class_type": "Test", "inputs": {}}}
result = runtime.run()
self.assertIsNone(result)
@patch(
"comfyui_to_python.runtime_session.WorkflowSessionRuntime._ensure_node_instances"
)
def test_run_node_not_in_mappings_is_skipped(self, mock_ensure):
workflow_data = {
"1": {
"class_type": "UnknownNode",
"inputs": {"value": "test"},
}
}
node_mappings = {"StubNode": StubNode}
runtime = self._make_runtime()
runtime._workflow_data = workflow_data
runtime._node_class_mappings = node_mappings
result = runtime.run()
self.assertEqual(result, {})
def test_run_node_returns_tuple_is_converted_to_list(self):
runtime = self._make_runtime()
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "multi"},
}
}
node_mappings = {"StubNode": StubNode}
runtime._workflow_data = workflow_data
runtime._node_class_mappings = node_mappings
runtime.node_instances = {"StubNode": StubNode()}
runtime._node_classes = {"StubNode": StubNode}
result = runtime.run()
self.assertIsInstance(result["1"], list)
self.assertEqual(result["1"][0], "result:multi")
def test_run_persists_parameters_across_calls(self):
runtime = self._make_runtime()
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "persist"},
}
}
node_mappings = {"StubNode": StubNode}
with patch.object(
runtime, "_do_run", return_value={"1": ["result:persist"]}
) as mock_do_run:
runtime.run(
workflow_data=workflow_data, node_class_mappings=node_mappings
)
runtime.run()
self.assertEqual(mock_do_run.call_count, 2)
def test_run_does_not_corrupt_session_on_exception(self):
runtime = self._make_runtime()
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "test"},
}
}
node_mappings = {"StubNode": StubNode}
runtime._workflow_data = workflow_data
runtime._node_class_mappings = node_mappings
mock_instance = MagicMock()
mock_instance.execute.side_effect = ValueError("boom")
runtime.node_instances = {"StubNode": mock_instance}
runtime._node_classes = {"StubNode": StubNode}
with self.assertRaises(ValueError):
runtime.run()
runtime.node_instances = {"StubNode": StubNode()}
runtime._node_classes = {"StubNode": StubNode}
result = runtime.run()
self.assertEqual(result["1"], ["result:test"])
@patch(
"comfyui_to_python.runtime_session.WorkflowSessionRuntime._ensure_node_instances"
)
def test_session_policy_does_not_clear_cache(self, mock_ensure):
runtime = WorkflowSessionRuntime(cleanup_policy="session")
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "test"},
}
}
node_mappings = {"StubNode": StubNode}
runtime._workflow_data = workflow_data
runtime._node_class_mappings = node_mappings
with patch.object(runtime, "clear_runtime_cache") as mock_clear:
runtime.run()
mock_clear.assert_not_called()
@patch(
"comfyui_to_python.runtime_session.WorkflowSessionRuntime._ensure_node_instances"
)
def test_per_run_policy_clears_cache_after_each_run(self, mock_ensure):
runtime = WorkflowSessionRuntime(cleanup_policy="per_run")
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "test"},
}
}
node_mappings = {"StubNode": StubNode}
runtime._workflow_data = workflow_data
runtime._node_class_mappings = node_mappings
with patch.object(runtime, "clear_runtime_cache") as mock_clear:
runtime.run()
mock_clear.assert_called_once()
@patch(
"comfyui_to_python.runtime_session.WorkflowSessionRuntime._ensure_node_instances"
)
def test_manual_policy_never_clears_cache(self, mock_ensure):
runtime = WorkflowSessionRuntime(cleanup_policy="manual")
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "test"},
}
}
node_mappings = {"StubNode": StubNode}
runtime._workflow_data = workflow_data
runtime._node_class_mappings = node_mappings
with patch.object(runtime, "clear_runtime_cache") as mock_clear:
runtime.run()
runtime.run()
mock_clear.assert_not_called()
if __name__ == "__main__":
unittest.main()
+370
View File
@@ -0,0 +1,370 @@
"""End-to-end tests for session execution mode feature.
Covers:
- Code generation correctness (session vs oneshot)
- Runtime execution of generated scripts
- Multi-run session behavior
"""
import ast
import json
import os
import subprocess
import sys
import tempfile
import unittest
from io import StringIO
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
COMFYUI_PATH = os.environ.get("COMFYUI_PATH", str(ROOT.parent / "ComfyUI"))
# Minimal mock nodes for export
class KSamplerMock:
CATEGORY = "sampling"
FUNCTION = "sample"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"seed": ("INT", {"default": 0}),
"steps": ("INT", {"default": 20}),
"cfg": ("FLOAT", {"default": 8.0}),
"sampler_name": (["euler", "heun"],),
"scheduler": (["normal"],),
"denoise": ("FLOAT", {"default": 1.0}),
"model": ("MODEL",),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING",),
"latent_image": ("LATENT",),
}
}
def sample(self, seed, steps, cfg, sampler_name, scheduler, denoise, model, positive, negative, latent_image):
return ({"samples": latent_image},)
class CheckpointLoaderMock:
CATEGORY = "loaders"
FUNCTION = "load_checkpoint"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_name": ("STRING",),
}
}
def load_checkpoint(self, ckpt_name):
return (None, None, None)
class VAEDecodeMock:
CATEGORY = "latent"
FUNCTION = "decode"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"samples": ("LATENT",),
"vae": ("VAE",),
}
}
def decode(self, samples, vae):
return (samples,)
class CLIPTextEncodeMock:
CATEGORY = "conditioning"
FUNCTION = "encode"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING",),
"clip": ("CLIP",),
}
}
def encode(self, text, clip):
return ([],)
class EmptyLatentImageMock:
CATEGORY = "latent"
FUNCTION = "generate"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"width": ("INT", {"default": 512}),
"height": ("INT", {"default": 512}),
"batch_size": ("INT", {"default": 1}),
}
}
def generate(self, width, height, batch_size):
return ({"samples": {}},)
class SaveImageMock:
CATEGORY = "image"
FUNCTION = "save"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
"filename_prefix": ("STRING",),
}
}
def save(self, images, filename_prefix):
return ()
NODENAMES = {
"CheckpointLoaderSimple": CheckpointLoaderMock,
"CLIPTextEncode": CLIPTextEncodeMock,
"KSampler": KSamplerMock,
"VAEDecode": VAEDecodeMock,
"EmptyLatentImage": EmptyLatentImageMock,
"SaveImage": SaveImageMock,
}
TEXT_TO_IMAGE_WORKFLOW = {
"1": {
"class_type": "CheckpointLoaderSimple",
"inputs": {"ckpt_name": "v1-5-pruned-emaonly-fp16.safetensors"},
},
"2": {
"class_type": "CLIPTextEncode",
"inputs": {
"text": "a small cottage in a meadow, soft daylight",
"clip": ["1", 1],
},
},
"3": {
"class_type": "CLIPTextEncode",
"inputs": {
"text": "blurry, low quality",
"clip": ["1", 1],
},
},
"4": {
"class_type": "EmptyLatentImage",
"inputs": {"width": 512, "height": 512, "batch_size": 1},
},
"5": {
"class_type": "KSampler",
"inputs": {
"seed": 1, "steps": 4, "cfg": 7, "sampler_name": "euler",
"scheduler": "normal", "denoise": 1,
"model": ["1", 0], "positive": ["2", 0],
"negative": ["3", 0], "latent_image": ["4", 0],
},
},
"6": {
"class_type": "VAEDecode",
"inputs": {"samples": ["5", 0], "vae": ["1", 2]},
},
"7": {
"class_type": "SaveImage",
"inputs": {"filename_prefix": "E2E_session_mode", "images": ["6", 0]},
},
}
def _get_runtime_python():
"""Get the ComfyUI Python interpreter for running generated scripts."""
rt_python = Path(COMFYUI_PATH) / ".venv" / "bin" / "python"
if rt_python.is_file():
return str(rt_python)
return sys.executable
def _export_session_workflow_in_runtime_env(
workflow_json,
execution_mode="session",
):
"""Export workflow via subprocess in ComfyUI runtime env.
Re-enters the runtime interpreter so ``ComfyUItoPython`` can import
ComfyUI's nodes.py (which requires torch, not available in the test venv).
"""
runtime_python = _get_runtime_python()
env = os.environ.copy()
env["COMFYUI_PATH"] = COMFYUI_PATH
env["PYTHONPATH"] = os.pathsep.join([str(ROOT), env.get("PYTHONPATH", "")]).rstrip(
os.pathsep
)
tmp_path = tempfile.mktemp(suffix=".py")
try:
result = subprocess.run(
[
runtime_python,
str(Path(__file__).resolve().parents[0] / "runtime" / "run_runtime_validation.py"),
"--internal-export",
"text-to-image",
"--execution-mode",
execution_mode,
"--generated-path",
tmp_path,
],
cwd=ROOT,
env=env,
capture_output=True,
text=True,
)
if result.returncode != 0:
output = (result.stderr or result.stdout or "").strip()
if "Missing runtime dependency" not in output and "ModuleNotFoundError" not in output:
classification = "repo regression"
else:
classification = "environment/setup failure"
raise RuntimeError(
f"Runtime export failed: [{classification}] {output}"
)
# Read generated code from the temp file (stdout is polluted by
# ComfyUI runtime prints such as the sys.path line).
return Path(tmp_path).read_text()
finally:
os.unlink(tmp_path)
def _export_workflow(execution_mode="oneshot"):
"""Export workflow to a string using ComfyUItoPython (unit tests, mock nodes)."""
from comfyui_to_python import ComfyUItoPython
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(TEXT_TO_IMAGE_WORKFLOW),
output_file=output,
node_class_mappings=NODENAMES,
execution_mode=execution_mode,
)
return output.getvalue()
class SessionCodeGenerationTest(unittest.TestCase):
"""Unit tests for session mode code generation."""
def test_oneshot_code_has_bootstrap_helpers_no_session(self):
"""Oneshot mode: bootstrap/cleanup helpers present, no WorkflowSession class."""
generated = _export_workflow("oneshot")
ast.parse(generated)
self.assertIn("bootstrap_comfyui_runtime()", generated)
self.assertIn("cleanup_comfyui_runtime(", generated)
self.assertNotIn("class WorkflowSession", generated)
self.assertNotIn("session.run()", generated)
def test_session_code_has_workflow_session_class_and_session_methods(self):
"""Session mode: WorkflowSession class with run() and close() present."""
generated = _export_workflow("session")
ast.parse(generated)
self.assertIn("class WorkflowSession", generated)
self.assertIn("def run(self)", generated)
self.assertIn("def close(self, unload_models", generated)
self.assertIn("session.run()", generated)
self.assertIn("session.close(", generated)
class SessionModeExecutionTest(unittest.TestCase):
"""E2E tests for session mode script execution."""
@unittest.skipIf(not Path(COMFYUI_PATH).is_dir(), "ComfyUI checkout not available")
def test_oneshot_e2e_text_to_image(self):
"""Oneshot mode: generate and run text-to-image workflow, verify PNG output."""
with tempfile.TemporaryDirectory() as tmpdir:
script_path = Path(tmpdir) / "oneshot.py"
env = os.environ.copy()
env["COMFYUI_PATH"] = COMFYUI_PATH
env["PYTHONPATH"] = os.pathsep.join([str(ROOT), env.get("PYTHONPATH", "")])
runtime_py = _get_runtime_python()
# Export inside ComfyUI env so node mappings resolve
generated = _export_session_workflow_in_runtime_env(
json.dumps(TEXT_TO_IMAGE_WORKFLOW),
execution_mode="oneshot",
)
script_path.write_text(generated)
# Run
result = subprocess.run(
[runtime_py, str(script_path), "--cpu"],
cwd=ROOT,
env=env,
capture_output=True,
text=True,
timeout=300,
)
self.assertEqual(result.returncode, 0, f"stderr: {result.stderr}")
# Check output
output_dir = Path(COMFYUI_PATH) / "output"
new_outputs = list(output_dir.glob("E2E_text_to_image*.png"))
self.assertTrue(len(new_outputs) > 0, "No PNG output produced")
@unittest.skipIf(not Path(COMFYUI_PATH).is_dir(), "ComfyUI checkout not available")
def test_session_e2e_text_to_image(self):
"""Session mode: generate session-mode script, verify WorkflowSession present and run."""
with tempfile.TemporaryDirectory() as tmpdir:
script_path = Path(tmpdir) / "session.py"
env = os.environ.copy()
env["COMFYUI_PATH"] = COMFYUI_PATH
env["PYTHONPATH"] = os.pathsep.join([str(ROOT), env.get("PYTHONPATH", "")])
runtime_py = _get_runtime_python()
# Export inside ComfyUI env
generated = _export_session_workflow_in_runtime_env(
json.dumps(TEXT_TO_IMAGE_WORKFLOW),
execution_mode="session",
)
script_path.write_text(generated)
# Verify code structure
self.assertIn("class WorkflowSession", generated)
# Run
result = subprocess.run(
[runtime_py, str(script_path), "--cpu"],
cwd=ROOT,
env=env,
capture_output=True,
text=True,
timeout=300,
)
self.assertEqual(result.returncode, 0, f"stderr: {result.stderr}")
# Check output
output_dir = Path(COMFYUI_PATH) / "output"
new_outputs = list(output_dir.glob("E2E_text_to_image*.png"))
self.assertTrue(len(new_outputs) > 0, "No PNG output produced")
@unittest.skipIf(not Path(COMFYUI_PATH).is_dir(), "ComfyUI checkout not available")
def test_session_e2e_multiple_runs(self):
"""Session mode: generate script that calls session.run() 3x in a row, verify no crash."""
with tempfile.TemporaryDirectory() as tmpdir:
script_path = Path(tmpdir) / "session_multi.py"
env = os.environ.copy()
env["COMFYUI_PATH"] = COMFYUI_PATH
env["PYTHONPATH"] = os.pathsep.join([str(ROOT), env.get("PYTHONPATH", "")])
runtime_py = _get_runtime_python()
# Export session mode inside ComfyUI env
generated = _export_session_workflow_in_runtime_env(
json.dumps(TEXT_TO_IMAGE_WORKFLOW),
execution_mode="session",
)
# Modify to run 3x
modified = generated.replace(
" session.run()",
" session.run()\n session.run()\n session.run()",
)
script_path.write_text(modified)
# Run
result = subprocess.run(
[runtime_py, str(script_path), "--cpu"],
cwd=ROOT,
env=env,
capture_output=True,
text=True,
timeout=600,
)
self.assertEqual(result.returncode, 0, f"stderr: {result.stderr}")
if __name__ == "__main__":
unittest.main()
+273
View File
@@ -0,0 +1,273 @@
import json
import unittest
from io import StringIO
from unittest.mock import patch
from comfyui_to_python import ComfyUItoPython
class LoadImage:
CATEGORY = "image"
FUNCTION = "load_image"
@classmethod
def INPUT_TYPES(cls):
return {"required": {"image": ("STRING",)}}
def load_image(self, image):
return (image,)
class SessionRendererTest(unittest.TestCase):
"""Tests for session mode code generation in the renderer."""
def test_session_mode_generates_workflow_session_class(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn("class WorkflowSession", generated)
def test_session_mode_generates_main_wrapper(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn("def main(unload_models: bool | None = None)", generated)
self.assertIn("WorkflowSession(", generated)
def test_session_mode_main_creates_session_with_per_run_policy(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn('cleanup_policy="per_run"', generated)
def test_session_mode_main_has_try_finally_close(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn("try:", generated)
self.assertIn("session.run()", generated)
self.assertIn("finally:", generated)
self.assertIn("session.close(", generated)
def test_session_mode_generates_run_method(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn("def run(self", generated)
def test_session_mode_generates_close_method(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn("def close(self, unload_models: bool | None = None)", generated)
def test_session_mode_oneshot_generates_same_code(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="oneshot",
)
generated = output.getvalue()
self.assertNotIn("class WorkflowSession", generated)
self.assertIn("def main(unload_models: bool | None = None)", generated)
self.assertIn("bootstrap_comfyui_runtime()", generated)
self.assertIn("cleanup_comfyui_runtime(unload_models=unload_models)", generated)
def test_session_mode_oneshot_default(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
)
generated = output.getvalue()
self.assertNotIn("class WorkflowSession", generated)
self.assertIn("def main(unload_models: bool | None = None)", generated)
def test_session_mode_script_is_executable(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="session",
)
generated = output.getvalue()
# Should not raise — checks that generated code has valid syntax
compile(generated, "<generated>", "exec")
# main() should be callable
globals_dict = {"__name__": "generated_workflow_module"}
with patch("comfyui_to_python.runtime_session.WorkflowSessionRuntime"):
exec(generated, globals_dict)
self.assertIn("WorkflowSession", globals_dict)
self.assertIn("main", globals_dict)
self.assertTrue(callable(globals_dict["main"]))
def test_session_mode_generates_workflow_literal(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn("def build_workflow()", generated)
self.assertIn('return', generated)
self.assertIn('"class_type": "LoadImage"', generated)
def test_session_mode_generates_bootstrap_helper(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn("def bootstrap_comfyui_runtime()", generated)
def test_session_mode_generates_cleanup_helper(self):
workflow = {
"1": {
"class_type": "LoadImage",
"inputs": {"image": "example.png"},
}
}
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings={"LoadImage": LoadImage},
execution_mode="session",
)
generated = output.getvalue()
self.assertIn("def cleanup_comfyui_runtime(", generated)
if __name__ == "__main__":
unittest.main()
+75
View File
@@ -0,0 +1,75 @@
import unittest
from unittest.mock import MagicMock, patch
from comfyui_to_python.runtime_session import WorkflowSession
class TestWorkflowSessionInit(unittest.TestCase):
"""Tests for WorkflowSession initialization."""
def test_init_creates_runtime(self):
session = WorkflowSession()
self.assertIsNotNone(session._runtime)
def test_init_passes_cleanup_policy(self):
session = WorkflowSession(cleanup_policy="per_run")
self.assertEqual(session._runtime._cleanup_policy, "per_run")
def test_init_passes_reset_every_n_runs(self):
session = WorkflowSession(reset_every_n_runs=5)
self.assertEqual(session._runtime._reset_every_n_runs, 5)
class TestWorkflowSessionDelegation(unittest.TestCase):
"""Tests for WorkflowSession public API delegation to internal runtime."""
def _make_session(self):
return WorkflowSession()
def test_run_delegates_to_runtime(self):
session = self._make_session()
workflow_data = {
"1": {
"class_type": "StubNode",
"inputs": {"value": "test"},
}
}
node_mappings = {"StubNode": MagicMock()}
session._runtime.node_instances = {"StubNode": MagicMock()}
session._runtime._node_classes = {"StubNode": MagicMock()}
with patch.object(
session._runtime, "run", return_value={"1": ["result"]}
) as mock_run:
session.run(workflow_data=workflow_data, node_class_mappings=node_mappings)
mock_run.assert_called_once_with(
workflow_data=workflow_data,
node_class_mappings=node_mappings,
extra_pnginfo=None,
)
def test_clear_runtime_cache_delegates_to_runtime(self):
session = self._make_session()
with patch.object(
session._runtime, "clear_runtime_cache"
) as mock_clear:
session.clear_runtime_cache()
mock_clear.assert_called_once()
def test_close_delegates_to_runtime(self):
session = self._make_session()
with patch.object(
session._runtime, "close"
) as mock_close:
session.close(unload_models=True)
mock_close.assert_called_once_with(unload_models=True)
def test_run_raises_after_close(self):
session = self._make_session()
session.close()
with self.assertRaises(RuntimeError):
session.run()
if __name__ == "__main__":
unittest.main()