mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-06 17:57:49 +02:00
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:
@@ -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.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user