simplify delta overwrite fix

This commit is contained in:
Sydney Runkle
2026-06-17 14:07:19 -04:00
parent 941c170c58
commit 4210feccd9
7 changed files with 18 additions and 194 deletions
@@ -74,49 +74,6 @@ async def test_history_excludes_target_pending_writes(
assert "extra" not in values, f"Target's writes should be excluded, got {values}"
async def test_history_overwrite_bypasses_same_step_writes(
saver: BaseCheckpointSaver,
) -> None:
tid = str(uuid4())
channel = "ch"
from langgraph.checkpoint.base import Checkpoint
from langgraph.checkpoint.base.id import uuid6
from langgraph.checkpoint.conformance.test_utils import generate_metadata
parent_cfg = None
configs: list = []
for step in range(2):
config = {"configurable": {"thread_id": tid, "checkpoint_ns": ""}}
if parent_cfg:
config["configurable"]["checkpoint_id"] = parent_cfg["configurable"][
"checkpoint_id"
]
cp = Checkpoint(
v=1,
id=str(uuid6(clock_seq=-1)),
ts="",
channel_values={},
channel_versions={},
versions_seen={},
updated_channels=None,
)
parent_cfg = await saver.aput(config, cp, generate_metadata(step=step), {})
configs.append(parent_cfg)
await saver.aput_writes(
configs[0],
[(channel, [1]), (channel, {"__overwrite__": [50]}), (channel, [2])],
str(uuid4()),
)
result = await saver.aget_delta_channel_history(
config=configs[1], channels=[channel]
)
values = [w[2] for w in result[channel]["writes"]]
assert values == [{"__overwrite__": [50]}], f"Expected overwrite only, got {values}"
async def test_history_multi_channel(
saver: BaseCheckpointSaver,
) -> None:
@@ -255,7 +212,6 @@ ALL_DELTA_CHANNEL_HISTORY_TESTS = [
test_history_returns_writes_oldest_first,
test_history_seed_is_nearest_snapshot,
test_history_excludes_target_pending_writes,
test_history_overwrite_bypasses_same_step_writes,
test_history_multi_channel,
test_history_empty_channels_returns_empty,
test_history_walk_to_root_no_seed,
@@ -13,7 +13,6 @@ from langgraph.checkpoint.base import (
ChannelVersions,
DeltaChannelHistory,
PendingWrite,
_apply_delta_history_overwrite_semantics,
get_checkpoint_id,
)
from langgraph.checkpoint.serde.types import TASKS
@@ -451,10 +450,10 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
"tuple[str, bytes]", (r["type"], r["blob"])
)
# Sort writes per (channel, cid) oldest-first by (task_id, idx).
# Sort writes per (channel, cid) newest-first by (task_id, idx)
for cid_map in writes_by_ch_by_cid.values():
for ws in cid_map.values():
ws.sort(key=lambda w: (w[2], w[3]))
ws.sort(key=lambda w: (w[2], w[3]), reverse=True)
result: dict[str, DeltaChannelHistory] = {}
for ch in channels:
@@ -463,13 +462,11 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
collected: list[PendingWrite] = []
cid_writes = writes_by_ch_by_cid.get(ch, {})
# Chain is newest-first; iterate oldest-first for the public order.
for cid in reversed(chain_cids):
step_writes: list[PendingWrite] = [
(task_id, ch, self.serde.loads_typed((type_tag, write_blob)))
for type_tag, write_blob, task_id, _idx in cid_writes.get(cid, [])
]
collected.extend(_apply_delta_history_overwrite_semantics(step_writes))
for cid in chain_cids:
for type_tag, write_blob, task_id, _idx in cid_writes.get(cid, []):
val = self.serde.loads_typed((type_tag, write_blob))
collected.append((task_id, ch, val))
collected.reverse()
entry: DeltaChannelHistory = {"writes": collected}
if seed_version is not None:
@@ -24,11 +24,7 @@ from __future__ import annotations
from collections.abc import Mapping, Sequence
from typing import Any
from langgraph.checkpoint.base import (
DeltaChannelHistory,
PendingWrite,
_apply_delta_history_overwrite_semantics,
)
from langgraph.checkpoint.base import DeltaChannelHistory, PendingWrite
# Stage 1 streams ancestors of `target_cid` newest-first. The `<=`
# predicate keeps target itself in the stream so we can read its
@@ -165,11 +161,10 @@ def build_delta_channels_writes_history(
collected: list[PendingWrite] = []
# Chain is newest-first; iterate oldest-first for the public order.
for cid in reversed(chain_cids):
step_writes: list[PendingWrite] = [
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
for type_tag, value_blob, task_id, _idx in cid_writes.get(cid, [])
]
collected.extend(_apply_delta_history_overwrite_semantics(step_writes))
for type_tag, value_blob, task_id, _idx in cid_writes.get(cid, []):
collected.append(
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
)
entry: DeltaChannelHistory = {"writes": collected}
if ch in seeded:
entry["seed"] = seed_val_by_ch[ch]
@@ -33,8 +33,6 @@ pytest.importorskip("langgraph.channels.delta", reason="langgraph core not insta
pytest.importorskip("langgraph.graph", reason="langgraph core not installed")
from langgraph.channels.delta import DeltaChannel # type: ignore[import-untyped] # noqa: E402,I001
from langgraph.checkpoint.base import Checkpoint # noqa: E402
from langgraph.checkpoint.base.id import uuid6 # noqa: E402
from langgraph.checkpoint.serde.types import _DeltaSnapshot # noqa: E402
from langgraph.graph import END, START, StateGraph # type: ignore[import-untyped] # noqa: E402
from typing_extensions import TypedDict # noqa: E402
@@ -215,42 +213,6 @@ def test_seed_omitted_when_walk_reaches_root_sync() -> None:
assert entry["writes"] == []
def test_overwrite_bypasses_same_step_writes_sync() -> None:
with SqliteSaver.from_conn_string(":memory:") as saver:
config: RunnableConfig = {
"configurable": {"thread_id": "overwrite-sync", "checkpoint_ns": ""}
}
cp1 = Checkpoint(
v=1,
id=str(uuid6(clock_seq=-1)),
ts="",
channel_values={},
channel_versions={},
versions_seen={},
updated_channels=None,
)
cfg1 = saver.put(config, cp1, {"source": "loop", "step": 0}, {})
cp2 = Checkpoint(
v=1,
id=str(uuid6(clock_seq=-1)),
ts="",
channel_values={},
channel_versions={},
versions_seen={},
updated_channels=None,
)
cfg2 = saver.put(cfg1, cp2, {"source": "loop", "step": 1}, {})
saver.put_writes(
cfg1,
[("items", [1]), ("items", {"__overwrite__": [50]}), ("items", [2])],
"task",
)
result = saver.get_delta_channel_history(config=cfg2, channels=["items"])
values = [w[2] for w in result["items"]["writes"]]
assert values == [{"__overwrite__": [50]}]
# ---------------------------------------------------------------------------
# Async: AsyncSqliteSaver
# ---------------------------------------------------------------------------
@@ -34,24 +34,6 @@ PendingWrite = tuple[str, str, Any]
logger = logging.getLogger(__name__)
_OVERWRITE_KEY = "__overwrite__"
def _is_overwrite_value(value: Any) -> bool:
if isinstance(value, dict) and len(value) == 1 and _OVERWRITE_KEY in value:
return True
return value.__class__.__name__ == "Overwrite" and hasattr(value, "value")
def _apply_delta_history_overwrite_semantics(
writes: Sequence[PendingWrite],
) -> list[PendingWrite]:
for write in writes:
if _is_overwrite_value(write[2]):
return [write]
return list(writes)
# Marked as total=False to allow for future expansion.
class CheckpointMetadata(TypedDict, total=False):
"""Metadata associated with a checkpoint."""
@@ -649,17 +631,10 @@ class BaseCheckpointSaver(Generic[V]):
if tup is None:
break
if tup.pending_writes:
step_writes_by_ch: dict[str, list[PendingWrite]] = {
ch: [] for ch in remaining
}
for write in tup.pending_writes:
for write in reversed(tup.pending_writes):
ch = write[1]
if ch in remaining:
step_writes_by_ch[ch].append(write)
for ch, step_writes in step_writes_by_ch.items():
collected_by_ch[ch].extend(
reversed(_apply_delta_history_overwrite_semantics(step_writes))
)
collected_by_ch[ch].append(write)
for ch in list(remaining):
if ch in tup.checkpoint["channel_values"]:
seed_by_ch[ch] = tup.checkpoint["channel_values"][ch]
@@ -697,17 +672,10 @@ class BaseCheckpointSaver(Generic[V]):
if tup is None:
break
if tup.pending_writes:
step_writes_by_ch: dict[str, list[PendingWrite]] = {
ch: [] for ch in remaining
}
for write in tup.pending_writes:
for write in reversed(tup.pending_writes):
ch = write[1]
if ch in remaining:
step_writes_by_ch[ch].append(write)
for ch, step_writes in step_writes_by_ch.items():
collected_by_ch[ch].extend(
reversed(_apply_delta_history_overwrite_semantics(step_writes))
)
collected_by_ch[ch].append(write)
for ch in list(remaining):
if ch in tup.checkpoint["channel_values"]:
seed_by_ch[ch] = tup.checkpoint["channel_values"][ch]
@@ -23,7 +23,6 @@ from langgraph.checkpoint.base import (
DeltaChannelHistory,
PendingWrite,
SerializerProtocol,
_apply_delta_history_overwrite_semantics,
get_checkpoint_id,
get_checkpoint_metadata,
)
@@ -201,11 +200,8 @@ class InMemorySaver(
terminated_here.add(ch)
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
step_writes_by_ch: dict[str, list[PendingWrite]] = {
ch: [] for ch in remaining
}
for (_task_id, _idx), (tid, ch, serialized, _) in sorted(
step_writes.items()
step_writes.items(), reverse=True
):
if ch not in remaining:
continue
@@ -214,13 +210,9 @@ class InMemorySaver(
blob_value, _DeltaSnapshot
):
continue
step_writes_by_ch[ch].append(
collected_by_ch[ch].append(
(tid, ch, self.serde.loads_typed(serialized))
)
for ch, writes in step_writes_by_ch.items():
collected_by_ch[ch].extend(
reversed(_apply_delta_history_overwrite_semantics(writes))
)
for ch in terminated_here:
seed_by_ch[ch] = blob_value_by_ch[ch]
-46
View File
@@ -387,52 +387,6 @@ class TestInMemorySaverDeltaChannel:
values = [v for _, _, v in result["writes"]]
assert values == [{"content": "hi"}]
def test_get_channel_writes_overwrite_bypasses_same_step_writes(self) -> None:
saver = InMemorySaver()
serde = JsonPlusSerializer()
thread_id, ns, channel = "t1", "", "messages"
cp1 = empty_checkpoint()
cp1["id"] = "cp1"
cp2 = empty_checkpoint()
cp2["id"] = "cp2"
saver.storage[thread_id][ns] = {
"cp1": (serde.dumps_typed(cp1), serde.dumps_typed({}), None),
"cp2": (serde.dumps_typed(cp2), serde.dumps_typed({}), "cp1"),
}
saver.writes[(thread_id, ns, "cp1")][("task1", 0)] = (
"task1",
channel,
serde.dumps_typed([1]),
"",
)
saver.writes[(thread_id, ns, "cp1")][("task2", 0)] = (
"task2",
channel,
serde.dumps_typed({"__overwrite__": [50]}),
"",
)
saver.writes[(thread_id, ns, "cp1")][("task3", 0)] = (
"task3",
channel,
serde.dumps_typed([2]),
"",
)
config: RunnableConfig = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": ns,
"checkpoint_id": "cp2",
}
}
result = saver.get_delta_channel_history(config=config, channels=[channel])[
channel
]
values = [v for _, _, v in result["writes"]]
assert values == [{"__overwrite__": [50]}]
def test_get_channel_writes_at_root_returns_empty(self) -> None:
"""Reconstructing the root checkpoint's state: no ancestors → []."""
saver = InMemorySaver()