mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-07 02:07:52 +02:00
simplify delta overwrite fix
This commit is contained in:
-44
@@ -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]
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user