Compare commits

...
Author SHA1 Message Date
Elior Nataf Lackritz 7d601b1b82 fix(langgraph): share one saver resolution between runs and state methods
`_defaults` picked the run's saver on its own, and disagreed with the state
methods on `checkpointer=False`: a run ignored a lent saver while
`update_state` wrote a checkpoint with it that no run reads. A
`checkpointer=True` root also got `True` back from `_state_checkpointer`
and failed with an AttributeError instead of the run's error. Both now go
through `_resolve_checkpointer`, and the state methods require an actual
saver.
2026-09-30 12:24:47 -04:00
Elior Nataf Lackritz 314cb19da8 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.
2026-09-29 10:52:52 -04:00
Elior Nataf Lackritz 0207ec7cff fix(langgraph): raise when a written DeltaChannel is hydrated without a saver
A missing saver used to hydrate the channel as empty, silently. Raise when a
DeltaChannel has a version but no stored value and no saver was passed.
Keying on the version keeps a never-written channel, which has no value
either, reading as empty without a saver.
2026-09-29 10:48:45 -04:00
Elior Nataf Lackritz a4a028a254 refactor(langgraph): make _prepare_state_snapshot recurse a bool
recurse only ever carried the same saver as saver, or None. Take a flag and
hand saver to the child's CONFIG_KEY_CHECKPOINTER, so the snapshot helpers
hold the saver in one argument.
2026-09-29 10:48:45 -04:00
0afdb91adc fix(langgraph): hydrate subgraph delta channels with the resolved saver
A subgraph has no saver of its own: it uses the parent's, passed through
`CONFIG_KEY_CHECKPOINTER`, or holds `checkpointer=True`. Every state
reader resolved that saver, then called `_prepare_state_snapshot`, which
re-derived it from `self.checkpointer` and got `None`. A `DeltaChannel`
stores nothing in `channel_values`, so without a saver it hydrated empty
while plain channels in the same snapshot read correctly.

`bulk_update_state` had the same expression, and once an update reached
the channel's snapshot cadence it persisted the empty value as a
`_DeltaSnapshot`, losing history on disk.

Pass the caller's saver in as a required keyword argument.

Fixes #8470

