feat(checkpoint-sqlite): override get_delta_channel_history with streaming walk (#7702)

## Summary

Adds a sqlite-specific override of
`BaseCheckpointSaver.get_delta_channel_history` (and async). Before this
PR, `SqliteSaver` / `AsyncSqliteSaver` inherited the default impl, which
calls `get_tuple` once per ancestor — N round-trips, full pending-writes
fetch per step regardless of channel relevance.

The override mirrors the postgres two-stage shape (ancestor walk +
per-channel UNION ALL writes fetch) but adapted for sqlite:

- **No JSONB** → stage 1 streams the cursor row-by-row in
`checkpoint_id` DESC order. The merged walk advances one row at a time,
deserializing only on-path checkpoints and dropping each before
advancing — peak in-flight is one deserialized checkpoint, no
`fetchall()` materialization.
- **No separate blob table** → `channel_values` lives inline in the
checkpoint blob, so seeds come back from stage 1 with no second fetch.
- **Single merged walk (not K independent walks)**: each visited cid is
deserialized exactly once, regardless of how many channels are still
seeking their seed.
- **Stage 2** stays per-channel UNION ALL to avoid over-fetching writes
when channels have different chain depths — same rationale as postgres.

`AsyncSqliteSaver.get_delta_channel_history` bridges to its async form
via `run_coroutine_threadsafe`, matching the same cross-thread guard
used by `get_tuple` / `delete_thread`.

## Tests

- New `tests/test_delta_channel_migration.py`: covers the
`BinaryOperatorAggregate -> DeltaChannel` migration path on sqlite (sync
round-trip, sync continuation with post-migration delta folding, async
round-trip). Mirrors
`libs/langgraph/tests/test_delta_channel_migration.py` (which covered
`InMemorySaver`); without these, the override's behavior on
pre-migration threads was unverified — the override has to identify a
plain accumulated `channel_values[ch]` at a pre-migration ancestor as a
valid `seed`, not just `_DeltaSnapshot` sentinels.
- Existing `tests/test_get_delta_channel_history.py` (7 tests) continues
to pass and now exercises the optimized override end-to-end (previously
hit the inherited default impl).
- `make format`, `make lint`, `make test`: clean. 97/97 in the non-flaky
sqlite suite (the one ignored test, `test_async_asearch_refresh_ttl`, is
a known TTL-store timing flake on a separate module unrelated to this
PR).

## Benchmarks

### `get_delta_channel_history` micro-bench (override vs inherited
default impl)

1000-turn synthetic threads with sentinel snapshots + per-step writes;
`bench_sqlite_delta_history.py`. Per-call latency in microseconds.

| Scenario | min | median | mean |
|---|---:|---:|---:|
| S1 single channel, root-only snapshot | **4.60x** | **4.90x** |
**5.13x** |
| S2 mixed cadence (every-50 + root-only), 2 channels | **6.08x** |
**6.37x** | **6.84x** |
| S3 K=8 channels, root-only snapshot | 1.23x | 1.27x | 0.90x |

S2 wins biggest because per-channel UNION ALL avoids over-fetching
writes for the shallow channel. S3 is the worst case for sqlite (8
channels all walking to root, 1000 deserializations either way) — the
override still wins on min/median.

### Long-running thread mem/storage bench (delta vs no-delta)

`bench_sqlite_delta_memory.py`. `delta` mode uses `DeltaChannel` + the
override; `no_delta` uses `Annotated[list, _messages_delta_reducer]`
(full state in every blob). Same workload, file-backed sqlite. Latency
measured untraced (30 iterations); peak heap measured separately under
tracemalloc.

| Scenario | Turns | Storage Δ | Peak heap Δ | Read latency Δ |
|---|---:|---|---|---|
| K=1, freq=50 | 200 | **-96%** (942 KB vs 25.1 MB) | +21% (504 KB vs
418 KB) | **+13%** |
| K=1, freq=50 | 500 | **-98%** (2.9 MB vs 152.3 MB) | +20% (1.2 MB vs
1.0 MB) | **-6%** (delta wins) |
| K=3, freq=50 uniform | 200 | **-98%** (1.7 MB vs 73.5 MB) | +7% (1.3
MB vs 1.2 MB) | **+10%** |
| K=3, freq=50 uniform | 500 | **-99%** (6.0 MB vs 452.5 MB) | +7% (3.3
MB vs 3.0 MB) | **+6%** |
| K=3, freq=mixed | 200 | **-98%** (1.4 MB vs 73.5 MB) | +5% (1.3 MB vs
1.2 MB) | +190% (5.1 ms vs 1.7 ms abs) |
| K=3, freq=mixed | 500 | **-99%** (4.1 MB vs 452.5 MB) | +8% (3.3 MB vs
3.0 MB) | +377% (20.9 ms vs 4.4 ms abs) |

