mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-22 23:52:23 +02:00
@@ -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;
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user