Compare commits

..
Author SHA1 Message Date
Elior Nataf Lackritz 0d91f9ed5b fix(langgraph): don't cache a node that failed in the sync loop
The async loop skipped caching an interrupted or failed task, but the
sync loop didn't, so a cached error made the next run with the same input
skip the node without raising. Both loops now share the check, which
looks at every write, since a cancelled task's error comes after any
writes it already made.

Also drops `arun_with_retry`'s `match_cached_writes`, unused since #4691
moved cache matching into the loop.
2026-10-10 07:47:26 -04:00
Elior Nataf Lackritz d8d02f052f fix(langgraph): save a cached node's writes like a node that ran
A cache hit's writes were applied but never saved, so a DeltaChannel,
rebuilt from saved writes, lost them for good under every durability.
The loop now saves a hit with put_writes like a node that ran, in all
six places it matches cached writes. A cached flag keeps it streamed as
cached and stops put_writes from writing it back to the cache.
2026-10-09 10:53:45 -04:00
21 changed files with 272 additions and 815 deletions
@@ -14,8 +14,8 @@ async def memory_checkpointer():
@pytest.mark.asyncio
async def test_validate_memory():
"""InMemorySaver passes the tests of every capability it implements."""
async def test_validate_memory_base():
"""InMemorySaver passes all base capability tests."""
report = await validate(memory_checkpointer)
report.print_report()
assert report.passed_all(), f"Capability tests failed: {report.to_dict()}"
assert report.passed_all_base(), f"Base 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 not checkpoint_id:
if checkpoint_id is None:
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 not checkpoint_id:
if checkpoint_id is None:
target = await self.aget_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
@@ -164,40 +164,6 @@ 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 not checkpoint_id:
if checkpoint_id is None:
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 not checkpoint_id:
if checkpoint_id is None:
target = await self.aget_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
@@ -65,39 +65,6 @@ 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)
@@ -0,0 +1,37 @@
"""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,19 +420,6 @@ 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
+6 -35
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
import uuid
from collections.abc import Callable, Iterable, Mapping, Sequence
from collections.abc import Callable, Iterable, Mapping
from datetime import datetime, timezone
from inspect import signature
from typing import Any, Literal, cast
@@ -32,7 +32,6 @@ 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
@@ -153,23 +152,6 @@ 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:
@@ -274,7 +256,6 @@ 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.
@@ -333,9 +314,7 @@ 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, trigger_to_nodes or {}
),
versions_seen=_mark_bumps_seen(checkpoint["versions_seen"], bumped),
updated_channels=None if updated_channels is None else sorted(updated_channels),
)
@@ -343,27 +322,19 @@ 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. 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.
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.
"""
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 out.items():
for node, seen in versions_seen.items():
if node == SNAPSHOT_BUMPS:
continue
for k, (old, new) in bumped.items():
+30 -38
View File
@@ -19,7 +19,6 @@ from typing import (
TypeVar,
cast,
)
from uuid import UUID, uuid5
from langchain_core.callbacks import AsyncParentRunManager, ParentRunManager
from langchain_core.runnables import RunnableConfig
@@ -161,6 +160,12 @@ 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
@@ -270,9 +275,6 @@ 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
@@ -443,7 +445,9 @@ class PregelLoop:
return None
return self._graph_lifecycle_events.popleft()
def put_writes(self, task_id: str, writes: WritesT) -> None:
def put_writes(
self, task_id: str, writes: WritesT, *, cached: bool = False
) -> None:
"""Put writes for a task, to be read by the next tick."""
if not writes:
return
@@ -536,7 +540,7 @@ class PregelLoop:
self._error_handler_write_futs.append(fut)
# output writes
if hasattr(self, "tasks"):
self.output_writes(task_id, writes)
self.output_writes(task_id, writes, cached=cached)
def _put_pending_writes(self) -> None:
if self.checkpointer_put_writes is None:
@@ -1121,26 +1125,16 @@ class PregelLoop:
self._exit_delta_writes.append(
(self.step, NULL_TASK_ID, "", c, v)
)
# 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.
# Persist delta-channel input writes so sub-freq inputs are
# recoverable via ancestor walk (mirrors the Command input path).
if self.durability != "exit":
delta_input = [
(c, v)
for c, v in input_writes
if isinstance(self.specs.get(c), DeltaChannel)
]
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
)
if delta_input:
self.put_writes(NULL_TASK_ID, delta_input)
# save input checkpoint
self.updated_channels = updated_channels
self._put_checkpoint({"source": "input"})
@@ -1270,7 +1264,6 @@ 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)
@@ -1699,7 +1692,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
) -> PregelExecutableTask | None:
if pushed := super().accept_push(task, write_idx, call):
for task in self.match_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
self.put_writes(task.id, task.writes, cached=True)
return pushed
def schedule_error_handler(
@@ -1736,16 +1729,18 @@ 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.output_writes(task.id, task.writes, cached=True)
self.put_writes(task.id, task.writes, cached=True)
return handler_task
def put_writes(self, task_id: str, writes: WritesT) -> None:
def put_writes(
self, task_id: str, writes: WritesT, *, cached: bool = False
) -> None:
"""Put writes for a task, to be read by the next tick."""
super().put_writes(task_id, writes)
if not writes or self.cache is None or not hasattr(self, "tasks"):
super().put_writes(task_id, writes, cached=cached)
if cached or 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:
if task is None or task.cache_key is None or not _cacheable(writes):
return
self.submit(
self.cache.set,
@@ -1789,7 +1784,6 @@ 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)
@@ -1953,7 +1947,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.output_writes(task.id, task.writes, cached=True)
self.put_writes(task.id, task.writes, cached=True)
return pushed
async def aschedule_error_handler(
@@ -1990,19 +1984,18 @@ 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.output_writes(task.id, task.writes, cached=True)
self.put_writes(task.id, task.writes, cached=True)
return handler_task
def put_writes(self, task_id: str, writes: WritesT) -> None:
def put_writes(
self, task_id: str, writes: WritesT, *, cached: bool = False
) -> None:
"""Put writes for a task, to be read by the next tick."""
super().put_writes(task_id, writes)
if not writes or self.cache is None or not hasattr(self, "tasks"):
super().put_writes(task_id, writes, cached=cached)
if cached or 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:
return
if writes[0][0] in (INTERRUPT, ERROR):
# only cache successful tasks
if task is None or task.cache_key is None or not _cacheable(writes):
return
self.submit(
self.cache.aset,
@@ -2046,7 +2039,6 @@ 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)
+1 -8
View File
@@ -7,7 +7,7 @@ import sys
import threading
import time
import weakref
from collections.abc import Awaitable, Callable, Sequence
from collections.abc import Callable, Sequence
from contextlib import suppress
from dataclasses import dataclass, replace
from datetime import datetime, timedelta, timezone
@@ -686,8 +686,6 @@ 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."""
@@ -711,11 +709,6 @@ 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):
+64 -162
View File
@@ -137,7 +137,6 @@ 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,
@@ -153,7 +152,6 @@ 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
@@ -1768,7 +1766,6 @@ 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,
@@ -1791,22 +1788,6 @@ 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,
@@ -1814,9 +1795,6 @@ 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 = (
@@ -1844,7 +1822,6 @@ 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,
@@ -1856,6 +1833,13 @@ 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
)
@@ -2049,22 +2033,30 @@ class Pregel(
),
)
updated_channels = get_updated_channels_from_tasks(run_tasks)
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,
)
# 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),
)
apply_writes(
checkpoint,
channels,
@@ -2094,7 +2086,6 @@ 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,
@@ -2267,7 +2258,6 @@ 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,
@@ -2290,22 +2280,6 @@ 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,
@@ -2313,9 +2287,6 @@ 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 = (
@@ -2343,7 +2314,6 @@ 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,
@@ -2355,6 +2325,13 @@ 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
)
@@ -2547,22 +2524,30 @@ class Pregel(
),
)
updated_channels = get_updated_channels_from_tasks(run_tasks)
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,
)
# 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),
)
apply_writes(
checkpoint,
channels,
@@ -2592,7 +2577,6 @@ 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,
@@ -3054,7 +3038,7 @@ class Pregel(
# with channel updates applied only at the transition between steps.
while loop.tick():
for task in loop.match_cached_writes():
loop.output_writes(task.id, task.writes, cached=True)
loop.put_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,
@@ -3525,7 +3509,7 @@ class Pregel(
try:
while loop.tick():
for task in await loop.amatch_cached_writes():
loop.output_writes(task.id, task.writes, cached=True)
loop.put_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,
@@ -4271,88 +4255,6 @@ 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.
@@ -0,0 +1,122 @@
"""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"]
@@ -1,103 +0,0 @@
"""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", [])
@@ -1,78 +0,0 @@
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]]
@@ -1,110 +0,0 @@
"""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,12 +25,9 @@ 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
@@ -541,189 +538,3 @@ 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):
def spy(self, task_id, writes, **kwargs):
put_writes_task_ids.append(task_id)
return orig(self, task_id, writes)
return orig(self, task_id, writes, **kwargs)
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):
def spy(self, task_id, writes, **kwargs):
put_writes_task_ids.append(task_id)
return orig(self, task_id, writes)
return orig(self, task_id, writes, **kwargs)
with patch.object(PregelLoop, "put_writes", spy):
result = await workflow.ainvoke(Command(resume="ans"), config=config)