From cc3ed61e097640dbcda55d98ac4b4a32b99065f2 Mon Sep 17 00:00:00 2001 From: dacert Date: Sat, 25 Jul 2026 14:17:09 +0100 Subject: [PATCH] Added start node validation --- examples/dialog_demo.py | 4 +- src/comfy_graph_bind/api.py | 31 +++- src/comfy_graph_bind/port_item.py | 24 ++- src/comfy_graph_bind/scene.py | 60 ++++++- src/comfy_graph_bind/start_node.py | 53 +++++- tests/test_dialog.py | 275 ++++++++++++++++++++++++++++- tests/test_start_node.py | 57 ++++++ 7 files changed, 489 insertions(+), 15 deletions(-) diff --git a/examples/dialog_demo.py b/examples/dialog_demo.py index 45596f7..733fbdc 100644 --- a/examples/dialog_demo.py +++ b/examples/dialog_demo.py @@ -58,9 +58,9 @@ def _config(self) -> EditorConfig: workflow_name=self._workflow_path.rsplit("/", 1)[-1], start_node_title="Krita", start_outputs=[ - StartOutputConfig("prompt", "STRING"), + StartOutputConfig("prompt", "STRING", required=True), StartOutputConfig("negative", "STRING"), - StartOutputConfig("seed", "INT"), + StartOutputConfig("seed", "INT", required=True), StartOutputConfig("image", "IMAGE"), ], ) diff --git a/src/comfy_graph_bind/api.py b/src/comfy_graph_bind/api.py index 3e4d0ac..4b2fe25 100644 --- a/src/comfy_graph_bind/api.py +++ b/src/comfy_graph_bind/api.py @@ -11,6 +11,7 @@ QApplication, QDialog, QDialogButtonBox, + QMessageBox, QToolBar, QVBoxLayout, QWidget, @@ -133,6 +134,18 @@ def get_result(self) -> dict: def clear_user_links(self) -> None: self._scene.clear_user_links() + def unwired_required_start_outputs(self) -> list: + """Return the start-node output ports that are required but unwired.""" + return self._scene.unwired_required_start_outputs() + + def required_start_outputs_with_bad_links(self) -> list: + """Return required start outputs that are unwired or wired to a missing target.""" + return self._scene.required_start_outputs_with_bad_links() + + def apply_start_output_errors(self) -> list: + """Paint required ports with bad links red; clear the rest. Returns errored ports.""" + return self._scene.apply_start_output_errors() + def _build_nodes(self, graph: WorkflowGraph) -> None: for node in graph.nodes: item = self._make_node_item(node, port_filter=self._config.port_filter) @@ -251,7 +264,7 @@ def __init__( layout.addWidget(self._editor.view()) buttons = QDialogButtonBox(QDialogButtonBox.Ok | QDialogButtonBox.Cancel) - buttons.accepted.connect(self.accept) + buttons.accepted.connect(self._on_ok) buttons.rejected.connect(self.reject) layout.addWidget(buttons) spacing = layout.spacing() @@ -259,6 +272,7 @@ def __init__( if initial_result is not None: self._apply_initial_result(initial_result) + self._editor.apply_start_output_errors() @classmethod def from_api_workflow( @@ -281,6 +295,21 @@ def editor_result(self) -> dict | None: return None return self._editor.get_result() + def _on_ok(self) -> None: + self._editor.apply_start_output_errors() + errored = self._editor.required_start_outputs_with_bad_links() + if errored: + names = ", ".join(sorted(p.ref.name for p in errored)) + QMessageBox.warning( + self, + "Missing required inputs", + "The following required start-node outputs are not wired to a valid target:\n\n" + f" {names}\n\n" + "Connect them to a target before accepting.", + ) + return + self.accept() + def _apply_initial_result(self, initial: dict) -> None: if not isinstance(initial, dict): return diff --git a/src/comfy_graph_bind/port_item.py b/src/comfy_graph_bind/port_item.py index 301563c..c9e7cc0 100644 --- a/src/comfy_graph_bind/port_item.py +++ b/src/comfy_graph_bind/port_item.py @@ -29,11 +29,16 @@ class PortRef: is_input: bool is_widget: bool = False is_start: bool = False + required: bool = False -def _port_color(is_input: bool, is_start: bool) -> QColor: +_START_OUTPUT_COLOR = QColor("#3fb950") +_START_OUTPUT_ERROR_COLOR = QColor("#d63a3a") + + +def _port_color(is_input: bool, is_start: bool, has_error: bool = False) -> QColor: if is_start and not is_input: - return QColor("#3fb950") + return _START_OUTPUT_ERROR_COLOR if has_error else _START_OUTPUT_COLOR return QColor("#4a9eff") if not is_input else QColor("#ff9a4a") @@ -49,7 +54,8 @@ def __init__( self._ref = ref self._radius = radius self._hovered = False - color = _port_color(ref.is_input, ref.is_start) + self._has_error = False + color = _port_color(ref.is_input, ref.is_start, self._has_error) self.setBrush(QBrush(color)) self.setPen(QPen(color.darker(140), 1.2)) self.setAcceptHoverEvents(True) @@ -62,6 +68,16 @@ def __init__( def ref(self) -> PortRef: return self._ref + def set_error(self, has_error: bool) -> None: + """Mark or clear the port's error state and repaint it.""" + if has_error == self._has_error: + return + self._has_error = has_error + color = _port_color(self._ref.is_input, self._ref.is_start, self._has_error) + self.setBrush(QBrush(color)) + self.setPen(QPen(color.darker(160), 1.4)) + self.update() + def scene_center(self) -> QPointF: return self.scenePos() @@ -84,7 +100,7 @@ def paint( # noqa: D401 - Qt API ) -> None: # type: ignore[override] del option, widget painter.setRenderHint(painter.Antialiasing, True) - color = _port_color(self._ref.is_input, self._ref.is_start) + color = _port_color(self._ref.is_input, self._ref.is_start, self._has_error) r = PORT_HOVER_RADIUS if self._hovered else self._radius painter.setBrush(QBrush(color)) painter.setPen(QPen(color.darker(160), 1.4)) diff --git a/src/comfy_graph_bind/scene.py b/src/comfy_graph_bind/scene.py index 56bf94a..10fbf1a 100644 --- a/src/comfy_graph_bind/scene.py +++ b/src/comfy_graph_bind/scene.py @@ -98,18 +98,69 @@ def create_link(self, source: PortRef, target: PortRef) -> LinkItem | None: self._links[link_id] = link self.addItem(link) self._refresh_link(link) + if source.is_start: + self._refresh_required_errors() return link def remove_link(self, link_id: int) -> None: link = self._links.pop(link_id, None) - if link is not None: - self.removeItem(link) + if link is None: + return + was_start = link.source.is_start + self.removeItem(link) + if was_start: + self._refresh_required_errors() def clear_user_links(self) -> None: + had_start = any(link.source.is_start for link in self._links.values()) for link in list(self._links.values()): if link.source.is_start: self.removeItem(link) self._links.pop(link.link_id, None) + if had_start: + self._refresh_required_errors() + + def required_start_outputs_with_bad_links(self) -> list[PortItem]: + """Required start outputs that are unwired or wired to a missing target.""" + required_names = self._start_node.required_output_names + if not required_names: + return [] + bad: list[PortItem] = [] + for port in self._start_node.output_ports(): + if port.ref.name not in required_names: + continue + link = self._find_link_for_start_output(port.ref) + if link is None or not self._target_input_resolves(link): + bad.append(port) + return bad + + def unwired_required_start_outputs(self) -> list[PortItem]: + """Return required start outputs that have no outgoing link at all.""" + return [ + p + for p in self.required_start_outputs_with_bad_links() + if self._find_link_for_start_output(p.ref) is None + ] + + def apply_start_output_errors(self) -> list[PortItem]: + """Paint all required start outputs with bad links red; clear the rest.""" + bad = self.required_start_outputs_with_bad_links() + bad_set = {id(p) for p in bad} + for port in self._start_node.output_ports(): + port.set_error(id(port) in bad_set) + return bad + + def _target_input_resolves(self, link: LinkItem) -> bool: + node = self._node_items.get(link.target.node_id) + if node is None: + return False + return any(p.ref.name == link.target.name for p in node.input_ports()) + + def _refresh_required_errors(self) -> None: + bad = self.required_start_outputs_with_bad_links() + bad_set = {id(p) for p in bad} + for port in self._start_node.output_ports(): + port.set_error(id(port) in bad_set) def mousePressEvent(self, event: QGraphicsSceneMouseEvent) -> None: view = next(iter(self.views()), None) @@ -133,14 +184,19 @@ def mouseReleaseEvent(self, event: QGraphicsSceneMouseEvent) -> None: view = next(iter(self.views()), None) transform = view.transform() if view is not None else QTransform() target_item = self.itemAt(event.scenePos(), transform) + touched_start = False if isinstance(target_item, PortItem) and target_item.ref.is_input: self.create_link(self._drag_source.ref, target_item.ref) + touched_start = True elif self._drag_source.ref.is_start: existing = self._find_link_for_start_output(self._drag_source.ref) if existing is not None: self.removeItem(existing) self._links.pop(existing.link_id, None) + touched_start = True self._end_link_drag() + if touched_start: + self._refresh_required_errors() event.accept() return super().mouseReleaseEvent(event) diff --git a/src/comfy_graph_bind/start_node.py b/src/comfy_graph_bind/start_node.py index 90b034e..346ce01 100644 --- a/src/comfy_graph_bind/start_node.py +++ b/src/comfy_graph_bind/start_node.py @@ -9,7 +9,7 @@ from collections.abc import Sequence from dataclasses import dataclass -from PyQt5.QtCore import QPointF +from PyQt5.QtCore import QPointF, Qt from PyQt5.QtWidgets import QGraphicsScene from .layout import node_metrics @@ -23,9 +23,11 @@ class StartOutputConfig: name: str type: str = "*" + required: bool = False def label(self) -> str: - return f"{self.name} : {self.type}" + suffix = " *" if self.required else "" + return f"{self.name}{suffix} : {self.type}" class StartNodeItem(GraphNodeItem): @@ -35,6 +37,8 @@ class StartNodeItem(GraphNodeItem): START_POS = QPointF(-260, 0) def __init__(self, outputs: Sequence[StartOutputConfig], title: str = "Start") -> None: + self._output_names: list[str] = [o.name for o in outputs] + self._required_outputs: set[str] = {o.name for o in outputs if o.required} inputs: list[tuple[int, str, str, bool]] = [] output_specs: list[tuple[int, str, str, bool]] = [ (i, o.name, o.type, False) for i, o in enumerate(outputs) @@ -49,12 +53,20 @@ def __init__(self, outputs: Sequence[StartOutputConfig], title: str = "Start") - outputs=output_specs, is_start=True, ) - self._output_names: list[str] = [o.name for o in outputs] + self._propagate_required_to_ports() + self._refresh_output_labels() @property def output_names(self) -> list[str]: return list(self._output_names) + @property + def required_output_names(self) -> set[str]: + return set(self._required_outputs) + + def required_output_label(self, name: str) -> str: + return f"{name} *" if name in self._required_outputs else name + def rebuild_outputs(self, outputs: Sequence[StartOutputConfig]) -> None: """Replace the current outputs and reposition the ports.""" from .scene import GraphScene @@ -64,11 +76,13 @@ def rebuild_outputs(self, outputs: Sequence[StartOutputConfig]) -> None: scene.detach_links_for_node(self) self._port_items.clear() self._output_names = [o.name for o in outputs] - self._inputs = [] + self._required_outputs = {o.name for o in outputs if o.required} stub = _StubNode("Start", self.title, outputs) self._outputs = build_node_outputs(stub) - output_labels = [o.name for o in outputs] + output_labels = [ + self.required_output_label(o.name) if o.required else o.name for o in outputs + ] self._metrics = node_metrics( stub, input_count=0, @@ -84,6 +98,7 @@ def rebuild_outputs(self, outputs: Sequence[StartOutputConfig]) -> None: is_input=False, is_widget=False, is_start=True, + required=o.required, ) port = PortItem(ref) port.setParentItem(self) @@ -93,6 +108,34 @@ def rebuild_outputs(self, outputs: Sequence[StartOutputConfig]) -> None: if isinstance(scene, QGraphicsScene): scene.update() + def _propagate_required_to_ports(self) -> None: + for port in self._port_items: + ref = port.ref + if ref.is_input: + continue + if ref.name in self._required_outputs and not ref.required: + port.setData( + Qt.UserRole, + PortRef( + node_id=ref.node_id, + slot_index=ref.slot_index, + name=ref.name, + type=ref.type, + is_input=ref.is_input, + is_widget=ref.is_widget, + is_start=ref.is_start, + required=True, + ), + ) + + def _refresh_output_labels(self) -> None: + for i, (_idx, _name, _type, _is_widget) in enumerate(self._outputs): + if i < len(self._output_names): + name = self._output_names[i] + decorated = self.required_output_label(name) + if decorated != _name: + self._outputs[i] = (i, decorated, _type, _is_widget) + class _StubNode: def __init__(self, type_name: str, title: str, outputs: Sequence[StartOutputConfig]) -> None: diff --git a/tests/test_dialog.py b/tests/test_dialog.py index 62d7dbd..bf576f4 100644 --- a/tests/test_dialog.py +++ b/tests/test_dialog.py @@ -2,8 +2,10 @@ from __future__ import annotations +from unittest.mock import patch + import pytest -from PyQt5.QtWidgets import QToolBar +from PyQt5.QtWidgets import QDialog, QToolBar from comfy_graph_bind import EditorConfig, GraphEditorDialog, StartOutputConfig @@ -182,3 +184,274 @@ def test_dialog_clear_links_action_wires_to_editor(qtbot, sample_graph): clear = next(a for a in toolbar.actions() if a.text() == "Clear start links") clear.trigger() assert dialog.editor().get_result()["inputs"] == {} + + +def test_dialog_marks_unwired_required_ports_on_load(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + StartOutputConfig("seed", "INT", required=True), + ], + initial_result={ + "workflow_name": "sample", + "inputs": {"prompt": {"node_id": "1", "property": "positive"}}, + }, + ) + qtbot.addWidget(dialog) + prompt = dialog.editor().start_node().output_port(0) + seed = dialog.editor().start_node().output_port(1) + assert not prompt._has_error + assert seed._has_error + + +def test_dialog_clears_errors_once_all_required_wired(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + StartOutputConfig("seed", "INT", required=True), + ], + initial_result={ + "workflow_name": "sample", + "inputs": { + "prompt": {"node_id": "1", "property": "positive"}, + "seed": {"node_id": "1", "property": "seed"}, + }, + }, + ) + qtbot.addWidget(dialog) + prompt = dialog.editor().start_node().output_port(0) + seed = dialog.editor().start_node().output_port(1) + assert not prompt._has_error + assert not seed._has_error + + +def test_dialog_ok_blocks_when_required_missing(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + StartOutputConfig("seed", "INT", required=True), + ], + initial_result=None, + ) + qtbot.addWidget(dialog) + with patch("comfy_graph_bind.api.QMessageBox.warning") as warn: + dialog._on_ok() + assert warn.called + assert dialog.result() != QDialog.Accepted + assert dialog.editor_result() is None + + +def test_dialog_ok_accepts_when_all_required_wired(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + StartOutputConfig("seed", "INT", required=True), + ], + initial_result={ + "workflow_name": "sample", + "inputs": { + "prompt": {"node_id": "1", "property": "positive"}, + "seed": {"node_id": "1", "property": "seed"}, + }, + }, + ) + qtbot.addWidget(dialog) + with patch("comfy_graph_bind.api.QMessageBox.warning") as warn: + dialog._on_ok() + assert not warn.called + assert dialog.result() == QDialog.Accepted + result = dialog.editor_result() + assert result is not None + assert set(result["inputs"].keys()) == {"prompt", "seed"} + + +def test_dialog_ok_accepts_when_no_required_outputs(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [StartOutputConfig("prompt", "STRING")], + initial_result=None, + ) + qtbot.addWidget(dialog) + with patch("comfy_graph_bind.api.QMessageBox.warning") as warn: + dialog._on_ok() + assert not warn.called + assert dialog.result() == QDialog.Accepted + + +def test_unwired_required_helper_returns_only_required(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + StartOutputConfig("negative", "STRING"), + StartOutputConfig("seed", "INT", required=True), + ], + initial_result=None, + ) + qtbot.addWidget(dialog) + scene = dialog.editor().scene() + missing = scene.unwired_required_start_outputs() + names = sorted(p.ref.name for p in missing) + assert names == ["prompt", "seed"] + + +def test_unwired_required_helper_after_partial_wiring(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + StartOutputConfig("seed", "INT", required=True), + ], + initial_result={ + "workflow_name": "sample", + "inputs": {"prompt": {"node_id": "1", "property": "positive"}}, + }, + ) + qtbot.addWidget(dialog) + scene = dialog.editor().scene() + missing = scene.unwired_required_start_outputs() + names = [p.ref.name for p in missing] + assert names == ["seed"] + + +def test_dialog_does_not_mark_required_red_on_open_without_initial_result(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + StartOutputConfig("seed", "INT", required=True), + ], + initial_result=None, + ) + qtbot.addWidget(dialog) + prompt = dialog.editor().start_node().output_port(0) + seed = dialog.editor().start_node().output_port(1) + assert not prompt._has_error + assert not seed._has_error + + +def test_dialog_marks_required_red_on_load_when_target_missing_in_graph(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + ], + initial_result={ + "workflow_name": "sample", + "inputs": {"prompt": {"node_id": "1", "property": "does-not-exist"}}, + }, + ) + qtbot.addWidget(dialog) + prompt = dialog.editor().start_node().output_port(0) + assert prompt._has_error + + +def test_dialog_clears_error_after_valid_link_replaces_bad_one(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + ], + initial_result={ + "workflow_name": "sample", + "inputs": {"prompt": {"node_id": "1", "property": "does-not-exist"}}, + }, + ) + qtbot.addWidget(dialog) + scene = dialog.editor().scene() + start = dialog.editor().start_node() + prompt = start.output_port(0) + assert prompt._has_error + + target_node = scene.get_node_item(1) + valid_port = next(p for p in target_node.input_ports() if p.ref.name == "positive") + scene.create_link(prompt.ref, valid_port.ref) + assert not prompt._has_error + + +def test_dialog_turns_red_after_required_link_removed(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + ], + initial_result={ + "workflow_name": "sample", + "inputs": {"prompt": {"node_id": "1", "property": "positive"}}, + }, + ) + qtbot.addWidget(dialog) + scene = dialog.editor().scene() + start = dialog.editor().start_node() + prompt = start.output_port(0) + assert not prompt._has_error + + start_links = [lnk for lnk in scene.links() if lnk.source.is_start] + assert len(start_links) == 1 + scene.remove_link(start_links[0].link_id) + assert prompt._has_error + + +def test_required_with_bad_links_helper_resolves_targets(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + StartOutputConfig("seed", "INT", required=True), + StartOutputConfig("negative", "STRING"), + ], + initial_result={ + "workflow_name": "sample", + "inputs": { + "prompt": {"node_id": "1", "property": "positive"}, + "seed": {"node_id": "1", "property": "does-not-exist"}, + "negative": {"node_id": "1", "property": "negative"}, + }, + }, + ) + qtbot.addWidget(dialog) + scene = dialog.editor().scene() + bad = scene.required_start_outputs_with_bad_links() + bad_names = sorted(p.ref.name for p in bad) + assert bad_names == ["seed"] + + +def test_dialog_clear_links_turns_required_red(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + ], + initial_result={ + "workflow_name": "sample", + "inputs": {"prompt": {"node_id": "1", "property": "positive"}}, + }, + ) + qtbot.addWidget(dialog) + prompt = dialog.editor().start_node().output_port(0) + assert not prompt._has_error + dialog.editor().clear_user_links() + assert prompt._has_error + + +def test_dialog_ok_blocks_when_required_linked_to_missing_target(qtbot, sample_graph): + dialog = _make_dialog( + sample_graph, + [ + StartOutputConfig("prompt", "STRING", required=True), + ], + initial_result={ + "workflow_name": "sample", + "inputs": {"prompt": {"node_id": "1", "property": "does-not-exist"}}, + }, + ) + qtbot.addWidget(dialog) + with patch("comfy_graph_bind.api.QMessageBox.warning") as warn: + dialog._on_ok() + assert warn.called + assert dialog.result() != QDialog.Accepted + assert dialog.editor_result() is None diff --git a/tests/test_start_node.py b/tests/test_start_node.py index e7c7f12..417dad3 100644 --- a/tests/test_start_node.py +++ b/tests/test_start_node.py @@ -50,3 +50,60 @@ def _fake_graph(): from comfy_graph_bind.workflow_parser import WorkflowGraph return WorkflowGraph(nodes=[]) + + +def test_required_flag_defaults_to_false(): + _ensure_qapp() + cfg = StartOutputConfig("prompt", "STRING") + assert cfg.required is False + assert cfg.label() == "prompt : STRING" + + +def test_required_flag_renders_star_in_label(): + _ensure_qapp() + cfg = StartOutputConfig("prompt", "STRING", required=True) + assert cfg.required is True + assert cfg.label() == "prompt * : STRING" + + +def test_start_node_required_names_exposed(): + _ensure_qapp() + editor = GraphEditor( + _fake_graph(), + EditorConfig( + workflow_name="x", + start_outputs=[ + StartOutputConfig("a", required=True), + StartOutputConfig("b"), + StartOutputConfig("c", required=True), + ], + ), + ) + assert editor.start_node().required_output_names == {"a", "c"} + + +def test_rebuild_outputs_preserves_required_flag(): + _ensure_qapp() + editor = GraphEditor( + _fake_graph(), + EditorConfig(workflow_name="x", start_outputs=[StartOutputConfig("a")]), + ) + editor.set_start_outputs([ + StartOutputConfig("x", required=True), + StartOutputConfig("y"), + ]) + assert editor.start_node().required_output_names == {"x"} + + +def test_required_label_appears_in_node_outputs(): + _ensure_qapp() + editor = GraphEditor( + _fake_graph(), + EditorConfig( + workflow_name="x", + start_outputs=[StartOutputConfig("prompt", "STRING", required=True)], + ), + ) + start = editor.start_node() + (_idx, name, _type, _is_widget) = start._outputs[0] + assert name == "prompt *"