Compare commits

..
Author SHA1 Message Date
Elior Nataf LackritzandGitHub b4991f1ba3 fix(langgraph): store an update_state input's writes on the checkpoint it updates (#9258)
`update_state(..., as_node="__input__")` stored its writes on the
checkpoint it saved. A DeltaChannel reads the writes stored on a
checkpoint's ancestors, not its own, so that checkpoint read back
without the input unless the channel snapshotted there, and once a later
checkpoint snapshotted, the input was gone for good. Only a raw `Pregel`
graph whose input channel is a DeltaChannel hits it, since a
`StateGraph`'s input goes to its start channel.

The input's writes now go on the checkpoint the update builds on, as a
node update's do. An update on a checkpoint the thread has moved past
stores nothing there and snapshots the DeltaChannels it writes on its
own checkpoint instead, the rule a node update already follows (#9165).

Both kinds of update now go through one helper, which also gives a
message an id before saving it, as the loop's `put_writes` does with
`ensure_message_ids`. Before, a message saved through `update_state`
without an id read back with id `None`, on the node path too, while the
same message from a node got one.

JS twin: langchain-ai/langgraphjs#2973, where review found it.

## Tests

`test_delta_channel_update_state.py`, sync and async: an update as input
reads back on its checkpoint and after the next run, with
`snapshot_frequency` 1 and 2, and an update as input to an older
checkpoint stays out of its other branch. The frequency 2 and
older-checkpoint cases fail on main. Five more check that a message
saved through `update_state` gets an id every read keeps, on the node
path (latest and older checkpoint) and the input path, sync and async;
they fail on main. The langgraph suite passes, and `lint_package` and
`lint_tests` are clean.
2026-10-10 11:57:46 +00:00
Elior Nataf LackritzandGitHub 6aa0afba68 fix(checkpoint): default InMemorySaver delta history to the latest checkpoint (#9257)
`InMemorySaver.get_delta_channel_history` looked its target up by
`config["configurable"].get("checkpoint_id", "")`, so a config without a
`checkpoint_id` matched no checkpoint, and every channel came back with
no seed and no writes, which a `DeltaChannel` reads as empty. The base
implementation, `PostgresSaver` and `SqliteSaver` read from the thread's
latest checkpoint instead, as `get_tuple` does.

A missing id now resolves to the namespace's newest checkpoint id, with
the same `max` that `InMemorySaver.get_tuple` uses. Graphs always pass a
checkpoint id here, so only direct callers of the saver API were
affected, and a call with an id does the same work as before.

Thanks @Hotragn, who pinned this down on #8242, where @longquanzheng's
branch had already fixed it.

JS twin: langchain-ai/langgraphjs#2979

## Tests

- `test_get_channel_writes_without_checkpoint_id_reads_the_latest` fails
on main (`{'messages': {'writes': []}}`).
- The `checkpoint`, `checkpoint-conformance`, `checkpoint-sqlite`,
`checkpoint-postgres` (Postgres 16), `prebuilt` and `langgraph` suites
pass, the last two without their Redis cases. The DeltaChannel
conformance suite also passes against `InMemorySaver`; CI skips that run
in `libs/checkpoint`, where the conformance package isn't installed, so
the new test lives in `test_memory.py`.
2026-10-10 00:59:20 +00:00
Elior Nataf LackritzandGitHub 12aeb0fddb fix(langgraph): read a run's DeltaChannel input back only on its own checkpoints (#9260)
A raw `Pregel` graph whose input channel is a DeltaChannel saved a run's
input to it as a `NULL_TASK_ID` write on the checkpoint the run started
from, under `"sync"` and `"async"` durability. Readers apply a
checkpoint's `NULL_TASK_ID` writes as its own state, so:

- on a new thread that checkpoint is never saved, and the first run's
input was lost for good: after `{"log": [0], "plain": [0], ...}`, `log`
read `[2]` next to `plain` `[0, 2]`;
- the last checkpoint of a run read the next run's input;
- a run with input from an older checkpoint (`checkpoint_id`) leaked its
input into the branch that already grew from it, and read that branch's
input when it had one.

The input's writes now go on the checkpoint the run starts from under a
task id of their own, `uuid5(checkpoint_id, INPUT)`, the one
`update_state` uses for an input update. On a new thread, or a
checkpoint a `checkpoint_id` addressed, which may already have children
that would read them, nothing is stored there and the input checkpoint
snapshots the channel instead. A normal run stores the same write as
before, and `durability="exit"`, which already kept the input right, is
unchanged. Threads saved before this keep the input writes already
stored on their checkpoints.

A `StateGraph` routes its input through its start channel, so only raw
`Pregel` graphs hit this.

JS twin: langchain-ai/langgraphjs#2981

## Tests

`test_delta_channel_run_input.py`, under every durability, with `invoke`
and `ainvoke`: two runs on a new thread, and a run with input from an
older checkpoint whose other branch had DeltaChannel input or none.
Every checkpoint reads the same in the DeltaChannel and a plain channel.
The 12 sync and async cases fail on main, and dropping the task id, the
new-thread snapshot or the addressed-checkpoint snapshot each fails its
own cases. The `langgraph` and `prebuilt` suites pass (without Redis),
and `lint_package` and `lint_tests` are clean.
2026-10-09 15:59:10 -04:00
Elior Nataf LackritzandGitHub d05236f805 fix(langgraph): snapshot a DeltaChannel that an update_state Overwrite resets (#9263)
A node that writes an `Overwrite` to a DeltaChannel makes the loop
snapshot the channel on the checkpoint that superstep saves, so no read
replays across the reset. The same write through `update_state`, as a
node or as input, didn't snapshot: the value read back right, but reads
of the update's checkpoint and the ones after it walked back past the
reset to the last snapshot. `update_state` now snapshots a DeltaChannel
an `Overwrite` resets, as the loop does.

Raised in review of langchain-ai/langgraphjs#2973. JS twin:
langchain-ai/langgraphjs#2982

## Tests

`test_delta_channel_update_overwrite.py`, sync and async: an `Overwrite`
through `update_state` as a node, and as input to a raw `Pregel` graph,
snapshots the channel on the checkpoint the update saves, and the value
reads back as the overwrite. All four fail on main. The `langgraph`
suite passes (without Redis), and `lint_package` and `lint_tests` are
clean.
2026-10-09 15:26:16 -04:00
Elior Nataf LackritzandGitHub 9d92f33cca fix(langgraph): keep a fork's DeltaChannel snapshot from starting nodes that never ran (#9264)
A fork from a checkpoint with pending writes to a DeltaChannel snapshots
the channel on the fork's first checkpoint, so the fork doesn't replay
the other branch's writes. If the channel had never been written there,
the snapshot also gives it its first version, and `_mark_bumps_seen`
only marked the bump seen for nodes that already have a `versions_seen`
entry. A node subscribed to the channel that never ran has none, so the
bump started it, and it read the empty value before the node that writes
the channel had run:

```python
graph.invoke("go", config)  # writer writes [1] to d, reader reads [1]
fork = graph.update_state(first_checkpoint, {"a": "go"}, as_node="__input__")
graph.get_state(fork).next  # ("writer", "reader"): resuming runs reader on [] first
```

A replay from that checkpoint saves the same kind of fork from the loop.
A channel bumped from no version was never written, so
`create_checkpoint` now also marks the bump seen for the nodes the
channel triggers, from the graph's `trigger_to_nodes`, at every call
that can bump. Only raw `Pregel` graphs hit this: `StateGraph` nodes
trigger on their `branch:to:*` channels, not on state keys. An
`as_node="__copy__"` fork still starts such a node, with a plain list
channel as well, so that case isn't DeltaChannel-specific and isn't
changed here.

Raised in review of langchain-ai/langgraphjs#2974. JS twin:
langchain-ai/langgraphjs#2983

## Tests

`test_delta_channel_seal_subscribers.py`: an update as input (sync and
async) and a replay from the checkpoint before a DeltaChannel's first
write leave only the writer next, and the reader then reads the written
value once. All three fail on main. The `langgraph` suite passes
(without Redis), and `lint_package` and `lint_tests` are clean.
2026-10-09 14:36:27 -04:00
Elior Nataf LackritzandGitHub 26356227c4 test(checkpoint): run the delta conformance suite against InMemorySaver in CI (#9261)
The conformance self-test in `libs/checkpoint-conformance` ran every
capability against `InMemorySaver` but asserted only the base ones.
`libs/checkpoint/tests/test_conformance_delta.py`, which ran the delta
conformance tests against it, always skips in CI: `libs/checkpoint`
can't depend on the conformance package, which depends on it. So no CI
job checked `InMemorySaver.get_delta_channel_history` against the delta
conformance tests.

The self-test now asserts every capability `InMemorySaver` implements,
and the skipped copy is gone.

## Tests

- With `InMemorySaver.get_delta_channel_history` broken to return no
history, the self-test now fails on 7 of the 10 delta conformance tests;
on main it still passes.
- `make format`, `make lint` and `make test` pass in
`libs/checkpoint-conformance`.
2026-10-09 13:44:55 -04:00
Elior Nataf LackritzandGitHub a5dbacae0d fix(checkpoint-postgres,checkpoint-sqlite): treat an empty checkpoint_id as the latest in delta history (#9262)
`PostgresSaver` and `SqliteSaver`, sync and async, read a DeltaChannel's
history from the latest checkpoint only when the config's
`checkpoint_id` was missing or `None`. An empty one was looked up as an
id, matched nothing, and came back with no seed and no writes, while
their own `get_tuple` treats an empty id as the latest checkpoint, as
the default implementation does. They now check for a missing or empty
id the way `get_tuple` does. The in-memory saver gets the same in #9257.

Graphs always pass the id of a saved checkpoint, so only direct callers
of the saver API hit this.

## Tests

-
`test_{sync,async}_empty_checkpoint_id_reads_from_the_latest_checkpoint`
in `checkpoint-postgres/tests/test_delta_pagination.py` and
`checkpoint-sqlite/tests/test_delta_parent_walk.py`: all four fail on
main.
- The `checkpoint-postgres` (Postgres 16) and `checkpoint-sqlite` suites
pass, and `lint_package` and `lint_tests` are clean in both.
2026-10-09 13:44:51 -04:00
21 changed files with 815 additions and 272 deletions
@@ -14,8 +14,8 @@ async def memory_checkpointer():
@pytest.mark.asyncio
async def test_validate_memory_base():
"""InMemorySaver passes all base capability tests."""
async def test_validate_memory():
"""InMemorySaver passes the tests of every capability it implements."""
report = await validate(memory_checkpointer)
report.print_report()
assert report.passed_all_base(), f"Base tests failed: {report.to_dict()}"
assert report.passed_all(), f"Capability tests failed: {report.to_dict()}"
@@ -469,7 +469,7 @@ class PostgresSaver(BasePostgresSaver):
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = get_checkpoint_id(config)
if checkpoint_id is None:
if not checkpoint_id:
target = self.get_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
@@ -417,7 +417,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = get_checkpoint_id(config)
if checkpoint_id is None:
if not checkpoint_id:
target = await self.aget_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
@@ -164,6 +164,40 @@ def test_sync_walk_reads_nothing_newer_than_the_target(
assert read == _ids_from_target_down(configs)
def _empty_checkpoint_id(config: dict) -> dict:
return {"configurable": {**config["configurable"], "checkpoint_id": ""}}
async def test_async_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver:
await saver.setup()
configs = await _abuild_chain(saver)
latest = await saver.aget_delta_channel_history(
config=configs[-1], channels=[CHANNEL]
)
result = await saver.aget_delta_channel_history(
config=_empty_checkpoint_id(configs[-1]), channels=[CHANNEL]
)
assert latest[CHANNEL]["writes"]
assert result == latest
def test_sync_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
with PostgresSaver.from_conn_string(DEFAULT_URI) as saver:
saver.setup()
configs = _build_chain(saver)
latest = saver.get_delta_channel_history(config=configs[-1], channels=[CHANNEL])
result = saver.get_delta_channel_history(
config=_empty_checkpoint_id(configs[-1]), channels=[CHANNEL]
)
assert latest[CHANNEL]["writes"]
assert result == latest
async def test_root_target_has_no_history_and_still_terminates(
monkeypatch: pytest.MonkeyPatch,
) -> None:
@@ -539,7 +539,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
thread_id = str(config["configurable"]["thread_id"])
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = get_checkpoint_id(config)
if checkpoint_id is None:
if not checkpoint_id:
target = self.get_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
@@ -653,7 +653,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
thread_id = str(config["configurable"]["thread_id"])
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = get_checkpoint_id(config)
if checkpoint_id is None:
if not checkpoint_id:
target = await self.aget_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
@@ -65,6 +65,39 @@ async def test_async_walk_reaches_parent_whatever_the_id_order(
assert got[CHANNEL] == EXPECTED
EMPTY_CHECKPOINT_ID: dict[str, Any] = {
"configurable": {**CONFIG["configurable"], "checkpoint_id": ""}
}
def test_sync_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
with SqliteSaver.from_conn_string(":memory:") as saver:
root = saver.put(CONFIG, _checkpoint("a-older", {CHANNEL: "seed"}), {}, {})
saver.put_writes(root, [(CHANNEL, "write-root")], "task")
saver.put(root, _checkpoint("z-newer", {}), {}, {})
got = saver.get_delta_channel_history(
config=EMPTY_CHECKPOINT_ID, channels=[CHANNEL]
)
assert got[CHANNEL] == EXPECTED
async def test_async_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
root = await saver.aput(
CONFIG, _checkpoint("a-older", {CHANNEL: "seed"}), {}, {}
)
await saver.aput_writes(root, [(CHANNEL, "write-root")], "task")
await saver.aput(root, _checkpoint("z-newer", {}), {}, {})
got = await saver.aget_delta_channel_history(
config=EMPTY_CHECKPOINT_ID, channels=[CHANNEL]
)
assert got[CHANNEL] == EXPECTED
def test_walk_reaches_root_of_long_chain_with_descending_ids() -> None:
steps = 40
with SqliteSaver.from_conn_string(":memory:") as saver:
@@ -170,8 +170,8 @@ class InMemorySaver(
return {}
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = config["configurable"].get("checkpoint_id", "")
ns_storage = self.storage.get(thread_id, {}).get(checkpoint_ns, {})
checkpoint_id = get_checkpoint_id(config) or max(ns_storage, default="")
chain: list[str] = []
target_entry = ns_storage.get(checkpoint_id)
@@ -1,37 +0,0 @@
"""Run delta-channel conformance capabilities against InMemorySaver."""
from __future__ import annotations
import pytest
conformance = pytest.importorskip(
"langgraph.checkpoint.conformance",
reason="langgraph-checkpoint-conformance not installed",
)
@pytest.mark.asyncio
async def test_delta_channel_conformance():
# Imported inside the test: the module-level importorskip above is what
# makes these safe, so they cannot move to the top of the file.
from langgraph.checkpoint.conformance import validate # noqa: PLC0415
from langgraph.checkpoint.conformance.initializer import ( # noqa: PLC0415
checkpointer_test,
)
from langgraph.checkpoint.memory import InMemorySaver # noqa: PLC0415
@checkpointer_test(name="InMemorySaver")
async def mem_saver():
yield InMemorySaver()
report = await validate(
mem_saver,
capabilities={
"delta_channel_history",
},
)
for cap, result in report.results.items():
if result.passed is False:
details = "\n".join(result.failures or [])
pytest.fail(f"Capability {cap} failed:\n{details}")
+13
View File
@@ -420,6 +420,19 @@ class TestInMemorySaverDeltaChannel:
assert "seed" not in result
assert result["writes"] == []
def test_get_channel_writes_without_checkpoint_id_reads_the_latest(self) -> None:
saver = InMemorySaver()
thread: RunnableConfig = {
"configurable": {"thread_id": "t1", "checkpoint_ns": ""}
}
parent = saver.put(thread, empty_checkpoint(), {}, {})
saver.put_writes(parent, [("messages", "hi")], "task1")
saver.put(parent, empty_checkpoint(), {}, {})
result = saver.get_delta_channel_history(config=thread, channels=["messages"])
assert result == {"messages": {"writes": [("task1", "messages", "hi")]}}
class TestBaseFallbackGetChannelWrites:
"""Exercises the `BaseCheckpointSaver.get_delta_channel_history` default
+35 -6
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
import uuid
from collections.abc import Callable, Iterable, Mapping
from collections.abc import Callable, Iterable, Mapping, Sequence
from datetime import datetime, timezone
from inspect import signature
from typing import Any, Literal, cast
@@ -32,6 +32,7 @@ from langgraph._internal._constants import (
)
from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel
from langgraph.channels.binop import _get_overwrite
from langgraph.channels.delta import DeltaChannel
from langgraph.managed.base import ManagedValueMapping, ManagedValueSpec
@@ -152,6 +153,23 @@ def delta_channels_with_pending_writes(
}
def delta_channels_overwritten(
specs: Mapping[str, Any], writes: Iterable[tuple[str, Any]]
) -> set[str]:
"""Return the names of the DeltaChannels that `writes` set with an `Overwrite`.
`update_state` saves a full snapshot of these channels in the checkpoint it
creates, like the loop does when a node returns an `Overwrite`. Otherwise,
reading the channel later starts from an older snapshot and replays the
writes the `Overwrite` threw away.
"""
return {
ch
for ch, value in writes
if isinstance(specs.get(ch), DeltaChannel) and _get_overwrite(value)[0]
}
def checkpoint_superseded(
saver: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple
) -> bool:
@@ -256,6 +274,7 @@ def create_checkpoint(
get_next_version: GetNextVersion | None = None,
channels_to_snapshot: set[str] | None = None,
stored_versions: ChannelVersions | None = None,
trigger_to_nodes: Mapping[str, Sequence[str]] | None = None,
) -> Checkpoint:
"""Build a new Checkpoint from the previous one and live channel state.
@@ -314,7 +333,9 @@ def create_checkpoint(
id=id or str(uuid6(clock_seq=step)),
channel_values=values,
channel_versions=channel_versions,
versions_seen=_mark_bumps_seen(checkpoint["versions_seen"], bumped),
versions_seen=_mark_bumps_seen(
checkpoint["versions_seen"], bumped, trigger_to_nodes or {}
),
updated_channels=None if updated_channels is None else sorted(updated_channels),
)
@@ -322,19 +343,27 @@ def create_checkpoint(
def _mark_bumps_seen(
versions_seen: dict[str, ChannelVersions],
bumped: Mapping[str, tuple[Any, Any]],
trigger_to_nodes: Mapping[str, Sequence[str]],
) -> dict[str, ChannelVersions]:
"""Advance whoever had seen a bumped channel's old version to the new one.
A bump that only stores a snapshot is not a write. Left unseen, it would
re-fire `interrupt_before` and rerun the channel's subscribers. For each
entry it advances, `SNAPSHOT_BUMPS` keeps the new version and the one the
node really read, so `versions_seen_without_bumps` can put the read back.
re-fire `interrupt_before` and rerun the channel's subscribers. A channel
bumped from no version was never written, so it also goes to the
subscribers that never ran: they have no entry, and would start on the bump.
For each entry it advances, `SNAPSHOT_BUMPS` keeps the new version and the
one the node really read, so `versions_seen_without_bumps` can put the read
back.
"""
if not bumped:
return versions_seen
out = dict(versions_seen)
for k, (old, _) in bumped.items():
if old is None:
for node in trigger_to_nodes.get(k, ()):
out.setdefault(node, {})
marks = dict(versions_seen.get(SNAPSHOT_BUMPS, {}))
for node, seen in versions_seen.items():
for node, seen in out.items():
if node == SNAPSHOT_BUMPS:
continue
for k, (old, new) in bumped.items():
+38 -30
View File
@@ -19,6 +19,7 @@ from typing import (
TypeVar,
cast,
)
from uuid import UUID, uuid5
from langchain_core.callbacks import AsyncParentRunManager, ParentRunManager
from langchain_core.runnables import RunnableConfig
@@ -160,12 +161,6 @@ def DuplexStream(*streams: StreamProtocol) -> StreamProtocol:
return StreamProtocol(__call__, {mode for s in streams for mode in s.modes})
def _cacheable(writes: WritesT) -> bool:
# A cache hit skips the node, so a cached error would skip it without
# raising: only a node that finished goes to the cache.
return not any(c in (INTERRUPT, ERROR) for c, _ in writes)
class PregelLoop:
config: RunnableConfig
store: BaseStore | None
@@ -275,6 +270,9 @@ class PregelLoop:
# `_put_exit_delta_writes` uses this to decide between anchoring on
# the existing parent (True) or creating a lazy stub (False).
_has_persisted_parent: bool = False
# True iff `__enter__` loaded the thread's latest checkpoint, not one a
# `checkpoint_id` addressed, so nothing has been built on it yet.
_loaded_latest: bool = False
managed: ManagedValueMapping
checkpoint: Checkpoint
@@ -445,9 +443,7 @@ class PregelLoop:
return None
return self._graph_lifecycle_events.popleft()
def put_writes(
self, task_id: str, writes: WritesT, *, cached: bool = False
) -> None:
def put_writes(self, task_id: str, writes: WritesT) -> None:
"""Put writes for a task, to be read by the next tick."""
if not writes:
return
@@ -540,7 +536,7 @@ class PregelLoop:
self._error_handler_write_futs.append(fut)
# output writes
if hasattr(self, "tasks"):
self.output_writes(task_id, writes, cached=cached)
self.output_writes(task_id, writes)
def _put_pending_writes(self) -> None:
if self.checkpointer_put_writes is None:
@@ -1125,16 +1121,26 @@ class PregelLoop:
self._exit_delta_writes.append(
(self.step, NULL_TASK_ID, "", c, v)
)
# Persist delta-channel input writes so sub-freq inputs are
# recoverable via ancestor walk (mirrors the Command input path).
# A DeltaChannel reads its input from the writes stored on the
# checkpoint this run starts from, under a task id of their own:
# readers apply a checkpoint's NULL_TASK_ID writes as its own state.
# A new thread has no checkpoint to store them on, and one a
# `checkpoint_id` addressed may have children that would read them,
# so then the input checkpoint snapshots the channel instead.
if self.durability != "exit":
delta_input = [
(c, v)
for c, v in input_writes
if isinstance(self.specs.get(c), DeltaChannel)
]
if delta_input:
self.put_writes(NULL_TASK_ID, delta_input)
if delta_input and self._has_persisted_parent and self._loaded_latest:
self.put_writes(
str(uuid5(UUID(self.checkpoint["id"]), INPUT)), delta_input
)
else:
self._delta_channels_forced_snapshot.update(
c for c, _ in delta_input
)
# save input checkpoint
self.updated_channels = updated_channels
self._put_checkpoint({"source": "input"})
@@ -1264,6 +1270,7 @@ class PregelLoop:
else None,
channels_to_snapshot=channels_to_snapshot,
stored_versions=self.checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
for k in channels_to_snapshot:
new_counters[k] = (0, 0)
@@ -1692,7 +1699,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
) -> PregelExecutableTask | None:
if pushed := super().accept_push(task, write_idx, call):
for task in self.match_cached_writes():
self.put_writes(task.id, task.writes, cached=True)
self.output_writes(task.id, task.writes, cached=True)
return pushed
def schedule_error_handler(
@@ -1729,18 +1736,16 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
if self._reapplies_pending_writes:
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
for task in self.match_cached_writes():
self.put_writes(task.id, task.writes, cached=True)
self.output_writes(task.id, task.writes, cached=True)
return handler_task
def put_writes(
self, task_id: str, writes: WritesT, *, cached: bool = False
) -> None:
def put_writes(self, task_id: str, writes: WritesT) -> None:
"""Put writes for a task, to be read by the next tick."""
super().put_writes(task_id, writes, cached=cached)
if cached or not writes or self.cache is None or not hasattr(self, "tasks"):
super().put_writes(task_id, writes)
if not writes or self.cache is None or not hasattr(self, "tasks"):
return
task = self.tasks.get(task_id)
if task is None or task.cache_key is None or not _cacheable(writes):
if task is None or task.cache_key is None:
return
self.submit(
self.cache.set,
@@ -1784,6 +1789,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
# Normal case: fetch the most recent checkpoint for this
# graph/thread. Returns None on first invocation.
saved = self.checkpointer.get_tuple(self.checkpoint_config)
self._loaded_latest = True
# Capture before the synthetic-empty fallback below overwrites `saved`.
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
@@ -1947,7 +1953,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
) -> PregelExecutableTask | None:
if pushed := super().accept_push(task, write_idx, call):
for task in await self.amatch_cached_writes():
self.put_writes(task.id, task.writes, cached=True)
self.output_writes(task.id, task.writes, cached=True)
return pushed
async def aschedule_error_handler(
@@ -1984,18 +1990,19 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
if self._reapplies_pending_writes:
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
for task in await self.amatch_cached_writes():
self.put_writes(task.id, task.writes, cached=True)
self.output_writes(task.id, task.writes, cached=True)
return handler_task
def put_writes(
self, task_id: str, writes: WritesT, *, cached: bool = False
) -> None:
def put_writes(self, task_id: str, writes: WritesT) -> None:
"""Put writes for a task, to be read by the next tick."""
super().put_writes(task_id, writes, cached=cached)
if cached or not writes or self.cache is None or not hasattr(self, "tasks"):
super().put_writes(task_id, writes)
if not writes or self.cache is None or not hasattr(self, "tasks"):
return
task = self.tasks.get(task_id)
if task is None or task.cache_key is None or not _cacheable(writes):
if task is None or task.cache_key is None:
return
if writes[0][0] in (INTERRUPT, ERROR):
# only cache successful tasks
return
self.submit(
self.cache.aset,
@@ -2039,6 +2046,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
# Normal case: fetch the most recent checkpoint for this
# graph/thread. Returns None on first invocation.
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
self._loaded_latest = True
# Capture before the synthetic-empty fallback below overwrites `saved`.
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
+8 -1
View File
@@ -7,7 +7,7 @@ import sys
import threading
import time
import weakref
from collections.abc import Callable, Sequence
from collections.abc import Awaitable, Callable, Sequence
from contextlib import suppress
from dataclasses import dataclass, replace
from datetime import datetime, timedelta, timezone
@@ -686,6 +686,8 @@ async def arun_with_retry(
task: PregelExecutableTask,
retry_policy: Sequence[RetryPolicy] | None,
stream: bool = False,
match_cached_writes: Callable[[], Awaitable[Sequence[PregelExecutableTask]]]
| None = None,
configurable: dict[str, Any] | None = None,
) -> None:
"""Run a task asynchronously with retries."""
@@ -709,6 +711,11 @@ async def arun_with_retry(
)
},
)
if match_cached_writes is not None and task.cache_key is not None:
for t in await match_cached_writes():
if t is task:
# if the task is already cached, return
return
while True:
runtime = config.get(CONF, {}).get(CONFIG_KEY_RUNTIME)
if isinstance(runtime, Runtime):
+162 -64
View File
@@ -137,6 +137,7 @@ from langgraph.pregel._checkpoint import (
copy_checkpoint,
create_checkpoint,
create_checkpoint_plan_for_update_state_api,
delta_channels_overwritten,
delta_channels_with_pending_writes,
empty_checkpoint,
get_updated_channels_from_tasks,
@@ -152,6 +153,7 @@ from langgraph.pregel._loop import (
from langgraph.pregel._messages import (
StreamMessagesHandler,
StreamMessagesHandlerV2,
ensure_message_ids,
)
from langgraph.pregel._read import DEFAULT_BOUND, PregelNode
from langgraph.pregel._retry import RetryPolicy
@@ -1766,6 +1768,7 @@ class Pregel(
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = checkpointer.put(
checkpoint_config,
@@ -1788,6 +1791,22 @@ class Pregel(
)
if input_writes := deque(map_input(self.input_channels, values)):
_store_or_fork_delta_writes(
checkpointer,
config,
saved,
checkpoint_config,
self.channels,
[
(
str(uuid5(UUID(checkpoint["id"]), INPUT)),
input_writes,
None,
)
],
fork_pending,
is_first=is_first,
)
updated_channels = apply_writes(
checkpoint,
channels,
@@ -1795,6 +1814,9 @@ class Pregel(
checkpointer.get_next_version,
self.trigger_to_nodes,
)
fork_pending |= delta_channels_overwritten(
self.channels, input_writes
)
# apply input write to channels
next_step = (
@@ -1822,6 +1844,7 @@ class Pregel(
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = checkpointer.put(
checkpoint_config,
@@ -1833,13 +1856,6 @@ class Pregel(
),
)
# store the writes
checkpointer.put_writes(
next_config,
input_writes,
str(uuid5(UUID(checkpoint["id"]), INPUT)),
)
return patch_checkpoint_map(
next_config, saved.metadata if saved else None
)
@@ -2033,30 +2049,22 @@ class Pregel(
),
)
updated_channels = get_updated_channels_from_tasks(run_tasks)
# The base's other children replay whatever is stored on it, so an
# edit of an older checkpoint stores none of its writes there: the
# checkpoint written here carries them, its delta channels
# snapshotted. Later supersteps address the checkpoint just written.
if (
is_first
and saved is not None
and checkpoint_superseded(checkpointer, config, saved)
):
fork_pending.update(
ch
for ch in updated_channels
if isinstance(self.channels.get(ch), DeltaChannel)
)
elif saved is not None:
for task_id, task in zip(run_task_ids, run_tasks):
channel_writes = [w for w in task.writes if w[0] != PUSH]
if channel_writes:
checkpointer.put_writes(
checkpoint_config,
channel_writes,
task_id,
**_task_path_kwarg(checkpointer.put_writes, task),
)
fork_pending |= delta_channels_overwritten(
self.channels, (w for t in run_tasks for w in t.writes)
)
_store_or_fork_delta_writes(
checkpointer,
config,
saved,
checkpoint_config,
self.channels,
[
(task_id, [w for w in task.writes if w[0] != PUSH], task)
for task_id, task in zip(run_task_ids, run_tasks)
],
fork_pending,
is_first=is_first,
)
apply_writes(
checkpoint,
channels,
@@ -2086,6 +2094,7 @@ class Pregel(
else None,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = checkpointer.put(
checkpoint_config,
@@ -2258,6 +2267,7 @@ class Pregel(
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = await checkpointer.aput(
checkpoint_config,
@@ -2280,6 +2290,22 @@ class Pregel(
)
if input_writes := deque(map_input(self.input_channels, values)):
await _astore_or_fork_delta_writes(
checkpointer,
config,
saved,
checkpoint_config,
self.channels,
[
(
str(uuid5(UUID(checkpoint["id"]), INPUT)),
input_writes,
None,
)
],
fork_pending,
is_first=is_first,
)
updated_channels = apply_writes(
checkpoint,
channels,
@@ -2287,6 +2313,9 @@ class Pregel(
checkpointer.get_next_version,
self.trigger_to_nodes,
)
fork_pending |= delta_channels_overwritten(
self.channels, input_writes
)
# apply input write to channels
next_step = (
@@ -2314,6 +2343,7 @@ class Pregel(
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = await checkpointer.aput(
checkpoint_config,
@@ -2325,13 +2355,6 @@ class Pregel(
),
)
# store the writes
await checkpointer.aput_writes(
next_config,
input_writes,
str(uuid5(UUID(checkpoint["id"]), INPUT)),
)
return patch_checkpoint_map(
next_config, saved.metadata if saved else None
)
@@ -2524,30 +2547,22 @@ class Pregel(
),
)
updated_channels = get_updated_channels_from_tasks(run_tasks)
# The base's other children replay whatever is stored on it, so an
# edit of an older checkpoint stores none of its writes there: the
# checkpoint written here carries them, its delta channels
# snapshotted. Later supersteps address the checkpoint just written.
if (
is_first
and saved is not None
and await acheckpoint_superseded(checkpointer, config, saved)
):
fork_pending.update(
ch
for ch in updated_channels
if isinstance(self.channels.get(ch), DeltaChannel)
)
elif saved is not None:
for task_id, task in zip(run_task_ids, run_tasks):
channel_writes = [w for w in task.writes if w[0] != PUSH]
if channel_writes:
await checkpointer.aput_writes(
checkpoint_config,
channel_writes,
task_id,
**_task_path_kwarg(checkpointer.aput_writes, task),
)
fork_pending |= delta_channels_overwritten(
self.channels, (w for t in run_tasks for w in t.writes)
)
await _astore_or_fork_delta_writes(
checkpointer,
config,
saved,
checkpoint_config,
self.channels,
[
(task_id, [w for w in task.writes if w[0] != PUSH], task)
for task_id, task in zip(run_task_ids, run_tasks)
],
fork_pending,
is_first=is_first,
)
apply_writes(
checkpoint,
channels,
@@ -2577,6 +2592,7 @@ class Pregel(
else None,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = await checkpointer.aput(
checkpoint_config,
@@ -3038,7 +3054,7 @@ class Pregel(
# with channel updates applied only at the transition between steps.
while loop.tick():
for task in loop.match_cached_writes():
loop.put_writes(task.id, task.writes, cached=True)
loop.output_writes(task.id, task.writes, cached=True)
for _ in runner.tick(
[t for t in loop.tasks.values() if not t.writes],
timeout=self.step_timeout,
@@ -3509,7 +3525,7 @@ class Pregel(
try:
while loop.tick():
for task in await loop.amatch_cached_writes():
loop.put_writes(task.id, task.writes, cached=True)
loop.output_writes(task.id, task.writes, cached=True)
async for _ in runner.atick(
[t for t in loop.tasks.values() if not t.writes],
timeout=self.step_timeout,
@@ -4255,6 +4271,88 @@ def _task_path_kwarg(put_writes: Callable[..., Any], task: PregelTaskWrites) ->
return {"task_path": task_path_str(task.path)}
_UpdateWrites = Sequence[tuple[str, Sequence[tuple[str, Any]], PregelTaskWrites | None]]
def _delta_writes(
channels: Mapping[str, BaseChannel | ManagedValueSpec], writes: _UpdateWrites
) -> list[tuple[str, Any]]:
"""Give the update's DeltaChannel writes message ids, as the loop's
`put_writes` does, so every read of them returns the same ids."""
delta = [
(ch, value)
for _, task_writes, _ in writes
for ch, value in task_writes
if isinstance(channels.get(ch), DeltaChannel)
]
for _, value in delta:
ensure_message_ids(value)
return delta
def _store_or_fork_delta_writes(
checkpointer: BaseCheckpointSaver,
config: RunnableConfig,
saved: CheckpointTuple | None,
checkpoint_config: RunnableConfig,
channels: Mapping[str, BaseChannel | ManagedValueSpec],
writes: _UpdateWrites,
fork_pending: set[str],
*,
is_first: bool,
) -> None:
"""Save an update's writes where its DeltaChannels will read them.
A DeltaChannel rebuilds its value from the writes saved on a checkpoint's
ancestors, so the writes go on the checkpoint the update builds on. If the
thread already moved past that checkpoint, its other children would read
them too, so the new checkpoint snapshots those channels instead. `writes`
holds `(task_id, writes, task)` per task; `task` is `None` for input.
"""
delta = _delta_writes(channels, writes)
if saved is None:
return
if is_first and checkpoint_superseded(checkpointer, config, saved):
fork_pending.update(ch for ch, _ in delta)
return
for task_id, task_writes, task in writes:
if task_writes:
checkpointer.put_writes(
checkpoint_config,
task_writes,
task_id,
**(_task_path_kwarg(checkpointer.put_writes, task) if task else {}),
)
async def _astore_or_fork_delta_writes(
checkpointer: BaseCheckpointSaver,
config: RunnableConfig,
saved: CheckpointTuple | None,
checkpoint_config: RunnableConfig,
channels: Mapping[str, BaseChannel | ManagedValueSpec],
writes: _UpdateWrites,
fork_pending: set[str],
*,
is_first: bool,
) -> None:
"""Async `_store_or_fork_delta_writes`."""
delta = _delta_writes(channels, writes)
if saved is None:
return
if is_first and await acheckpoint_superseded(checkpointer, config, saved):
fork_pending.update(ch for ch, _ in delta)
return
for task_id, task_writes, task in writes:
if task_writes:
await checkpointer.aput_writes(
checkpoint_config,
task_writes,
task_id,
**(_task_path_kwarg(checkpointer.aput_writes, task) if task else {}),
)
def _update_task_id(checkpoint_id: str, i: int) -> str:
"""Task id for the `i`th update of a superstep that has no task to reuse.
@@ -1,122 +0,0 @@
"""A node served from the node cache saves its writes like a node that ran,
without writing them back to the cache, and a node that fails isn't cached."""
import operator
from typing import Annotated, Any
import pytest
from langgraph.cache.memory import InMemoryCache
from langgraph.checkpoint.memory import InMemorySaver
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.types import CachePolicy, Durability
pytestmark = pytest.mark.anyio
INPUT = {"log": [], "plain": []}
def _append(current: list, writes: list) -> list:
return [*current, *(item for write in writes for item in write)]
class _State(TypedDict):
log: Annotated[list, DeltaChannel(_append)]
plain: Annotated[list, operator.add]
class _CountsSets(InMemoryCache):
sets = 0
def set(self, keys: Any) -> None:
self.sets += 1
super().set(keys)
def _a_then_cached_b_then_c(runs: list[str], cache: InMemoryCache) -> Any:
def node(name: str) -> Any:
def run(state: _State) -> dict:
runs.append(name)
return {"log": [name], "plain": [name]}
return run
builder = StateGraph(_State)
builder.add_node("a", node("a"))
builder.add_node("b", node("b"), cache_policy=CachePolicy())
builder.add_node("c", node("c"))
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("b", "c")
return builder.compile(checkpointer=InMemorySaver(), cache=cache)
def test_a_cache_hit_saves_its_writes_without_caching_them_again(
durability: Durability,
) -> None:
runs: list[str] = []
cache = _CountsSets()
graph = _a_then_cached_b_then_c(runs, cache)
graph.invoke(INPUT, {"configurable": {"thread_id": "1"}}, durability=durability)
config = {"configurable": {"thread_id": "2"}}
graph.invoke(INPUT, config, durability=durability)
assert runs == ["a", "b", "c", "a", "c"]
assert cache.sets == 1, "the cache hit was written back to the cache"
for state in graph.get_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
async def test_a_cache_hit_saves_its_writes_without_caching_them_again_async(
durability: Durability,
) -> None:
runs: list[str] = []
cache = _CountsSets()
graph = _a_then_cached_b_then_c(runs, cache)
await graph.ainvoke(
INPUT, {"configurable": {"thread_id": "1"}}, durability=durability
)
config = {"configurable": {"thread_id": "2"}}
await graph.ainvoke(INPUT, config, durability=durability)
assert runs == ["a", "b", "c", "a", "c"]
assert cache.sets == 1, "the cache hit was written back to the cache"
async for state in graph.aget_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
def _cached_node_that_fails(runs: list[str]) -> Any:
def fail(state: _State) -> dict:
runs.append("b")
raise ValueError("b failed")
builder = StateGraph(_State)
builder.add_node("b", fail, cache_policy=CachePolicy())
builder.add_edge(START, "b")
return builder.compile(checkpointer=InMemorySaver(), cache=InMemoryCache())
def test_a_node_that_fails_is_not_cached() -> None:
runs: list[str] = []
graph = _cached_node_that_fails(runs)
for thread in ("1", "2"):
with pytest.raises(ValueError, match="b failed"):
graph.invoke(INPUT, {"configurable": {"thread_id": thread}})
assert runs == ["b", "b"]
async def test_a_node_that_fails_is_not_cached_async() -> None:
runs: list[str] = []
graph = _cached_node_that_fails(runs)
for thread in ("1", "2"):
with pytest.raises(ValueError, match="b failed"):
await graph.ainvoke(INPUT, {"configurable": {"thread_id": thread}})
assert runs == ["b", "b"]
@@ -0,0 +1,103 @@
"""A run's input to a DeltaChannel input channel reads back on the checkpoints
built from it, and on no others."""
import pytest
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.last_value import LastValue
from langgraph.pregel import NodeBuilder, Pregel
from langgraph.types import Durability
pytestmark = pytest.mark.anyio
def _sorted_extend(current: list, writes: list) -> list:
return sorted([*current, *(item for write in writes for item in write)])
def _delta_input_graph() -> Pregel:
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log", "plain")
return Pregel(
nodes={"n": node},
channels={
"log": DeltaChannel(_sorted_extend),
"plain": BinaryOperatorAggregate(list, lambda a, b: sorted(a + b)),
"go": LastValue(int),
},
input_channels=["log", "plain", "go"],
output_channels=["log", "plain"],
checkpointer=InMemorySaver(),
)
def test_each_run_input_reads_back_on_its_own_checkpoints(
durability: Durability,
) -> None:
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"log": [0], "plain": [0], "go": 1}, config, durability=durability)
graph.invoke({"log": [5], "plain": [5], "go": 1}, config, durability=durability)
for state in graph.get_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
async def test_each_run_input_reads_back_on_its_own_checkpoints_async(
durability: Durability,
) -> None:
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke(
{"log": [0], "plain": [0], "go": 1}, config, durability=durability
)
await graph.ainvoke(
{"log": [5], "plain": [5], "go": 1}, config, durability=durability
)
async for state in graph.aget_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
OTHER_BRANCH_INPUTS = pytest.mark.parametrize(
"other_branch_input",
[{"go": 1}, {"log": [5], "plain": [5], "go": 1}],
ids=["other-branch-without-delta-input", "other-branch-with-delta-input"],
)
@OTHER_BRANCH_INPUTS
def test_run_input_from_an_older_checkpoint_stays_out_of_its_other_branch(
durability: Durability, other_branch_input: dict
) -> None:
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"go": 1}, config, durability=durability)
older = graph.get_state(config).config
graph.invoke(other_branch_input, config, durability=durability)
graph.invoke({"log": [7], "plain": [7], "go": 1}, older, durability=durability)
for state in graph.get_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
@OTHER_BRANCH_INPUTS
async def test_run_input_from_an_older_checkpoint_stays_out_of_its_other_branch_async(
durability: Durability, other_branch_input: dict
) -> None:
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"go": 1}, config, durability=durability)
older = (await graph.aget_state(config)).config
await graph.ainvoke(other_branch_input, config, durability=durability)
await graph.ainvoke(
{"log": [7], "plain": [7], "go": 1}, older, durability=durability
)
async for state in graph.aget_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
@@ -0,0 +1,78 @@
import pytest
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.last_value import LastValue
from langgraph.pregel import NodeBuilder, Pregel
pytestmark = pytest.mark.anyio
CONFIG = {"configurable": {"thread_id": "t"}}
def _extend(current: list, writes: list) -> list:
return [*current, *(item for write in writes for item in write)]
def _graph(reads: list) -> Pregel:
writer = NodeBuilder().subscribe_only("a").do(lambda _: [1]).write_to("d")
reader = NodeBuilder().subscribe_only("d").do(lambda d: reads.append(list(d)))
return Pregel(
nodes={"writer": writer, "reader": reader},
channels={"a": LastValue(str), "d": DeltaChannel(_extend)},
input_channels="a",
output_channels=["d"],
checkpointer=InMemorySaver(),
)
def test_an_input_update_from_before_the_first_write_starts_only_the_writer() -> None:
reads: list = []
graph = _graph(reads)
graph.invoke("go", CONFIG)
first = next(
s.config for s in graph.get_state_history(CONFIG) if s.metadata["step"] == -1
)
fork = graph.update_state(first, {"a": "go"}, as_node="__input__")
assert graph.get_state(fork).next == ("writer",)
graph.invoke(None, fork)
assert reads == [[1], [1]]
def test_a_replay_from_before_the_first_write_forks_with_only_the_writer_next() -> None:
reads: list = []
graph = _graph(reads)
graph.invoke("go", CONFIG)
first = next(
s.config for s in graph.get_state_history(CONFIG) if s.metadata["step"] == -1
)
graph.invoke(None, first, durability="sync")
fork = next(
s for s in graph.get_state_history(CONFIG) if s.metadata["source"] == "fork"
)
assert fork.next == ("writer",)
graph.invoke(None, fork.config)
assert reads == [[1], [1], [1]]
async def test_an_ainput_update_from_before_the_first_write_starts_only_the_writer() -> (
None
):
reads: list = []
graph = _graph(reads)
await graph.ainvoke("go", CONFIG)
first = [
s.config
async for s in graph.aget_state_history(CONFIG)
if s.metadata["step"] == -1
][0]
fork = await graph.aupdate_state(first, {"a": "go"}, as_node="__input__")
assert (await graph.aget_state(fork)).next == ("writer",)
await graph.ainvoke(None, fork)
assert reads == [[1], [1]]
@@ -0,0 +1,110 @@
"""An `Overwrite` through `update_state` snapshots its DeltaChannel on the
checkpoint the update saves, as a node's `Overwrite` does on the loop's."""
from typing import Annotated, Any
import pytest
from langchain_core.messages import HumanMessage
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.last_value import LastValue
from langgraph.graph import START, StateGraph
from langgraph.graph.message import _messages_delta_reducer
from langgraph.pregel import NodeBuilder, Pregel
from langgraph.types import Overwrite
pytestmark = pytest.mark.anyio
class _State(TypedDict):
messages: Annotated[list, DeltaChannel(_messages_delta_reducer)]
def _messages_graph(saver: InMemorySaver) -> Any:
builder = StateGraph(_State)
builder.add_node("model", lambda state: {})
builder.add_edge(START, "model")
return builder.compile(checkpointer=saver)
def _extend(current: list, writes: list) -> list:
return [*current, *(item for write in writes for item in write)]
def _delta_input_graph(saver: InMemorySaver) -> Pregel:
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log")
return Pregel(
nodes={"n": node},
channels={"log": DeltaChannel(_extend), "go": LastValue(int)},
input_channels=["log", "go"],
output_channels=["log"],
checkpointer=saver,
)
def test_update_state_with_an_overwrite_snapshots_the_channel() -> None:
saver = InMemorySaver()
graph = _messages_graph(saver)
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"messages": [HumanMessage(content="a", id="1")]}, config)
graph.update_state(
config,
{"messages": Overwrite([HumanMessage(content="b", id="2")])},
as_node="model",
)
head = saver.get_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
assert [m.content for m in graph.get_state(config).values["messages"]] == ["b"]
async def test_aupdate_state_with_an_overwrite_snapshots_the_channel() -> None:
saver = InMemorySaver()
graph = _messages_graph(saver)
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"messages": [HumanMessage(content="a", id="1")]}, config)
await graph.aupdate_state(
config,
{"messages": Overwrite([HumanMessage(content="b", id="2")])},
as_node="model",
)
head = await saver.aget_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
values = (await graph.aget_state(config)).values
assert [m.content for m in values["messages"]] == ["b"]
def test_update_state_as_input_with_an_overwrite_snapshots_the_channel() -> None:
saver = InMemorySaver()
graph = _delta_input_graph(saver)
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"log": [0], "go": 1}, config)
graph.update_state(config, {"log": Overwrite([1])}, as_node="__input__")
head = saver.get_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
assert graph.get_state(config).values["log"] == [1]
async def test_aupdate_state_as_input_with_an_overwrite_snapshots_the_channel() -> None:
saver = InMemorySaver()
graph = _delta_input_graph(saver)
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"log": [0], "go": 1}, config)
await graph.aupdate_state(config, {"log": Overwrite([1])}, as_node="__input__")
head = await saver.aget_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
assert (await graph.aget_state(config)).values["log"] == [1]
@@ -25,9 +25,12 @@ from langgraph.checkpoint.memory import InMemorySaver
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from typing_extensions import TypedDict
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.last_value import LastValue
from langgraph.graph import START, StateGraph
from langgraph.graph.message import _messages_delta_reducer
from langgraph.pregel import NodeBuilder, Pregel
from langgraph.types import StateSnapshot, StateUpdate
pytestmark = pytest.mark.anyio
@@ -538,3 +541,189 @@ def test_update_state_that_snapshots_keeps_a_deferred_node_pending() -> None:
assert [m.content for m in final["messages"]] == ["s", "a", "u", "b"]
assert graph.get_state(config).next == ()
def _sorted_extend(current: list, writes: list) -> list:
return sorted([*current, *(item for write in writes for item in write)])
def _delta_input_graph(snapshot_frequency: int = 1000) -> Any:
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log", "plain")
return Pregel(
nodes={"n": node},
channels={
"log": DeltaChannel(_sorted_extend, snapshot_frequency=snapshot_frequency),
"plain": BinaryOperatorAggregate(list, lambda a, b: sorted(a + b)),
"go": LastValue(int),
},
input_channels=["log", "plain", "go"],
output_channels=["log", "plain"],
checkpointer=InMemorySaver(),
)
@pytest.mark.parametrize("snapshot_frequency", [1, 2])
def test_update_as_input_reads_back_on_its_checkpoint_and_after_the_next_run(
snapshot_frequency: int,
) -> None:
graph = _delta_input_graph(snapshot_frequency)
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"log": [0], "plain": [0], "go": 1}, config)
graph.update_state(config, {"log": [1], "plain": [1], "go": 1}, as_node="__input__")
after_update = graph.get_state(config).values
graph.invoke(None, config)
after_run = graph.get_state(config).values
assert after_update["log"] == after_update["plain"]
assert after_run["log"] == after_run["plain"]
@pytest.mark.parametrize("snapshot_frequency", [1, 2])
async def test_aupdate_as_input_reads_back_on_its_checkpoint_and_after_the_next_run(
snapshot_frequency: int,
) -> None:
graph = _delta_input_graph(snapshot_frequency)
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"log": [0], "plain": [0], "go": 1}, config)
await graph.aupdate_state(
config, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
)
after_update = (await graph.aget_state(config)).values
await graph.ainvoke(None, config)
after_run = (await graph.aget_state(config)).values
assert after_update["log"] == after_update["plain"]
assert after_run["log"] == after_run["plain"]
def test_update_as_input_to_an_older_checkpoint_stays_out_of_its_other_branch() -> None:
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"go": 1}, config)
older = graph.get_state(config).config
graph.invoke({"go": 1}, config)
other_branch = graph.get_state(config)
edited = graph.update_state(
older, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
)
values = graph.get_state(edited).values
assert values["log"] == values["plain"]
assert graph.get_state(other_branch.config).values == other_branch.values
async def test_aupdate_as_input_to_an_older_checkpoint_stays_out_of_its_other_branch() -> (
None
):
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"go": 1}, config)
older = (await graph.aget_state(config)).config
await graph.ainvoke({"go": 1}, config)
other_branch = await graph.aget_state(config)
edited = await graph.aupdate_state(
older, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
)
values = (await graph.aget_state(edited)).values
assert values["log"] == values["plain"]
assert (await graph.aget_state(other_branch.config)).values == other_branch.values
def _message_ids(graph: Any, config: dict) -> list[str | None]:
return [m.id for m in graph.get_state(config).values["messages"]]
async def _amessage_ids(graph: Any, config: dict) -> list[str | None]:
return [m.id for m in (await graph.aget_state(config)).values["messages"]]
def _messages_input_graph() -> Any:
node = (
NodeBuilder()
.subscribe_only("go")
.do(lambda _: [HumanMessage("n", id="n")])
.write_to("messages")
)
return Pregel(
nodes={"n": node},
channels={
"messages": DeltaChannel(_messages_delta_reducer),
"go": LastValue(int),
},
input_channels=["messages", "go"],
output_channels=["messages"],
checkpointer=InMemorySaver(),
)
def test_update_state_gives_a_message_an_id_that_every_read_keeps() -> None:
graph = _build_graph(InMemorySaver())
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"messages": [HumanMessage("a", id="a")]}, config)
graph.update_state(config, {"messages": [HumanMessage("b")]})
first, second = _message_ids(graph, config), _message_ids(graph, config)
assert first[-1] is not None
assert first == second
async def test_aupdate_state_gives_a_message_an_id_that_every_read_keeps() -> None:
graph = _build_graph(InMemorySaver())
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"messages": [HumanMessage("a", id="a")]}, config)
await graph.aupdate_state(config, {"messages": [HumanMessage("b")]})
first, second = (
await _amessage_ids(graph, config),
await _amessage_ids(graph, config),
)
assert first[-1] is not None
assert first == second
def test_update_state_on_an_older_checkpoint_gives_a_message_an_id() -> None:
graph = _build_graph(InMemorySaver())
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"messages": [HumanMessage("a", id="a")]}, config)
older = graph.get_state(config).config
graph.invoke({"messages": [HumanMessage("c", id="c")]}, config)
branch = graph.update_state(older, {"messages": [HumanMessage("b")]})
assert _message_ids(graph, branch)[-1] is not None
def test_update_as_input_gives_a_message_an_id_that_every_read_keeps() -> None:
graph = _messages_input_graph()
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"messages": [HumanMessage("a", id="a")], "go": 1}, config)
graph.update_state(config, {"messages": [HumanMessage("b")]}, as_node="__input__")
first, second = _message_ids(graph, config), _message_ids(graph, config)
assert first[-1] is not None
assert first == second
async def test_aupdate_as_input_gives_a_message_an_id_that_every_read_keeps() -> None:
graph = _messages_input_graph()
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"messages": [HumanMessage("a", id="a")], "go": 1}, config)
await graph.aupdate_state(
config, {"messages": [HumanMessage("b")]}, as_node="__input__"
)
first, second = (
await _amessage_ids(graph, config),
await _amessage_ids(graph, config),
)
assert first[-1] is not None
assert first == second
+2 -2
View File
@@ -5932,9 +5932,9 @@ def test_no_redundant_put_writes_for_cached_task(
put_writes_task_ids: list[str] = []
orig = PregelLoop.put_writes
def spy(self, task_id, writes, **kwargs):
def spy(self, task_id, writes):
put_writes_task_ids.append(task_id)
return orig(self, task_id, writes, **kwargs)
return orig(self, task_id, writes)
with patch.object(PregelLoop, "put_writes", spy):
result = workflow.invoke(Command(resume="ans"), config=config)
+2 -2
View File
@@ -8157,9 +8157,9 @@ async def test_no_redundant_put_writes_for_cached_task(
put_writes_task_ids: list[str] = []
orig = PregelLoop.put_writes
def spy(self, task_id, writes, **kwargs):
def spy(self, task_id, writes):
put_writes_task_ids.append(task_id)
return orig(self, task_id, writes, **kwargs)
return orig(self, task_id, writes)
with patch.object(PregelLoop, "put_writes", spy):
result = await workflow.ainvoke(Command(resume="ans"), config=config)