From b20130d4e4f50b95d3a97e648d461956d4a35445 Mon Sep 17 00:00:00 2001 From: vbarda Date: Wed, 14 May 2025 09:40:02 -0400 Subject: [PATCH] langgraph: fix graph drawing for self-loops --- libs/langgraph/langgraph/pregel/draw.py | 7 ++- .../tests/__snapshots__/test_pregel.ambr | 62 +++++++++++++++++++ libs/langgraph/tests/test_pregel.py | 27 ++++++++ 3 files changed, 93 insertions(+), 3 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/draw.py b/libs/langgraph/langgraph/pregel/draw.py index 6e6b7440c..f2d8e2631 100644 --- a/libs/langgraph/langgraph/pregel/draw.py +++ b/libs/langgraph/langgraph/pregel/draw.py @@ -229,9 +229,10 @@ def draw_graph( first, last = graph.extend(subgraph, prefix=name) for idx, edge in enumerate(graph.edges): if edge.source == name: - graph.edges[idx] = edge.copy(source=cast(Node, last).id) - elif edge.target == name: - graph.edges[idx] = edge.copy(target=cast(Node, first).id) + edge = edge.copy(source=cast(Node, last).id) + if edge.target == name: + edge = edge.copy(target=cast(Node, first).id) + graph.edges[idx] = edge return graph diff --git a/libs/langgraph/tests/__snapshots__/test_pregel.ambr b/libs/langgraph/tests/__snapshots__/test_pregel.ambr index feef534c9..4caad8801 100644 --- a/libs/langgraph/tests/__snapshots__/test_pregel.ambr +++ b/libs/langgraph/tests/__snapshots__/test_pregel.ambr @@ -396,6 +396,68 @@ ''' # --- +# name: test_get_graph_self_loop + ''' + { + "nodes": [ + { + "id": "__start__", + "type": "runnable", + "data": { + "id": [ + "langchain", + "schema", + "runnable", + "RunnablePassthrough" + ], + "name": "__start__" + } + }, + { + "id": "worker_node", + "type": "runnable", + "data": { + "id": [ + "langgraph", + "utils", + "runnable", + "RunnableCallable" + ], + "name": "worker_node" + } + }, + { + "id": "__end__" + } + ], + "edges": [ + { + "source": "__start__", + "target": "worker_node" + }, + { + "source": "worker_node", + "target": "__end__", + "conditional": true + }, + { + "source": "worker_node", + "target": "worker_node", + "conditional": true + } + ] + } + ''' +# --- +# name: test_get_graph_self_loop.1 + ''' + graph TD; + __start__ --> worker_node; + worker_node -.-> __end__; + worker_node -.-> worker_node; + + ''' +# --- # name: test_in_one_fan_out_state_graph_defer_node[memory-False] ''' graph TD; diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index f0c11b36f..b251cf465 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -8727,3 +8727,30 @@ def test_get_graph_loop(snapshot: SnapshotAssertion) -> None: app = workflow.compile() assert json.dumps(app.get_graph().to_json(), indent=2) == snapshot assert app.get_graph().draw_mermaid(with_styles=False) == snapshot + + +def test_get_graph_self_loop(snapshot: SnapshotAssertion) -> None: + import random + + subgraph_builder = StateGraph(MessagesState) + subgraph_builder.add_node("agent", lambda x: x) + subgraph_builder.add_edge(START, "agent") + subgraph = subgraph_builder.compile() + + def worker_node(state: MessagesState) -> Command[Literal["worker_node", "__end__"]]: + subgraph_result = subgraph.invoke(state) + + if random.choice([True, False]): + next_node_name = "worker_node" + else: + next_node_name = END + + return Command(update=subgraph_result, goto=next_node_name) + + self_loop_builder = StateGraph(MessagesState) + self_loop_builder.add_node("worker_node", worker_node) + self_loop_builder.add_edge(START, "worker_node") + self_loop_graph = self_loop_builder.compile() + + assert json.dumps(self_loop_graph.get_graph().to_json(), indent=2) == snapshot + assert self_loop_graph.get_graph().draw_mermaid(with_styles=False) == snapshot