Co-authored-by: gururafiki <22777967+gururafiki@users.noreply.github.com>
Co-authored-by: Yuan Gao <119447586+DavidGao520@users.noreply.github.com>
2026-09-29 10:48:45 -04:00
3 changed files with 485 additions and 102 deletions
@@ -226,6 +226,20 @@ def _needs_replay(spec: BaseChannel, stored: object) -> bool:
return stored is MISSING
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 or config is None):
raise ValueError(
f"DeltaChannel {written} has history to replay but no checkpointer "
"or config was passed to read it"
)
def channels_from_checkpoint(
specs: Mapping[str, BaseChannel | ManagedValueSpec],
checkpoint: Checkpoint,
@@ -256,6 +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, config)
histories: Mapping[str, Any] = {}
if delta_channels and saver is not None and config is not None:
histories = saver.get_delta_channel_history(
@@ -298,6 +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, 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(
+74 -102
View File
@@ -835,6 +835,41 @@ class Pregel(
if auto_validate:
self.validate()
def _resolve_checkpointer(
self, config: RunnableConfig
) -> BaseCheckpointSaver | None:
"""The saver runs and state methods use: none for `checkpointer=False`,
else the one a parent lends a subgraph through the config, else this
graph's own."""
if self.checkpointer is False:
return None
conf = config.get(CONF, {})
if CONFIG_KEY_CHECKPOINTER in conf:
checkpointer = conf[CONFIG_KEY_CHECKPOINTER]
elif self.checkpointer is True:
raise RuntimeError("checkpointer=True cannot be used for root graphs.")
else:
checkpointer = self.checkpointer
if isinstance(checkpointer, BaseCheckpointSaver):
checkpointer = self._apply_checkpointer_allowlist(checkpointer)
return checkpointer
def _state_checkpointer(self, config: RunnableConfig) -> BaseCheckpointSaver:
checkpointer = self._resolve_checkpointer(ensure_config(config))
if not isinstance(checkpointer, BaseCheckpointSaver):
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:
@@ -1146,7 +1181,9 @@ class Pregel(
self,
config: RunnableConfig,
saved: CheckpointTuple | None,
recurse: BaseCheckpointSaver | None = None,
*,
saver: BaseCheckpointSaver,
recurse: bool = False,
apply_pending_writes: bool = False,
) -> StateSnapshot:
if not saved:
@@ -1169,9 +1206,7 @@ class Pregel(
channels, managed = channels_from_checkpoint(
self.channels,
saved.checkpoint,
saver=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
saver=saver,
config=saved.config,
)
# tasks for this checkpoint
@@ -1186,11 +1221,7 @@ class Pregel(
stop,
for_execution=True,
store=self.store,
checkpointer=(
self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None
),
checkpointer=saver,
manager=None,
)
# get the subgraphs
@@ -1217,7 +1248,7 @@ class Pregel(
# get the state of the subgraph
config = {
CONF: {
CONFIG_KEY_CHECKPOINTER: recurse,
CONFIG_KEY_CHECKPOINTER: saver,
"thread_id": saved.config[CONF]["thread_id"],
CONFIG_KEY_CHECKPOINT_NS: task_ns,
}
@@ -1269,7 +1300,9 @@ class Pregel(
self,
config: RunnableConfig,
saved: CheckpointTuple | None,
recurse: BaseCheckpointSaver | None = None,
*,
saver: BaseCheckpointSaver,
recurse: bool = False,
apply_pending_writes: bool = False,
) -> StateSnapshot:
if not saved:
@@ -1292,9 +1325,7 @@ class Pregel(
channels, managed = await achannels_from_checkpoint(
self.channels,
saved.checkpoint,
saver=self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
saver=saver,
config=saved.config,
)
# tasks for this checkpoint
@@ -1309,11 +1340,7 @@ class Pregel(
stop,
for_execution=True,
store=self.store,
checkpointer=(
self.checkpointer
if isinstance(self.checkpointer, BaseCheckpointSaver)
else None
),
checkpointer=saver,
manager=None,
)
# get the subgraphs
@@ -1340,7 +1367,7 @@ class Pregel(
# get the state of the subgraph
config = {
CONF: {
CONFIG_KEY_CHECKPOINTER: recurse,
CONFIG_KEY_CHECKPOINTER: saver,
"thread_id": saved.config[CONF]["thread_id"],
CONFIG_KEY_CHECKPOINT_NS: task_ns,
}
@@ -1393,13 +1420,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, "")
@@ -1416,11 +1437,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)
@@ -1429,7 +1446,8 @@ class Pregel(
return self._prepare_state_snapshot(
config,
saved,
recurse=checkpointer if subgraphs else None,
saver=checkpointer,
recurse=subgraphs,
apply_pending_writes=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
)
@@ -1437,13 +1455,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, "")
@@ -1460,11 +1472,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)
@@ -1473,7 +1481,8 @@ class Pregel(
return await self._aprepare_state_snapshot(
config,
saved,
recurse=checkpointer if subgraphs else None,
saver=checkpointer,
recurse=subgraphs,
apply_pending_writes=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
)
@@ -1487,13 +1496,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, "")
@@ -1522,12 +1525,13 @@ 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)
):
yield self._prepare_state_snapshot(
checkpoint_tuple.config, checkpoint_tuple
checkpoint_tuple.config, checkpoint_tuple, saver=checkpointer
)
async def aget_state_history(
@@ -1540,13 +1544,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, "")
@@ -1576,6 +1574,7 @@ class Pregel(
}
},
)
config = self._own_checkpoint_config(config)
# eagerly consume list() to avoid holding up the db cursor
for checkpoint_tuple in [
c
@@ -1584,7 +1583,7 @@ class Pregel(
)
]:
yield await self._aprepare_state_snapshot(
checkpoint_tuple.config, checkpoint_tuple
checkpoint_tuple.config, checkpoint_tuple, saver=checkpointer
)
def bulk_update_state(
@@ -1608,13 +1607,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")
@@ -1641,7 +1634,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)
@@ -1666,10 +1661,7 @@ class Pregel(
channels, managed = channels_from_checkpoint(
self.channels,
checkpoint,
saver=self.checkpointer
if saved is not None
and isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
saver=checkpointer,
config=saved.config if saved is not None else None,
)
values, as_node = updates[0][:2]
@@ -2074,13 +2066,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")
@@ -2107,7 +2093,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)
@@ -2132,10 +2120,7 @@ class Pregel(
channels, managed = await achannels_from_checkpoint(
self.channels,
checkpoint,
saver=self.checkpointer
if saved is not None
and isinstance(self.checkpointer, BaseCheckpointSaver)
else None,
saver=checkpointer,
config=saved.config if saved is not None else None,
)
values, as_node = updates[0][:2]
@@ -2576,16 +2561,7 @@ class Pregel(
stream_modes.add(print_mode)
else:
stream_modes.update(print_mode)
if self.checkpointer is False:
checkpointer: BaseCheckpointSaver | None = None
elif CONFIG_KEY_CHECKPOINTER in config.get(CONF, {}):
checkpointer = config[CONF][CONFIG_KEY_CHECKPOINTER]
elif self.checkpointer is True:
raise RuntimeError("checkpointer=True cannot be used for root graphs.")
else:
checkpointer = self.checkpointer
if isinstance(checkpointer, BaseCheckpointSaver):
checkpointer = self._apply_checkpointer_allowlist(checkpointer)
checkpointer = self._resolve_checkpointer(config)
if checkpointer and not config.get(CONF):
raise ValueError(
"Checkpointer requires one or more of the following 'configurable' "
@@ -2804,9 +2780,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))
@@ -3231,9 +3205,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?
@@ -0,0 +1,395 @@
import operator
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._internal._constants import CONFIG_KEY_CHECKPOINTER
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import END, START, StateGraph
from langgraph.pregel._checkpoint import (
achannels_from_checkpoint,
channels_from_checkpoint,
empty_checkpoint,
)
pytestmark = pytest.mark.anyio
def _extend(state: list | None, writes: list[Any]) -> list:
out = list(state or [])
for write in writes:
out.extend(write if isinstance(write, list) else [write])
return out
def _state_schema(snapshot_frequency: int = 1000) -> type:
class State(TypedDict, total=False):
delta: Annotated[
list, DeltaChannel(_extend, snapshot_frequency=snapshot_frequency)
]
plain: Annotated[list, operator.add]
return State
def _both(*items: str) -> dict:
return {"delta": list(items), "plain": list(items)}
def _child_builder(*, snapshot_frequency: int = 1000) -> StateGraph:
builder = StateGraph(_state_schema(snapshot_frequency))
builder.add_node("a", lambda state: _both("a1"))
builder.add_node("b", lambda state: _both("b1", "b2"))
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("b", END)
return builder
def _wrap(
inner: StateGraph,
*,
checkpointer: bool | None = None,
interrupt_before: list[str] | None = None,
) -> StateGraph:
builder = StateGraph(inner.state_schema)
builder.add_node(
"child",
inner.compile(checkpointer=checkpointer, interrupt_before=interrupt_before),
)
builder.add_edge(START, "child")
builder.add_edge("child", END)
return builder
def _nested_app(
checkpointer: BaseCheckpointSaver,
*,
depth: int = 1,
snapshot_frequency: int = 1000,
pause_before_b: bool = False,
subgraph_checkpointer: bool | None = None,
) -> Any:
graph = _child_builder(snapshot_frequency=snapshot_frequency)
for _ in range(depth):
graph = _wrap(
graph,
checkpointer=subgraph_checkpointer,
interrupt_before=["b"] if pause_before_b else None,
)
return graph.compile(checkpointer=checkpointer)
def _scoped(config: dict, namespace: str) -> dict:
return {"configurable": {**config["configurable"], "checkpoint_ns": namespace}}
def _child_namespace(app: Any, config: dict, *, depth: int = 1) -> str:
namespace = ""
for level in range(depth):
scoped = _scoped(config, namespace) if namespace else config
namespace = next(
(
task.state["configurable"]["checkpoint_ns"]
for snapshot in app.get_state_history(scoped)
for task in snapshot.tasks
if task.name == "child" and isinstance(task.state, dict)
),
"",
)
assert namespace, f"no `child` subgraph task at nesting level {level}"
return namespace
async def _achild_namespace(app: Any, config: dict) -> str:
async for snapshot in app.aget_state_history(config):
for task in snapshot.tasks:
if task.name == "child" and isinstance(task.state, dict):
return task.state["configurable"]["checkpoint_ns"]
raise AssertionError("no `child` subgraph task")
HISTORY = [_both("a1", "b1", "b2"), _both("a1"), _both(), _both()]
def test_subgraph_get_state(sync_checkpointer: BaseCheckpointSaver) -> None:
app = _nested_app(sync_checkpointer)
config = {"configurable": {"thread_id": "1"}}
app.invoke({}, config)
child = _scoped(config, _child_namespace(app, config))
assert app.get_state(config).values == _both("a1", "b1", "b2")
assert app.get_state(child).values == _both("a1", "b1", "b2")
async def test_subgraph_aget_state(async_checkpointer: BaseCheckpointSaver) -> None:
app = _nested_app(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
await app.ainvoke({}, config)
child = _scoped(config, await _achild_namespace(app, config))
assert (await app.aget_state(child)).values == _both("a1", "b1", "b2")
def test_subgraph_get_state_history(sync_checkpointer: BaseCheckpointSaver) -> None:
app = _nested_app(sync_checkpointer)
config = {"configurable": {"thread_id": "1"}}
app.invoke({}, config)
child = _scoped(config, _child_namespace(app, config))
assert [s.values for s in app.get_state_history(child)] == HISTORY
async def test_subgraph_aget_state_history(
async_checkpointer: BaseCheckpointSaver,
) -> None:
app = _nested_app(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
await app.ainvoke({}, config)
child = _scoped(config, await _achild_namespace(app, config))
assert [s.values async for s in app.aget_state_history(child)] == HISTORY
def test_doubly_nested_subgraph_get_state(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
app = _nested_app(sync_checkpointer, depth=2)
config = {"configurable": {"thread_id": "1"}}
app.invoke({}, config)
child = _scoped(config, _child_namespace(app, config, depth=2))
assert app.get_state(child).values == _both("a1", "b1", "b2")
@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"])
def test_interrupted_subgraph_task_state(
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, subgraphs=True).tasks
assert task.state.values == _both("a1")
@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"])
async def test_interrupted_subgraph_task_state_async(
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, subgraphs=True)).tasks
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:
app = _nested_app(sync_checkpointer, snapshot_frequency=2)
config = {"configurable": {"thread_id": "1"}}
app.invoke({}, config)
child = _scoped(config, _child_namespace(app, config))
app.update_state(child, _both("manual"))
assert app.get_state(child).values == _both("a1", "b1", "b2", "manual")
async def test_subgraph_aupdate_state_keeps_history(
async_checkpointer: BaseCheckpointSaver,
) -> None:
app = _nested_app(async_checkpointer, snapshot_frequency=2)
config = {"configurable": {"thread_id": "1"}}
await app.ainvoke({}, config)
child = _scoped(config, await _achild_namespace(app, config))
await app.aupdate_state(child, _both("manual"))
assert (await app.aget_state(child)).values == _both("a1", "b1", "b2", "manual")
def test_stateless_subgraph_persists_nothing(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
app = _nested_app(sync_checkpointer, subgraph_checkpointer=False)
config = {"configurable": {"thread_id": "1"}}
app.invoke({}, config)
child_tasks = [
task
for snapshot in app.get_state_history(config)
for task in snapshot.tasks
if task.name == "child" and isinstance(task.state, dict)
]
assert child_tasks == []
assert app.get_state(config).values == _both("a1", "b1", "b2")
def test_completed_subgraph_exposes_no_task_state(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
app = _nested_app(sync_checkpointer)
config = {"configurable": {"thread_id": "1"}}
app.invoke({}, config)
assert app.get_state(config, subgraphs=True).tasks == ()
def _written_delta_checkpoint() -> Any:
checkpoint = empty_checkpoint()
checkpoint["channel_versions"]["delta"] = 1
return checkpoint
def test_hydrating_written_delta_channel_without_saver_raises() -> None:
with pytest.raises(ValueError, match="no checkpointer"):
channels_from_checkpoint(
{"delta": DeltaChannel(_extend)}, _written_delta_checkpoint()
)
async def test_ahydrating_written_delta_channel_without_saver_raises() -> None:
with pytest.raises(ValueError, match="no checkpointer"):
await achannels_from_checkpoint(
{"delta": DeltaChannel(_extend)}, _written_delta_checkpoint()
)
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()
)
assert channels["delta"].get() == []
def test_root_checkpointer_true_graph_state_read_raises() -> None:
app = _child_builder().compile(checkpointer=True)
with pytest.raises(RuntimeError, match="checkpointer=True cannot be used"):
app.get_state({"configurable": {"thread_id": "1"}})
def test_stateless_graph_update_state_ignores_lent_saver() -> None:
saver = InMemorySaver()
app = _child_builder().compile(checkpointer=False)
config = {
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child:1",
CONFIG_KEY_CHECKPOINTER: saver,
}
}
with pytest.raises(ValueError, match="No checkpointer set"):
app.update_state(config, _both("x"), as_node="a")
assert list(saver.list(None)) == []