Fix graph repr

This commit is contained in:
Nuno Campos
2024-03-21 16:53:59 -07:00
parent ccea395002
commit e11153cf09
3 changed files with 49 additions and 9 deletions
+6 -3
View File
@@ -50,6 +50,10 @@ class Graph:
self.entry_point: Optional[str] = None
self.entry_point_branch: Optional[Branch] = None
@property
def _all_edges(self) -> set[tuple[str, str]]:
return self.edges
def add_node(self, key: str, action: RunnableLike) -> None:
if self.compiled:
logger.warning(
@@ -149,9 +153,8 @@ class Graph:
def validate(
self,
interrupt: Optional[Sequence[str]] = None,
additional_edges: Optional[set[tuple[str, str]]] = None,
) -> None:
edges = self.edges.union(additional_edges or [])
edges = self._all_edges
all_starts = {src for src, _ in edges} | {src for src in self.branches}
for node in self.nodes:
if node not in all_starts:
@@ -280,7 +283,7 @@ class CompiledGraph(Pregel):
n = graph.add_node(node, key)
start_nodes[key] = n
end_nodes[key] = n
for start, end in self.graph.edges:
for start, end in self.graph._all_edges:
graph.add_edge(start_nodes[start], end_nodes[end])
for start, branches in self.graph.branches.items():
for i, branch in enumerate(branches):
+7 -6
View File
@@ -31,6 +31,12 @@ class StateGraph(Graph):
self.support_multiple_edges = True
self.w_edges: set[tuple[tuple[str, ...], str]] = set()
@property
def _all_edges(self) -> set[tuple[str, str]]:
return self.edges | {
(start, end) for starts, end in self.w_edges for start in starts
}
def add_node(self, key: str, action: RunnableLike) -> None:
if key in self.channels:
raise ValueError(
@@ -66,12 +72,7 @@ class StateGraph(Graph):
) -> CompiledGraph:
interrupt_before = interrupt_before or []
interrupt_after = interrupt_after or []
self.validate(
interrupt=interrupt_before + interrupt_after,
additional_edges={
(start, end) for starts, end in self.w_edges for start in starts
},
)
self.validate(interrupt=interrupt_before + interrupt_after)
state_keys = list(self.channels)
state_keys_read = state_keys[0] if state_keys == ["__root__"] else state_keys
+36
View File
@@ -2832,6 +2832,42 @@ def test_in_one_fan_out_state_graph_waiting_edge() -> None:
app = workflow.compile()
assert app.get_graph().draw_ascii() == (
""" +-----------+
| __start__ |
+-----------+
*
*
*
+---------------+
| rewrite_query |
+---------------+
*** ***
* *
** ***
+--------------+ *
| analyzer_one | *
+--------------+ *
* *
* *
* *
+---------------+ +---------------+
| retriever_one | | retriever_two |
+---------------+ +---------------+
*** ***
* *
** **
+----+
| qa |
+----+
*
*
*
+---------+
| __end__ |
+---------+ """
)
assert app.invoke({"query": "what is weather in sf"}) == {
"query": "analyzed: query: what is weather in sf",
"docs": ["doc1", "doc2", "doc3", "doc4"],