mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 20:45:05 +02:00
Fix graph repr
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"],
|
||||
|
||||
Reference in New Issue
Block a user