mirror of
https://github.com/Comfy-Org/ComfyUI.git
synced 2026-10-01 18:38:01 -05:00
feat: add DynamicGroup widget input (#16260)
This commit is contained in:
+173
-17
@@ -1316,6 +1316,125 @@ class DynamicSlot(ComfyTypeI):
|
||||
out_dict[input_type][finalized_id] = value
|
||||
out_dict["dynamic_paths"][finalized_id] = finalize_prefix(curr_prefix, curr_prefix[-1])
|
||||
|
||||
class DynamicInputError(ValueError):
|
||||
def __init__(self, input_name: str, message: str):
|
||||
super().__init__(message)
|
||||
self.input_name = input_name
|
||||
|
||||
|
||||
@comfytype(io_type="COMFY_DYNAMICGROUP_V3")
|
||||
class DynamicGroup(ComfyTypeI):
|
||||
"""Repeat a widget template and pass its values to execute as a list of row dicts.
|
||||
|
||||
Template fields must be widget inputs without force_input or nested dynamic inputs.
|
||||
Submit fields as '<group>.<index>.<field>', using indices below max without leading zeros.
|
||||
The '<group>.' prefix is reserved for group fields, not separate sibling inputs.
|
||||
min/max are integers that count submitted rows (defaults: 0/20), even when optional=True.
|
||||
max also bounds the reconstructed list length and cannot exceed 20.
|
||||
Each submitted row follows the template's required/optional field declarations.
|
||||
|
||||
Missing positions are dicts whose fields are None. Missing optional fields are
|
||||
also None; widget defaults are not injected. With INPUT_IS_LIST, missing fields
|
||||
are [None], matching the lists supplied for submitted fields.
|
||||
|
||||
Empty groups are [] in execute and check_lazy_status. Nonempty lazy groups
|
||||
contain (value, original_key) tuples at each field.
|
||||
PriceBadge widget dependencies use the same indexed field names as prompt keys.
|
||||
"""
|
||||
|
||||
Type = list[dict[str, Any]]
|
||||
_MaxRows = 20
|
||||
|
||||
class Input(DynamicInput):
|
||||
def __init__(self, id: str, template: list[WidgetInput], min: int=0, max: int=20,
|
||||
display_name: str=None, optional: bool=False, tooltip: str=None,
|
||||
lazy: bool=None, extra_dict=None, group_name: str="Group"):
|
||||
super().__init__(id, display_name, optional, tooltip, lazy, extra_dict)
|
||||
if not template:
|
||||
raise ValueError("DynamicGroup template must have at least one field.")
|
||||
for t in template:
|
||||
if isinstance(t, DynamicInput):
|
||||
raise ValueError("Nesting dynamic inputs inside DynamicGroup is not supported.")
|
||||
if not isinstance(t, WidgetInput):
|
||||
raise TypeError("DynamicGroup template fields must be WidgetInputs.")
|
||||
if t.force_input:
|
||||
raise ValueError(f"DynamicGroup template field '{t.id}' must not use force_input.")
|
||||
if not t.id or "." in t.id:
|
||||
raise ValueError(f"DynamicGroup template field id must be nonempty and must not contain '.'. Got: '{t.id}'")
|
||||
field_ids = [t.id for t in template]
|
||||
if len(field_ids) != len(set(field_ids)):
|
||||
raise ValueError("DynamicGroup template field ids must be unique within a row.")
|
||||
if not id or "." in id:
|
||||
raise ValueError(f"DynamicGroup id must be nonempty and must not contain '.'. Got: '{id}'")
|
||||
if isinstance(min, bool) or not isinstance(min, int) or isinstance(max, bool) or not isinstance(max, int):
|
||||
raise TypeError("DynamicGroup min and max must be integers.")
|
||||
if not 1 <= max <= DynamicGroup._MaxRows:
|
||||
raise ValueError(f"DynamicGroup max must be between 1 and {DynamicGroup._MaxRows}.")
|
||||
if not 0 <= min <= max:
|
||||
raise ValueError("DynamicGroup min must be between 0 and max.")
|
||||
self.template = template
|
||||
self.min = min
|
||||
self.max = max
|
||||
self.group_name = group_name
|
||||
|
||||
def get_all(self) -> list[Input]:
|
||||
return [self] + list(self.template)
|
||||
|
||||
def as_dict(self):
|
||||
return super().as_dict() | prune_dict({
|
||||
"template": create_input_dict_v1(self.template),
|
||||
"min": self.min,
|
||||
"max": self.max,
|
||||
"group_name": self.group_name,
|
||||
})
|
||||
|
||||
def validate(self):
|
||||
for t in self.template:
|
||||
t.validate()
|
||||
|
||||
@staticmethod
|
||||
def _expand_schema_for_dynamic(out_dict: dict[str, Any], live_inputs: dict[str, Any], value: tuple[str, dict[str, Any]], input_type: str, curr_prefix: list[str] | None):
|
||||
info = value[1]
|
||||
min_rows = info.get("min", 0)
|
||||
max_rows = info.get("max", DynamicGroup._MaxRows)
|
||||
template = info.get("template", {})
|
||||
field_specs = {
|
||||
field_id: (field_value, category)
|
||||
for category in ("required", "optional")
|
||||
for field_id, field_value in template.get(category, {}).items()
|
||||
}
|
||||
finalized_prefix = finalize_prefix(curr_prefix)
|
||||
prefix = finalized_prefix + "."
|
||||
present_rows = set()
|
||||
for key in live_inputs:
|
||||
if not key.startswith(prefix):
|
||||
continue
|
||||
index, separator, field_id = key[len(prefix):].partition(".")
|
||||
if (not separator or not index.isascii() or not index.isdecimal()
|
||||
or (len(index) > 1 and index.startswith("0")) or field_id not in field_specs):
|
||||
raise DynamicInputError(key, f"Invalid DynamicGroup input key '{key}'; expected '{finalized_prefix}.<index>.<template field>'.")
|
||||
row = int(index)
|
||||
if row >= max_rows:
|
||||
raise DynamicInputError(key, f"DynamicGroup input '{key}' exceeds the index limit of {max_rows - 1} (max={max_rows}).")
|
||||
present_rows.add(row)
|
||||
|
||||
if not min_rows <= len(present_rows) <= max_rows:
|
||||
raise DynamicInputError(finalized_prefix, f"DynamicGroup input '{finalized_prefix}' received {len(present_rows)} rows; expected between {min_rows} and {max_rows}.")
|
||||
|
||||
for row in range(max(present_rows, default=-1) + 1):
|
||||
for field_id, (field_value, category) in field_specs.items():
|
||||
slot_id = f"{finalized_prefix}.{row}.{field_id}"
|
||||
if row in present_rows:
|
||||
out_dict[category][slot_id] = field_value
|
||||
# Preserve gaps without adding validation inputs for them.
|
||||
out_dict["dynamic_paths"][slot_id] = slot_id
|
||||
|
||||
out_dict["list_paths"].add(finalized_prefix)
|
||||
if not present_rows:
|
||||
out_dict["dynamic_paths"][finalized_prefix] = finalized_prefix
|
||||
out_dict["dynamic_paths_default_value"][finalized_prefix] = DynamicPathsDefaultValue.EMPTY_LIST
|
||||
|
||||
|
||||
@comfytype(io_type="IMAGECOMPARE")
|
||||
class ImageCompare(ComfyTypeI):
|
||||
Type = dict
|
||||
@@ -1529,6 +1648,8 @@ def setup_dynamic_input_funcs():
|
||||
register_dynamic_input_func(DynamicCombo.io_type, DynamicCombo._expand_schema_for_dynamic)
|
||||
# DynamicSlot.Input
|
||||
register_dynamic_input_func(DynamicSlot.io_type, DynamicSlot._expand_schema_for_dynamic)
|
||||
# DynamicGroup.Input
|
||||
register_dynamic_input_func(DynamicGroup.io_type, DynamicGroup._expand_schema_for_dynamic)
|
||||
|
||||
if len(DYNAMIC_INPUT_LOOKUP) == 0:
|
||||
setup_dynamic_input_funcs()
|
||||
@@ -1540,6 +1661,8 @@ class V3Data(TypedDict):
|
||||
'Dictionary where the keys are the input ids and the values dictate how to turn the inputs into a nested dictionary.'
|
||||
dynamic_paths_default_value: dict[str, Any]
|
||||
'Dictionary where the keys are the input ids and the values are a string from DynamicPathsDefaultValue for the inputs if value is None.'
|
||||
list_paths: set[str]
|
||||
'Paths whose index-keyed row dictionaries are reconstructed as lists.'
|
||||
create_dynamic_tuple: bool
|
||||
'When True, the value of the dynamic input will be in the format (value, path_key).'
|
||||
|
||||
@@ -1651,16 +1774,19 @@ class PriceBadgeDepends:
|
||||
raise ValueError("PriceBadgeDepends.input_groups must be a list[str].")
|
||||
|
||||
def as_dict(self, schema_inputs: list["Input"]) -> dict[str, Any]:
|
||||
# Build lookup: widget_id -> io_type
|
||||
input_types: dict[str, str] = {}
|
||||
for inp in schema_inputs:
|
||||
all_inputs = inp.get_all()
|
||||
input_types[inp.id] = inp.get_io_type() # First input is always the parent itself
|
||||
for nested_inp in all_inputs[1:]:
|
||||
# For DynamicCombo/DynamicSlot, nested inputs are prefixed with parent ID
|
||||
# to match frontend naming convention (e.g., "should_texture.enable_pbr")
|
||||
prefixed_id = f"{inp.id}.{nested_inp.id}"
|
||||
input_types[prefixed_id] = nested_inp.get_io_type()
|
||||
|
||||
def collect_inputs(inputs: list[Input], prefix: str = "") -> None:
|
||||
for inp in inputs:
|
||||
name = prefix + inp.id
|
||||
input_types[name] = inp.get_io_type()
|
||||
if isinstance(inp, DynamicGroup.Input):
|
||||
for row in range(inp.max):
|
||||
collect_inputs(inp.template, f"{name}.{row}.")
|
||||
else:
|
||||
collect_inputs(inp.get_all()[1:], name + ".")
|
||||
|
||||
collect_inputs(schema_inputs)
|
||||
|
||||
# Enrich widgets with type information, raising error for unknown widgets
|
||||
widgets_data: list[dict[str, str]] = []
|
||||
@@ -1888,6 +2014,7 @@ def get_finalized_class_inputs(d: dict[str, Any], live_inputs: dict[str, Any], i
|
||||
"optional": {},
|
||||
"dynamic_paths": {},
|
||||
"dynamic_paths_default_value": {},
|
||||
"list_paths": set(),
|
||||
}
|
||||
d = d.copy()
|
||||
# ignore hidden for parsing
|
||||
@@ -1899,10 +2026,12 @@ def get_finalized_class_inputs(d: dict[str, Any], live_inputs: dict[str, Any], i
|
||||
dynamic_paths = out_dict.pop("dynamic_paths", None)
|
||||
if dynamic_paths is not None and len(dynamic_paths) > 0:
|
||||
v3_data["dynamic_paths"] = dynamic_paths
|
||||
# this list is used for autogrow, in the case all inputs are optional and no values are passed
|
||||
dynamic_paths_default_value = out_dict.pop("dynamic_paths_default_value", None)
|
||||
if dynamic_paths_default_value is not None and len(dynamic_paths_default_value) > 0:
|
||||
v3_data["dynamic_paths_default_value"] = dynamic_paths_default_value
|
||||
list_paths = out_dict.pop("list_paths")
|
||||
if list_paths:
|
||||
v3_data["list_paths"] = list_paths
|
||||
return out_dict, hidden, v3_data
|
||||
|
||||
def parse_class_inputs(out_dict: dict[str, Any], live_inputs: dict[str, Any], curr_dict: dict[str, Any], curr_prefix: list[str] | None=None) -> None:
|
||||
@@ -1921,11 +2050,23 @@ def parse_class_inputs(out_dict: dict[str, Any], live_inputs: dict[str, Any], cu
|
||||
if curr_prefix:
|
||||
out_dict["dynamic_paths"][finalized_id] = finalized_id
|
||||
|
||||
def _dynamic_group_prefixes(inputs: list[Input]) -> Iterable[str]:
|
||||
for inp in inputs:
|
||||
if isinstance(inp, DynamicGroup.Input):
|
||||
yield inp.id + "."
|
||||
elif isinstance(inp, DynamicInput):
|
||||
for prefix in _dynamic_group_prefixes(inp.get_all()[1:]):
|
||||
yield f"{inp.id}.{prefix}"
|
||||
|
||||
|
||||
def create_input_dict_v1(inputs: list[Input]) -> dict:
|
||||
input = {
|
||||
"required": {}
|
||||
}
|
||||
group_prefixes = tuple(_dynamic_group_prefixes(inputs))
|
||||
for i in inputs:
|
||||
if group_prefixes and i.id.startswith(group_prefixes):
|
||||
raise ValueError(f"Input '{i.id}' conflicts with a DynamicGroup field prefix.")
|
||||
add_to_dict_v1(i, input)
|
||||
return input
|
||||
|
||||
@@ -1938,10 +2079,12 @@ def add_to_dict_v1(i: Input, d: dict):
|
||||
|
||||
class DynamicPathsDefaultValue:
|
||||
EMPTY_DICT = "empty_dict"
|
||||
EMPTY_LIST = "empty_list"
|
||||
|
||||
def build_nested_inputs(values: dict[str, Any], v3_data: V3Data):
|
||||
def build_nested_inputs(values: dict[str, Any], v3_data: V3Data, *, input_is_list: bool = False):
|
||||
paths = v3_data.get("dynamic_paths", None)
|
||||
default_value_dict = v3_data.get("dynamic_paths_default_value", {})
|
||||
list_paths = v3_data.get("list_paths", set())
|
||||
if paths is None:
|
||||
return values
|
||||
values = values.copy()
|
||||
@@ -1958,19 +2101,31 @@ def build_nested_inputs(values: dict[str, Any], v3_data: V3Data):
|
||||
is_last = (i == len(parts) - 1)
|
||||
|
||||
if is_last:
|
||||
missing = key not in values
|
||||
value = values.pop(key, None)
|
||||
if value is None:
|
||||
# see if a default value was provided for this key
|
||||
default_option = default_value_dict.get(key, None)
|
||||
if default_option == DynamicPathsDefaultValue.EMPTY_DICT:
|
||||
value = {}
|
||||
if create_tuple:
|
||||
default_option = default_value_dict.get(key, None)
|
||||
if default_option == DynamicPathsDefaultValue.EMPTY_LIST:
|
||||
# No row keys were submitted, regardless of the root input value.
|
||||
value = []
|
||||
elif value is None and default_option == DynamicPathsDefaultValue.EMPTY_DICT:
|
||||
value = {}
|
||||
elif missing and input_is_list and ".".join(parts[:-2]) in list_paths:
|
||||
value = [None]
|
||||
if create_tuple and default_option != DynamicPathsDefaultValue.EMPTY_LIST:
|
||||
value = (value, key)
|
||||
current[p] = value
|
||||
else:
|
||||
current = current.setdefault(p, {})
|
||||
|
||||
values.update(result)
|
||||
for list_path in list_paths:
|
||||
parts = list_path.split(".")
|
||||
container = values
|
||||
for part in parts[:-1]:
|
||||
container = container[part]
|
||||
rows = container[parts[-1]]
|
||||
if isinstance(rows, dict):
|
||||
container[parts[-1]] = [rows[key] for key in sorted(rows, key=int)]
|
||||
return values
|
||||
|
||||
|
||||
@@ -2538,6 +2693,7 @@ __all__ = [
|
||||
"MatchType",
|
||||
"DynamicCombo",
|
||||
"Autogrow",
|
||||
"DynamicGroup",
|
||||
# Other classes
|
||||
"HiddenHolder",
|
||||
"Hidden",
|
||||
|
||||
+12
-2
@@ -293,7 +293,7 @@ async def _async_map_node_over_list(prompt_id, unique_id, obj, input_data_all, f
|
||||
f = make_locked_method_func(type_obj, func, class_clone)
|
||||
# in case of dynamic inputs, restructure inputs to expected nested dict
|
||||
if v3_data is not None:
|
||||
inputs = _io.build_nested_inputs(inputs, v3_data)
|
||||
inputs = _io.build_nested_inputs(inputs, v3_data, input_is_list=input_is_list)
|
||||
# V1
|
||||
else:
|
||||
f = getattr(obj, func)
|
||||
@@ -880,7 +880,17 @@ async def validate_inputs(prompt_id, prompt, item, validated, visiting=None):
|
||||
if issubclass(obj_class, _ComfyNodeInternal):
|
||||
obj_class: _io._ComfyNodeBaseInternal
|
||||
class_inputs = obj_class.INPUT_TYPES()
|
||||
class_inputs, _, v3_data = _io.get_finalized_class_inputs(class_inputs, inputs)
|
||||
try:
|
||||
class_inputs, _, v3_data = _io.get_finalized_class_inputs(class_inputs, inputs)
|
||||
except _io.DynamicInputError as ex:
|
||||
errors.append({
|
||||
"type": "invalid_dynamic_input",
|
||||
"message": "Invalid dynamic input",
|
||||
"details": str(ex),
|
||||
"extra_info": {"input_name": ex.input_name},
|
||||
})
|
||||
validated[unique_id] = (False, errors, unique_id)
|
||||
return validated[unique_id]
|
||||
validate_function_name = "validate_inputs"
|
||||
validate_function = first_real_override(obj_class, validate_function_name)
|
||||
else:
|
||||
|
||||
@@ -0,0 +1,266 @@
|
||||
import pytest
|
||||
|
||||
from comfy_api.latest import io
|
||||
from comfy_api.latest._io import DynamicSlot, build_nested_inputs, create_input_dict_v1, get_finalized_class_inputs
|
||||
|
||||
|
||||
def _reconstruct(group, values, *, lazy=False):
|
||||
_, _, v3_data = get_finalized_class_inputs(create_input_dict_v1([group]), values)
|
||||
v3_data["create_dynamic_tuple"] = lazy
|
||||
return build_nested_inputs(values, v3_data)
|
||||
|
||||
|
||||
def test_serializes_one_template_with_field_requirements():
|
||||
group = io.DynamicGroup.Input(
|
||||
"rows",
|
||||
template=[io.String.Input("name"), io.Float.Input("weight", default=1.0, optional=True)],
|
||||
min=0, max=5, group_name="Item",
|
||||
)
|
||||
schema = create_input_dict_v1([group])
|
||||
assert schema == {"required": {"rows": ("COMFY_DYNAMICGROUP_V3", {
|
||||
"template": {
|
||||
"required": {"name": ("STRING", {"multiline": False})},
|
||||
"optional": {"weight": ("FLOAT", {"default": 1.0})},
|
||||
},
|
||||
"min": 0, "max": 5, "group_name": "Item",
|
||||
})}}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("group_id,template,limits", [
|
||||
("rows", [], {}),
|
||||
("rows", [io.Float.Input("x"), io.Float.Input("x")], {}),
|
||||
("rows.bad", [io.Float.Input("x")], {}),
|
||||
("rows", [io.Float.Input("x.bad")], {}),
|
||||
("", [io.Float.Input("x")], {}),
|
||||
("rows", [io.Float.Input("")], {}),
|
||||
("rows", [io.Float.Input("x", force_input=True)], {}),
|
||||
("rows", [io.DynamicGroup.Input("nested", template=[io.Float.Input("x")])], {}),
|
||||
("rows", [io.Float.Input("x")], {"min": -1}),
|
||||
("rows", [io.Float.Input("x")], {"min": 2, "max": 1}),
|
||||
("rows", [io.Float.Input("x")], {"max": 0}),
|
||||
("rows", [io.Float.Input("x")], {"max": 21}),
|
||||
])
|
||||
def test_rejects_invalid_template_or_limits(group_id, template, limits):
|
||||
with pytest.raises(ValueError):
|
||||
io.DynamicGroup.Input(group_id, template=template, **limits)
|
||||
|
||||
|
||||
def test_rejects_socket_template():
|
||||
with pytest.raises(TypeError, match="WidgetInputs"):
|
||||
io.DynamicGroup.Input("rows", template=[io.Image.Input("image")])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("limit", ["min", "max"])
|
||||
@pytest.mark.parametrize("value", [1.5, True, False])
|
||||
def test_rejects_non_integer_limits(limit, value):
|
||||
with pytest.raises(TypeError, match="min and max must be integers"):
|
||||
io.DynamicGroup.Input("rows", template=[io.Float.Input("x")], **{limit: value})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("lazy", [False, True])
|
||||
@pytest.mark.parametrize("values", [{}, {"rows": 0}, {"rows": 7}, {"rows": [1, 2, 3]}, {"rows": {"bad": "data"}}])
|
||||
def test_empty_group_is_an_empty_list(lazy, values):
|
||||
group = io.DynamicGroup.Input("rows", template=[io.Float.Input("x", default=1.0)], min=0)
|
||||
assert _reconstruct(group, values, lazy=lazy) == {"rows": []}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sibling_id", ["rows.summary", "rows.0.x"])
|
||||
@pytest.mark.parametrize("nested", [False, True])
|
||||
def test_rejects_sibling_input_in_group_namespace(sibling_id, nested):
|
||||
inputs = [
|
||||
io.Float.Input(sibling_id, optional=True),
|
||||
io.DynamicGroup.Input("rows", template=[io.Float.Input("x")]),
|
||||
]
|
||||
if nested:
|
||||
inputs = [io.DynamicCombo.Input("mode", options=[io.DynamicCombo.Option("on", inputs)])]
|
||||
with pytest.raises(ValueError, match="conflicts with a DynamicGroup field prefix"):
|
||||
create_input_dict_v1(inputs)
|
||||
|
||||
|
||||
def test_other_dotted_input_ids_are_unchanged():
|
||||
inputs = [io.DynamicGroup.Input("rows", template=[io.Float.Input("x")]), io.Float.Input("rows_summary.value")]
|
||||
schema = create_input_dict_v1(inputs)
|
||||
assert schema["required"]["rows_summary.value"] == ("FLOAT", {})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sibling_id", ["mode.rows.summary", "mode.rows.0.x"])
|
||||
@pytest.mark.parametrize("sibling_first", [False, True])
|
||||
def test_rejects_outer_input_in_nested_group_namespace(sibling_id, sibling_first):
|
||||
inputs = [
|
||||
io.DynamicCombo.Input("mode", options=[io.DynamicCombo.Option("on", [
|
||||
io.DynamicGroup.Input("rows", template=[io.Float.Input("x")]),
|
||||
])]),
|
||||
io.Float.Input(sibling_id),
|
||||
]
|
||||
if sibling_first:
|
||||
inputs.reverse()
|
||||
with pytest.raises(ValueError, match="conflicts with a DynamicGroup field prefix"):
|
||||
create_input_dict_v1(inputs)
|
||||
|
||||
|
||||
def test_group_namespace_is_scoped_to_its_combo_option():
|
||||
combo = io.DynamicCombo.Input("mode", options=[
|
||||
io.DynamicCombo.Option("on", [io.DynamicGroup.Input("rows", template=[io.Float.Input("x")])]),
|
||||
io.DynamicCombo.Option("off", [io.Float.Input("rows.summary")]),
|
||||
])
|
||||
assert _reconstruct(combo, {"mode": "off", "mode.rows.summary": 0.5}) == {
|
||||
"mode": {"mode": "off", "rows": {"summary": 0.5}},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("nested", [False, True])
|
||||
def test_price_badge_resolves_indexed_group_fields(nested):
|
||||
group = io.DynamicGroup.Input("rows", template=[io.Float.Input("weight"), io.String.Input("name")], max=2)
|
||||
inputs = [group, io.Float.Input("fixed")]
|
||||
prefix = "rows"
|
||||
if nested:
|
||||
inputs = [io.DynamicCombo.Input("mode", options=[io.DynamicCombo.Option("on", inputs)])]
|
||||
prefix = "mode.rows"
|
||||
outer = "mode." if nested else ""
|
||||
badge = io.PriceBadgeDepends(widgets=[f"{prefix}.0.weight", f"{prefix}.1.name", outer + "fixed"])
|
||||
assert badge.as_dict(inputs)["widgets"] == [
|
||||
{"name": f"{prefix}.0.weight", "type": "FLOAT"},
|
||||
{"name": f"{prefix}.1.name", "type": "STRING"},
|
||||
{"name": outer + "fixed", "type": "FLOAT"},
|
||||
]
|
||||
for invalid in (f"{prefix}.weight", f"{prefix}.2.weight"):
|
||||
with pytest.raises(ValueError, match="unknown widget"):
|
||||
io.PriceBadgeDepends(widgets=[invalid]).as_dict(inputs)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("minimum", [0, 1, 2])
|
||||
@pytest.mark.parametrize("optional_group", [False, True])
|
||||
def test_every_submitted_row_keeps_template_requirements(minimum, optional_group):
|
||||
group = io.DynamicGroup.Input("rows", template=[
|
||||
io.String.Input("name"),
|
||||
io.Float.Input("weight", default=1.0),
|
||||
io.Boolean.Input("enabled", optional=True),
|
||||
], min=minimum, optional=optional_group)
|
||||
values = {"rows.0.name": "A", "rows.0.weight": 0.8, "rows.2.name": "C"}
|
||||
schema, _, _ = get_finalized_class_inputs(create_input_dict_v1([group]), values)
|
||||
assert set(schema["required"]) == {
|
||||
"rows.0.name", "rows.0.weight", "rows.2.name", "rows.2.weight",
|
||||
}
|
||||
assert set(schema["optional"]) == {"rows.0.enabled", "rows.2.enabled"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("optional_group", [False, True])
|
||||
@pytest.mark.parametrize("values", [{}, {"rows.0.x": 1.0}, {"rows.2.x": 1.0}])
|
||||
def test_min_counts_submitted_rows_without_padding(optional_group, values):
|
||||
group = io.DynamicGroup.Input("rows", template=[io.Float.Input("x", optional=True)], min=2, optional=optional_group)
|
||||
with pytest.raises(ValueError, match="expected between 2 and"):
|
||||
_reconstruct(group, values)
|
||||
|
||||
|
||||
def test_sparse_rows_preserve_positions_without_defaults():
|
||||
group = io.DynamicGroup.Input("rows", template=[
|
||||
io.String.Input("name"), io.Float.Input("weight", default=1.0, optional=True),
|
||||
], min=2, max=3)
|
||||
values = {"rows.2.name": "C", "rows.2.weight": 0.5, "rows.0.name": "A"}
|
||||
assert _reconstruct(group, values) == {"rows": [
|
||||
{"name": "A", "weight": None},
|
||||
{"name": None, "weight": None},
|
||||
{"name": "C", "weight": 0.5},
|
||||
]}
|
||||
assert values == {"rows.2.name": "C", "rows.2.weight": 0.5, "rows.0.name": "A"}
|
||||
|
||||
|
||||
def test_max_counts_rows_not_fields():
|
||||
group = io.DynamicGroup.Input("rows", template=[io.String.Input("name"), io.Float.Input("weight")], max=1)
|
||||
assert _reconstruct(group, {"rows.0.name": "A", "rows.0.weight": 0.8}) == {
|
||||
"rows": [{"name": "A", "weight": 0.8}],
|
||||
}
|
||||
with pytest.raises(ValueError, match="exceeds the index limit of 0"):
|
||||
_reconstruct(group, {"rows.0.name": "A", "rows.2.name": "C"})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("key", [
|
||||
"rows.foo.x", "rows.-1.x", "rows.01.x", "rows.+1.x", "rows.١.x",
|
||||
"rows.0", "rows..x", "rows.0.unknown", "rows.0.x.extra",
|
||||
])
|
||||
def test_rejects_malformed_row_keys(key):
|
||||
group = io.DynamicGroup.Input("rows", template=[io.Float.Input("x")])
|
||||
with pytest.raises(ValueError) as error:
|
||||
_reconstruct(group, {key: 1.0})
|
||||
assert key in str(error.value)
|
||||
|
||||
|
||||
def test_default_max_is_exposed_and_enforced():
|
||||
group = io.DynamicGroup.Input("rows", template=[io.Float.Input("x")])
|
||||
assert create_input_dict_v1([group])["required"]["rows"][1]["max"] == 20
|
||||
with pytest.raises(ValueError, match="exceeds the index limit of 19"):
|
||||
_reconstruct(group, {"rows.20.x": 0.5})
|
||||
|
||||
|
||||
def test_largest_supported_index_preserves_position():
|
||||
group = io.DynamicGroup.Input("rows", template=[io.Float.Input("x")], max=20)
|
||||
rows = _reconstruct(group, {"rows.19.x": 0.5})["rows"]
|
||||
assert rows == [{"x": None}] * 19 + [{"x": 0.5}]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("maximum,index", [(1, 1), (1, 99), (2, 2), (20, 20), (20, 100), (1, 1_000_000)])
|
||||
def test_rejects_out_of_range_index_before_registering_padding(maximum, index):
|
||||
group = io.DynamicGroup.Input("rows", template=[io.Float.Input("x")], max=maximum)
|
||||
expanded = {"required": {}, "optional": {}, "dynamic_paths": {}, "dynamic_paths_default_value": {}, "list_paths": set()}
|
||||
with pytest.raises(ValueError, match=f"exceeds the index limit of {maximum - 1}"):
|
||||
io.DynamicGroup._expand_schema_for_dynamic(
|
||||
expanded, {f"rows.{index}.x": 0.5}, (group.io_type, group.as_dict()), "required", ["rows"],
|
||||
)
|
||||
assert expanded["dynamic_paths"] == {}
|
||||
|
||||
|
||||
def test_lazy_rows_keep_original_field_keys_and_positions():
|
||||
group = io.DynamicGroup.Input("rows", template=[io.Float.Input("x")], max=3)
|
||||
assert _reconstruct(group, {"rows.2.x": 0.5, "rows.0.x": 0.8}, lazy=True) == {"rows": [
|
||||
{"x": (0.8, "rows.0.x")},
|
||||
{"x": (None, "rows.1.x")},
|
||||
{"x": (0.5, "rows.2.x")},
|
||||
]}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("lazy", [False, True])
|
||||
def test_group_inside_dynamic_combo_preserves_other_inputs(lazy):
|
||||
group = io.DynamicGroup.Input("rows", template=[io.Float.Input("x")])
|
||||
combo = io.DynamicCombo.Input("mode", options=[io.DynamicCombo.Option("on", [group])])
|
||||
values = {"mode": "on", "mode.rows.0.x": 0.8, "fixed": "untouched"}
|
||||
assert _reconstruct(combo, values, lazy=lazy) == {
|
||||
"mode": {
|
||||
"mode": ("on", "mode") if lazy else "on",
|
||||
"rows": [{"x": (0.8, "mode.rows.0.x") if lazy else 0.8}],
|
||||
},
|
||||
"fixed": "untouched",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("kind", ["combo", "slot"])
|
||||
@pytest.mark.parametrize("lazy", [False, True])
|
||||
def test_list_mode_padding_is_limited_to_group_fields(kind, lazy):
|
||||
inputs = [
|
||||
io.Float.Input("optional", optional=True),
|
||||
io.DynamicGroup.Input("rows", template=[io.Float.Input("x")], max=2),
|
||||
]
|
||||
if kind == "combo":
|
||||
container = io.DynamicCombo.Input("mode", options=[io.DynamicCombo.Option("on", inputs)])
|
||||
selection = "on"
|
||||
else:
|
||||
container = DynamicSlot.Input(io.Float.Input("mode"), inputs=inputs)
|
||||
selection = 1.0
|
||||
values = {"mode": selection, "mode.rows.1.x": 0.5}
|
||||
_, _, metadata = get_finalized_class_inputs(create_input_dict_v1([container]), values)
|
||||
metadata["create_dynamic_tuple"] = lazy
|
||||
result = build_nested_inputs({key: [value] for key, value in values.items()}, metadata, input_is_list=True)
|
||||
|
||||
assert result == {"mode": {
|
||||
"mode": ([selection], "mode") if lazy else [selection],
|
||||
"optional": (None, "mode.optional") if lazy else None,
|
||||
"rows": [
|
||||
{"x": ([None], "mode.rows.0.x") if lazy else [None]},
|
||||
{"x": ([0.5], "mode.rows.1.x") if lazy else [0.5]},
|
||||
],
|
||||
}}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("lazy", [False, True])
|
||||
def test_autogrow_empty_value_is_unchanged(lazy):
|
||||
group = io.Autogrow.Input("items", template=io.Autogrow.TemplatePrefix(io.Float.Input("x"), prefix="item", min=0))
|
||||
assert _reconstruct(group, {}, lazy=lazy) == {"items": ({}, "items") if lazy else {}}
|
||||
@@ -0,0 +1,95 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
args.cpu = True
|
||||
|
||||
import execution
|
||||
import nodes
|
||||
from comfy_api.latest import io
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
@pytest.mark.parametrize("is_input_list", [False, True])
|
||||
@pytest.mark.parametrize("lazy", [False, True])
|
||||
@pytest.mark.parametrize("empty", [False, True])
|
||||
async def test_group_missing_fields_follow_execution_list_mode(is_input_list, lazy, empty):
|
||||
received = []
|
||||
|
||||
class Group(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(node_id=cls.__name__, is_input_list=is_input_list, inputs=[
|
||||
io.DynamicGroup.Input("rows", template=[
|
||||
io.Float.Input("x"), io.Float.Input("optional", optional=True, default=1.0),
|
||||
], max=3),
|
||||
], outputs=[])
|
||||
|
||||
@classmethod
|
||||
def execute(cls, rows):
|
||||
received.append(rows)
|
||||
return io.NodeOutput()
|
||||
|
||||
@classmethod
|
||||
def check_lazy_status(cls, rows):
|
||||
received.append(rows)
|
||||
return []
|
||||
|
||||
values = {} if empty else {"rows.0.x": 0.5, "rows.2.x": 0.8}
|
||||
inputs, _, metadata = execution.get_input_data(values, Group, "group")
|
||||
metadata["create_dynamic_tuple"] = lazy
|
||||
await execution._async_map_node_over_list(
|
||||
"test", "group", Group, inputs, "check_lazy_status" if lazy else "execute", v3_data=metadata,
|
||||
)
|
||||
expected = []
|
||||
if not empty:
|
||||
for index, value in enumerate([0.5, None, 0.8]):
|
||||
row = {"x": [value] if is_input_list else value, "optional": [None] if is_input_list else None}
|
||||
if lazy:
|
||||
row = {name: (value, f"rows.{index}.{name}") for name, value in row.items()}
|
||||
expected.append(row)
|
||||
assert received == [expected]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("values,input_name", [
|
||||
({"rows.3.x": 0.5}, "rows.3.x"),
|
||||
({"rows.bad.x": 0.5}, "rows.bad.x"),
|
||||
({}, "rows"),
|
||||
])
|
||||
@pytest.mark.parametrize("downstream", [False, True])
|
||||
async def test_group_validation_errors_identify_original_input(monkeypatch, values, input_name, downstream):
|
||||
class Group(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(node_id=cls.__name__, inputs=[
|
||||
io.DynamicGroup.Input("rows", template=[io.Float.Input("x")], min=1, max=3),
|
||||
], outputs=[io.Float.Output()], is_output_node=not downstream)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, rows):
|
||||
raise AssertionError("Invalid group must not execute")
|
||||
|
||||
class Sink(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(node_id=cls.__name__, inputs=[io.Float.Input("source")], outputs=[], is_output_node=True)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, source):
|
||||
return io.NodeOutput()
|
||||
|
||||
monkeypatch.setitem(nodes.NODE_CLASS_MAPPINGS, "Group", Group)
|
||||
monkeypatch.setitem(nodes.NODE_CLASS_MAPPINGS, "Sink", Sink)
|
||||
prompt = {"group": {"class_type": "Group", "inputs": values}}
|
||||
if downstream:
|
||||
prompt["sink"] = {"class_type": "Sink", "inputs": {"source": ["group", 0]}}
|
||||
valid, _, _, errors = await execution.validate_prompt("test", prompt, None)
|
||||
assert not valid
|
||||
error, = errors["group"]["errors"]
|
||||
assert error["type"] == "invalid_dynamic_input"
|
||||
assert error["extra_info"] == {"input_name": input_name}
|
||||
assert input_name in error["details"]
|
||||
Reference in New Issue
Block a user