mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-06 17:57:49 +02:00
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:
@@ -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))
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user