diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py index 3259ff150..0fe04de67 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py @@ -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, diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/_delta.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/_delta.py index 1fe617ff7..a28313bdb 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/_delta.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/_delta.py @@ -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 diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py index 1ad0777c1..c1253e1cb 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py @@ -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, diff --git a/libs/checkpoint-sqlite/tests/test_delta_parent_walk.py b/libs/checkpoint-sqlite/tests/test_delta_parent_walk.py new file mode 100644 index 000000000..bdd513366 --- /dev/null +++ b/libs/checkpoint-sqlite/tests/test_delta_parent_walk.py @@ -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": []}