mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 21:45:08 +02:00
fix(checkpoint-postgres): derive the delta walk cursor once the target loads (#8556)
## Summary `get_delta_channel_history` on Postgres returns an empty history for any `DeltaChannel` on a target checkpoint that is not within the first stage-1 pagination page (1024 rows) of the thread. No exception, no warning: the channel just hydrates empty. Fixes #8448 ## Problem Stage 1 pages `checkpoints` newest-first from the head of the thread, and after each page `_try_advance_walks` tries to move every not-yet-seeded channel's walk along the partial `parent_of` map accumulated so far. The walk starts at the target's parent: ```python if ch not in walk_cursor_by_ch: walk_cursor_by_ch[ch] = parent_of.get(target_id) ``` The target can be any checkpoint in the thread, not just the head, so on the first page `parent_of` frequently has no row for it yet. `.get` then returns `None`, which is also what a target with no parent returns, and the two are stored identically. Because the initialisation is guarded by `ch not in walk_cursor_by_ch`, it never runs again: once the walk is parked at `None` it stays there even after the target's real row and real parent load on a later page. The result is an empty chain and no seed. Downstream `channels_from_checkpoint` does ```python replay_ch = delta_spec.from_checkpoint(history.get("seed", MISSING)) replay_ch.replay_writes(history["writes"]) ``` so `get_state`, `get_state_history` and `update_state` against an older checkpoint reconstruct a `messages` channel as `[]` on a thread with hundreds of real messages. ## Fix Start the walk only once `target_id` is actually present in `parent_of`, so "the target has not loaded yet" stops sharing a representation with "the target is a root": ```python if ch not in walk_cursor_by_ch: if target_id not in parent_of: continue walk_cursor_by_ch[ch] = parent_of[target_id] ``` `_try_advance_walks` is a static method on `BasePostgresSaver`, so `PostgresSaver` and `AsyncPostgresSaver` are both covered by the one change. ## Why it's safe `continue` leaves the channel exactly as it was, so a later page retries. The three existing stop conditions are untouched: a channel that finds its seed still seeds, one that reaches a real root still parks at `None`, and one waiting on an ancestor still keeps its cursor. Paging still terminates on a short page, which is what ends the run for a target that really is a root. ## Long-term The sibling sqlite implementation avoids this class of bug differently, by starting its stage-1 scan at the target (`checkpoint_id <= ?`) instead of at the head. Postgres could adopt the same bound and would then never fetch a checkpoint newer than the target at all, which looks like the bigger win on a long thread. It makes the read path depend on ancestors always sorting below their descendants, though, which sqlite already assumes but the Postgres fast path currently does not. #8550 now reports that assumption as a bug in sqlite, on the grounds that ancestry is defined by `parent_checkpoint_id` and the contract does not require ids to be monotonic, so the bound is the wrong direction to move Postgres in. Paging the full thread and following parent pointers is what keeps this path correct when ids are not monotonic, and with this fix Postgres returns the right history for #8550's scenario at every page size. ## Test plan New `libs/checkpoint-postgres/tests/test_delta_pagination.py`. Page size is monkeypatched rather than writing 1024+ real checkpoints per case, since the only thing that decides the behaviour is which page the target lands on. - [x] `test_async_target_older_than_the_first_page` and its sync twin, parametrised over page sizes `[_DELTA_PAGE_SIZE, 3, 2, 1]`. The thread has 8 checkpoints with a snapshot at step 1 and the target at step 4, so every size at or below 3 leaves the target off the first page. The real page size is the control. - [x] `test_root_target_has_no_history_and_still_terminates` covers the case where a `None` cursor is the correct answer, at page size 1 so the paging loop runs the length of the thread. - [x] 6 of the 9 fail on `main` (`expected a snapshot seed, got '<missing>'`); the 3 that pass are the two controls and the root case. - [x] `make format`, `make lint_package`, `make lint_tests` clean. - [x] Full `libs/checkpoint-postgres` suite, rebased on current `main`: 279 passed, 3 skipped on Postgres 16. - [x] Graph-level repro with `_DELTA_PAGE_SIZE = 5`: 10 invocations, then `get_state` on the 8th-newest checkpoint returns `[]` on `main` and the full history on this branch. Thanks to @Navneet-Scaler for the report, the mechanism write-up, and the fix in #8453, which this matches. Co-authored-by: Navneet-Scaler <147032454+Navneet-Scaler@users.noreply.github.com>
This commit is contained in:
co-authored by
Navneet-Scaler
parent
f5804a5bf5
commit
c0279f0910
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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', '<missing>')!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": []}
|
||||
Reference in New Issue
Block a user