From e11153cf0959427634a75483273619bd314b0fb3 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Thu, 21 Mar 2024 16:53:59 -0700 Subject: [PATCH] Fix graph repr --- langgraph/graph/graph.py | 9 ++++++--- langgraph/graph/state.py | 13 +++++++------ tests/test_pregel.py | 36 ++++++++++++++++++++++++++++++++++++ 3 files changed, 49 insertions(+), 9 deletions(-) diff --git a/langgraph/graph/graph.py b/langgraph/graph/graph.py index c7b6a2074..097e75af0 100644 --- a/langgraph/graph/graph.py +++ b/langgraph/graph/graph.py @@ -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): diff --git a/langgraph/graph/state.py b/langgraph/graph/state.py index c12f2f9d7..2c766c410 100644 --- a/langgraph/graph/state.py +++ b/langgraph/graph/state.py @@ -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 diff --git a/tests/test_pregel.py b/tests/test_pregel.py index ab190751f..0cd238c66 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -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"],