- **Storage**: -96 to -99% on long threads (a 500-turn K=3 thread
shrinks from 452 MB to 6 MB on disk). This is the headline win.
- **Peak heap**: within +5 to +21% of the no-delta path — the streaming
cursor + merged walk + drop-after-deserialize keep peak in-flight at one
checkpoint at a time.
- **Read latency**: equivalent-ish (within ~15%) on uniform-cadence
scenarios; at K=1/500 turns delta even wins by 6%. The mixed-cadence
rows have one channel with `snapshot_frequency=1000` walking to root on
a 500-turn thread — by configuration. Absolute mixed-delta latency is
still 5-21 ms per read.

Bench scripts (not committed; workspace-root convention matches other
`bench_*.py` files):
- `bench_sqlite_delta_history.py`
- `bench_sqlite_delta_memory.py`

## Test plan

- [x] `cd libs/checkpoint-sqlite && make format` clean
- [x] `cd libs/checkpoint-sqlite && make lint` clean
- [x] `cd libs/checkpoint-sqlite && make test` — 97 passed (1 known
flake unrelated)
- [x] `tests/test_get_delta_channel_history.py` — 7/7 (now exercises the
override)
- [x] `tests/test_delta_channel_migration.py` — 3/3 (new)
This commit is contained in:
Sydney Runkle
2026-05-05 15:30:57 -04:00
committed by GitHub
parent 0ae01c7366
commit 86baa5d08e
8 changed files with 545 additions and 7 deletions
@@ -4,7 +4,7 @@ import json
import random
import sqlite3
import threading
from collections.abc import AsyncIterator, Iterator, Sequence
from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
from contextlib import closing, contextmanager
from typing import Any, cast
@@ -16,12 +16,19 @@ from langgraph.checkpoint.base import (
Checkpoint,
CheckpointMetadata,
CheckpointTuple,
DeltaChannelHistory,
SerializerProtocol,
get_checkpoint_id,
get_checkpoint_metadata,
)
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
from langgraph.checkpoint.sqlite._delta import (
DELTA_STAGE1_SQL,
build_delta_channels_writes_history,
build_delta_stage2_sql,
step_walk_with_row,
)
from langgraph.checkpoint.sqlite.utils import search_where
_AIO_ERROR_MSG = (
@@ -493,6 +500,88 @@ class SqliteSaver(BaseCheckpointSaver[str]):
(str(thread_id),),
)
def get_delta_channel_history(
self, *, config: RunnableConfig, channels: Sequence[str]
) -> Mapping[str, DeltaChannelHistory]:
"""Fast-path override of `BaseCheckpointSaver.get_delta_channel_history`.
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 2 (per-channel UNION ALL): one branch per channel reading
`writes` filtered to that channel's specific `chain_cids`. No
separate seed-blob fetch — sqlite stores `channel_values` inline
in the checkpoint blob, so seeds come back from stage 1.
"""
if not channels:
return {}
channels = list(channels)
thread_id = str(config["configurable"]["thread_id"])
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = get_checkpoint_id(config)
if checkpoint_id is None:
target = self.get_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
checkpoint_id = target.config["configurable"]["checkpoint_id"]
chain_by_ch: dict[str, list[str]] = {ch: [] for ch in channels}
seed_val_by_ch: dict[str, Any] = {}
walk_state: dict[str, Any] = {}
seeded: set[str] = set()
with self.cursor(transaction=False) as cur:
cur.execute(DELTA_STAGE1_SQL, (thread_id, checkpoint_ns, checkpoint_id))
for row in cur:
cid, parent_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,
serde=self.serde,
chain_by_ch=chain_by_ch,
seed_val_by_ch=seed_val_by_ch,
walk_state=walk_state,
seeded=seeded,
channels=channels,
):
break
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
stage2_sql = build_delta_stage2_sql(
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
)
if stage2_sql:
stage2_params: list[Any] = []
for ch in channels_with_chain:
stage2_params.extend(
[thread_id, checkpoint_ns, ch, *chain_by_ch[ch]]
)
cur.execute(stage2_sql, stage2_params)
stage2_rows = cast(
"list[tuple[str, str, str, int, str, bytes]]", cur.fetchall()
)
else:
stage2_rows = []
return build_delta_channels_writes_history(
channels=channels,
chain_by_ch=chain_by_ch,
seed_val_by_ch=seed_val_by_ch,
seeded=seeded,
stage2_rows=stage2_rows,
serde=self.serde,
)
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
"""Get a checkpoint tuple from the database asynchronously.