diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index 534c2dcb7..3a51ec67e 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -230,12 +230,13 @@ def _require_saver_for_history( checkpoint: Checkpoint, delta_channels: list[str], saver: BaseCheckpointSaver | None, + config: RunnableConfig | None, ) -> None: written = [k for k in delta_channels if k in checkpoint["channel_versions"]] - if written and saver is None: + if written and (saver is None or config is None): raise ValueError( f"DeltaChannel {written} has history to replay but no checkpointer " - "was passed to read it" + "and config were passed to read it" ) @@ -269,7 +270,7 @@ def channels_from_checkpoint( for k, spec in channel_specs.items() if _needs_replay(spec, checkpoint["channel_values"].get(k, MISSING)) ] - _require_saver_for_history(checkpoint, delta_channels, saver) + _require_saver_for_history(checkpoint, delta_channels, saver, config) histories: Mapping[str, Any] = {} if delta_channels and saver is not None and config is not None: histories = saver.get_delta_channel_history( @@ -312,7 +313,7 @@ async def achannels_from_checkpoint( for k, spec in channel_specs.items() if _needs_replay(spec, checkpoint["channel_values"].get(k, MISSING)) ] - _require_saver_for_history(checkpoint, delta_channels, saver) + _require_saver_for_history(checkpoint, delta_channels, saver, config) histories: Mapping[str, Any] = {} if delta_channels and saver is not None and config is not None: histories = await saver.aget_delta_channel_history( diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index 39a196d42..01e67c740 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -835,6 +835,28 @@ class Pregel( if auto_validate: self.validate() + def _state_checkpointer(self, config: RunnableConfig) -> BaseCheckpointSaver: + """The checkpointer state reads and writes use: the one a parent lends a + subgraph through the config, else this graph's own.""" + checkpointer = ensure_config(config)[CONF].get( + CONFIG_KEY_CHECKPOINTER, self.checkpointer + ) + if isinstance(checkpointer, BaseCheckpointSaver): + checkpointer = self._apply_checkpointer_allowlist(checkpointer) + if not checkpointer: + raise ValueError("No checkpointer set") + return checkpointer + + def _own_checkpoint_config(self, config: RunnableConfig) -> RunnableConfig: + """A `checkpointer=True` subgraph keeps one history per thread, stored + under its namespace with the task ids removed.""" + if self.checkpointer is not True: + return config + ns = config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") + return patch_configurable( + config, {CONFIG_KEY_CHECKPOINT_NS: recast_checkpoint_ns(ns)} + ) + def _apply_checkpointer_allowlist( self, checkpointer: BaseCheckpointSaver | None ) -> BaseCheckpointSaver | None: @@ -1385,13 +1407,7 @@ class Pregel( self, config: RunnableConfig, *, subgraphs: bool = False ) -> StateSnapshot: """Get the current state of the graph.""" - checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get( - CONFIG_KEY_CHECKPOINTER, self.checkpointer - ) - if isinstance(checkpointer, BaseCheckpointSaver): - checkpointer = self._apply_checkpointer_allowlist(checkpointer) - if not checkpointer: - raise ValueError("No checkpointer set") + checkpointer = self._state_checkpointer(config) if ( checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") @@ -1408,11 +1424,7 @@ class Pregel( raise ValueError(f"Subgraph {recast} not found") config = merge_configs(self.config, config) if self.config else config - if self.checkpointer is True: - ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS]) - config = merge_configs( - config, {CONF: {CONFIG_KEY_CHECKPOINT_NS: recast_checkpoint_ns(ns)}} - ) + config = self._own_checkpoint_config(config) thread_id = config[CONF][CONFIG_KEY_THREAD_ID] if not isinstance(thread_id, str): config[CONF][CONFIG_KEY_THREAD_ID] = str(thread_id) @@ -1430,13 +1442,7 @@ class Pregel( self, config: RunnableConfig, *, subgraphs: bool = False ) -> StateSnapshot: """Get the current state of the graph.""" - checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get( - CONFIG_KEY_CHECKPOINTER, self.checkpointer - ) - if isinstance(checkpointer, BaseCheckpointSaver): - checkpointer = self._apply_checkpointer_allowlist(checkpointer) - if not checkpointer: - raise ValueError("No checkpointer set") + checkpointer = self._state_checkpointer(config) if ( checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") @@ -1453,11 +1459,7 @@ class Pregel( raise ValueError(f"Subgraph {recast} not found") config = merge_configs(self.config, config) if self.config else config - if self.checkpointer is True: - ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS]) - config = merge_configs( - config, {CONF: {CONFIG_KEY_CHECKPOINT_NS: recast_checkpoint_ns(ns)}} - ) + config = self._own_checkpoint_config(config) thread_id = config[CONF][CONFIG_KEY_THREAD_ID] if not isinstance(thread_id, str): config[CONF][CONFIG_KEY_THREAD_ID] = str(thread_id) @@ -1481,13 +1483,7 @@ class Pregel( ) -> Iterator[StateSnapshot]: """Get the history of the state of the graph.""" config = ensure_config(config) - checkpointer: BaseCheckpointSaver | None = config[CONF].get( - CONFIG_KEY_CHECKPOINTER, self.checkpointer - ) - if isinstance(checkpointer, BaseCheckpointSaver): - checkpointer = self._apply_checkpointer_allowlist(checkpointer) - if not checkpointer: - raise ValueError("No checkpointer set") + checkpointer = self._state_checkpointer(config) if ( checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") @@ -1516,6 +1512,7 @@ class Pregel( } }, ) + config = self._own_checkpoint_config(config) # eagerly consume list() to avoid holding up the db cursor for checkpoint_tuple in list( checkpointer.list(config, before=before, limit=limit, filter=filter) @@ -1534,13 +1531,7 @@ class Pregel( ) -> AsyncIterator[StateSnapshot]: """Asynchronously get the history of the state of the graph.""" config = ensure_config(config) - checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get( - CONFIG_KEY_CHECKPOINTER, self.checkpointer - ) - if isinstance(checkpointer, BaseCheckpointSaver): - checkpointer = self._apply_checkpointer_allowlist(checkpointer) - if not checkpointer: - raise ValueError("No checkpointer set") + checkpointer = self._state_checkpointer(config) if ( checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "") @@ -1570,6 +1561,7 @@ class Pregel( } }, ) + config = self._own_checkpoint_config(config) # eagerly consume list() to avoid holding up the db cursor for checkpoint_tuple in [ c @@ -1602,13 +1594,7 @@ class Pregel( RunnableConfig: The updated config. """ - checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get( - CONFIG_KEY_CHECKPOINTER, self.checkpointer - ) - if isinstance(checkpointer, BaseCheckpointSaver): - checkpointer = self._apply_checkpointer_allowlist(checkpointer) - if not checkpointer: - raise ValueError("No checkpointer set") + checkpointer = self._state_checkpointer(config) if len(supersteps) == 0: raise ValueError("No supersteps provided") @@ -1635,7 +1621,9 @@ class Pregel( input_config: RunnableConfig, updates: Sequence[StateUpdate] ) -> RunnableConfig: # get last checkpoint - config = ensure_config(self.config, input_config) + config = self._own_checkpoint_config( + ensure_config(self.config, input_config) + ) saved = checkpointer.get_tuple(config) if saved is not None: self._migrate_checkpoint(saved.checkpoint) @@ -2065,13 +2053,7 @@ class Pregel( RunnableConfig: The updated config. """ - checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get( - CONFIG_KEY_CHECKPOINTER, self.checkpointer - ) - if isinstance(checkpointer, BaseCheckpointSaver): - checkpointer = self._apply_checkpointer_allowlist(checkpointer) - if not checkpointer: - raise ValueError("No checkpointer set") + checkpointer = self._state_checkpointer(config) if len(supersteps) == 0: raise ValueError("No supersteps provided") @@ -2098,7 +2080,9 @@ class Pregel( input_config: RunnableConfig, updates: Sequence[StateUpdate] ) -> RunnableConfig: # get last checkpoint - config = ensure_config(self.config, input_config) + config = self._own_checkpoint_config( + ensure_config(self.config, input_config) + ) saved = await checkpointer.aget_tuple(config) if saved is not None: self._migrate_checkpoint(saved.checkpoint) @@ -2792,9 +2776,7 @@ class Pregel( "`durability` has no effect when no checkpointer is present.", ) # set up subgraph checkpointing - if self.checkpointer is True: - ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS]) - config[CONF][CONFIG_KEY_CHECKPOINT_NS] = recast_checkpoint_ns(ns) + config = self._own_checkpoint_config(config) # set up messages stream mode if "messages" in stream_modes: ns_ = cast(str | None, config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)) @@ -3219,9 +3201,7 @@ class Pregel( "`durability` has no effect when no checkpointer is present.", ) # set up subgraph checkpointing - if self.checkpointer is True: - ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS]) - config[CONF][CONFIG_KEY_CHECKPOINT_NS] = recast_checkpoint_ns(ns) + config = self._own_checkpoint_config(config) # set up messages stream mode if "messages" in stream_modes: # namespace can be None in a root level graph? diff --git a/libs/langgraph/tests/test_delta_channel_subgraph.py b/libs/langgraph/tests/test_delta_channel_subgraph.py index c873ae759..b5fcaeda0 100644 --- a/libs/langgraph/tests/test_delta_channel_subgraph.py +++ b/libs/langgraph/tests/test_delta_channel_subgraph.py @@ -3,6 +3,7 @@ from typing import Annotated, Any, Literal import pytest from langgraph.checkpoint.base import BaseCheckpointSaver +from langgraph.checkpoint.memory import InMemorySaver from typing_extensions import TypedDict from langgraph.channels.delta import DeltaChannel @@ -204,6 +205,84 @@ async def test_interrupted_subgraph_task_state_async( assert task.state.values == _both("a1") +@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"]) +def test_interrupted_subgraph_history_from_its_task_config( + sync_checkpointer: BaseCheckpointSaver, + persistence: Literal["per-invocation", "per-thread"], +) -> None: + app = _nested_app( + sync_checkpointer, + pause_before_b=True, + subgraph_checkpointer=True if persistence == "per-thread" else None, + ) + config = {"configurable": {"thread_id": "1"}} + app.invoke({}, config) + (task,) = app.get_state(config).tasks + + history = list(app.get_state_history(task.state)) + + assert [snapshot.values for snapshot in history][:1] == [_both("a1")] + + +@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"]) +async def test_interrupted_subgraph_ahistory_from_its_task_config( + async_checkpointer: BaseCheckpointSaver, + persistence: Literal["per-invocation", "per-thread"], +) -> None: + app = _nested_app( + async_checkpointer, + pause_before_b=True, + subgraph_checkpointer=True if persistence == "per-thread" else None, + ) + config = {"configurable": {"thread_id": "1"}} + await app.ainvoke({}, config) + (task,) = (await app.aget_state(config)).tasks + + history = [snapshot async for snapshot in app.aget_state_history(task.state)] + + assert [snapshot.values for snapshot in history][:1] == [_both("a1")] + + +@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"]) +def test_interrupted_subgraph_update_state_from_its_task_config( + sync_checkpointer: BaseCheckpointSaver, + persistence: Literal["per-invocation", "per-thread"], +) -> None: + app = _nested_app( + sync_checkpointer, + pause_before_b=True, + subgraph_checkpointer=True if persistence == "per-thread" else None, + ) + config = {"configurable": {"thread_id": "1"}} + app.invoke({}, config) + (task,) = app.get_state(config).tasks + + app.update_state(task.state, _both("edit"), as_node="a") + + (task,) = app.get_state(config, subgraphs=True).tasks + assert task.state.values == _both("a1", "edit") + + +@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"]) +async def test_interrupted_subgraph_aupdate_state_from_its_task_config( + async_checkpointer: BaseCheckpointSaver, + persistence: Literal["per-invocation", "per-thread"], +) -> None: + app = _nested_app( + async_checkpointer, + pause_before_b=True, + subgraph_checkpointer=True if persistence == "per-thread" else None, + ) + config = {"configurable": {"thread_id": "1"}} + await app.ainvoke({}, config) + (task,) = (await app.aget_state(config)).tasks + + await app.aupdate_state(task.state, _both("edit"), as_node="a") + + (task,) = (await app.aget_state(config, subgraphs=True)).tasks + assert task.state.values == _both("a1", "edit") + + def test_subgraph_update_state_keeps_history( sync_checkpointer: BaseCheckpointSaver, ) -> None: @@ -278,6 +357,15 @@ async def test_ahydrating_written_delta_channel_without_saver_raises() -> None: ) +def test_hydrating_written_delta_channel_without_config_raises() -> None: + with pytest.raises(ValueError, match="no checkpointer"): + channels_from_checkpoint( + {"delta": DeltaChannel(_extend)}, + _written_delta_checkpoint(), + saver=InMemorySaver(), + ) + + def test_hydrating_unwritten_delta_channel_without_saver_is_empty() -> None: channels, _ = channels_from_checkpoint( {"delta": DeltaChannel(_extend)}, empty_checkpoint()