diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 352cae7e2..235606929 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -370,9 +370,9 @@ class Pregel( ) async def aget_subgraphs( - self, recursive: bool = False + self, recurse: bool = False ) -> AsyncIterator[tuple[str, Pregel]]: - for name, node in self.get_subgraphs(recurse=recursive): + for name, node in self.get_subgraphs(recurse=recurse): yield name, node def _prepare_state_snapshot( @@ -482,7 +482,7 @@ class Pregel( for_execution=False, ) # get the subgraphs - subgraphs = dict(self.get_subgraphs()) + subgraphs = {n: g async for n, g in self.aget_subgraphs()} parent_ns = saved.config["configurable"].get("checkpoint_ns", "") task_states: dict[str, Union[RunnableConfig, StateSnapshot]] = {} for task in next_tasks: @@ -534,6 +534,28 @@ class Pregel( if not checkpointer: raise ValueError("No checkpointer set") + if ( + checkpoint_ns := config["configurable"].get("checkpoint_ns", "") + ) and CONFIG_KEY_CHECKPOINTER not in config["configurable"]: + # remove task_ids from checkpoint_ns + recast_checkpoint_ns = NS_SEP.join( + part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP) + ) + # find the subgraph with the matching name + for name, pregel in self.get_subgraphs(recurse=True): + if name == recast_checkpoint_ns: + return pregel.get_state( + { + "configurable": { + **config["configurable"], + CONFIG_KEY_CHECKPOINTER: checkpointer, + } + }, + subgraphs=subgraphs, + ) + else: + raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") + config = merge_configs(self.config, config) if self.config else config saved = checkpointer.get_tuple(config) return self._prepare_state_snapshot( @@ -550,6 +572,28 @@ class Pregel( if not checkpointer: raise ValueError("No checkpointer set") + if ( + checkpoint_ns := config["configurable"].get("checkpoint_ns", "") + ) and CONFIG_KEY_CHECKPOINTER not in config["configurable"]: + # remove task_ids from checkpoint_ns + recast_checkpoint_ns = NS_SEP.join( + part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP) + ) + # find the subgraph with the matching name + async for name, pregel in self.aget_subgraphs(recurse=True): + if name == recast_checkpoint_ns: + return await pregel.aget_state( + { + "configurable": { + **config["configurable"], + CONFIG_KEY_CHECKPOINTER: checkpointer, + } + }, + subgraphs=subgraphs, + ) + else: + raise ValueError(f"Subgraph {recast_checkpoint_ns} not found") + config = merge_configs(self.config, config) if self.config else config saved = await checkpointer.aget_tuple(config) return await self._aprepare_state_snapshot( @@ -630,7 +674,7 @@ class Pregel( part.split(NS_END)[0] for part in checkpoint_ns.split(NS_SEP) ) # find the subgraph with the matching name - for name, pregel in self.get_subgraphs(recurse=True): + async for name, pregel in self.aget_subgraphs(recurse=True): if name == recast_checkpoint_ns: async for state in pregel.aget_state_history( { diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index e5d263121..9b6a1967d 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -10899,7 +10899,8 @@ def test_doubly_nested_graph_state( config = {"configurable": {"thread_id": "1"}} app.invoke({"my_key": "my value"}, config, debug=True) # get state without subgraphs - assert app.get_state(config) == StateSnapshot( + outer_state = app.get_state(config) + assert outer_state == StateSnapshot( values={"my_key": "hi my value"}, tasks=( PregelTask( @@ -10936,6 +10937,84 @@ def test_doubly_nested_graph_state( } }, ) + child_state = app.get_state(outer_state.tasks[0].state) + assert ( + child_state.tasks[0] + == StateSnapshot( + values={"my_key": "hi my value"}, + tasks=( + PregelTask( + AnyStr(), + "child_1", + state={ + "configurable": { + "thread_id": "1", + "checkpoint_ns": AnyStr(), + } + }, + ), + ), + next=("child_1",), + config={ + "configurable": { + "thread_id": "1", + "checkpoint_ns": AnyStr("child:"), + "checkpoint_id": AnyStr(), + } + }, + metadata={ + "parents": {"": AnyStr()}, + "source": "loop", + "writes": None, + "step": 0, + }, + created_at=AnyStr(), + parent_config={ + "configurable": { + "thread_id": "1", + "checkpoint_ns": AnyStr("child:"), + "checkpoint_id": AnyStr(), + } + }, + ).tasks[0] + ) + grandchild_state = app.get_state(child_state.tasks[0].state) + assert grandchild_state == StateSnapshot( + values={"my_key": "hi my value here"}, + tasks=( + PregelTask( + AnyStr(), + "grandchild_2", + ), + ), + next=("grandchild_2",), + config={ + "configurable": { + "thread_id": "1", + "checkpoint_ns": AnyStr(), + "checkpoint_id": AnyStr(), + } + }, + metadata={ + "parents": AnyDict( + { + "": AnyStr(), + AnyStr("child:"): AnyStr(), + } + ), + "source": "loop", + "writes": {"grandchild_1": {"my_key": "hi my value here"}}, + "step": 1, + }, + created_at=AnyStr(), + parent_config={ + "configurable": { + "thread_id": "1", + "checkpoint_ns": AnyStr(), + "checkpoint_id": AnyStr(), + } + }, + ) # get state with subgraphs assert app.get_state(config, subgraphs=True) == StateSnapshot( values={"my_key": "hi my value"}, diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 94b60a7f5..9163bfe7f 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -9340,7 +9340,8 @@ async def test_doubly_nested_graph_state( config = {"configurable": {"thread_id": "1"}} await app.ainvoke({"my_key": "my value"}, config, debug=True) # get state without subgraphs - assert await app.aget_state(config) == StateSnapshot( + outer_state = await app.aget_state(config) + assert outer_state == StateSnapshot( values={"my_key": "hi my value"}, tasks=( PregelTask( @@ -9377,6 +9378,84 @@ async def test_doubly_nested_graph_state( } }, ) + child_state = await app.aget_state(outer_state.tasks[0].state) + assert ( + child_state.tasks[0] + == StateSnapshot( + values={"my_key": "hi my value"}, + tasks=( + PregelTask( + AnyStr(), + "child_1", + state={ + "configurable": { + "thread_id": "1", + "checkpoint_ns": AnyStr(), + } + }, + ), + ), + next=("child_1",), + config={ + "configurable": { + "thread_id": "1", + "checkpoint_ns": AnyStr("child:"), + "checkpoint_id": AnyStr(), + } + }, + metadata={ + "parents": {"": AnyStr()}, + "source": "loop", + "writes": None, + "step": 0, + }, + created_at=AnyStr(), + parent_config={ + "configurable": { + "thread_id": "1", + "checkpoint_ns": AnyStr("child:"), + "checkpoint_id": AnyStr(), + } + }, + ).tasks[0] + ) + grandchild_state = await app.aget_state(child_state.tasks[0].state) + assert grandchild_state == StateSnapshot( + values={"my_key": "hi my value here"}, + tasks=( + PregelTask( + AnyStr(), + "grandchild_2", + ), + ), + next=("grandchild_2",), + config={ + "configurable": { + "thread_id": "1", + "checkpoint_ns": AnyStr(), + "checkpoint_id": AnyStr(), + } + }, + metadata={ + "parents": AnyDict( + { + "": AnyStr(), + AnyStr("child:"): AnyStr(), + } + ), + "source": "loop", + "writes": {"grandchild_1": {"my_key": "hi my value here"}}, + "step": 1, + }, + created_at=AnyStr(), + parent_config={ + "configurable": { + "thread_id": "1", + "checkpoint_ns": AnyStr(), + "checkpoint_id": AnyStr(), + } + }, + ) # get state with subgraphs assert await app.aget_state(config, subgraphs=True) == StateSnapshot( values={"my_key": "hi my value"},