langgraph: fix graph drawing for self-loops (#4688)

Fixes #4685
This commit is contained in:
Nuno Campos
2025-05-14 11:44:44 -07:00
committed by GitHub
3 changed files with 93 additions and 3 deletions
+4 -3
View File
@@ -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
@@ -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;
+27
View File
@@ -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