diff --git a/libs/langgraph/langgraph/pregel/_retry.py b/libs/langgraph/langgraph/pregel/_retry.py index 554a90a59..b42b63644 100644 --- a/libs/langgraph/langgraph/pregel/_retry.py +++ b/libs/langgraph/langgraph/pregel/_retry.py @@ -9,7 +9,7 @@ from collections.abc import Awaitable, Callable, Sequence from dataclasses import replace from typing import Any -from langgraph._internal._config import patch_configurable +from langgraph._internal._config import patch_configurable, recast_checkpoint_ns from langgraph._internal._constants import ( CONF, CONFIG_KEY_CHECKPOINT_NS, @@ -43,7 +43,8 @@ def run_with_retry( except ParentCommand as exc: ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS] cmd = exc.args[0] - if cmd.graph in (ns, task.name): + # strip task_ids from namespace for comparison (ns format: "node1|node2:task_id") + if cmd.graph in (ns, recast_checkpoint_ns(ns), task.name): # this command is for the current graph, handle it for w in task.writers: w.invoke(cmd, config) @@ -138,7 +139,8 @@ async def arun_with_retry( except ParentCommand as exc: ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS] cmd = exc.args[0] - if cmd.graph in (ns, task.name): + # strip task_ids from namespace for comparison (ns format: "node1|node2:task_id") + if cmd.graph in (ns, recast_checkpoint_ns(ns), task.name): # this command is for the current graph, handle it for w in task.writers: w.invoke(cmd, config) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 97928ec6f..14388f159 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -7927,6 +7927,83 @@ def test_parent_command_goto( } +@pytest.mark.parametrize("subgraph_persist", [True, False]) +def test_parent_command_goto_deeply_nested( + sync_checkpointer: BaseCheckpointSaver, + subgraph_persist: bool, +) -> None: + """Test Command.PARENT in a 3-level nested subgraph. + + Command.PARENT should jump to sub_child_3 in the immediate parent (sub_graph). + + Note: With operator.add, subgraph state (including its input) is merged with + parent state, causing the input to appear multiple times. This is expected. + """ + + class State(TypedDict): + dialog_state: Annotated[list[str], operator.add] + + # Level 3: Deepest subgraph that issues Command.PARENT + def sub_sub_child_node(state): + # Jump to immediate parent (sub_graph) + return Command( + graph=Command.PARENT, + goto="sub_child_3", + update={"dialog_state": ["sub_sub_child"]}, + ) + + sub_sub_builder = StateGraph(State) + sub_sub_builder.add_node("sub_sub_child", sub_sub_child_node) + sub_sub_builder.add_edge(START, "sub_sub_child") + sub_sub_graph = sub_sub_builder.compile( + name="sub_sub_graph", checkpointer=subgraph_persist + ) + + # Level 2: Middle subgraph containing Level 3 + def sub_child_1(state): + return {"dialog_state": ["sub_child_1"]} + + def sub_child_3(state): + return {"dialog_state": ["sub_child_3"]} + + sub_builder = StateGraph(State) + sub_builder.add_node("sub_child_1", sub_child_1) + sub_builder.add_node("sub_child_2", sub_sub_graph, destinations=("sub_child_3",)) + sub_builder.add_node("sub_child_3", sub_child_3) + sub_builder.add_edge(START, "sub_child_1") + sub_builder.add_edge("sub_child_1", "sub_child_2") + sub_graph = sub_builder.compile(name="sub_graph", checkpointer=subgraph_persist) + + # Level 1: Main graph containing Level 2 + def child_1(state): + return {"dialog_state": ["child_1"]} + + builder = StateGraph(State) + builder.add_node("child_1", child_1) + builder.add_node("child_2", sub_graph) + builder.add_edge(START, "child_1") + builder.add_edge("child_1", "child_2") + graph = builder.compile(name="main_graph", checkpointer=sync_checkpointer) + + config = {"configurable": {"thread_id": 1}} + + result = graph.invoke(input={"dialog_state": ["init"]}, config=config) + + # Command.PARENT from sub_sub_child jumps to sub_child_3 in immediate parent + # State duplication occurs due to operator.add merging behavior + assert result == { + "dialog_state": [ + "init", + "child_1", + "init", + "child_1", + "sub_child_1", + "sub_sub_child", + "sub_child_3", + ] + } + + @pytest.mark.parametrize("with_timeout", [True, False]) def test_timeout_with_parent_command( sync_checkpointer: BaseCheckpointSaver, with_timeout: bool