feat: public get_writes_history saver API + delta cadence rework (#7699)

## Summary

- Promotes the private K-channel batched ancestor-walk to a stable
public `get_delta_channel_history` / `aget_delta_channel_history` API on
`BaseCheckpointSaver` (returns `Mapping[str, DeltaChannelHistory]`, a
TypedDict with `writes` always present and `seed` `NotRequired`)
- Removes `DELTA_SENTINEL` / `_DeltaSentinel` entirely — the saver layer
is now delta-agnostic on both write and read paths
- Reworks `DeltaChannel` snapshot cadence from "every Nth superstep" to
"every N updates to this channel," persisted in
`CheckpointMetadata.delta_updates_since_snapshot`
- Adds Postgres optimizations: paged stage-1 with cursor (1024-row
pages) and per-channel UNION ALL stage-2 (no over-fetch when channels
have different chain depths)
- Default `snapshot_frequency` becomes a positive int (default `1000`);
the previous `None` opt-out is removed

## Public API

```python
class DeltaChannelHistory(TypedDict):
    writes: list[PendingWrite]            # always present, possibly empty
    seed: NotRequired[Any]                # absent if walk reached root

def get_delta_channel_history(
    self, *, config: RunnableConfig, channels: Sequence[str]
) -> Mapping[str, DeltaChannelHistory]: ...

async def aget_delta_channel_history(
    self, *, config: RunnableConfig, channels: Sequence[str]
) -> Mapping[str, DeltaChannelHistory]: ...
```

`config` and `channels` are keyword-only so later additions (e.g.
`page_size`) don't shift the positional API.

The TypedDict-with-`NotRequired[seed]` shape matches the existing
checkpoint-package convention (`CheckpointMetadata` is
`TypedDict(total=False)`) — absence-via-key-omission rather than
introducing a new sentinel. Pregel translates `"seed" not in hist` to
`MISSING` on its side at consume time.

The default impl walks `get_tuple` + `parent_config` correctly but is
slow on long chains; savers that care override (`InMemorySaver`,
`PostgresSaver`).

## Sentinel removal

`DELTA_SENTINEL` and `_DeltaSentinel` are deleted entirely. The saver
layer becomes delta-agnostic:

- `DeltaChannel.checkpoint()` returns `MISSING` for non-snapshot steps;
pregel's `create_checkpoint` skips MISSING so delta channels without a
snapshot simply don't appear in `channel_values`
- `InMemorySaver.put` and Postgres `put` no longer filter sentinels
(they have nothing to filter)
- `_needs_replay` becomes `stored is MISSING`
- `DeltaChannel.from_checkpoint` accepts: `MISSING` → empty,
`_DeltaSnapshot(value)` → snapshot value, plain value → pre-migration
legacy

## Snapshot cadence

`DeltaChannel.snapshot_frequency: int` (default `1000`, positive). The
previous `None` opt-out is gone.

```python
def should_snapshot(ch_name, ch):
    if force_delta_snapshot:                                  # durability="exit"
        return True
    return updates_since_snapshot.get(ch_name, 0) >= ch.snapshot_frequency
```

Per-channel update counters are persisted in
`CheckpointMetadata.delta_updates_since_snapshot` (`NotRequired`,
`total=False`). The counter is incremented by `_put_checkpoint` for any
delta channel in `updated_channels` and reset to `0` by
`create_checkpoint` for channels that fire a snapshot this step.
Version-format-independent — works for `int`, `float`, and `str`
versioning schemes alike.

## Postgres optimization

Two improvements internal to the override:

**Stage-1 paged with cursor** (`LIMIT 1024` internal const, `AND
checkpoint_id < ?` for subsequent pages). The previous unpaged form
scanned every checkpoint in `(thread_id, ns)` and was pathological at
high thread depths.

**Stage-2 per-channel UNION ALL**: one `WHERE channel='X' AND
checkpoint_id = ANY(chain_X)` branch per channel plus one seed-blob
branch per channel with a seed. The previous form filtered by `channel =
ANY(channels) AND checkpoint_id = ANY(union_chain_cids)`, over-fetching
writes when channels had different chain depths (`K ×
max(chain_lengths)` vs the correct `sum(chain_lengths)`).

Both improvements stay internal to `PostgresSaver`/`AsyncPostgresSaver`;
the public contract returns a single `Mapping`.

## Benchmarks

`libs/langgraph/tests/test_delta_channel_benchmark.py`. Run via `python
libs/langgraph/tests/test_delta_channel_benchmark.py`. Postgres against
local pg:5441.

Results below trimmed to the high-signal cells. Sub-millisecond /
sub-100-turn rows omitted as warmup-bound; freq=1 omitted (chain depth =
1, nothing to optimize); peak read-time memory and Postgres storage are
flat between branches and omitted. Deep-thread reads and the
cadence-rework storage win are the load-bearing numbers.

### Postgres reads, 500 turns

| Scenario | main | branch | Δ |
|---|---:|---:|---:|
| Single-channel deep read | 17.7 ms | **6.1 ms** | **-66%** |
| Single-channel, 1000 turns | 35.0 ms | **14.3 ms** | **-59%** |
| K=3 channels, freq=50 uniform | 70.5 ms | **41.4 ms** | **-41%** |
| K=8 channels, freq=50 uniform | 214.2 ms | **139.4 ms** | **-35%** |
| K=8 channels, mixed freq (25/50/100/.../1000) | 295.6 ms | **214.4
ms** | **-27%** |

K-channel batching + paged stage-1 + per-channel UNION ALL stage-2 doing
exactly what they should at depth.

### InMemory reads, 500 turns

| Scenario | main | branch | Δ |
|---|---:|---:|---:|
| Single-channel deep read | 7.9 ms | **3.8 ms** | **-52%** |
| Single-channel, 1000 turns | 15.6 ms | **7.2 ms** | **-54%** |
| K=8 channels, freq=50 uniform | 112.3 ms | 94.6 ms | -16% |
| K=8 channels, mixed freq | 184.9 ms | **134.5 ms** | **-27%** |

### InMemory storage, 500 turns (cadence-rework win)

| Scenario | main | branch | Δ |
|---|---:|---:|---:|
| K=3, freq=50 uniform | 8.7 MB | **3.3 MB** | **-62%** |
| K=3 mixed freq | 3.8 MB | **1.3 MB** | **-66%** |
| K=8, freq=50 uniform | 23.1 MB | **8.7 MB** | **-62%** |
| K=8 mixed freq | 11.5 MB | **4.2 MB** | **-64%** |

Snapshot frequency now counts **channel updates** instead of
**supersteps**. On graphs where supersteps outpace per-channel updates
(e.g., input/end steps that don't write to channels), branch stores ~3×
fewer snapshot blobs.

### Tradeoff worth flagging

InMemory K=3 with mixed frequencies (50/200/1000) at 500 turns: **+64%
read latency** (46.6 → 76.5 ms). The mixed scenario has a channel with
`freq=1000` that goes the entire 500-turn run with no snapshot. On main,
the old superstep-counted cadence happened to fire at step=500 anyway.
New cadence gives users explicit control over walk depth via
`snapshot_frequency`. The K=8 mixed case still wins overall (-27%); this
regression is specific to the K=3 mixed shape.

Default `snapshot_frequency=1000` is the upper bound on walk depth —
it's a tunable knob.

## Tests

- New sqlite smoke test (`test_get_delta_channel_history.py`) exercises
the inherited default `BaseCheckpointSaver` impl via `SqliteSaver` /
`AsyncSqliteSaver` end-to-end with a real `DeltaChannel`-backed graph.
Sqlite uses the default unchanged — this validates the default path
actually works on a real second saver, not just on the optimized
override.
- Module-level `pytest.importorskip("langgraph.channels.delta")` guards
the test for sqlite's standalone CI environment (matches the postgres
pattern).

## Test plan

- [x] `libs/checkpoint`: 150 passed, 16 skipped
- [x] `libs/langgraph` (channels + delta migration): 41/41 (post-merge)
- [x] `libs/langgraph` (full pregel suite): 1784 passing — 6 "failures"
verified via `env -i` clean shell are local LangSmith env vars + `git
describe revision_id` polluting LangChain metadata fixtures; CI is
unaffected
- [x] `libs/checkpoint-postgres`: 40/40 saver tests + 3/3 delta channel
reconstruction tests against local Postgres
- [x] `libs/checkpoint-sqlite`: 105/105 (incl. retry-passed flake
`test_ttl_refresh`, unrelated to this PR)
- [x] Lint clean across all four libs (`ruff format`, `ruff check`,
`mypy`)
- [x] Branch-vs-main benchmarks — see results above

---------

Co-authored-by: Quanzheng Long <long@langchain.dev>
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sydney Runkle
2026-05-04 15:18:43 -04:00
committed by GitHub
co-authored by Quanzheng Long Cursor
parent 35cea28707
commit 0a53c385b2
16 changed files with 1352 additions and 901 deletions
@@ -2,19 +2,18 @@ from __future__ import annotations
import threading
from collections import defaultdict
from collections.abc import Iterator, Sequence
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from typing import Any, cast
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
DELTA_SENTINEL,
WRITES_IDX_MAP,
ChannelVersions,
Checkpoint,
CheckpointMetadata,
CheckpointTuple,
_ChannelWritesHistory,
DeltaChannelHistory,
get_checkpoint_id,
get_serializable_checkpoint_metadata,
)
@@ -27,10 +26,10 @@ from psycopg_pool import ConnectionPool
from langgraph.checkpoint.postgres import _internal
from langgraph.checkpoint.postgres.base import (
SELECT_DELTA_STAGE1_SQL,
SELECT_DELTA_STAGE2_SQL,
_DELTA_PAGE_SIZE,
BasePostgresSaver,
_DeltaStage1Row,
_build_delta_stage1_sql,
_build_delta_stage2_sql,
_DeltaStage2Row,
)
from langgraph.checkpoint.postgres.shallow import ShallowPostgresSaver
@@ -311,9 +310,7 @@ class PostgresSaver(BasePostgresSaver):
# others are stored in blobs table
blob_values = {}
for k, v in checkpoint["channel_values"].items():
if v is DELTA_SENTINEL:
copy["channel_values"].pop(k)
elif isinstance(v, _DeltaSnapshot):
if isinstance(v, _DeltaSnapshot):
blob_values[k] = copy["channel_values"].pop(k)
copy["channel_values"][k] = True
elif v is None or isinstance(v, (str, int, float, bool)):
@@ -444,53 +441,111 @@ class PostgresSaver(BasePostgresSaver):
with conn.cursor(binary=True, row_factory=dict_row) as cur:
yield cur
def _get_channel_writes_history(
self, config: RunnableConfig, channel: str
) -> _ChannelWritesHistory:
"""Fast-path override of `BaseCheckpointSaver._get_channel_writes_history`.
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 scans checkpoint metadata to walk the parent
chain and locate the nearest snapshot; stage 2 fetches only the
chain-limited writes and single seed blob.
Two-stage query, both stages cover ALL requested channels:
* Stage 1 (paged): dynamic SELECT over `checkpoints` with K parallel
JSONB key lookups (one column pair per channel) — no subquery, no
aggregation. Pages newest-first by `checkpoint_id` with a cursor;
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
`checkpoint_writes` filtered to that channel's specific
`chain_cids`, plus one branch per channel that has a seed reading
`checkpoint_blobs` for that channel + version. Avoids the
over-fetch of a single `channel = ANY(channels)` filter when
channels have different chain depths.
"""
if not channels:
return {}
channels = list(channels)
thread_id = 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 _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=[])
return {ch: {"writes": []} for ch in channels}
checkpoint_id = target.config["configurable"]["checkpoint_id"]
# Stage 1: paged K-JSONB-lookup scan, walking the parent chain in
# Python after each page. Stops as soon as every channel has its seed.
stage1_sql = _build_delta_stage1_sql(channels, paged=True)
parent_of: dict[str, str | None] = {}
ver_by_i_by_cid: list[dict[str, str | None]] = [{} for _ in channels]
hs_by_i_by_cid: list[dict[str, bool]] = [{} for _ in channels]
chain_by_ch: dict[str, list[str]] = {ch: [] for ch in channels}
seed_ver_by_ch: dict[str, str | None] = {ch: None for ch in channels}
walk_cursor_by_ch: dict[str, str | None] = {}
seeded: set[str] = set()
cursor: str | None = None
with self._cursor() as cur:
cur.execute(
SELECT_DELTA_STAGE1_SQL,
(channel, channel, thread_id, checkpoint_ns),
)
stage1_rows = cur.fetchall()
chain_cids, seed_version = self._walk_stage1(
cast("list[_DeltaStage1Row]", stage1_rows), checkpoint_id
while True:
stage1_params: list[Any] = []
for ch in channels:
stage1_params.extend([ch, ch])
stage1_params.extend(
[thread_id, checkpoint_ns, cursor, cursor, _DELTA_PAGE_SIZE]
)
cur.execute(stage1_sql, stage1_params)
page = cur.fetchall()
if not page:
break
oldest = self._ingest_stage1_page(
cast("list[Mapping[str, Any]]", page),
channels,
parent_of,
ver_by_i_by_cid,
hs_by_i_by_cid,
)
self._try_advance_walks(
checkpoint_id,
channels,
parent_of,
ver_by_i_by_cid,
hs_by_i_by_cid,
chain_by_ch,
seed_ver_by_ch,
walk_cursor_by_ch,
seeded,
)
# Stop if every channel is seeded, or the page was short
# (chain exhausted — no more rows to fetch).
if len(seeded) == len(channels) or len(page) < _DELTA_PAGE_SIZE:
break
cursor = oldest
# Stage 2: per-channel UNION ALL — one writes branch per channel
# with non-empty chain, plus one blob branch per seeded channel.
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
channels_with_seed = [ch for ch in channels if seed_ver_by_ch[ch] is not None]
stage2_sql = _build_delta_stage2_sql(
channels_with_chain=channels_with_chain,
channels_with_seed=channels_with_seed,
)
seed_versions = [seed_version] if seed_version else []
with self._cursor() as cur:
cur.execute(
SELECT_DELTA_STAGE2_SQL,
(
thread_id,
checkpoint_ns,
channel,
chain_cids,
thread_id,
checkpoint_ns,
channel,
seed_versions,
),
)
stage2_rows = cur.fetchall()
return self._build_delta_channel_writes_history(
channel=channel,
chain_cids=chain_cids,
seed_version=seed_version,
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]])
for ch in channels_with_seed:
stage2_params.extend([thread_id, checkpoint_ns, ch, seed_ver_by_ch[ch]])
with self._cursor() as cur:
cur.execute(stage2_sql, stage2_params)
stage2_rows = cur.fetchall()
else:
stage2_rows = []
return self._build_delta_channels_writes_history(
channels=channels,
chain_by_ch=chain_by_ch,
seed_ver_by_ch=seed_ver_by_ch,
stage2_rows=cast("list[_DeltaStage2Row]", stage2_rows),
)