diff --git a/README.rst b/README.rst
index 0b64c74..8fad43f 100644
--- a/README.rst
+++ b/README.rst
@@ -3,7 +3,8 @@ Trame Flow
A node editor for `trame `_ based on `VueFlow `_.
-.. image:: screenshot.png
+.. image:: https://raw.githubusercontent.com/Kitware/trame-flow/refs/heads/main/screenshot.png
+ :alt: Usage example of trame-flow with custom node and edge styles.
License
----------------------------------------
diff --git a/src/trame_flow/widgets/flow/node_editor.py b/src/trame_flow/widgets/flow/node_editor.py
index 4a23e01..4c15b11 100644
--- a/src/trame_flow/widgets/flow/node_editor.py
+++ b/src/trame_flow/widgets/flow/node_editor.py
@@ -1,4 +1,5 @@
from ast import literal_eval
+from collections.abc import Iterable
from typing import Callable, Literal, Optional
from trame_client.widgets.core import Template
@@ -275,6 +276,30 @@ def remove_edge(
def graph(self) -> Graph:
return Graph(nodes=self._nodes, edges=self._edges)
+ @graph.setter
+ def graph(self, graph: Graph):
+ self._nodes = graph["nodes"]
+ self._edges = graph["edges"]
+ self._sync()
+
+ @property
+ def nodes(self) -> list[Node]:
+ return self._nodes
+
+ @nodes.setter
+ def nodes(self, nodes: Iterable[Node]):
+ self._nodes = list(nodes)
+ self._sync()
+
+ @property
+ def edges(self) -> list[Edge]:
+ return self._edges
+
+ @edges.setter
+ def edges(self, edges: Iterable[Edge]):
+ self._edges = list(edges)
+ self._sync()
+
def serialize_graph(self) -> str:
"""Returns graph as a string representing a `Graph` object."""
return str(self.graph)
@@ -286,10 +311,7 @@ def deserialize_graph(self, graph_str: str) -> bool:
Returns False if deserialization produced any error, else True.
"""
try:
- graph = literal_eval(graph_str)
- self._nodes = graph["nodes"]
- self._edges = graph["edges"]
- self._sync()
+ self.graph = literal_eval(graph_str)
except Exception:
return False
return True
@@ -313,3 +335,9 @@ def update_edge(self, source: str, target: str, **kwargs):
def fit_view(self):
"""Fit VueFlow's view to show the entire graph (excluding hidden nodes)"""
self.server.js_call(self.__ref, "fitView")
+
+ def clear_graph(self):
+ """Remove every edges and then every nodes."""
+ self._nodes = []
+ self._edges = []
+ self._sync()