fix(langgraph): snapshot DeltaChannel overwrite supersteps (#8125)

When a DeltaChannel receives an Overwrite, force that checkpoint to
store a snapshot so sparse replay starts from the post-overwrite value.
This also aligns live DeltaChannel overwrite handling with
BinaryOperatorAggregate by letting Overwrite bypass other reducer writes
in the same superstep.
This commit is contained in:
Sydney Runkle
2026-06-29 17:06:10 -07:00
committed by GitHub
parent 1b5ca0a1b1
commit 9a27693c64
3 changed files with 113 additions and 4 deletions
+1 -3
View File
@@ -172,13 +172,11 @@ class DeltaChannel(Generic[Value], BaseChannel[Any, Any, Any]):
overwrite_idx = i
if overwrite_idx is not None:
_, overwrite_value = _get_overwrite(values[overwrite_idx])
base = (
self.value = (
_copy.copy(overwrite_value)
if overwrite_value is not None
else self.typ()
)
remaining = [v for i, v in enumerate(values) if i != overwrite_idx]
self.value = self.reducer(base, remaining) if remaining else base
return True
base = self.typ() if self.value is MISSING else self.value
self.value = self.reducer(base, list(values))
+25 -1
View File
@@ -70,6 +70,7 @@ from langgraph.callbacks import (
GraphResumeEvent,
)
from langgraph.channels.base import BaseChannel
from langgraph.channels.binop import _get_overwrite
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.untracked_value import UntrackedValue
from langgraph.constants import TAG_HIDDEN
@@ -221,6 +222,11 @@ class PregelLoop:
# under the saver's `ORDER BY task_id, idx` sorting.
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
# Delta channels that saw an Overwrite since the last checkpoint. These
# channels must snapshot after live update applies overwrite semantics so
# sparse replay starts from the same post-overwrite value.
_delta_channels_with_overwrite: set[str]
# The checkpoint_config that points at the parent loaded at `__enter__`
# (or the synthetic-empty checkpoint, on first run). We capture it
# eagerly because every `_put_checkpoint` advances `self.checkpoint_config`
@@ -677,6 +683,11 @@ class PregelLoop:
def after_tick(self) -> None:
# finish superstep
writes = [w for t in self.tasks.values() for w in t.writes]
self._delta_channels_with_overwrite.update(
ch
for ch, v in writes
if isinstance(self.specs.get(ch), DeltaChannel) and _get_overwrite(v)[0]
)
# all tasks have finished
self.updated_channels = apply_writes(
self.checkpoint,
@@ -980,6 +991,11 @@ class PregelLoop:
manager=None,
updated_channels=updated_channels,
)
self._delta_channels_with_overwrite.update(
c
for c, v in input_writes
if isinstance(self.specs.get(c), DeltaChannel) and _get_overwrite(v)[0]
)
# apply input writes
updated_channels = apply_writes(
self.checkpoint,
@@ -1120,6 +1136,7 @@ class PregelLoop:
# create new checkpoint
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, new_counters)
| self._delta_channels_with_overwrite
if do_checkpoint
else set()
)
@@ -1136,6 +1153,8 @@ class PregelLoop:
)
for k in channels_to_snapshot:
new_counters[k] = (0, 0)
if do_checkpoint:
self._delta_channels_with_overwrite.difference_update(channels_to_snapshot)
non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}
if non_zero:
self.checkpoint_metadata["counters_since_delta_snapshot"] = non_zero
@@ -1218,7 +1237,10 @@ class PregelLoop:
counters = dict(
self.checkpoint_metadata.get("counters_since_delta_snapshot") or {}
)
channels_to_snapshot = delta_channels_to_snapshot(self.channels, counters)
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, counters)
| self._delta_channels_with_overwrite
)
pending = [
(step, tid, ch, v)
@@ -1662,6 +1684,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
)
self._delta_write_futs = []
self._error_handler_write_futs = []
self._delta_channels_with_overwrite = set()
self._exit_delta_writes = (
[] if self.durability == "exit" and self.checkpointer is not None else None
)
@@ -1919,6 +1942,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
)
self._delta_write_futs = []
self._error_handler_write_futs = []
self._delta_channels_with_overwrite = set()
self._exit_delta_writes = (
[] if self.durability == "exit" and self.checkpointer is not None else None
)
+87
View File
@@ -441,6 +441,93 @@ def test_delta_channel_inmemory_saver_assembles_writes() -> None:
assert len(state.values["messages"]) == 4 # 2 human + 2 AI
def test_delta_channel_overwrite_superstep_snapshots() -> None:
def reducer(state: list[str], writes: Sequence[list[str]]) -> list[str]:
result = list(state)
for write in writes:
result.extend(write)
return result
class State(TypedDict):
items: Annotated[
list[str], DeltaChannel(reducer, list, snapshot_frequency=1000)
]
def node_a(state: State) -> dict:
return {"items": ["a"]}
def node_b(state: State) -> dict:
return {"items": Overwrite(["b"])}
def node_c(state: State) -> dict:
return {"items": ["c"]}
builder = StateGraph(State)
builder.add_node("node_a", node_a)
builder.add_node("node_b", node_b)
builder.add_node("node_c", node_c)
builder.add_edge(START, "node_a")
builder.add_edge("node_a", "node_b")
builder.add_edge("node_a", "node_c")
saver = InMemorySaver()
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "overwrite-snapshot"}}
result = graph.invoke({"items": ["START"]}, config)
assert result == {"items": ["b"]}
saved = saver.get_tuple(config)
assert saved is not None
snapshot = saved.checkpoint["channel_values"].get("items")
assert isinstance(snapshot, _DeltaSnapshot)
assert snapshot.value == ["b"]
assert saved.metadata.get("counters_since_delta_snapshot", {}).get("items") is None
def test_delta_channel_replay_after_overwrite_snapshot() -> None:
def reducer(state: list[str], writes: Sequence[list[str]]) -> list[str]:
result = list(state)
for write in writes:
result.extend(write)
return result
class State(TypedDict):
items: Annotated[
list[str], DeltaChannel(reducer, list, snapshot_frequency=1000)
]
calls = 0
def node(state: State) -> dict:
nonlocal calls
calls += 1
if calls == 1:
return {"items": Overwrite(["reset"])}
return {"items": ["after"]}
builder = StateGraph(State)
builder.add_node("node", node)
builder.add_edge(START, "node")
saver = InMemorySaver()
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "overwrite-replay"}}
assert graph.invoke({"items": ["before"]}, config) == {"items": ["reset"]}
first_saved = saver.get_tuple(config)
assert first_saved is not None
assert isinstance(
first_saved.checkpoint["channel_values"].get("items"), _DeltaSnapshot
)
assert graph.invoke({"items": []}, config) == {"items": ["reset", "after"]}
second_saved = saver.get_tuple(config)
assert second_saved is not None
assert "items" not in second_saved.checkpoint["channel_values"]
assert graph.get_state(config).values == {"items": ["reset", "after"]}
# ---------------------------------------------------------------------------
# DeltaChannel — dict reducer
# ---------------------------------------------------------------------------