Compare commits

...
7 changed files with 447 additions and 8 deletions
+41 -8
View File
@@ -48,6 +48,7 @@ class WorkflowPlanner:
for idx, data, is_special_function in load_order:
inputs, class_type = data["inputs"], data["class_type"]
input_types = self.node_class_mappings[class_type].INPUT_TYPES()
input_value_types = self.get_input_value_types(input_types)
class_def = self.node_class_mappings[class_type]()
missing_required_variable = False
@@ -106,7 +107,7 @@ class WorkflowPlanner:
)
inputs = self.update_inputs(inputs, executed_variables)
seed_sync_code = self.create_prompt_seed_sync_code(
idx, inputs, is_special_function
idx, inputs, input_value_types, is_special_function
)
target_lines = special_functions_code if is_special_function else code
@@ -118,6 +119,7 @@ class WorkflowPlanner:
class_def.FUNCTION,
executed_variables[idx],
is_special_function,
input_value_types=input_value_types,
**inputs,
)
)
@@ -138,16 +140,24 @@ class WorkflowPlanner:
func: str,
variable_name: str,
is_special_function: bool,
input_value_types: dict[str, str] | None = None,
**kwargs,
) -> str:
args = ", ".join(self.format_arg(key, value) for key, value in kwargs.items())
args = ", ".join(
self.format_arg(key, value, (input_value_types or {}).get(key))
for key, value in kwargs.items()
)
code = f"{variable_name} = {obj_name}.{func}({args})\n"
if not is_special_function:
code = f"\t{code}"
return code
def create_prompt_seed_sync_code(
self, node_id: str, inputs: dict, is_special_function: bool
self,
node_id: str,
inputs: dict,
input_value_types: dict[str, str],
is_special_function: bool,
) -> list[str]:
seed_sync_lines = []
for key in ("seed", "noise_seed"):
@@ -156,8 +166,11 @@ class WorkflowPlanner:
randomized_seed_variable = (
f"node_{self.sanitize_node_id(str(node_id))}_{self.clean_variable_name(key)}"
)
randomized_seed_code = self.get_randomized_seed_code(
input_value_types.get(key)
)
seed_sync_lines.append(
f'{randomized_seed_variable} = prompt["{node_id}"]["inputs"]["{key}"] = random.randint(1, 2**64)'
f'{randomized_seed_variable} = prompt["{node_id}"]["inputs"]["{key}"] = {randomized_seed_code}'
)
inputs[key] = {"variable_name": randomized_seed_variable}
@@ -167,22 +180,42 @@ class WorkflowPlanner:
indentation = "" if is_special_function else "\t"
return [f"{indentation}{line}\n" for line in seed_sync_lines]
def format_arg(self, key: str, value: Any) -> str:
value_code = self.format_arg_value(key, value)
def format_arg(self, key: str, value: Any, input_value_type: str | None = None) -> str:
value_code = self.format_arg_value(key, value, input_value_type)
if key.isidentifier() and not keyword.iskeyword(key):
return f"{key}={value_code}"
return f"**{{{json.dumps(key)}: {value_code}}}"
@staticmethod
def format_arg_value(key: str, value: Any) -> str:
def format_arg_value(
key: str, value: Any, input_value_type: str | None = None
) -> str:
if isinstance(value, dict) and "variable_name" in value:
return value["variable_name"]
if key == "noise_seed" or key == "seed":
return "random.randint(1, 2**64)"
return WorkflowPlanner.get_randomized_seed_code(input_value_type)
if isinstance(value, str):
return json.dumps(value)
return repr(value)
@staticmethod
def get_input_value_types(input_types: dict) -> dict[str, str]:
value_types = {}
for section in ("required", "optional", "hidden"):
for key, value in input_types.get(section, {}).items():
if isinstance(value, tuple) and value:
value_types[key] = value[0]
elif isinstance(value, str):
value_types[key] = value
return value_types
@staticmethod
def get_randomized_seed_code(input_value_type: str | None) -> str:
randomized_seed_code = "random.randint(1, 2**64)"
if input_value_type == "STRING":
return f"str({randomized_seed_code})"
return randomized_seed_code
def get_class_info(self, class_type: str) -> tuple[str, tuple[str, str], str]:
class_obj = self.base_node_class_mappings.get(class_type)
module_name = "nodes"
@@ -0,0 +1,8 @@
{
"1": {
"class_type": "StringSeedNode",
"inputs": {
"seed": "seed-placeholder"
}
}
}
@@ -0,0 +1,27 @@
{
"33": {
"class_type": "VaeDecode",
"inputs": {
"samples": "latent-placeholder"
}
},
"42:0": {
"class_type": "UpscaleModelLoader",
"inputs": {
"model_name": "4x-ultrasharp.safetensors"
}
},
"42:1": {
"class_type": "ImageUpscaleWithModel",
"inputs": {
"upscale_model": [
"42:0",
0
],
"image": [
"33",
0
]
}
}
}
@@ -0,0 +1,10 @@
{
"1": {
"class_type": "TextConcatenateNode",
"inputs": {
"delimiter": "",
"clean_whitespace": "true",
"text_b": "\\"
}
}
}
@@ -0,0 +1,41 @@
{
"11": {
"class_type": "DualClipLoader",
"inputs": {
"clip_name": "clip.safetensors"
}
},
"633": {
"class_type": "AnySwitchRgthree",
"inputs": {
"model": "model-placeholder"
}
},
"631": {
"class_type": "PowerLoraLoaderRgthree",
"inputs": {
"PowerLoraLoaderHeaderWidget": {
"type": "PowerLoraLoaderHeaderWidget"
},
"lora_1": {
"on": false,
"lora": "lora.safetensors",
"strength": 1.2
},
"lora_2": {
"on": false,
"lora": "lora2.safetensors",
"strength": 0.7
},
"\u2795 Add Lora": "",
"model": [
"633",
0
],
"clip": [
"11",
0
]
}
}
}
@@ -0,0 +1,8 @@
{
"1": {
"class_type": "WindowsPathNode",
"inputs": {
"path": "C:\\ComfyUI\\models\\upscale_models\\RealESRGAN_x4plus.safetensors"
}
}
}
@@ -0,0 +1,312 @@
import json
import unittest
from io import StringIO
from pathlib import Path
from comfyui_to_python import ComfyUItoPython
FIXTURE_DIR = Path(__file__).parent / "fixtures" / "unit" / "generator_codegen"
class AnySwitchRgthree:
CATEGORY = "utils"
FUNCTION = "switch"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
}
}
def switch(self, model):
return (model,)
class DualClipLoader:
CATEGORY = "loaders"
FUNCTION = "load_clip"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"clip_name": ("STRING",),
}
}
def load_clip(self, clip_name):
return (clip_name,)
class PowerLoraLoaderRgthree:
CATEGORY = "loaders"
FUNCTION = "load_loras"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"PowerLoraLoaderHeaderWidget": ("DICT",),
"model": ("MODEL",),
"clip": ("CLIP",),
}
}
def load_loras(self, **kwargs):
return (kwargs,)
class UpscaleModelLoader:
CATEGORY = "loaders"
FUNCTION = "load_model"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_name": ("STRING",),
}
}
def load_model(self, model_name):
return (model_name,)
class ImageUpscaleWithModel:
CATEGORY = "image/upscaling"
FUNCTION = "upscale"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"upscale_model": ("UPSCALE_MODEL",),
"image": ("IMAGE",),
}
}
def upscale(self, upscale_model, image):
return (image,)
class VaeDecode:
CATEGORY = "latent"
FUNCTION = "decode"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"samples": ("LATENT",),
}
}
def decode(self, samples):
return (samples,)
class WindowsPathNode:
CATEGORY = "paths"
FUNCTION = "open_path"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"path": ("STRING",),
}
}
def open_path(self, path):
return (path,)
class TextConcatenateNode:
CATEGORY = "text"
FUNCTION = "text_concatenate"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"delimiter": ("STRING",),
"clean_whitespace": ("STRING",),
"text_b": ("STRING",),
}
}
def text_concatenate(self, delimiter, clean_whitespace, text_b):
return (delimiter, clean_whitespace, text_b)
class StringSeedNode:
CATEGORY = "sampling"
FUNCTION = "sample"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"seed": ("STRING",),
},
"hidden": {
"prompt": "PROMPT",
},
}
def sample(self, seed, prompt):
return (seed, prompt)
def load_fixture(name: str) -> dict:
return json.loads((FIXTURE_DIR / name).read_text(encoding="utf-8"))
def export_workflow(workflow: dict, node_class_mappings: dict) -> str:
output = StringIO()
ComfyUItoPython(
workflow=json.dumps(workflow),
output_file=output,
node_class_mappings=node_class_mappings,
)
return output.getvalue()
class GeneratorCodegenIssueRegressionTest(unittest.TestCase):
def test_export_uses_dictionary_expansion_for_rgthree_symbol_heavy_input_names(self):
generated = export_workflow(
load_fixture("unsafe-rgthree-kwargs.json"),
{
"AnySwitchRgthree": AnySwitchRgthree,
"DualClipLoader": DualClipLoader,
"PowerLoraLoaderRgthree": PowerLoraLoaderRgthree,
},
)
self.assertIn(
'powerloraloaderrgthree_631 = powerloraloaderrgthree.load_loras(',
generated,
)
self.assertIn(
'PowerLoraLoaderHeaderWidget={"type": "PowerLoraLoaderHeaderWidget"}',
generated,
)
self.assertIn('**{"\\u2795 Add Lora": ""}', generated)
self.assertNotIn('➕ Add Lora=""', generated)
def test_export_sanitizes_subgraph_identifiers_for_upscaler_workflows(self):
generated = export_workflow(
load_fixture("subgraph-upscaler-identifiers.json"),
{
"VaeDecode": VaeDecode,
"UpscaleModelLoader": UpscaleModelLoader,
"ImageUpscaleWithModel": ImageUpscaleWithModel,
},
)
self.assertIn("upscalemodelloader_42_0 = upscalemodelloader.load_model(", generated)
self.assertIn(
"imageupscalewithmodel_42_1 = imageupscalewithmodel.upscale(",
generated,
)
self.assertIn(
"upscale_model=get_value_at_index(upscalemodelloader_42_0, 0)",
generated,
)
self.assertNotIn("imageupscalewithmodel_42:1", generated)
self.assertNotIn("upscalemodelloader_42:0", generated)
def test_export_preserves_windows_style_model_paths(self):
generated = export_workflow(
load_fixture("windows-path-string.json"),
{
"WindowsPathNode": WindowsPathNode,
},
)
globals_dict = {"__name__": "generated_workflow_module"}
exec(generated, globals_dict)
self.assertEqual(
globals_dict["build_workflow"]()["1"]["inputs"]["path"],
r"C:\ComfyUI\models\upscale_models\RealESRGAN_x4plus.safetensors",
)
def test_export_preserves_trailing_backslash_string_literals(self):
generated = export_workflow(
load_fixture("trailing-backslash-string.json"),
{
"TextConcatenateNode": TextConcatenateNode,
},
)
globals_dict = {"__name__": "generated_workflow_module"}
exec(generated, globals_dict)
self.assertEqual(
globals_dict["build_workflow"]()["1"]["inputs"]["text_b"],
"\\",
)
def test_export_randomizes_string_seed_inputs_as_strings(self):
generated = export_workflow(
load_fixture("string-seed-node.json"),
{
"StringSeedNode": StringSeedNode,
},
)
self.assertIn(
'node_1_seed = prompt["1"]["inputs"]["seed"] = str(random.randint(1, 2**64))',
generated,
)
self.assertIn("seed=node_1_seed", generated)
def test_issue_cluster_regressions_render_parseable_python(self):
workflows = [
(
load_fixture("unsafe-rgthree-kwargs.json"),
{
"AnySwitchRgthree": AnySwitchRgthree,
"DualClipLoader": DualClipLoader,
"PowerLoraLoaderRgthree": PowerLoraLoaderRgthree,
},
),
(
load_fixture("subgraph-upscaler-identifiers.json"),
{
"VaeDecode": VaeDecode,
"UpscaleModelLoader": UpscaleModelLoader,
"ImageUpscaleWithModel": ImageUpscaleWithModel,
},
),
(
load_fixture("trailing-backslash-string.json"),
{
"TextConcatenateNode": TextConcatenateNode,
},
),
(
load_fixture("windows-path-string.json"),
{
"WindowsPathNode": WindowsPathNode,
},
),
(
load_fixture("string-seed-node.json"),
{
"StringSeedNode": StringSeedNode,
},
),
]
for workflow, mapping in workflows:
generated = export_workflow(workflow, mapping)
compile(generated, "<generated_workflow>", "exec")
if __name__ == "__main__":
unittest.main()