Compare commits

...
Author SHA1 Message Date
Christian Bromann 95225b4171 cr 2026-06-02 16:43:46 -07:00
Christian Bromann b60ef4adb0 fix(checkpoint): replay migrated delta writes through a plain seed
InMemorySaver.get_delta_channel_history skipped on-path writes when the
terminating ancestor's blob was a plain (pre-delta migration) value
instead of a _DeltaSnapshot, dropping post-migration writes on reload.
Collect writes regardless of seed type, matching the base implementation.
2026-06-01 12:54:11 -07:00
2 changed files with 64 additions and 8 deletions
@@ -148,15 +148,15 @@ class InMemorySaver(
whose stored blob is non-empty. Other channels keep walking until whose stored blob is non-empty. Other channels keep walking until
they find their own terminator or hit the root. they find their own terminator or hit the root.
Pre-delta plain-value blobs subsume their ancestor's pending The seed value (whether a `_DeltaSnapshot` or a plain pre-delta
writes (the value already includes them); `_DeltaSnapshot` blobs migration blob) is the value AT that ancestor. Pending writes on a
do not (snapshot is the value AT that ancestor, prior to its own plain seed checkpoint are replayed when they are the first deltas
pending writes that produce the child). after migration (no delta-era ancestor already contributed writes).
Stale pre-delta writes on the same checkpoint as the blob are skipped
once a delta-era ancestor in the walk has already contributed writes.
""" """
if not channels: if not channels:
return {} return {}
# Imported lazily to avoid a hard checkpoint→serde-types coupling at
# module import; only this override needs the runtime check.
from langgraph.checkpoint.serde.types import _DeltaSnapshot from langgraph.checkpoint.serde.types import _DeltaSnapshot
thread_id = config["configurable"]["thread_id"] thread_id = config["configurable"]["thread_id"]
@@ -178,6 +178,8 @@ class InMemorySaver(
collected_by_ch: dict[str, list[PendingWrite]] = {c: [] for c in channels} collected_by_ch: dict[str, list[PendingWrite]] = {c: [] for c in channels}
seed_by_ch: dict[str, Any] = {} seed_by_ch: dict[str, Any] = {}
remaining: set[str] = set(channels) remaining: set[str] = set(channels)
# True once a delta-era (empty-blob) ancestor has contributed a write.
seen_delta_write: dict[str, bool] = dict.fromkeys(channels, False)
for cp_id in chain: for cp_id in chain:
if not remaining: if not remaining:
@@ -187,14 +189,17 @@ class InMemorySaver(
terminated_here: set[str] = set() terminated_here: set[str] = set()
blob_value_by_ch: dict[str, Any] = {} blob_value_by_ch: dict[str, Any] = {}
delta_era_by_ch: dict[str, bool] = {}
if ckpt is not None: if ckpt is not None:
versions = ckpt.get("channel_versions", {}) versions = ckpt.get("channel_versions", {})
for ch in remaining: for ch in remaining:
ver = versions.get(ch) ver = versions.get(ch)
if ver is None: if ver is None:
delta_era_by_ch[ch] = True
continue continue
blob_entry = self.blobs.get((thread_id, checkpoint_ns, ch, ver)) blob_entry = self.blobs.get((thread_id, checkpoint_ns, ch, ver))
if blob_entry is None or blob_entry[0] == "empty": if blob_entry is None or blob_entry[0] == "empty":
delta_era_by_ch[ch] = True
continue continue
blob_value_by_ch[ch] = self.serde.loads_typed(blob_entry) blob_value_by_ch[ch] = self.serde.loads_typed(blob_entry)
terminated_here.add(ch) terminated_here.add(ch)
@@ -206,13 +211,21 @@ class InMemorySaver(
if ch not in remaining: if ch not in remaining:
continue continue
blob_value = blob_value_by_ch.get(ch) blob_value = blob_value_by_ch.get(ch)
if blob_value is not None and not isinstance( if (
blob_value, _DeltaSnapshot blob_value is not None
and not isinstance(blob_value, _DeltaSnapshot)
and seen_delta_write.get(ch, False)
): ):
# Pre-delta blob already settled; stale writes on the same
# checkpoint are subsumed. Post-migration writes on the
# migration boundary are collected when no delta-era
# ancestor has contributed writes yet.
continue continue
collected_by_ch[ch].append( collected_by_ch[ch].append(
(tid, ch, self.serde.loads_typed(serialized)) (tid, ch, self.serde.loads_typed(serialized))
) )
if delta_era_by_ch.get(ch, False):
seen_delta_write[ch] = True
for ch in terminated_here: for ch in terminated_here:
seed_by_ch[ch] = blob_value_by_ch[ch] seed_by_ch[ch] = blob_value_by_ch[ch]
@@ -616,3 +616,46 @@ async def test_add_messages_to_delta_migration_preserves_message_history_async()
assert [m.id for m in snap.values["messages"]] == ["h1", "a1"], ( assert [m.id for m in snap.values["messages"]] == ["h1", "a1"], (
f"async tip hydration mismatch: got {[m.id for m in snap.values['messages']]}" f"async tip hydration mismatch: got {[m.id for m in snap.values['messages']]}"
) )
def test_post_migration_write_survives_reload_through_plain_seed() -> None:
"""A write made on the first post-migration super-step must be preserved
when the checkpoint is reloaded (reconstructed via the plain seed)."""
checkpointer = InMemorySaver()
config = {"configurable": {"thread_id": "post-migration-reload"}}
binop = _binop_graph(checkpointer)
_drive(binop, config, "u", 3)
delta = _delta_graph(checkpointer)
live_result = delta.invoke({"items": ["POST"]}, config)
reloaded = delta.get_state(config)
assert list(live_result.get("items", [])) == ["u0", "u1", "u2", "POST"], (
f"sanity: live invoke should include POST, got {live_result.get('items')}"
)
assert list(reloaded.values.get("items", [])) == ["u0", "u1", "u2", "POST"], (
"post-migration write dropped on reload through plain seed: "
f"got {reloaded.values.get('items')}"
)
def test_post_migration_reload_base_matches_optimized_override() -> None:
"""The reference `BaseCheckpointSaver` path and the optimized
`InMemorySaver` override must agree once a post-migration write is
reconstructed through a plain seed (the scenario that triggers the
write-collection guard)."""
def _run(saver: Any) -> list:
config = {"configurable": {"thread_id": "parity"}}
_drive(_binop_graph(saver), config, "u", 2)
delta = _delta_graph(saver)
delta.invoke({"items": ["POST"]}, config)
return list(delta.get_state(config).values.get("items", []))
fast = _run(InMemorySaver())
slow = _run(_ThirdPartyStyleSaver())
assert fast == ["u0", "u1", "POST"], f"optimized override wrong: {fast}"
assert slow == fast, f"base fallback diverged from override: {slow} != {fast}"