fix(langgraph): resolve the saver and checkpointer=True namespace in one place

The six state methods each resolved the saver themselves, and only
get_state recast a checkpointer=True subgraph's namespace. So
get_state_history and update_state given the subgraph task's config
looked under `name:<task_id>`, where the run never stores anything: the
history came back empty and the update was silently lost. The state
methods and the run loop now share one saver lookup and one namespace
rule. The hydration guard also raises when a saver comes without a
config, which skipped the replay the same way.
This commit is contained in:
Elior Nataf Lackritz
2026-09-29 10:52:52 -04:00
parent 0207ec7cff
commit 314cb19da8
3 changed files with 133 additions and 64 deletions
@@ -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(
+40 -60
View File
@@ -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?
@@ -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()