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
5 changes: 3 additions & 2 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
106 changes: 106 additions & 0 deletions examples/handles.py
Original file line number Diff line number Diff line change
@@ -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()
6 changes: 4 additions & 2 deletions src/trame_flow/module/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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],
},
Expand Down
66 changes: 49 additions & 17 deletions src/trame_flow/widgets/flow/node_editor.py
Original file line number Diff line number Diff line change
@@ -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

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

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