Compare commits

...
Author SHA1 Message Date
Elior Nataf Lackritzandiroiro147 b791c1f1d5 fix(langgraph): don't save a checkpoint past a failed DeltaChannel write
A DeltaChannel is rebuilt from its writes along the parent chain, so a
checkpoint saved without one of them reads back short for good. Each save
now gets the delta writes submitted before it, and fails if one of them
or the previous save failed.

Co-authored-by: iroiro147 <265728356+iroiro147@users.noreply.github.com>
2026-10-07 18:28:47 -04:00
2 changed files with 145 additions and 27 deletions
+30 -27
View File
@@ -191,6 +191,7 @@ class PregelLoop:
Callable[
[
concurrent.futures.Future | None,
Sequence[Any],
RunnableConfig,
Checkpoint,
str,
@@ -204,11 +205,13 @@ class PregelLoop:
submit: Submit
channels: Mapping[str, BaseChannel]
# Futures from `checkpointer.put_writes` calls that produced delta-channel
# writes. `_checkpointer_put_after_previous` drains this list (swap to a
# local `futs` then reset to `[]` and wait/gather) before putting the
# next checkpoint, so a checkpoint never becomes durable before the
# writes that produced it. Initialised to `[]` in both sync and async
# `__enter__`; stays `None` only when no checkpointer.
# writes. `_put_checkpoint` hands this list to the save it submits, which
# waits for them first, so a checkpoint never becomes durable before the
# writes that produced it. If a write or the previous save failed, the
# save fails too: a DeltaChannel is rebuilt from its writes along the
# parent chain, so a checkpoint saved past either gap reads back short
# for good. Initialised to `[]` in both sync and async `__enter__`;
# stays `None` only when no checkpointer.
_delta_write_futs: list[Any] | None = None
# Same pattern as `_delta_write_futs` but for error-handler writes.
@@ -1299,12 +1302,17 @@ class PregelLoop:
)
self.checkpoint_previous_versions = channel_versions
# Take this checkpoint's writes now: saves run in the background
# and can start out of order, so a save that took them itself
# could get another checkpoint's writes.
delta_write_futs, self._delta_write_futs = self._delta_write_futs, []
# save it, without blocking
# if there's a previous checkpoint save in progress, wait for it
# ensuring checkpointers receive checkpoints in order
self._put_checkpoint_fut = self.submit(
self._checkpointer_put_after_previous,
getattr(self, "_put_checkpoint_fut", None),
delta_write_futs,
self.checkpoint_config,
copy_checkpoint(self.checkpoint),
self.checkpoint_metadata,
@@ -1374,6 +1382,7 @@ class PregelLoop:
self._put_checkpoint_fut = self.submit(
self._checkpointer_put_after_previous,
getattr(self, "_put_checkpoint_fut", None),
(),
stub_put_config,
stub_cp,
{"step": -2},
@@ -1641,21 +1650,19 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
def _checkpointer_put_after_previous(
self,
prev: concurrent.futures.Future | None,
delta_write_futs: Sequence[concurrent.futures.Future],
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> RunnableConfig:
if self._delta_write_futs:
futs, self._delta_write_futs = self._delta_write_futs, []
concurrent.futures.wait(futs)
try:
if prev is not None:
prev.result()
finally:
cast(BaseCheckpointSaver, self.checkpointer).put(
config, checkpoint, metadata, new_versions
)
for fut in delta_write_futs:
fut.result()
if prev is not None:
prev.result()
cast(BaseCheckpointSaver, self.checkpointer).put(
config, checkpoint, metadata, new_versions
)
def match_cached_writes(self) -> Sequence[PregelExecutableTask]:
if self.cache is None:
@@ -1896,23 +1903,19 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
async def _checkpointer_put_after_previous(
self,
prev: asyncio.Task | None,
delta_write_futs: Sequence[asyncio.Future],
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> RunnableConfig:
# Drain DeltaChannel write futures before committing the checkpoint so
# ancestor walks never see a checkpoint without its backing writes.
if self._delta_write_futs:
futs, self._delta_write_futs = self._delta_write_futs, []
await asyncio.gather(*futs)
try:
if prev is not None:
await prev
finally:
await cast(BaseCheckpointSaver, self.checkpointer).aput(
config, checkpoint, metadata, new_versions
)
if delta_write_futs:
await asyncio.gather(*delta_write_futs)
if prev is not None:
await prev
await cast(BaseCheckpointSaver, self.checkpointer).aput(
config, checkpoint, metadata, new_versions
)
async def amatch_cached_writes(self) -> Sequence[PregelExecutableTask]:
if self.cache is None:
@@ -0,0 +1,115 @@
"""A checkpoint must never be saved without the `DeltaChannel` writes it reads."""
import operator
import threading
from typing import Annotated, Any
import pytest
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 Durability
pytestmark = pytest.mark.anyio
INPUT = {"log": [], "plain": []}
FINAL = {"log": ["a", "b", "c"], "plain": ["a", "b", "c"]}
# Exit mode saves nothing before the failed write, so its retry starts over.
RETRIES = [
pytest.param("sync", None, id="sync"),
pytest.param("async", None, id="async"),
pytest.param("exit", INPUT, id="exit"),
]
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 _FailsTheWriteOfBOnce(InMemorySaver):
failed = False
def _fail_once(self, writes: Any) -> None:
if not self.failed and ("log", ["b"]) in writes:
self.failed = True
raise ConnectionError("b's write was not saved")
def put_writes(
self, config: Any, writes: Any, task_id: str, task_path: str = ""
) -> None:
self._fail_once(writes)
super().put_writes(config, writes, task_id, task_path)
async def aput_writes(
self, config: Any, writes: Any, task_id: str, task_path: str = ""
) -> None:
self._fail_once(writes)
await super().aput_writes(config, writes, task_id, task_path)
def _a_then_b_then_c(saver: InMemorySaver) -> Any:
builder = StateGraph(_State)
for name in "abc":
builder.add_node(
name, lambda state, name=name: {"log": [name], "plain": [name]}
)
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("b", "c")
return builder.compile(checkpointer=saver)
@pytest.mark.parametrize(("durability", "retry_input"), RETRIES)
def test_a_failed_delta_write_is_rerun_not_lost(
durability: Durability, retry_input: dict | None
) -> None:
graph = _a_then_b_then_c(_FailsTheWriteOfBOnce())
config = {"configurable": {"thread_id": "t"}}
with pytest.raises(ConnectionError):
graph.invoke(INPUT, config, durability=durability)
for state in graph.get_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
graph.invoke(retry_input, config, durability=durability)
assert graph.get_state(config).values == FINAL
@pytest.mark.parametrize(("durability", "retry_input"), RETRIES)
async def test_a_failed_delta_write_is_rerun_not_lost_async(
durability: Durability, retry_input: dict | None
) -> None:
graph = _a_then_b_then_c(_FailsTheWriteOfBOnce())
config = {"configurable": {"thread_id": "t"}}
with pytest.raises(ConnectionError):
await graph.ainvoke(INPUT, config, durability=durability)
async for state in graph.aget_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
await graph.ainvoke(retry_input, config, durability=durability)
assert (await graph.aget_state(config)).values == FINAL
def test_a_delta_graph_finishes_on_a_single_background_thread() -> None:
graph = _a_then_b_then_c(InMemorySaver())
config = {"configurable": {"thread_id": "t"}, "max_concurrency": 1}
result: dict = {}
run = threading.Thread(
target=lambda: result.update(graph.invoke(INPUT, config, durability="async")),
daemon=True,
)
run.start()
run.join(timeout=10)
assert not run.is_alive(), "invoke hung"
assert result == FINAL