diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 9f583b5b5..5c89ed27f 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -190,15 +190,12 @@ class Channel: ) -def _get_nodes_and_channels( - nodes: Mapping[str, PregelNode], - channels: Mapping[str, BaseChannel], - checkpoint_ns: str, -) -> tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]]: +def _get_subgraph(graph: Pregel, checkpoint_ns: str) -> Pregel: if checkpoint_ns == "": - return nodes, channels + return graph path = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR) + nodes = graph.nodes for subgraph_node_name in path: # if we have this separator it means we have a node that was triggered by Send if SEND_CHECKPOINT_NAMESPACE_SEPARATOR in subgraph_node_name: @@ -210,17 +207,17 @@ def _get_nodes_and_channels( if subgraph_node_name not in nodes: raise ValueError(f"Couldn't find node '{subgraph_node_name}'.") - subgraph_node = nodes[subgraph_node_name].get_node() - - if not isinstance(subgraph_node, RunnableSequence): - break - - first_step = subgraph_node.steps[0] - if isinstance(first_step, Pregel): - nodes = first_step.nodes - channels = first_step.channels - - return nodes, channels + subgraph_node = nodes[subgraph_node_name] + if isinstance(subgraph_node.bound, Pregel): + nodes = subgraph_node.bound.nodes + elif isinstance(subgraph_node.bound, RunnableSequence): + for runnable in subgraph_node.bound.steps: + if isinstance(runnable, Pregel): + nodes = runnable.nodes + break + else: + continue + return subgraph_node.bound def _assemble_state_snapshot_hierarchy( @@ -259,24 +256,21 @@ def _assemble_state_snapshot_hierarchy( def _prepare_state_snapshot( saved: CheckpointTuple, - nodes: Mapping[str, PregelNode], - channels: Mapping[str, BaseChannel], - managed_values_dict: dict[str, ManagedValueSpec], - select_channels: str | list[str], + graph: Pregel, ) -> StateSnapshot: with ChannelsManager( { k: LastValue(None) if isinstance(c, Context) else c - for k, c in channels.items() + for k, c in graph.channels.items() }, saved.checkpoint, saved.config, ) as channels, ManagedValuesManager( - managed_values_dict, ensure_config(saved.config) + graph.managed_values_dict, ensure_config(saved.config) ) as managed: next_tasks = prepare_next_tasks( saved.checkpoint, - nodes, + graph.nodes, channels, managed, saved.config, @@ -284,7 +278,7 @@ def _prepare_state_snapshot( for_execution=False, ) return StateSnapshot( - values=read_channels(channels, select_channels), + values=read_channels(channels, graph.stream_channels_asis), next=tuple(t.name for t in next_tasks), config=saved.config, metadata=saved.metadata, @@ -294,25 +288,21 @@ def _prepare_state_snapshot( async def _prepare_state_snapshot_async( - saved: CheckpointTuple, - nodes: Mapping[str, PregelNode], - channels: Mapping[str, BaseChannel], - managed_values_dict: dict[str, ManagedValueSpec], - select_channels: str | list[str], + saved: CheckpointTuple, graph: Pregel ) -> StateSnapshot: async with AsyncChannelsManager( { k: LastValue(None) if isinstance(c, Context) else c - for k, c in channels.items() + for k, c in graph.channels.items() }, saved.checkpoint, saved.config, ) as channels, AsyncManagedValuesManager( - managed_values_dict, ensure_config(saved.config) + graph.managed_values_dict, ensure_config(saved.config) ) as managed: next_tasks = prepare_next_tasks( saved.checkpoint, - nodes, + graph.nodes, channels, managed, saved.config, @@ -320,7 +310,7 @@ async def _prepare_state_snapshot_async( for_execution=False, ) return StateSnapshot( - values=read_channels(channels, select_channels), + values=read_channels(channels, graph.stream_channels_asis), next=tuple(t.name for t in next_tasks), config=saved.config, metadata=saved.metadata, @@ -534,9 +524,7 @@ class Pregel( checkpoint_id = checkpoint_config["configurable"].get("checkpoint_id") checkpoint_ns_to_checkpoint_id: dict[str, str] = {} checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot] = {} - checkpoint_ns_to_nodes_and_channels: dict[ - str, tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]] - ] = {} + checkpoint_ns_to_graph: dict[str, Pregel] = {} for checkpoint_tuple in checkpoint_tuples: saved_checkpoint_ns = checkpoint_tuple.config["configurable"][ "checkpoint_ns" @@ -555,22 +543,14 @@ class Pregel( existing_checkpoint_id is None or saved_checkpoint_id > existing_checkpoint_id ): - if saved_checkpoint_ns not in checkpoint_ns_to_nodes_and_channels: - checkpoint_ns_to_nodes_and_channels[ - saved_checkpoint_ns - ] = _get_nodes_and_channels( - self.nodes, self.channels, saved_checkpoint_ns + if saved_checkpoint_ns not in checkpoint_ns_to_graph: + checkpoint_ns_to_graph[saved_checkpoint_ns] = _get_subgraph( + self, saved_checkpoint_ns ) - nodes, channels = checkpoint_ns_to_nodes_and_channels[ - saved_checkpoint_ns - ] state_snapshot = _prepare_state_snapshot( checkpoint_tuple, - nodes, - channels, - self.managed_values_dict, - self.stream_channels_asis, + checkpoint_ns_to_graph[saved_checkpoint_ns], ) checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot checkpoint_ns_to_checkpoint_id[ @@ -610,9 +590,7 @@ class Pregel( checkpoint_id = checkpoint_config["configurable"].get("checkpoint_id") checkpoint_ns_to_checkpoint_id: dict[str, str] = {} checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot] = {} - checkpoint_ns_to_nodes_and_channels: dict[ - str, tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]] - ] = {} + checkpoint_ns_to_graph: dict[str, Pregel] = {} async for checkpoint_tuple in checkpoint_tuples: saved_checkpoint_ns = checkpoint_tuple.config["configurable"][ "checkpoint_ns" @@ -631,22 +609,14 @@ class Pregel( existing_checkpoint_id is None or saved_checkpoint_id > existing_checkpoint_id ): - if saved_checkpoint_ns not in checkpoint_ns_to_nodes_and_channels: - checkpoint_ns_to_nodes_and_channels[ - saved_checkpoint_ns - ] = _get_nodes_and_channels( - self.nodes, self.channels, saved_checkpoint_ns + if saved_checkpoint_ns not in checkpoint_ns_to_graph: + checkpoint_ns_to_graph[saved_checkpoint_ns] = _get_subgraph( + self, saved_checkpoint_ns ) - nodes, channels = checkpoint_ns_to_nodes_and_channels[ - saved_checkpoint_ns - ] state_snapshot = await _prepare_state_snapshot_async( checkpoint_tuple, - nodes, - channels, - self.managed_values_dict, - self.stream_channels_asis, + checkpoint_ns_to_graph[saved_checkpoint_ns], ) checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot checkpoint_ns_to_checkpoint_id[ @@ -698,17 +668,13 @@ class Pregel( ) yield state_snapshot else: - nodes, channels = _get_nodes_and_channels( - self.nodes, - self.channels, + graph = _get_subgraph( + self, checkpoint_tuple.config["configurable"]["checkpoint_ns"], ) yield _prepare_state_snapshot( checkpoint_tuple, - nodes, - channels, - self.managed_values_dict, - self.stream_channels_asis, + graph, ) async def aget_state_history( @@ -746,17 +712,13 @@ class Pregel( ) yield state_snapshot else: - nodes, channels = _get_nodes_and_channels( - self.nodes, - self.channels, + graph = _get_subgraph( + self, checkpoint_tuple.config["configurable"]["checkpoint_ns"], ) yield await _prepare_state_snapshot_async( checkpoint_tuple, - nodes, - channels, - self.managed_values_dict, - self.stream_channels_asis, + graph, ) def update_state( diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index aa1db664e..7bdc357d7 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -8602,7 +8602,7 @@ def test_nested_graph_interrupts( ] assert child_state_history == [ StateSnapshot( - values={"my_key": "hi my value here"}, + values={"my_key": "hi my value here", "my_other_key": "hi my value"}, next=("inner_2",), config={ "configurable": { @@ -9140,6 +9140,7 @@ def test_nested_graph_state( class State(TypedDict): my_key: str + other_parent_key: str def outer_1(state: State): return {"my_key": "hi " + state["my_key"]} @@ -9214,7 +9215,7 @@ def test_nested_graph_state( }, subgraph_state_snapshots={ "inner": StateSnapshot( - values={"my_key": "hi my value here"}, + values={"my_key": "hi my value here", "my_other_key": "hi my value"}, next=("inner_2",), config={ "configurable": { @@ -9271,7 +9272,10 @@ def test_nested_graph_state( }, subgraph_state_snapshots={ "inner": StateSnapshot( - values={"my_key": "hi my value here"}, + values={ + "my_key": "hi my value here", + "my_other_key": "hi my value", + }, next=("inner_2",), config={ "configurable": { @@ -9376,7 +9380,10 @@ def test_nested_graph_state( {"configurable": {"thread_id": "1", "checkpoint_ns": "inner"}} ) assert child_snapshot == StateSnapshot( - values={"my_key": "hi my value here and there"}, + values={ + "my_key": "hi my value here and there", + "my_other_key": "hi my value here", + }, next=(), config={ "configurable": { @@ -9519,7 +9526,10 @@ def test_nested_graph_state( }, subgraph_state_snapshots={ "inner": StateSnapshot( - values={"my_key": "hi my value here and there"}, + values={ + "my_key": "hi my value here and there", + "my_other_key": "hi my value here", + }, next=(), config={ "configurable": { @@ -9931,6 +9941,13 @@ def test_send_to_nested_graphs( for subgraph_node in subgraph_nodes: assert subgraph_node.split(":")[0] == "generate_joke" + subgraph_state_snapshots = { + subgraph_node: graph.get_state( + {"configurable": {"thread_id": "1", "checkpoint_ns": subgraph_node}} + ) + for subgraph_node in subgraph_nodes + } + expected_snapshot = StateSnapshot( values={"subjects": ["cats", "dogs"], "jokes": []}, next=("generate_joke", "generate_joke"), @@ -9950,50 +9967,7 @@ def test_send_to_nested_graphs( "checkpoint_id": AnyStr(), } }, - subgraph_state_snapshots={ - subgraph_nodes[0]: StateSnapshot( - values={"jokes": []}, - next=("generate",), - config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": subgraph_nodes[0], - "checkpoint_id": AnyStr(), - } - }, - metadata={"source": "loop", "writes": {"edit": None}, "step": 1}, - created_at=AnyStr(), - parent_config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": subgraph_nodes[0], - "checkpoint_id": AnyStr(), - } - }, - subgraph_state_snapshots=None, - ), - subgraph_nodes[1]: StateSnapshot( - values={"jokes": []}, - next=("generate",), - config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": subgraph_nodes[1], - "checkpoint_id": AnyStr(), - } - }, - metadata={"source": "loop", "writes": {"edit": None}, "step": 1}, - created_at=AnyStr(), - parent_config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": subgraph_nodes[1], - "checkpoint_id": AnyStr(), - } - }, - subgraph_state_snapshots=None, - ), - }, + subgraph_state_snapshots=subgraph_state_snapshots, ) assert actual_snapshot == expected_snapshot diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 3c881e630..87d9240e0 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -7104,7 +7104,7 @@ async def test_nested_graph_interrupts( ] assert child_state_history == [ StateSnapshot( - values={"my_key": "hi my value here"}, + values={"my_key": "hi my value here", "my_other_key": "hi my value"}, next=("inner_2",), config={ "configurable": { @@ -7647,6 +7647,7 @@ async def test_nested_graph_state( class State(TypedDict): my_key: str + other_parent_key: str async def outer_1(state: State): return {"my_key": "hi " + state["my_key"]} @@ -7721,7 +7722,10 @@ async def test_nested_graph_state( }, subgraph_state_snapshots={ "inner": StateSnapshot( - values={"my_key": "hi my value here"}, + values={ + "my_key": "hi my value here", + "my_other_key": "hi my value", + }, next=("inner_2",), config={ "configurable": { @@ -7780,7 +7784,10 @@ async def test_nested_graph_state( }, subgraph_state_snapshots={ "inner": StateSnapshot( - values={"my_key": "hi my value here"}, + values={ + "my_key": "hi my value here", + "my_other_key": "hi my value", + }, next=("inner_2",), config={ "configurable": { @@ -7885,7 +7892,10 @@ async def test_nested_graph_state( {"configurable": {"thread_id": "1", "checkpoint_ns": "inner"}} ) assert child_snapshot == StateSnapshot( - values={"my_key": "hi my value here and there"}, + values={ + "my_key": "hi my value here and there", + "my_other_key": "hi my value here", + }, next=(), config={ "configurable": { @@ -8030,7 +8040,10 @@ async def test_nested_graph_state( }, subgraph_state_snapshots={ "inner": StateSnapshot( - values={"my_key": "hi my value here and there"}, + values={ + "my_key": "hi my value here and there", + "my_other_key": "hi my value here", + }, next=(), config={ "configurable": { @@ -8442,6 +8455,12 @@ async def test_send_to_nested_graphs( for subgraph_node in subgraph_nodes: assert subgraph_node.split(":")[0] == "generate_joke" + subgraph_state_snapshots = { + subgraph_node: await graph.aget_state( + {"configurable": {"thread_id": "1", "checkpoint_ns": subgraph_node}} + ) + for subgraph_node in subgraph_nodes + } expected_snapshot = StateSnapshot( values={"subjects": ["cats", "dogs"], "jokes": []}, next=("generate_joke", "generate_joke"), @@ -8461,50 +8480,7 @@ async def test_send_to_nested_graphs( "checkpoint_id": AnyStr(), } }, - subgraph_state_snapshots={ - subgraph_nodes[0]: StateSnapshot( - values={"jokes": []}, - next=("generate",), - config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": subgraph_nodes[0], - "checkpoint_id": AnyStr(), - } - }, - metadata={"source": "loop", "writes": {"edit": None}, "step": 1}, - created_at=AnyStr(), - parent_config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": subgraph_nodes[0], - "checkpoint_id": AnyStr(), - } - }, - subgraph_state_snapshots=None, - ), - subgraph_nodes[1]: StateSnapshot( - values={"jokes": []}, - next=("generate",), - config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": subgraph_nodes[1], - "checkpoint_id": AnyStr(), - } - }, - metadata={"source": "loop", "writes": {"edit": None}, "step": 1}, - created_at=AnyStr(), - parent_config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": subgraph_nodes[1], - "checkpoint_id": AnyStr(), - } - }, - subgraph_state_snapshots=None, - ), - }, + subgraph_state_snapshots=subgraph_state_snapshots, ) assert actual_snapshot == expected_snapshot