Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
79312a25bb | ||
|
|
e844d508ac | ||
|
|
08ad731038 |
@@ -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",
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user