diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 5a90071..18d2659 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -76,10 +76,11 @@ repos: exclude: .pre-commit-config.yaml - repo: https://github.com/abravalheri/validate-pyproject - rev: "v0.23" + rev: "v0.25" hooks: - id: validate-pyproject - additional_dependencies: ["validate-pyproject-schema-store[all]"] + additional_dependencies: + ["validate-pyproject[all]", "validate-pyproject-schema-store"] - repo: https://github.com/python-jsonschema/check-jsonschema rev: "0.31.0" diff --git a/examples/handles.py b/examples/handles.py new file mode 100644 index 0000000..c326006 --- /dev/null +++ b/examples/handles.py @@ -0,0 +1,106 @@ +from trame.app import TrameApp +from trame.ui.vuetify3 import SinglePageLayout +from trame.widgets.html import Span +from trame.widgets.vuetify3 import ( + VBtn, + VIcon, + VSelect, +) +from trame_flow.module.core import create_node +from trame_flow.widgets.flow import ( + Background, + Controls, + CustomNode, + Handle, + NodeEditor, +) + + +class Example(TrameApp): + def __init__(self, server=None): + super().__init__(server) + self.ui = self.build_ui() + self.next_node_id = 0 + + @property + def state(self): + return self.server.state + + def add_node(self): + self.vueflow.add_node( + create_node( + id=str(self.next_node_id), + x=0, + y=0, + type=self.state.node_type, + label=f"Node {self.next_node_id}", + data={"subtitle": "subtitle"}, + ) + ) + self.next_node_id += 1 + + def build_ui(self): + with SinglePageLayout(self.server) as layout: + layout.title.set_text("trame-flow example") + with layout.toolbar: + VSelect( + label="Node type", + items=("['solver1', 'solver2', 'solver3', 'solver4']",), + v_model=("node_type", "solver1"), + density="compact", + hide_details="true", + max_width="120px", + ) + with VBtn( + "Add a node", + click=self.add_node, + ): + VIcon("mdi-plus") + + with NodeEditor() as self.vueflow: + Background(gap=10, size=1, pattern_color="#81818a") + Controls() + with CustomNode("solver1"): + Handle( + type="source", position="right", id="out1", style="top: 10px" + ) + Handle( + type="source", position="right", id="out2", style="top: 20px" + ) + Handle( + type="source", position="right", id="out3", style="top: 30px" + ) + Span("Solver 1") + + with CustomNode("solver2"): + Handle(type="source", position="right") + Handle(type="target", position="left", id="in1", style="top: 10px") + Handle(type="target", position="left", id="in2", style="top: 20px") + Span("Solver 2") + + with CustomNode("solver3"): + Handle(type="source", position="right") + Handle(type="target", position="left") + Span("Solver 3") + + with CustomNode("solver4"): + Handle(type="target", position="left", id="in1", style="top: 10px") + Handle(type="target", position="left", id="in2", style="top: 20px") + Span("Solver 4") + + def on_graph_change(nodes, edges): + with self.state: + self.state.nodes = nodes + self.state.edges = edges + self.state.selected_node_id = None + self.state.dirty("nodes") + self.state.dirty("edges") + + self.vueflow.graph_change = on_graph_change + + +# Main + +if __name__ == "__main__": + app = Example() + app.server.start() diff --git a/src/trame_flow/module/core.py b/src/trame_flow/module/core.py index 13ade82..45d6633 100644 --- a/src/trame_flow/module/core.py +++ b/src/trame_flow/module/core.py @@ -36,7 +36,7 @@ class Dimensions(TypedDict): Extent = Union[Literal["parent"], list[list[float]]] DEFAULT_EXTENT = [[float("-inf"), float("-inf")], [float("+inf"), float("+inf")]] -NodeType = Literal["default", "input", "output"] | str +NodeType = Union[Literal["default", "input", "output"], str] Node = TypedDict( @@ -100,7 +100,7 @@ def create_node( if style: node["style"] = style if data: - node["data"] = node["data"] | data + node["data"] = node["data"] or data # set default node style for custom node if type not in ["default", "input", "output"]: node["class"] = "vue-flow__node-default" @@ -145,8 +145,10 @@ class EdgeMarker(TypedDict): "markerStart": NotRequired[Union[EdgeMarkerType, EdgeMarker]], "selectable": NotRequired[bool], "source": str, + "sourceHandle": NotRequired[str], "style": NotRequired[dict], "target": str, + "targetHandle": NotRequired[str], "type": EdgeType, "zIndex": NotRequired[int], }, diff --git a/src/trame_flow/widgets/flow/node_editor.py b/src/trame_flow/widgets/flow/node_editor.py index b50216c..4a23e01 100644 --- a/src/trame_flow/widgets/flow/node_editor.py +++ b/src/trame_flow/widgets/flow/node_editor.py @@ -1,5 +1,5 @@ from ast import literal_eval -from typing import Callable, Literal +from typing import Callable, Literal, Optional from trame_client.widgets.core import Template @@ -141,17 +141,27 @@ def __init__(self, **kwargs): self.graph_change: Callable[[list[Node], list[Edge]], None] = lambda *_: None - def on_connect(self, event): - if not self.get_edge(source=event["source"], target=event["target"]): - self.add_edge( - Edge( - source=event["source"], - target=event["target"], - id=f"{event['source']}->{event['target']}", - type="default", - animated=False, - ) + def on_connect(self, event: dict): + event_source_handle = event.get("sourceHandle") + event_target_handle = event.get("targetHandle") + if not self.get_edge( + source=event["source"], + target=event["target"], + source_handle=event_source_handle, + target_handle=event_target_handle, + ): + edge = Edge( + source=event["source"], + target=event["target"], + id=f"{event['source']}{f'({event_source_handle})' if event_source_handle is not None else ''}->{event['target']}{f'({event_target_handle})' if event_target_handle is not None else ''}", + type="default", + animated=False, ) + if event_source_handle is not None: + edge["sourceHandle"] = event_source_handle + if event_target_handle is not None: + edge["targetHandle"] = event_target_handle + self.add_edge(edge) def on_nodes_change(self, events): need_sync = False @@ -164,11 +174,16 @@ def on_nodes_change(self, events): if need_sync: self._sync() - def on_edges_change(self, events): + def on_edges_change(self, events: list[dict]): need_sync = False for event in events: if event["type"] == "remove": - edge = self.get_edge(event["source"], event["target"]) + edge = self.get_edge( + event["source"], + event["target"], + event.get("sourceHandle"), + event.get("targetHandle"), + ) if edge: self._edges.remove(edge) need_sync = True @@ -216,10 +231,21 @@ def get_node(self, id: str): return node return None - def get_edge(self, source: str, target: str): + def get_edge( + self, + source: str, + target: str, + source_handle: Optional[str] = None, + target_handle: Optional[str] = None, + ): """Get an Edge from its source and target. Returns None if not found.""" for edge in self._edges: - if edge["source"] == source and edge["target"] == target: + if ( + edge["source"] == source + and edge["target"] == target + and source_handle == edge.get("sourceHandle") + and target_handle == edge.get("targetHandle") + ): return edge return None @@ -231,9 +257,15 @@ def remove_node(self, node_id: str): self._nodes.remove(node) self.graph_change(self._nodes, self._edges) - def remove_edge(self, source: str, target: str): + def remove_edge( + self, + source: str, + target: str, + source_handle: Optional[str] = None, + target_handle: Optional[str] = None, + ): """Remove an Edge from the graph. Does nothing if there is no edge from `source` to `target`.""" - edge = self.get_edge(source, target) + edge = self.get_edge(source, target, source_handle, target_handle) if edge is not None: self.server.js_call(self.__ref, "removeEdges", edge["id"]) self._edges.remove(edge)