mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-23 18:15:08 +02:00
fix(checkpoint-sqlite): walk delta ancestors by parent pointer
Stage 1 of the sqlite delta history filtered `checkpoint_id <= target` and streamed `ORDER BY checkpoint_id DESC`. Both encode an extra assumption: that every child's checkpoint id sorts above its parent's. Ancestry is defined by `parent_checkpoint_id`, and nothing in the contract requires ids to be monotonic. A parent whose id sorted above its child's was dropped from the stream, so its stored value and its writes were lost with no error raised. Removing the range filter alone would not help: in DESC order that parent arrives before the target, so the walk passes it before it has started. Replace it with a recursive CTE anchored at the target that follows `parent_checkpoint_id`. Rows now arrive in walk order, so the off-path skip and parent tracking in `step_walk_with_row` are dead and removed, and the query reads only true ancestors instead of every row at or below the target. Following pointers can loop where a bounded id scan could not, and a loop is reachable through `put` alone: it writes with `INSERT OR REPLACE`, so re-putting an existing checkpoint id under a descendant's config repoints that checkpoint at its own descendant. The walk therefore stops on a repeated checkpoint id. sqlite yields recursive rows lazily, so abandoning the cursor ends the recursion. Postgres needs no equivalent change: it pages the whole thread without an id bound and follows parent pointers in Python, and its upsert never rewrites `parent_checkpoint_id`, so it cannot form this loop. Fixes #8550 Co-authored-by: lylelllll <59271327+lylelllll@users.noreply.github.com>
This commit is contained in:
co-authored by
lylelllll
parent
bdb85b5aa8
commit
3c3a3dde5d
@@ -507,13 +507,12 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
|
||||
Two-stage query:
|
||||
|
||||
* Stage 1 (paged): newest-first slice of `checkpoints` returning
|
||||
`(checkpoint_id, parent_checkpoint_id, type, checkpoint)` per
|
||||
ancestor. Sqlite has no JSONB, so we ship the full serialized
|
||||
checkpoint blob and inspect `channel_values` in Python. Pages
|
||||
newest-first by `checkpoint_id` with a `< cursor` predicate;
|
||||
page size is `DELTA_PAGE_SIZE`. Stops paging when every channel
|
||||
has found its seed or the chain is exhausted.
|
||||
* Stage 1 (streamed): recursive CTE over `checkpoints` following
|
||||
`parent_checkpoint_id` from the target, returning
|
||||
`(checkpoint_id, type, checkpoint)` per ancestor. Sqlite has no
|
||||
JSONB, so we ship the full serialized checkpoint blob and inspect
|
||||
`channel_values` in Python. Stops reading when every channel has
|
||||
found its seed or the chain is exhausted.
|
||||
|
||||
* Stage 2 (per-channel UNION ALL): one branch per channel reading
|
||||
`writes` filtered to that channel's specific `chain_cids`. No
|
||||
@@ -538,12 +537,14 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
seeded: set[str] = set()
|
||||
|
||||
with self.cursor(transaction=False) as cur:
|
||||
cur.execute(DELTA_STAGE1_SQL, (thread_id, checkpoint_ns, checkpoint_id))
|
||||
cur.execute(
|
||||
DELTA_STAGE1_SQL,
|
||||
(thread_id, checkpoint_ns, checkpoint_id, thread_id, checkpoint_ns),
|
||||
)
|
||||
for row in cur:
|
||||
cid, parent_cid, type_tag, blob = row
|
||||
cid, type_tag, blob = row
|
||||
if step_walk_with_row(
|
||||
cid=cid,
|
||||
parent_cid=parent_cid,
|
||||
type_tag=type_tag,
|
||||
blob=blob,
|
||||
target_id=checkpoint_id,
|
||||
|
||||
@@ -26,16 +26,29 @@ from typing import Any
|
||||
|
||||
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
|
||||
# `parent_checkpoint_id` from the first row without a separate lookup;
|
||||
# the caller skips target's own writes/seed (matches the
|
||||
# `BaseCheckpointSaver` contract).
|
||||
# Stage 1 streams target, then its ancestors nearest-first, by following
|
||||
# `parent_checkpoint_id`. Ids carry no ordering guarantee, so a range scan by
|
||||
# id can miss a parent whose id sorts above its child's. Target is the anchor
|
||||
# row; its own writes/seed are skipped (matches the `BaseCheckpointSaver`
|
||||
# contract).
|
||||
#
|
||||
# `put` is `INSERT OR REPLACE`, so re-putting an existing id under a
|
||||
# descendant's config makes the chain a loop. `step_walk_with_row` stops on a
|
||||
# repeated id; sqlite yields recursive rows lazily, so abandoning the cursor
|
||||
# ends the recursion.
|
||||
DELTA_STAGE1_SQL = (
|
||||
"WITH RECURSIVE ancestors(checkpoint_id, parent_checkpoint_id, type, "
|
||||
"checkpoint) AS ("
|
||||
"SELECT checkpoint_id, parent_checkpoint_id, type, checkpoint "
|
||||
"FROM checkpoints "
|
||||
"WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id <= ? "
|
||||
"ORDER BY checkpoint_id DESC"
|
||||
"WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? "
|
||||
"UNION ALL "
|
||||
"SELECT c.checkpoint_id, c.parent_checkpoint_id, c.type, c.checkpoint "
|
||||
"FROM checkpoints c JOIN ancestors a "
|
||||
"ON c.checkpoint_id = a.parent_checkpoint_id "
|
||||
"WHERE c.thread_id = ? AND c.checkpoint_ns = ?"
|
||||
") "
|
||||
"SELECT checkpoint_id, type, checkpoint FROM ancestors"
|
||||
)
|
||||
|
||||
|
||||
@@ -68,7 +81,6 @@ def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
|
||||
def step_walk_with_row(
|
||||
*,
|
||||
cid: str,
|
||||
parent_cid: str | None,
|
||||
type_tag: str,
|
||||
blob: bytes,
|
||||
target_id: str,
|
||||
@@ -81,36 +93,32 @@ def step_walk_with_row(
|
||||
) -> bool:
|
||||
"""Process one streamed stage-1 row in the merged ancestor walk.
|
||||
|
||||
The cursor returns (cid, parent_cid, type, blob) rows in
|
||||
`checkpoint_id` DESC order starting at target. The first row is
|
||||
target itself; we read its parent_cid to seed the walk and otherwise
|
||||
skip it (target's own writes/seed are not part of the contract).
|
||||
The cursor returns (cid, type, blob) rows in walk order starting at
|
||||
target. The first row is target itself and is skipped (target's own
|
||||
writes/seed are not part of the contract).
|
||||
|
||||
For each subsequent row, if `cid` matches the walk's current
|
||||
position, we deserialize the blob, append the cid to every
|
||||
not-yet-seeded channel's chain, and check `channel_values` for
|
||||
For each subsequent row we deserialize the blob, append the cid to
|
||||
every not-yet-seeded channel's chain, and check `channel_values` for
|
||||
seeds. The deserialized checkpoint is dropped before advancing — no
|
||||
cross-row cache, so peak in-flight is one deserialized checkpoint.
|
||||
|
||||
Off-path rows (different branch on the same thread) advance the
|
||||
cursor without doing any work.
|
||||
|
||||
Returns True when every requested channel is seeded — the caller
|
||||
can stop iterating and close the cursor.
|
||||
Returns True when the caller can stop iterating and close the cursor:
|
||||
every requested channel is seeded, or the chain revisited a checkpoint.
|
||||
"""
|
||||
if "started" not in walk_state:
|
||||
if cid == target_id:
|
||||
walk_state["started"] = True
|
||||
walk_state["cur_cid"] = parent_cid
|
||||
walk_state["active"] = {ch for ch in channels if ch not in seeded}
|
||||
walk_state["walked"] = {cid}
|
||||
# Not target yet (or target not present): keep streaming.
|
||||
return False
|
||||
active: set[str] = walk_state["active"]
|
||||
if not active:
|
||||
return True
|
||||
if cid != walk_state["cur_cid"]:
|
||||
# Off-path row from a sibling branch — skip without deserializing.
|
||||
return False
|
||||
walked: set[str] = walk_state["walked"]
|
||||
if cid in walked:
|
||||
return True
|
||||
walked.add(cid)
|
||||
for ch in active:
|
||||
chain_by_ch[ch].append(cid)
|
||||
ckpt = serde.loads_typed((type_tag, blob))
|
||||
@@ -120,7 +128,6 @@ def step_walk_with_row(
|
||||
seeded.add(ch)
|
||||
active.discard(ch)
|
||||
del ckpt, channel_values
|
||||
walk_state["cur_cid"] = parent_cid
|
||||
return not active
|
||||
|
||||
|
||||
|
||||
@@ -625,8 +625,8 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
"""Fast-path override of `BaseCheckpointSaver.aget_delta_channel_history`.
|
||||
|
||||
See `SqliteSaver.get_delta_channel_history` for design notes; this
|
||||
is the async equivalent using `aiosqlite` cursors. Stage 1 pages
|
||||
the parent chain newest-first and Python-deserializes each
|
||||
is the async equivalent using `aiosqlite` cursors. Stage 1 streams
|
||||
the parent chain from the target and Python-deserializes each
|
||||
checkpoint blob to find per-channel snapshots; stage 2 fetches
|
||||
only the relevant writes via per-channel UNION ALL.
|
||||
"""
|
||||
@@ -650,13 +650,13 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
|
||||
async with self.lock, self.conn.cursor() as cur:
|
||||
await cur.execute(
|
||||
DELTA_STAGE1_SQL, (thread_id, checkpoint_ns, checkpoint_id)
|
||||
DELTA_STAGE1_SQL,
|
||||
(thread_id, checkpoint_ns, checkpoint_id, thread_id, checkpoint_ns),
|
||||
)
|
||||
async for row in cur:
|
||||
cid, parent_cid, type_tag, blob = row
|
||||
cid, type_tag, blob = row
|
||||
if step_walk_with_row(
|
||||
cid=cid,
|
||||
parent_cid=parent_cid,
|
||||
type_tag=type_tag,
|
||||
blob=blob,
|
||||
target_id=checkpoint_id,
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from langgraph.checkpoint.base import (
|
||||
BaseCheckpointSaver,
|
||||
Checkpoint,
|
||||
DeltaChannelHistory,
|
||||
empty_checkpoint,
|
||||
)
|
||||
|
||||
from langgraph.checkpoint.sqlite import SqliteSaver
|
||||
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
||||
|
||||
CHANNEL = "ch"
|
||||
CONFIG: dict[str, Any] = {"configurable": {"thread_id": "t", "checkpoint_ns": ""}}
|
||||
EXPECTED: DeltaChannelHistory = {
|
||||
"writes": [("task", CHANNEL, "write-root")],
|
||||
"seed": "seed",
|
||||
}
|
||||
|
||||
|
||||
def _checkpoint(checkpoint_id: str, values: dict[str, Any]) -> Checkpoint:
|
||||
value = empty_checkpoint()
|
||||
value["id"] = checkpoint_id
|
||||
value["channel_values"] = values
|
||||
return value
|
||||
|
||||
|
||||
PARENT_ID_ORDERS = [
|
||||
pytest.param("z-older", "a-newer", id="parent_id_sorts_above_child"),
|
||||
pytest.param("a-older", "z-newer", id="parent_id_sorts_below_child"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("root_id", "child_id"), PARENT_ID_ORDERS)
|
||||
def test_sync_walk_reaches_parent_whatever_the_id_order(
|
||||
root_id: str, child_id: str
|
||||
) -> None:
|
||||
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||
root = saver.put(CONFIG, _checkpoint(root_id, {CHANNEL: "seed"}), {}, {})
|
||||
saver.put_writes(root, [(CHANNEL, "write-root")], "task")
|
||||
child = saver.put(root, _checkpoint(child_id, {}), {}, {})
|
||||
|
||||
got = saver.get_delta_channel_history(config=child, channels=[CHANNEL])
|
||||
reference = BaseCheckpointSaver.get_delta_channel_history(
|
||||
saver, config=child, channels=[CHANNEL]
|
||||
)
|
||||
assert got[CHANNEL] == EXPECTED
|
||||
assert got[CHANNEL] == reference[CHANNEL], "fast path disagrees with base"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("root_id", "child_id"), PARENT_ID_ORDERS)
|
||||
async def test_async_walk_reaches_parent_whatever_the_id_order(
|
||||
root_id: str, child_id: str
|
||||
) -> None:
|
||||
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
|
||||
root = await saver.aput(CONFIG, _checkpoint(root_id, {CHANNEL: "seed"}), {}, {})
|
||||
await saver.aput_writes(root, [(CHANNEL, "write-root")], "task")
|
||||
child = await saver.aput(root, _checkpoint(child_id, {}), {}, {})
|
||||
|
||||
got = await saver.aget_delta_channel_history(config=child, channels=[CHANNEL])
|
||||
assert got[CHANNEL] == EXPECTED
|
||||
|
||||
|
||||
def test_walk_reaches_root_of_long_chain_with_descending_ids() -> None:
|
||||
steps = 40
|
||||
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||
parent = saver.put(
|
||||
CONFIG, _checkpoint(f"id-{steps:03d}", {CHANNEL: "seed"}), {}, {}
|
||||
)
|
||||
saver.put_writes(parent, [(CHANNEL, "write-root")], "task")
|
||||
for step in range(steps - 1, 0, -1):
|
||||
parent = saver.put(parent, _checkpoint(f"id-{step:03d}", {}), {}, {})
|
||||
|
||||
got = saver.get_delta_channel_history(config=parent, channels=[CHANNEL])
|
||||
assert got[CHANNEL] == EXPECTED
|
||||
|
||||
|
||||
def test_walk_terminates_when_put_makes_the_parent_chain_cycle() -> None:
|
||||
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||
a = saver.put(CONFIG, _checkpoint("cid-a", {}), {}, {})
|
||||
b = saver.put(a, _checkpoint("cid-b", {}), {}, {})
|
||||
repoint_a_under_b = _checkpoint("cid-a", {})
|
||||
saver.put(b, repoint_a_under_b, {}, {})
|
||||
|
||||
got = saver.get_delta_channel_history(config=b, channels=[CHANNEL])
|
||||
assert got[CHANNEL] == {"writes": []}
|
||||
Reference in New Issue
Block a user