mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-26 03:25:06 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
95225b4171 | ||
|
|
b60ef4adb0 |
@@ -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}"
|
||||||
|
|||||||
Reference in New Issue
Block a user