From cdd789956426eb8ff16e9f74d6529bd7a2cc91d4 Mon Sep 17 00:00:00 2001 From: Andrew Nguonly Date: Fri, 1 Nov 2024 15:21:39 -0700 Subject: [PATCH] RemoteGraph: If node name is not present, fallback to node id as the node name (#2304) --- libs/langgraph/langgraph/pregel/remote.py | 14 ++++++++++++-- libs/langgraph/tests/test_remote_graph.py | 20 ++++++++++---------- 2 files changed, 22 insertions(+), 12 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/remote.py b/libs/langgraph/langgraph/pregel/remote.py index d9ee53e8c..abe27eb28 100644 --- a/libs/langgraph/langgraph/pregel/remote.py +++ b/libs/langgraph/langgraph/pregel/remote.py @@ -133,10 +133,20 @@ class RemoteGraph(PregelProtocol): nodes = {} for node in graph["nodes"]: node_id = str(node["id"]) + node_data = node.get("data", {}) + + # Get node name from node_data if available. If not, use node_id. + node_name = node.get("name") + if node_name is None: + if isinstance(node_data, dict): + node_name = node_data.get("name", node_id) + else: + node_name = node_id + nodes[node_id] = DrawableNode( id=node_id, - name=node.get("name", ""), - data=node.get("data", {}), + name=node_name, + data=node_data, metadata=node.get("metadata"), ) return nodes diff --git a/libs/langgraph/tests/test_remote_graph.py b/libs/langgraph/tests/test_remote_graph.py index 83e5f913b..21fb8278f 100644 --- a/libs/langgraph/tests/test_remote_graph.py +++ b/libs/langgraph/tests/test_remote_graph.py @@ -54,7 +54,7 @@ def test_get_graph(): "type": "runnable", "data": { "id": ["langgraph", "utils", "RunnableCallable"], - "name": "agent", + "name": "agent_1", }, }, ], @@ -71,13 +71,13 @@ def test_get_graph(): assert drawable_graph.nodes == { "__start__": DrawableNode( - id="__start__", name="", data="__start__", metadata=None + id="__start__", name="__start__", data="__start__", metadata=None ), - "__end__": DrawableNode(id="__end__", name="", data="__end__", metadata=None), + "__end__": DrawableNode(id="__end__", name="__end__", data="__end__", metadata=None), "agent": DrawableNode( id="agent", - name="", - data={"id": ["langgraph", "utils", "RunnableCallable"], "name": "agent"}, + name="agent_1", + data={"id": ["langgraph", "utils", "RunnableCallable"], "name": "agent_1"}, metadata=None, ), } @@ -101,7 +101,7 @@ async def test_aget_graph(): "type": "runnable", "data": { "id": ["langgraph", "utils", "RunnableCallable"], - "name": "agent", + "name": "agent_1", }, }, ], @@ -118,13 +118,13 @@ async def test_aget_graph(): assert drawable_graph.nodes == { "__start__": DrawableNode( - id="__start__", name="", data="__start__", metadata=None + id="__start__", name="__start__", data="__start__", metadata=None ), - "__end__": DrawableNode(id="__end__", name="", data="__end__", metadata=None), + "__end__": DrawableNode(id="__end__", name="__end__", data="__end__", metadata=None), "agent": DrawableNode( id="agent", - name="", - data={"id": ["langgraph", "utils", "RunnableCallable"], "name": "agent"}, + name="agent_1", + data={"id": ["langgraph", "utils", "RunnableCallable"], "name": "agent_1"}, metadata=None, ), }