diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py index e9ca12305..038bb6677 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py @@ -253,9 +253,6 @@ class BasePostgresSaver(BaseCheckpointSaver): if config: wheres.append("thread_id = %s ") param_values.append(config["configurable"]["thread_id"]) - checkpoint_ns = config["configurable"].get("checkpoint_ns", "") - wheres.append("checkpoint_ns = %s") - param_values.append(checkpoint_ns) # construct predicate for metadata filter if filter: diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/utils.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/utils.py index 6e1baf5ae..56034ea34 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/utils.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/utils.py @@ -70,9 +70,6 @@ def search_where( if config is not None: wheres.append("thread_id = ?") param_values.append(config["configurable"]["thread_id"]) - checkpoint_ns = config["configurable"].get("checkpoint_ns", "") - wheres.append("checkpoint_ns = ?") - param_values.append(checkpoint_ns) # construct predicate for metadata filter if filter: diff --git a/libs/checkpoint/langgraph/checkpoint/base/__init__.py b/libs/checkpoint/langgraph/checkpoint/base/__init__.py index 4f0f08a84..86c8b0eec 100644 --- a/libs/checkpoint/langgraph/checkpoint/base/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/base/__init__.py @@ -259,7 +259,6 @@ class BaseCheckpointSaver(ABC): filter: Optional[Dict[str, Any]] = None, before: Optional[RunnableConfig] = None, limit: Optional[int] = None, - include_nested_checkpoints: bool = False, ) -> Iterator[CheckpointTuple]: """List checkpoints that match the given criteria. @@ -351,7 +350,6 @@ class BaseCheckpointSaver(ABC): filter: Optional[Dict[str, Any]] = None, before: Optional[RunnableConfig] = None, limit: Optional[int] = None, - include_nested_checkpoints: bool = False, ) -> AsyncIterator[CheckpointTuple]: """Asynchronously list checkpoints that match the given criteria. diff --git a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py index 5e60ebaf0..3bb7f8803 100644 --- a/libs/checkpoint/langgraph/checkpoint/memory/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/memory/__init__.py @@ -157,7 +157,6 @@ class MemorySaver( filter: Optional[Dict[str, Any]] = None, before: Optional[RunnableConfig] = None, limit: Optional[int] = None, - include_nested_checkpoints: bool = False, ) -> Iterator[CheckpointTuple]: """List checkpoints from the in-memory storage. @@ -178,16 +177,7 @@ class MemorySaver( config["configurable"].get("checkpoint_ns", "") if config else "" ) for thread_id in thread_ids: - checkpoint_ns_iter = ( - ( - key - for key in self.storage[thread_id].keys() - if key.startswith(checkpoint_ns) - ) - if include_nested_checkpoints - else [checkpoint_ns] - ) - for checkpoint_ns in checkpoint_ns_iter: + for checkpoint_ns in self.storage[thread_id].keys(): for checkpoint_id, ( checkpoint, metadata_b, @@ -330,7 +320,6 @@ class MemorySaver( filter: Optional[Dict[str, Any]] = None, before: Optional[RunnableConfig] = None, limit: Optional[int] = None, - include_nested_checkpoints: bool = False, ) -> AsyncIterator[CheckpointTuple]: """Asynchronous version of list. @@ -351,7 +340,6 @@ class MemorySaver( before=before, limit=limit, filter=filter, - include_nested_checkpoints=include_nested_checkpoints, ), config, ) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index dd105e0ba..9c591979f 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -360,70 +360,64 @@ class Pregel( if is_managed_value(v) } - def _prepare_state_snapshot( - self, saved: CheckpointTuple, config: RunnableConfig - ) -> StateSnapshot: - checkpoint = saved.checkpoint if saved else empty_checkpoint() - config = saved.config if saved else config + def _prepare_state_snapshot(self, saved: CheckpointTuple) -> StateSnapshot: with ChannelsManager( { k: LastValue(None) if isinstance(c, Context) else c for k, c in self.channels.items() }, - checkpoint, - config, + saved.checkpoint, + saved.config, ) as channels, ManagedValuesManager( - self.managed_values_dict, ensure_config(config) + self.managed_values_dict, ensure_config(saved.config) ) as managed: next_tasks = prepare_next_tasks( - checkpoint, + saved.checkpoint, self.nodes, channels, managed, - config, + saved.config, -1, for_execution=False, ) return StateSnapshot( values=read_channels(channels, self.stream_channels_asis), next=tuple(t.name for t in next_tasks), - config=saved.config if saved else config, - metadata=saved.metadata if saved else None, - created_at=saved.checkpoint["ts"] if saved else None, - parent_config=saved.parent_config if saved else None, + config=saved.config, + metadata=saved.metadata, + created_at=saved.checkpoint["ts"], + parent_config=saved.parent_config, ) async def _prepare_state_snapshot_async( - self, saved: CheckpointTuple, config: RunnableConfig + self, saved: CheckpointTuple ) -> StateSnapshot: - checkpoint = saved.checkpoint if saved else empty_checkpoint() - config = saved.config if saved else config async with AsyncChannelsManager( { k: LastValue(None) if isinstance(c, Context) else c for k, c in self.channels.items() }, - checkpoint, - config, + saved.checkpoint, + saved.config, ) as channels, AsyncManagedValuesManager( - self.managed_values_dict, ensure_config(config) + self.managed_values_dict, ensure_config(saved.config) ) as managed: next_tasks = prepare_next_tasks( - checkpoint, + saved.checkpoint, self.nodes, channels, managed, - config, + saved.config, -1, for_execution=False, ) return StateSnapshot( values=read_channels(channels, self.stream_channels_asis), next=tuple(t.name for t in next_tasks), - config=saved.config if saved else config, - metadata=saved.metadata if saved else None, - created_at=saved.checkpoint["ts"] if saved else None, - parent_config=saved.parent_config if saved else None, + config=saved.config, + metadata=saved.metadata, + created_at=saved.checkpoint["ts"], + parent_config=saved.parent_config, ) @staticmethod @@ -457,7 +451,9 @@ class Pregel( state_snapshot = checkpoint_ns_to_state_snapshots.pop(root_checkpoint_ns, None) if state_snapshot is None: - raise ValueError(f"Missing checkpoint for thread ID '{root_checkpoint_ns}'") + raise ValueError( + f"Missing checkpoint for checkpoint NS '{root_checkpoint_ns}'" + ) return state_snapshot def get_state( @@ -468,9 +464,7 @@ class Pregel( raise ValueError("No checkpointer set") if include_subgraph_state: - checkpoint_tuples = self.checkpointer.list( - config, include_nested_checkpoints=True - ) + checkpoint_tuples = self.checkpointer.list(config) else: checkpoint_tuples = iter([self.checkpointer.get_tuple(config)]) @@ -496,20 +490,16 @@ class Pregel( existing_checkpoint_id is None or saved_checkpoint_id > existing_checkpoint_id ): - state_snapshot = self._prepare_state_snapshot(checkpoint_tuple, config) + state_snapshot = self._prepare_state_snapshot(checkpoint_tuple) checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot checkpoint_ns_to_checkpoint_id[ saved_checkpoint_ns ] = saved_checkpoint_id if not checkpoint_ns_to_state_snapshots: - error_msg = ( - f"Could not find checkpoints for checkpoint NS '{checkpoint_ns}'" + return StateSnapshot( + values={}, next=(), config=config, checkpoint=empty_checkpoint() ) - if checkpoint_id: - error_msg += f" and checkpoint ID '{checkpoint_id}'" - - raise ValueError(error_msg) state_snapshot = self._assemble_state_snapshot_hierarchy( checkpoint_ns, checkpoint_ns_to_state_snapshots @@ -524,9 +514,7 @@ class Pregel( raise ValueError("No checkpointer set") if include_subgraph_state: - checkpoint_tuples = self.checkpointer.alist( - config, include_nested_checkpoints=True - ) + checkpoint_tuples = self.checkpointer.alist(config) else: async def alist_checkpoints(): @@ -556,20 +544,18 @@ class Pregel( existing_checkpoint_id is None or saved_checkpoint_id > existing_checkpoint_id ): - state_snapshot = self._prepare_state_snapshot(checkpoint_tuple, config) + state_snapshot = await self._prepare_state_snapshot_async( + checkpoint_tuple + ) checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot checkpoint_ns_to_checkpoint_id[ saved_checkpoint_ns ] = saved_checkpoint_id if not checkpoint_ns_to_state_snapshots: - error_msg = ( - f"Could not find checkpoints for checkpoint NS '{checkpoint_ns}'" + return StateSnapshot( + values={}, next=(), config=config, checkpoint=empty_checkpoint() ) - if checkpoint_id: - error_msg += f" and checkpoint ID '{checkpoint_id}'" - - raise ValueError(error_msg) state_snapshot = self._assemble_state_snapshot_hierarchy( checkpoint_ns, checkpoint_ns_to_state_snapshots @@ -593,41 +579,25 @@ class Pregel( and signature(self.checkpointer.list).parameters.get("filter") is None ): raise ValueError("Checkpointer does not support filtering") - for config, checkpoint, metadata, parent_config, _ in self.checkpointer.list( + + checkpoint_ns = config["configurable"].get("checkpoint_ns", "") + for checkpoint_tuple in self.checkpointer.list( config, before=before, limit=limit, filter=filter ): + if ( + checkpoint_tuple.config["configurable"]["checkpoint_ns"] + != checkpoint_ns + ): + # only list root checkpoints here + continue + if include_subgraph_state: - state_snapshot = self.get_state(config, include_subgraph_state=True) + state_snapshot = self.get_state( + checkpoint_tuple.config, include_subgraph_state=True + ) yield state_snapshot else: - with ChannelsManager( - { - k: LastValue(None) if isinstance(c, Context) else c - for k, c in self.channels.items() - }, - checkpoint, - config, - ) as channels, ManagedValuesManager( - self.managed_values_dict, ensure_config(config) - ) as managed: - next_tasks = prepare_next_tasks( - checkpoint, - self.nodes, - channels, - managed, - config, - -1, - for_execution=False, - ) - - yield StateSnapshot( - read_channels(channels, self.stream_channels_asis), - tuple(t.name for t in next_tasks), - config, - metadata, - checkpoint["ts"], - parent_config, - ) + yield self._prepare_state_snapshot(checkpoint_tuple) async def aget_state_history( self, @@ -646,46 +616,25 @@ class Pregel( and signature(self.checkpointer.list).parameters.get("filter") is None ): raise ValueError("Checkpointer does not support filtering") - async for ( - config, - checkpoint, - metadata, - parent_config, - _, - ) in self.checkpointer.alist(config, before=before, limit=limit, filter=filter): + + checkpoint_ns = config["configurable"].get("checkpoint_ns", "") + async for checkpoint_tuple in self.checkpointer.alist( + config, before=before, limit=limit, filter=filter + ): + if ( + checkpoint_tuple.config["configurable"]["checkpoint_ns"] + != checkpoint_ns + ): + # only list root checkpoints here + continue + if include_subgraph_state: state_snapshot = await self.aget_state( - config, include_subgraph_state=True + checkpoint_tuple.config, include_subgraph_state=True ) yield state_snapshot else: - async with AsyncChannelsManager( - { - k: LastValue(None) if isinstance(c, Context) else c - for k, c in self.channels.items() - }, - checkpoint, - config, - ) as channels, AsyncManagedValuesManager( - self.managed_values_dict, ensure_config(config) - ) as managed: - next_tasks = prepare_next_tasks( - checkpoint, - self.nodes, - channels, - managed, - config, - -1, - for_execution=False, - ) - yield StateSnapshot( - read_channels(channels, self.stream_channels_asis), - tuple(t.name for t in next_tasks), - config, - metadata, - checkpoint["ts"], - parent_config, - ) + yield await self._prepare_state_snapshot_async(checkpoint_tuple) def update_state( self, @@ -953,32 +902,36 @@ class Pregel( stream_mode = stream_mode if stream_mode is not None else self.stream_mode if not isinstance(stream_mode, list): stream_mode = [stream_mode] + + if config and config.get("configurable", {}).get(CONFIG_KEY_CHECKPOINTER): + parent_checkpointer = config["configurable"][CONFIG_KEY_CHECKPOINTER] + else: + parent_checkpointer = None + if config and config.get("configurable", {}).get(CONFIG_KEY_READ) is not None: # if being called as a node in another graph, always use values mode stream_mode = ["values"] - if self.checkpointer is None: + if parent_checkpointer is not None and self.checkpointer is None: raise ValueError( "Missing checkpointer for subgraph. " "Please compile the subgraph graph with checkpointer=INHERIT_CHECKPOINTER (from langgraph.pregel import INHERIT_CHECKPOINTER)." ) - if self.checkpointer != INHERIT_CHECKPOINTER: + if ( + parent_checkpointer is not None + and self.checkpointer != INHERIT_CHECKPOINTER + ): raise ValueError( "Custom checkpointers for subgraphs are not allowed. " "Please compile the subgraph graph with checkpointer=INHERIT_CHECKPOINTER (from langgraph.pregel import INHERIT_CHECKPOINTER)." ) - if ( - config is not None - and config.get("configurable", {}).get(CONFIG_KEY_CHECKPOINTER) - and self.checkpointer == INHERIT_CHECKPOINTER - ): - checkpointer: Optional[BaseCheckpointSaver] = config["configurable"][ - CONFIG_KEY_CHECKPOINTER - ] - else: - checkpointer = self.checkpointer + checkpointer = ( + parent_checkpointer + if parent_checkpointer is not None + else self.checkpointer + ) return ( debug, stream_mode, diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index ee5efaa37..b94a73494 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -7629,13 +7629,13 @@ async def test_nested_graph_state( my_key: str my_other_key: str - def inner_1(state: InnerState): + async def inner_1(state: InnerState): return { "my_key": state["my_key"] + " here", "my_other_key": state["my_key"], } - def inner_2(state: InnerState): + async def inner_2(state: InnerState): return { "my_key": state["my_key"] + " and there", "my_other_key": state["my_key"], @@ -7651,10 +7651,10 @@ async def test_nested_graph_state( class State(TypedDict): my_key: str - def outer_1(state: State): + async def outer_1(state: State): return {"my_key": "hi " + state["my_key"]} - def outer_2(state: State): + async def outer_2(state: State): return {"my_key": state["my_key"] + " and back again"} graph = StateGraph(State) @@ -7699,7 +7699,7 @@ async def test_nested_graph_state( }, subgraph_state_snapshots=None, ) - assert app.get_state(config, include_subgraph_state=True) == StateSnapshot( + assert await app.aget_state(config, include_subgraph_state=True) == StateSnapshot( values={"my_key": "hi my value"}, next=("inner",), config={ @@ -7755,7 +7755,9 @@ async def test_nested_graph_state( ) }, ) - assert list(app.get_state_history(config, include_subgraph_state=True)) == [ + assert [ + s async for s in app.aget_state_history(config, include_subgraph_state=True) + ] == [ StateSnapshot( values={"my_key": "hi my value"}, next=("inner",), @@ -8089,10 +8091,10 @@ async def test_doubly_nested_graph_state( class GrandChildState(TypedDict): my_key: str - def grandchild_1(state: ChildState): + async def grandchild_1(state: ChildState): return {"my_key": state["my_key"] + " here"} - def grandchild_2(state: ChildState): + async def grandchild_2(state: ChildState): return { "my_key": state["my_key"] + " and there", } @@ -8114,10 +8116,10 @@ async def test_doubly_nested_graph_state( child.set_entry_point("child_1") child.set_finish_point("child_1") - def parent_1(state: State): + async def parent_1(state: State): return {"my_key": "hi " + state["my_key"]} - def parent_2(state: State): + async def parent_2(state: State): return {"my_key": state["my_key"] + " and back again"} graph = StateGraph(State)