Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 31 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
@@ -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
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ MANIFEST

# Virtual environments
.venv/
.venv_win
venv/
env/
ENV/
Expand Down
10 changes: 8 additions & 2 deletions AGENTS.md
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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
Expand Down
6 changes: 5 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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).
[`examples/dialog_demo.py`](examples/dialog_demo.py).

Used in the krita_comfyui plugin <https://github.com/dacert/krita-comfyui>

![Select Workflow](media/example.png)
Binary file added media/example.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,8 @@ dev = [
"pytest-qt>=4",
"pytest-mock>=3",
"ruff>=0.6",
"ty>=0.0.61",
"PyQt5-stubs>=5.15",
]

[tool.setuptools]
Expand Down
37 changes: 15 additions & 22 deletions src/comfy_graph_bind/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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)

Expand All @@ -83,19 +80,20 @@ 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

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:
Expand All @@ -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]] = {}
Expand All @@ -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)
Expand All @@ -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

Expand All @@ -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):
Expand Down Expand Up @@ -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)

Expand Down
2 changes: 1 addition & 1 deletion src/comfy_graph_bind/layout.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)


Expand Down
7 changes: 1 addition & 6 deletions src/comfy_graph_bind/node_item.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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]
10 changes: 2 additions & 8 deletions src/comfy_graph_bind/port_item.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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:
Expand Down Expand Up @@ -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
49 changes: 29 additions & 20 deletions src/comfy_graph_bind/scene.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")))
Expand Down Expand Up @@ -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()):
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
14 changes: 5 additions & 9 deletions src/comfy_graph_bind/start_node.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand All @@ -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()


Expand Down
Loading
Loading