diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..c5a24d3 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,31 @@ +name: CI + +on: + pull_request: + branches: [master] + +jobs: + ci: + strategy: + matrix: + python-version: ["3.10", "3.12"] + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + - name: Install system deps (PyQt5) + run: sudo apt-get update && sudo apt-get install -y libxcb-cursor0 + - name: Install package + run: pip install -e ".[dev]" + - name: Install ty + run: pip install ty + - name: Lint + run: ruff check . + - name: Format check + run: ruff format --check . + - name: Type check + run: QT_QPA_PLATFORM=offscreen python -m ty check src tests + - name: Test + run: QT_QPA_PLATFORM=offscreen pytest -v diff --git a/.gitignore b/.gitignore index aef3281..e388711 100644 --- a/.gitignore +++ b/.gitignore @@ -23,6 +23,7 @@ MANIFEST # Virtual environments .venv/ +.venv_win venv/ env/ ENV/ diff --git a/AGENTS.md b/AGENTS.md index a02af43..2248ad4 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1,9 +1,15 @@ # Commands used in this repository +## Environment +```bash +python3 -m venv .venv +source .venv/bin/activate +``` + ## Install (editable + dev deps) ```bash -pip install -e . pip install -e ".[dev]" +pip install ty ``` ## Lint / format @@ -14,7 +20,7 @@ ruff format . ## Type check ```bash -PYTHONPATH=$(python3 -c "import site; print(site.getusersitepackages())") ty check src tests +QT_QPA_PLATFORM=offscreen python -m ty check src tests ``` ## Tests diff --git a/README.md b/README.md index ba2b3b3..8861648 100644 --- a/README.md +++ b/README.md @@ -173,4 +173,8 @@ so the user can immediately see the existing wiring. Any entry that does not resolve to a known node/property is silently ignored. A runnable version of this example lives at -[`examples/dialog_demo.py`](examples/dialog_demo.py). \ No newline at end of file +[`examples/dialog_demo.py`](examples/dialog_demo.py). + +Used in the krita_comfyui plugin + +![Select Workflow](media/example.png) \ No newline at end of file diff --git a/media/example.png b/media/example.png new file mode 100644 index 0000000..e6540ff Binary files /dev/null and b/media/example.png differ diff --git a/pyproject.toml b/pyproject.toml index d3096b6..c32bd9a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -19,6 +19,8 @@ dev = [ "pytest-qt>=4", "pytest-mock>=3", "ruff>=0.6", + "ty>=0.0.61", + "PyQt5-stubs>=5.15", ] [tool.setuptools] diff --git a/src/comfy_graph_bind/api.py b/src/comfy_graph_bind/api.py index 5893444..251fdaf 100644 --- a/src/comfy_graph_bind/api.py +++ b/src/comfy_graph_bind/api.py @@ -22,8 +22,8 @@ from .scene import GraphScene from .start_node import StartNodeItem, StartOutputConfig from .view import GraphView -from .workflow_parser import WorkflowGraph, WorkflowNode -from .workflow_parser import _is_link as _is_api_link +from .workflow_parser import WorkflowGraph, WorkflowNode, parse_workflow_dict +from .workflow_parser import is_link as _is_api_link @dataclass @@ -46,7 +46,6 @@ def __init__(self, graph: WorkflowGraph, config: EditorConfig) -> None: self._graph = graph self._config = config self._app: QApplication | None = None - self._owns_app = False self._ensure_qapp() self._start_node = StartNodeItem(config.start_outputs, title=config.start_node_title) self._scene = GraphScene([], self._start_node, port_filter=config.port_filter) @@ -60,8 +59,6 @@ def from_api_workflow( raw_api: dict, config: EditorConfig, ) -> GraphEditor: - from .workflow_parser import parse_workflow_dict - graph = parse_workflow_dict(raw_api) return cls(graph, config) @@ -83,8 +80,8 @@ def workflow_name(self) -> str: def set_start_node_title(self, title: str) -> None: self._config.start_node_title = title self._start_node.title = title - self._start_node.rebuild_outputs(self._config.start_outputs) - self._scene.setSceneRect(self._scene._compute_scene_rect()) + self._start_node.update() + self._scene.update_scene_rect() def start_node_title(self) -> str: return self._config.start_node_title @@ -92,10 +89,11 @@ def start_node_title(self) -> str: def set_start_outputs(self, outputs: Sequence[StartOutputConfig]) -> None: self._config.start_outputs = list(outputs) self._start_node.rebuild_outputs(outputs) - self._scene.setSceneRect(self._scene._compute_scene_rect()) + self._scene.update_scene_rect() - def add_output(self, name: str, type: str = "*") -> None: - outputs = list(self._config.start_outputs) + [StartOutputConfig(name=name, type=type)] + def add_output(self, name: str, output_type: str = "*") -> None: + new_output = StartOutputConfig(name=name, type=output_type) + outputs = list(self._config.start_outputs) + [new_output] self.set_start_outputs(outputs) def remove_output(self, name: str) -> None: @@ -115,8 +113,8 @@ def port_filter(self) -> PortFilter: def reset_view(self) -> None: for item in self._scene.node_items(): item.reset_position() + self._scene.update_scene_rect() self._view.reset_view() - self._scene.setSceneRect(self._scene._compute_scene_rect()) def get_result(self) -> dict: inputs: dict[str, dict[str, str]] = {} @@ -142,7 +140,7 @@ def _build_nodes(self, graph: WorkflowGraph) -> None: self._node_items[node.id] = item self._scene.add_node_item(item, data=node) self._create_existing_links(graph) - self._scene.setSceneRect(self._scene._compute_scene_rect()) + self._scene.update_scene_rect() def _make_node_item(self, node: WorkflowNode, port_filter: PortFilter) -> GraphNodeItem: inputs = build_node_inputs(node, port_filter) @@ -165,10 +163,9 @@ def _create_existing_links(self, graph: WorkflowGraph) -> None: for name, value in node.raw_inputs.items(): if not _is_api_link(value): continue - src_id = value[0] - src_idx = int(value[1]) with contextlib.suppress(ValueError, TypeError): - src_id = int(src_id) + src_id = int(value[0]) + src_idx = int(value[1]) if src_id not in self._node_items: continue @@ -194,11 +191,9 @@ def _create_existing_links(self, graph: WorkflowGraph) -> None: def _ensure_qapp(self) -> None: app = QApplication.instance() if app is None: - self._app = QApplication([]) - self._owns_app = True - else: - assert isinstance(app, QApplication) - self._app = app + app = QApplication([]) + assert isinstance(app, QApplication) + self._app = app class GraphEditorDialog(QDialog): @@ -272,8 +267,6 @@ def from_api_workflow( initial_result: dict | None = None, parent: QWidget | None = None, ) -> GraphEditorDialog: - from .workflow_parser import parse_workflow_dict - graph = parse_workflow_dict(raw_api) return cls(graph, config, initial_result=initial_result, parent=parent) diff --git a/src/comfy_graph_bind/layout.py b/src/comfy_graph_bind/layout.py index 238a2ff..e274242 100644 --- a/src/comfy_graph_bind/layout.py +++ b/src/comfy_graph_bind/layout.py @@ -41,7 +41,7 @@ def _max_label_width(labels: Sequence[str] | None) -> float: try: fm = QFontMetrics(_label_font()) return float(max(fm.horizontalAdvance(s) for s in labels)) - except Exception: + except (RuntimeError, ValueError): return max(len(s) * LABEL_CHAR_PX for s in labels) diff --git a/src/comfy_graph_bind/node_item.py b/src/comfy_graph_bind/node_item.py index b42c8c5..a36f828 100644 --- a/src/comfy_graph_bind/node_item.py +++ b/src/comfy_graph_bind/node_item.py @@ -2,8 +2,6 @@ from __future__ import annotations -from collections.abc import Callable - from PyQt5.QtCore import QPointF, QRectF, Qt from PyQt5.QtGui import QBrush, QColor, QFont, QPainter, QPen from PyQt5.QtWidgets import ( @@ -113,7 +111,7 @@ def itemChange(self, change, value): # noqa: N802 - Qt API scene = self.scene() if isinstance(scene, GraphScene): - scene.node_moved(self) + scene.node_moved(self, QPointF(value)) return super().itemChange(change, value) def mousePressEvent(self, event: QGraphicsSceneMouseEvent) -> None: # noqa: N802 - Qt API @@ -245,6 +243,3 @@ def build_node_outputs(node) -> list[tuple[int, str, str, bool]]: continue out.append((i, s.name or "", str(s.type) if s.type is not None else "*", False)) return out - - -CallableType = Callable[[QPointF, QPointF], None] diff --git a/src/comfy_graph_bind/port_item.py b/src/comfy_graph_bind/port_item.py index 95f65b8..301563c 100644 --- a/src/comfy_graph_bind/port_item.py +++ b/src/comfy_graph_bind/port_item.py @@ -15,6 +15,8 @@ from PyQt5.QtGui import QBrush, QColor, QPen from PyQt5.QtWidgets import QGraphicsEllipseItem, QStyleOptionGraphicsItem, QWidget +from .layout import PORT_HOVER_RADIUS + @dataclass(frozen=True) class PortRef: @@ -28,10 +30,6 @@ class PortRef: is_widget: bool = False is_start: bool = False - @property - def key(self) -> tuple[int | str, int, bool]: - return (self.node_id, self.slot_index, self.is_input) - def _port_color(is_input: bool, is_start: bool) -> QColor: if is_start and not is_input: @@ -91,7 +89,3 @@ def paint( # noqa: D401 - Qt API painter.setBrush(QBrush(color)) painter.setPen(QPen(color.darker(160), 1.4)) painter.drawEllipse(QPointF(0, 0), r, r) - - -# Late import to avoid a cycle between PortItem and layout helpers. -from .layout import PORT_HOVER_RADIUS # noqa: E402 diff --git a/src/comfy_graph_bind/scene.py b/src/comfy_graph_bind/scene.py index 31b24ca..56bf94a 100644 --- a/src/comfy_graph_bind/scene.py +++ b/src/comfy_graph_bind/scene.py @@ -24,7 +24,7 @@ def __init__( self, nodes: Iterable[GraphNodeItem], start_node: StartNodeItem, - port_filter: PortFilter = PortFilter.LINKED_ONLY, + port_filter: PortFilter = PortFilter.ALL_INCLUDE_WIDGETS, ) -> None: super().__init__() self.setBackgroundBrush(QBrush(QColor("#121212"))) @@ -67,12 +67,11 @@ def add_node_item(self, item: GraphNodeItem, data: Any | None = None) -> None: if data is not None: self._node_data[item.node_id] = data self.addItem(item) - for port in item.port_items: - del port - def node_moved(self, _node: GraphNodeItem) -> None: + def node_moved(self, _node: GraphNodeItem, new_scene_pos: QPointF) -> None: for link in self._links.values(): - self._refresh_link(link) + if link.source.node_id == _node.node_id or link.target.node_id == _node.node_id: + self._refresh_link(link, _node, new_scene_pos) def detach_links_for_node(self, node: GraphNodeItem) -> None: for link in list(self._links.values()): @@ -113,7 +112,7 @@ def clear_user_links(self) -> None: self._links.pop(link.link_id, None) def mousePressEvent(self, event: QGraphicsSceneMouseEvent) -> None: - view = self.views()[0] if self.views() else None + view = next(iter(self.views()), None) transform = view.transform() if view is not None else QTransform() item = self.itemAt(event.scenePos(), transform) if isinstance(item, PortItem) and not item.ref.is_input and item.ref.is_start: @@ -131,7 +130,7 @@ def mouseMoveEvent(self, event: QGraphicsSceneMouseEvent) -> None: def mouseReleaseEvent(self, event: QGraphicsSceneMouseEvent) -> None: if self._drag_source is not None and self._temp_link is not None: - view = self.views()[0] if self.views() else None + view = next(iter(self.views()), None) transform = view.transform() if view is not None else QTransform() target_item = self.itemAt(event.scenePos(), transform) if isinstance(target_item, PortItem) and target_item.ref.is_input: @@ -160,12 +159,29 @@ def _end_link_drag(self) -> None: self._temp_link = None self._drag_source = None - def _refresh_link(self, link: LinkItem) -> None: + def _refresh_link( + self, + link: LinkItem, + moved_node: GraphNodeItem | None = None, + moved_node_new_scene_pos: QPointF | None = None, + ) -> None: src_item = self._find_port_item(link.source) dst_item = self._find_port_item(link.target) if src_item is None or dst_item is None: return - link.set_endpoints(src_item.scene_center(), dst_item.scene_center()) + if moved_node is not None and moved_node_new_scene_pos is not None: + if link.source.node_id == moved_node.node_id: + src_center = src_item.pos() + moved_node_new_scene_pos + else: + src_center = src_item.scene_center() + if link.target.node_id == moved_node.node_id: + dst_center = dst_item.pos() + moved_node_new_scene_pos + else: + dst_center = dst_item.scene_center() + else: + src_center = src_item.scene_center() + dst_center = dst_item.scene_center() + link.set_endpoints(src_center, dst_center) def _find_port_item(self, ref: PortRef) -> PortItem | None: if ref.is_start and not ref.is_input: @@ -200,15 +216,6 @@ def _find_start_link_for_target(self, target: PortRef) -> LinkItem | None: return link return None - def _find_link_for_target(self, target: PortRef) -> LinkItem | None: - for link in self._links.values(): - if ( - link.target.node_id == target.node_id - and link.target.slot_index == target.slot_index - ): - return link - return None - def _rebuild_node_ports(self) -> None: for node_item in self._node_items.values(): data = self._node_data.get(node_item.node_id) @@ -236,6 +243,8 @@ def _rebuild_node_ports(self) -> None: node_item._build_ports() node_item.update() + def update_scene_rect(self) -> None: + self.setSceneRect(self._compute_scene_rect()) + def _compute_scene_rect(self) -> QRectF: - rect = self.itemsBoundingRect().adjusted(-200, -200, 200, 200) - return QRectF(rect) + return self.itemsBoundingRect().adjusted(-200, -200, 200, 200) diff --git a/src/comfy_graph_bind/start_node.py b/src/comfy_graph_bind/start_node.py index 0dbb3f1..90b034e 100644 --- a/src/comfy_graph_bind/start_node.py +++ b/src/comfy_graph_bind/start_node.py @@ -12,7 +12,9 @@ from PyQt5.QtCore import QPointF from PyQt5.QtWidgets import QGraphicsScene -from .node_item import GraphNodeItem +from .layout import node_metrics +from .node_item import GraphNodeItem, build_node_outputs +from .port_item import PortItem, PortRef @dataclass(frozen=True) @@ -60,15 +62,9 @@ def rebuild_outputs(self, outputs: Sequence[StartOutputConfig]) -> None: scene = self.scene() if isinstance(scene, GraphScene): scene.detach_links_for_node(self) - old_ports = list(self.output_ports()) - for p in old_ports: - if p.parentItem() is self: - scene.removeItem(p) if scene is not None else None + self._port_items.clear() self._output_names = [o.name for o in outputs] self._inputs = [] - from .layout import node_metrics - from .node_item import build_node_outputs # local import to avoid cycle - from .port_item import PortItem, PortRef stub = _StubNode("Start", self.title, outputs) self._outputs = build_node_outputs(stub) @@ -94,7 +90,7 @@ def rebuild_outputs(self, outputs: Sequence[StartOutputConfig]) -> None: port.setPos(QPointF(self._metrics.width, self._metrics.port_y(i))) self._port_items.append(port) self.update() - if scene is not None and isinstance(scene, QGraphicsScene): + if isinstance(scene, QGraphicsScene): scene.update() diff --git a/src/comfy_graph_bind/view.py b/src/comfy_graph_bind/view.py index 608cd4d..7ad4cec 100644 --- a/src/comfy_graph_bind/view.py +++ b/src/comfy_graph_bind/view.py @@ -6,6 +6,8 @@ from PyQt5.QtGui import QKeySequence from PyQt5.QtWidgets import QGraphicsView, QShortcut +from .scene import GraphScene + class GraphView(QGraphicsView): """A view that supports zooming, panning and resetting the camera.""" @@ -16,7 +18,6 @@ class GraphView(QGraphicsView): def __init__(self, scene) -> None: super().__init__(scene) - self.setRenderHints(self.renderHints() | self.renderHints()) from PyQt5.QtGui import QPainter self.setRenderHint(QPainter.Antialiasing, True) @@ -95,11 +96,10 @@ def mouseReleaseEvent(self, event) -> None: # noqa: N802 - Qt API # ------------------------------------------------------------------ def reset_view(self) -> None: """Restore the default zoom (1.0) and recenter the scene.""" - self.resetTransform() - from .scene import GraphScene - scene = self.scene() + self.resetTransform() if isinstance(scene, GraphScene) and scene.start_node() is not None: self.centerOn(scene.start_node()) else: self.centerOn(QPointF(0, 0)) + self.viewport().update() diff --git a/src/comfy_graph_bind/workflow_parser.py b/src/comfy_graph_bind/workflow_parser.py index fa61bce..2304d3f 100644 --- a/src/comfy_graph_bind/workflow_parser.py +++ b/src/comfy_graph_bind/workflow_parser.py @@ -33,7 +33,7 @@ class WorkflowGraph: def parse_workflow_json(path: str | pathlib.Path) -> WorkflowGraph: - with open(path) as f: + with open(path, encoding="utf-8") as f: data = json.load(f) return _parse_workflow_dict(data) @@ -43,6 +43,9 @@ def parse_workflow_dict(data: dict) -> WorkflowGraph: def _parse_workflow_dict(data: dict) -> WorkflowGraph: + if not isinstance(data, dict): + raise TypeError(f"Expected a dict, got {type(data).__name__}") + raw_inputs_map: dict[int | str, dict[str, Any]] = {} for node_id_str, node_data in data.items(): @@ -52,7 +55,7 @@ def _parse_workflow_dict(data: dict) -> WorkflowGraph: output_refs: dict[int | str, set[int]] = defaultdict(set) for _node_id, node_inputs in raw_inputs_map.items(): for value in node_inputs.values(): - if _is_link(value): + if is_link(value): src_id, src_idx = _parse_link(value) if src_id in raw_inputs_map: output_refs[src_id].add(src_idx) @@ -67,7 +70,7 @@ def _parse_workflow_dict(data: dict) -> WorkflowGraph: input_slots: list[Slot] = [] for name, value in node_inputs.items(): - if _is_link(value): + if is_link(value): slot = Slot(name=name, type="*", link=1, widget=None) else: slot = Slot(name=name, type=type(value).__name__, link=None, widget=value) @@ -106,8 +109,13 @@ def _parse_node_id(s: str) -> int | str: return s -def _is_link(value: Any) -> bool: - return isinstance(value, list) and len(value) == 2 +def is_link(value: Any) -> bool: + return ( + isinstance(value, list) + and len(value) == 2 + and isinstance(value[0], (int, str)) + and isinstance(value[1], (int, str)) + ) def _parse_link(value: list) -> tuple[int | str, int]: @@ -130,7 +138,7 @@ def _auto_layout( if nid not in node_ids: continue for value in n_inputs.values(): - if _is_link(value): + if is_link(value): src_id, _ = _parse_link(value) if src_id in node_ids: deps[nid].add(src_id) diff --git a/tests/test_link_creation.py b/tests/test_link_creation.py index 972aeb3..15e2c07 100644 --- a/tests/test_link_creation.py +++ b/tests/test_link_creation.py @@ -2,8 +2,13 @@ from __future__ import annotations +from PyQt5.QtCore import Qt + from comfy_graph_bind import EditorConfig, GraphEditor, PortFilter, StartOutputConfig +Qt_LeftButton = Qt.LeftButton +Qt_NoModifier = Qt.NoModifier + def _make_editor(qtbot, sample_graph): editor = GraphEditor( @@ -30,16 +35,17 @@ def test_create_link_via_scene_api(qtbot, sample_graph): def test_start_output_replaces_previous_link(qtbot, sample_graph): editor = _make_editor(qtbot, sample_graph) + initial_count = len(editor.scene().links()) prompt = editor.start_node().output_port(0) target_node = editor.scene().get_node_item(1) port_a = target_node.input_port(1) port_b = target_node.input_port(2) link_a = editor.scene().create_link(prompt.ref, port_a.ref) assert link_a is not None - assert len(editor.scene().links()) == 4 + assert len(editor.scene().links()) == initial_count + 1 link_b = editor.scene().create_link(prompt.ref, port_b.ref) assert link_b is not None - assert len(editor.scene().links()) == 4 + assert len(editor.scene().links()) == initial_count + 1 assert editor.get_result()["inputs"]["prompt"]["property"] == "negative" @@ -62,15 +68,17 @@ def test_drag_from_start_output_creates_link(qtbot, sample_graph): def test_different_start_output_replaces_on_same_target(qtbot, sample_graph): editor = _make_editor(qtbot, sample_graph) + initial_count = len(editor.scene().links()) output_a = editor.start_node().output_port(0) output_b = editor.start_node().output_port(1) target_node = editor.scene().get_node_item(1) target_port = target_node.input_port(1) link_a = editor.scene().create_link(output_a.ref, target_port.ref) assert link_a is not None + assert len(editor.scene().links()) == initial_count + 1 link_b = editor.scene().create_link(output_b.ref, target_port.ref) assert link_b is not None - assert len(editor.scene().links()) == 4 + assert len(editor.scene().links()) == initial_count + 1 result = editor.get_result() assert "prompt" not in result["inputs"] assert "seed" in result["inputs"] @@ -85,9 +93,3 @@ def test_link_to_widget_input(qtbot, sample_graph): link = editor.scene().create_link(prompt.ref, widget_port.ref) assert link is not None assert link.target.is_widget - - -from PyQt5.QtCore import Qt as _Qt # noqa: E402 - -Qt_LeftButton = _Qt.LeftButton -Qt_NoModifier = _Qt.NoModifier