From f65d9b2b7d05b4cc0d09a41c346b177ac81042f8 Mon Sep 17 00:00:00 2001 From: vbarda Date: Mon, 12 Aug 2024 19:50:07 -0400 Subject: [PATCH] pass subgraph nodes/channels --- libs/langgraph/langgraph/pregel/__init__.py | 85 ++++++++++++++++++--- libs/langgraph/tests/test_pregel.py | 10 +-- libs/langgraph/tests/test_pregel_async.py | 10 +-- 3 files changed, 85 insertions(+), 20 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index d2b996d08..13aacf0e5 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -360,11 +360,41 @@ class Pregel( if is_managed_value(v) } - def _prepare_state_snapshot(self, saved: CheckpointTuple) -> StateSnapshot: + def _get_nodes_and_channels( + self, checkpoint_ns: str + ) -> tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]]: + if checkpoint_ns == "": + return self.nodes, self.channels + + path = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR) + nodes = self.nodes + channels = self.channels + for subgraph_node_name in path: + 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 + + def _prepare_state_snapshot( + self, + saved: CheckpointTuple, + nodes: Mapping[str, PregelNode], + channels: Mapping[str, BaseChannel], + ) -> StateSnapshot: with ChannelsManager( { k: LastValue(None) if isinstance(c, Context) else c - for k, c in self.channels.items() + for k, c in channels.items() }, saved.checkpoint, saved.config, @@ -373,7 +403,7 @@ class Pregel( ) as managed: next_tasks = prepare_next_tasks( saved.checkpoint, - self.nodes, + nodes, channels, managed, saved.config, @@ -390,12 +420,15 @@ class Pregel( ) async def _prepare_state_snapshot_async( - self, saved: CheckpointTuple + self, + saved: CheckpointTuple, + nodes: Mapping[str, PregelNode], + channels: Mapping[str, BaseChannel], ) -> StateSnapshot: async with AsyncChannelsManager( { k: LastValue(None) if isinstance(c, Context) else c - for k, c in self.channels.items() + for k, c in channels.items() }, saved.checkpoint, saved.config, @@ -404,7 +437,7 @@ class Pregel( ) as managed: next_tasks = prepare_next_tasks( saved.checkpoint, - self.nodes, + nodes, channels, managed, saved.config, @@ -472,6 +505,9 @@ class Pregel( checkpoint_id = 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]] + ] = {} for checkpoint_tuple in checkpoint_tuples: saved_checkpoint_ns = checkpoint_tuple.config["configurable"][ "checkpoint_ns" @@ -490,7 +526,17 @@ class Pregel( existing_checkpoint_id is None or saved_checkpoint_id > existing_checkpoint_id ): - state_snapshot = self._prepare_state_snapshot(checkpoint_tuple) + if saved_checkpoint_ns not in checkpoint_ns_to_nodes_and_channels: + checkpoint_ns_to_nodes_and_channels[ + saved_checkpoint_ns + ] = self._get_nodes_and_channels(saved_checkpoint_ns) + + nodes, channels = checkpoint_ns_to_nodes_and_channels[ + saved_checkpoint_ns + ] + state_snapshot = self._prepare_state_snapshot( + checkpoint_tuple, nodes, channels + ) checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot checkpoint_ns_to_checkpoint_id[ saved_checkpoint_ns @@ -526,6 +572,9 @@ class Pregel( checkpoint_id = 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]] + ] = {} async for checkpoint_tuple in checkpoint_tuples: saved_checkpoint_ns = checkpoint_tuple.config["configurable"][ "checkpoint_ns" @@ -544,8 +593,16 @@ 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 + ] = self._get_nodes_and_channels(saved_checkpoint_ns) + + nodes, channels = checkpoint_ns_to_nodes_and_channels[ + saved_checkpoint_ns + ] state_snapshot = await self._prepare_state_snapshot_async( - checkpoint_tuple + checkpoint_tuple, nodes, channels ) checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot checkpoint_ns_to_checkpoint_id[ @@ -597,7 +654,10 @@ class Pregel( ) yield state_snapshot else: - yield self._prepare_state_snapshot(checkpoint_tuple) + nodes, channels = self._get_nodes_and_channels( + checkpoint_tuple.config["configurable"]["checkpoint_ns"] + ) + yield self._prepare_state_snapshot(checkpoint_tuple, nodes, channels) async def aget_state_history( self, @@ -634,7 +694,12 @@ class Pregel( ) yield state_snapshot else: - yield await self._prepare_state_snapshot_async(checkpoint_tuple) + nodes, channels = self._get_nodes_and_channels( + checkpoint_tuple.config["configurable"]["checkpoint_ns"] + ) + yield await self._prepare_state_snapshot_async( + checkpoint_tuple, nodes, channels + ) def update_state( self, diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 0f4fbb5b2..8780bc042 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -8604,7 +8604,7 @@ def test_nested_graph_interrupts( assert child_state_history == [ StateSnapshot( values={"my_key": "hi my value here"}, - next=(), + next=("inner_2",), config={ "configurable": { "thread_id": "6", @@ -9218,7 +9218,7 @@ def test_nested_graph_state( subgraph_state_snapshots={ "inner": StateSnapshot( values={"my_key": "hi my value here"}, - next=(), + next=("inner_2",), config={ "configurable": { "thread_id": "1", @@ -9275,7 +9275,7 @@ def test_nested_graph_state( subgraph_state_snapshots={ "inner": StateSnapshot( values={"my_key": "hi my value here"}, - next=(), + next=("inner_2",), config={ "configurable": { "thread_id": "1", @@ -9676,7 +9676,7 @@ def test_doubly_nested_graph_state( subgraph_state_snapshots={ "child": StateSnapshot( values={"my_key": "hi my value"}, - next=(), + next=("child_1",), config={ "configurable": { "thread_id": "1", @@ -9696,7 +9696,7 @@ def test_doubly_nested_graph_state( subgraph_state_snapshots={ "child_1": StateSnapshot( values={"my_key": "hi my value here"}, - next=(), + next=("grandchild_2",), config={ "configurable": { "thread_id": "1", diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index b94a73494..af624f639 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -7106,7 +7106,7 @@ async def test_nested_graph_interrupts( assert child_state_history == [ StateSnapshot( values={"my_key": "hi my value here"}, - next=(), + next=("inner_2",), config={ "configurable": { "thread_id": "6", @@ -7725,7 +7725,7 @@ async def test_nested_graph_state( subgraph_state_snapshots={ "inner": StateSnapshot( values={"my_key": "hi my value here"}, - next=(), + next=("inner_2",), config={ "configurable": { "thread_id": "1", @@ -7784,7 +7784,7 @@ async def test_nested_graph_state( subgraph_state_snapshots={ "inner": StateSnapshot( values={"my_key": "hi my value here"}, - next=(), + next=("inner_2",), config={ "configurable": { "thread_id": "1", @@ -8187,7 +8187,7 @@ async def test_doubly_nested_graph_state( subgraph_state_snapshots={ "child": StateSnapshot( values={"my_key": "hi my value"}, - next=(), + next=("child_1",), config={ "configurable": { "thread_id": "1", @@ -8207,7 +8207,7 @@ async def test_doubly_nested_graph_state( subgraph_state_snapshots={ "child_1": StateSnapshot( values={"my_key": "hi my value here"}, - next=(), + next=("grandchild_2",), config={ "configurable": { "thread_id": "1",