diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py index c2a8ea985..aa20f5339 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py @@ -448,11 +448,12 @@ class PostgresSaver(BasePostgresSaver): 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 1 (paged): dynamic SELECT over `checkpoints` with three + columns per channel: its version, an `EXISTS` probe for a stored + blob at that version, and its inline value. 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 a page comes + back short. * Stage 2 (per-channel UNION ALL): one branch per channel reading `checkpoint_writes` filtered to that channel's specific diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py index d58d2cc38..6cf8aa63a 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/base.py @@ -172,30 +172,8 @@ class _DeltaStage2Row(TypedDict, total=False): version: str | None # "b" rows only -# Multi-channel two-stage DeltaChannel reconstruction. -# -# Stage 1 scans checkpoint metadata (no blob bytes) and emits one row per -# checkpoint with K parallel JSONB key lookups (one column pair per -# requested delta channel: ver_i / hs_i). No subqueries, no aggregation. -# Python walks the parent chain once across all channels. -# -# Stage 2 fetches all writes and the seed blobs for ALL channels in a -# single roundtrip via `channel = ANY(%s)` and chain/seed-version -# filtering. -# -# Empirical comparison vs an alternative "ship full channel_versions / -# channel_values JSONB and let Python pick" form (1000 checkpoints, -# 8 total channels in graph, 3 delta channels requested): -# -# Postgres execution: A=0.24ms vs B=0.38ms (both negligible) -# End-to-end latency: A=6.83ms vs B=2.28ms (B is 3.0x faster) -# Wire payload: A=836KB vs B=330KB (61% smaller) -# Buffer hits: identical (167 blocks) -# -# B (this dynamic-columns design) wins because it avoids JSONB -# serialization on the wire and JSONB-to-dict deserialization in -# psycopg. Even at K=8 (8 delta channels = 16 dynamic columns), B -# still beats A end-to-end (4.2ms vs 6.8ms). +# Delta history is rebuilt in two queries; `_build_delta_stage1_sql` and +# `_build_delta_stage2_sql` document their shapes. def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str: @@ -335,10 +313,8 @@ def _build_delta_stage2_sql( return " UNION ALL ".join(branches) -# Stage 1 rows are dynamic-shape dicts: {checkpoint_id, parent_checkpoint_id, -# ver_0, hs_0, ver_1, hs_1, ...}. Walking is parameterized by the channel -# list to map indices back to channel names — no static TypedDict here. -# `dict[str, Any]` is the practical signature. +# Stage 1 rows are dicts keyed by the per-channel aliases +# `_build_delta_stage1_sql` emits, so there is no static TypedDict. class BasePostgresSaver(BaseCheckpointSaver[str]): @@ -431,9 +407,11 @@ class BasePostgresSaver(BaseCheckpointSaver[str]): (a) it found a stored value for its channel — a blob or an inline primitive (channel becomes seeded), (b) it reached a real root (parent_of[cid] is None — fully - materialized at this point), or + materialized at this point), (c) the next ancestor cid isn't in `parent_of` yet (waiting for - a later page; the cursor stays put). + a later page; the cursor stays put), or + (d) the target's own row isn't in `parent_of` yet (the walk has + not started; no cursor is set, so a later page retries). Mutates `chain_by_ch`, `seed_ver_by_ch`, `seed_inline_by_ch`, `walk_cursor_by_ch`, and `seeded` in place. @@ -441,9 +419,12 @@ class BasePostgresSaver(BaseCheckpointSaver[str]): for i, ch in enumerate(channels): if ch in seeded: continue - # First-time entry: cursor starts at the target's parent. + # Pages start at the thread head, so the target may not have + # loaded yet; a `None` cursor would read as "target is a root". if ch not in walk_cursor_by_ch: - walk_cursor_by_ch[ch] = parent_of.get(target_id) + if target_id not in parent_of: + continue + walk_cursor_by_ch[ch] = parent_of[target_id] cur_cid = walk_cursor_by_ch[ch] ch_chain = chain_by_ch[ch] hb_i = hb_by_i_by_cid[i] diff --git a/libs/checkpoint-postgres/tests/test_delta_pagination.py b/libs/checkpoint-postgres/tests/test_delta_pagination.py new file mode 100644 index 000000000..6f6cfeaff --- /dev/null +++ b/libs/checkpoint-postgres/tests/test_delta_pagination.py @@ -0,0 +1,132 @@ +from __future__ import annotations + +from typing import Any +from uuid import uuid4 + +import pytest +from langgraph.checkpoint.base import ( + Checkpoint, + DeltaChannelHistory, + empty_checkpoint, +) +from langgraph.checkpoint.base.id import uuid6 +from langgraph.checkpoint.serde.types import _DeltaSnapshot + +from langgraph.checkpoint.postgres import PostgresSaver +from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver +from langgraph.checkpoint.postgres.base import _DELTA_PAGE_SIZE +from tests.conftest import DEFAULT_URI + +CHANNEL = "items" +STEPS = 8 +SEED_STEP = 1 +SEED_VALUE = [10, 20] +TARGET_STEP = 4 + +# The real page size is the control; the rest leave the target off the first +# page (three checkpoints are newer than it). +PAGE_SIZES = [_DELTA_PAGE_SIZE, 3, 2, 1] + + +def _step_args( + thread_id: str, step: int, parent: dict | None +) -> tuple[dict, Checkpoint, dict[str, Any]]: + config: dict = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}} + if parent is not None: + config["configurable"]["checkpoint_id"] = parent["configurable"][ + "checkpoint_id" + ] + checkpoint: Checkpoint = empty_checkpoint() + checkpoint["id"] = str(uuid6(clock_seq=step)) + checkpoint["channel_versions"][CHANNEL] = f"v{step}" + if step == SEED_STEP: + checkpoint["channel_values"][CHANNEL] = _DeltaSnapshot(list(SEED_VALUE)) + return config, checkpoint, {CHANNEL: f"v{step}"} + return config, checkpoint, {} + + +async def _abuild_chain(saver: AsyncPostgresSaver) -> list[dict]: + thread_id = str(uuid4()) + parent: dict | None = None + configs: list[dict] = [] + for step in range(STEPS): + config, checkpoint, new_versions = _step_args(thread_id, step, parent) + parent = await saver.aput( + config, + checkpoint, + {"source": "loop", "step": step, "parents": {}}, + new_versions, + ) + await saver.aput_writes(parent, [(CHANNEL, f"w{step}")], str(uuid4())) + configs.append(parent) + return configs + + +def _build_chain(saver: PostgresSaver) -> list[dict]: + thread_id = str(uuid4()) + parent: dict | None = None + configs: list[dict] = [] + for step in range(STEPS): + config, checkpoint, new_versions = _step_args(thread_id, step, parent) + parent = saver.put( + config, + checkpoint, + {"source": "loop", "step": step, "parents": {}}, + new_versions, + ) + saver.put_writes(parent, [(CHANNEL, f"w{step}")], str(uuid4())) + configs.append(parent) + return configs + + +def _assert_history(entry: DeltaChannelHistory, page_size: int) -> None: + seed = entry.get("seed") + assert isinstance(seed, _DeltaSnapshot), ( + f"page_size={page_size}: expected a snapshot seed, " + f"got {entry.get('seed', '')!r}" + ) + assert seed.value == SEED_VALUE + assert [w[2] for w in entry["writes"]] == ["w1", "w2", "w3"], ( + f"page_size={page_size}: got {[w[2] for w in entry['writes']]}" + ) + + +@pytest.mark.parametrize("page_size", PAGE_SIZES) +async def test_async_target_older_than_the_first_page( + page_size: int, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("langgraph.checkpoint.postgres.aio._DELTA_PAGE_SIZE", page_size) + async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver: + await saver.setup() + configs = await _abuild_chain(saver) + result = await saver.aget_delta_channel_history( + config=configs[TARGET_STEP], channels=[CHANNEL] + ) + _assert_history(result[CHANNEL], page_size) + + +@pytest.mark.parametrize("page_size", PAGE_SIZES) +def test_sync_target_older_than_the_first_page( + page_size: int, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr("langgraph.checkpoint.postgres._DELTA_PAGE_SIZE", page_size) + with PostgresSaver.from_conn_string(DEFAULT_URI) as saver: + saver.setup() + configs = _build_chain(saver) + result = saver.get_delta_channel_history( + config=configs[TARGET_STEP], channels=[CHANNEL] + ) + _assert_history(result[CHANNEL], page_size) + + +async def test_root_target_has_no_history_and_still_terminates( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr("langgraph.checkpoint.postgres.aio._DELTA_PAGE_SIZE", 1) + async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver: + await saver.setup() + configs = await _abuild_chain(saver) + result = await saver.aget_delta_channel_history( + config=configs[0], channels=[CHANNEL] + ) + assert result[CHANNEL] == {"writes": []}