mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 05:55:14 +02:00
## 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>
648 lines
25 KiB
Python
648 lines
25 KiB
Python
from __future__ import annotations
|
|
|
|
import random
|
|
import warnings
|
|
from collections.abc import Mapping, Sequence
|
|
from importlib.metadata import version as get_version
|
|
from typing import Any, TypedDict, cast
|
|
|
|
from langchain_core.runnables import RunnableConfig
|
|
from langgraph.checkpoint.base import (
|
|
WRITES_IDX_MAP,
|
|
BaseCheckpointSaver,
|
|
ChannelVersions,
|
|
DeltaChannelHistory,
|
|
PendingWrite,
|
|
get_checkpoint_id,
|
|
)
|
|
from langgraph.checkpoint.serde.types import TASKS
|
|
from psycopg.types.json import Jsonb
|
|
|
|
# Page size for stage-1 paged scan in `get_delta_channel_history`. Internal
|
|
# constant — exposing this as a kwarg is left as a follow-up.
|
|
_DELTA_PAGE_SIZE = 1024
|
|
|
|
MetadataInput = dict[str, Any] | None
|
|
|
|
try:
|
|
major, minor = get_version("langgraph").split(".")[:2]
|
|
if int(major) == 0 and int(minor) < 5:
|
|
warnings.warn(
|
|
"You're using incompatible versions of langgraph and checkpoint-postgres. Please upgrade langgraph to avoid unexpected behavior.",
|
|
DeprecationWarning,
|
|
stacklevel=2,
|
|
)
|
|
except Exception:
|
|
# skip version check if running from source
|
|
pass
|
|
|
|
"""
|
|
To add a new migration, add a new string to the MIGRATIONS list.
|
|
The position of the migration in the list is the version number.
|
|
"""
|
|
MIGRATIONS = [
|
|
"""CREATE TABLE IF NOT EXISTS checkpoint_migrations (
|
|
v INTEGER PRIMARY KEY
|
|
);""",
|
|
"""CREATE TABLE IF NOT EXISTS checkpoints (
|
|
thread_id TEXT NOT NULL,
|
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
|
checkpoint_id TEXT NOT NULL,
|
|
parent_checkpoint_id TEXT,
|
|
type TEXT,
|
|
checkpoint JSONB NOT NULL,
|
|
metadata JSONB NOT NULL DEFAULT '{}',
|
|
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id)
|
|
);""",
|
|
"""CREATE TABLE IF NOT EXISTS checkpoint_blobs (
|
|
thread_id TEXT NOT NULL,
|
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
|
channel TEXT NOT NULL,
|
|
version TEXT NOT NULL,
|
|
type TEXT NOT NULL,
|
|
blob BYTEA,
|
|
PRIMARY KEY (thread_id, checkpoint_ns, channel, version)
|
|
);""",
|
|
"""CREATE TABLE IF NOT EXISTS checkpoint_writes (
|
|
thread_id TEXT NOT NULL,
|
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
|
checkpoint_id TEXT NOT NULL,
|
|
task_id TEXT NOT NULL,
|
|
idx INTEGER NOT NULL,
|
|
channel TEXT NOT NULL,
|
|
type TEXT,
|
|
blob BYTEA NOT NULL,
|
|
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
|
|
);""",
|
|
"ALTER TABLE checkpoint_blobs ALTER COLUMN blob DROP not null;",
|
|
# NOTE: this is a no-op migration to ensure that the versions in the migrations table are correct.
|
|
# This is necessary due to an empty migration previously added to the list.
|
|
"SELECT 1;",
|
|
"""
|
|
CREATE INDEX CONCURRENTLY IF NOT EXISTS checkpoints_thread_id_idx ON checkpoints(thread_id);
|
|
""",
|
|
"""
|
|
CREATE INDEX CONCURRENTLY IF NOT EXISTS checkpoint_blobs_thread_id_idx ON checkpoint_blobs(thread_id);
|
|
""",
|
|
"""
|
|
CREATE INDEX CONCURRENTLY IF NOT EXISTS checkpoint_writes_thread_id_idx ON checkpoint_writes(thread_id);
|
|
""",
|
|
"""ALTER TABLE checkpoint_writes ADD COLUMN IF NOT EXISTS task_path TEXT NOT NULL DEFAULT '';""",
|
|
]
|
|
|
|
SELECT_SQL = """
|
|
select
|
|
thread_id,
|
|
checkpoint,
|
|
checkpoint_ns,
|
|
checkpoint_id,
|
|
parent_checkpoint_id,
|
|
metadata,
|
|
(
|
|
select array_agg(array[bl.channel::bytea, bl.type::bytea, bl.blob])
|
|
from jsonb_each_text(checkpoint -> 'channel_versions')
|
|
inner join checkpoint_blobs bl
|
|
on bl.thread_id = checkpoints.thread_id
|
|
and bl.checkpoint_ns = checkpoints.checkpoint_ns
|
|
and bl.channel = jsonb_each_text.key
|
|
and bl.version = jsonb_each_text.value
|
|
) as channel_values,
|
|
(
|
|
select
|
|
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
|
|
from checkpoint_writes cw
|
|
where cw.thread_id = checkpoints.thread_id
|
|
and cw.checkpoint_ns = checkpoints.checkpoint_ns
|
|
and cw.checkpoint_id = checkpoints.checkpoint_id
|
|
) as pending_writes
|
|
from checkpoints """
|
|
|
|
SELECT_PENDING_SENDS_SQL = f"""
|
|
select
|
|
checkpoint_id,
|
|
array_agg(array[type::bytea, blob] order by task_path, task_id, idx) as sends
|
|
from checkpoint_writes
|
|
where thread_id = %s
|
|
and checkpoint_id = any(%s)
|
|
and channel = '{TASKS}'
|
|
group by checkpoint_id
|
|
"""
|
|
|
|
UPSERT_CHECKPOINT_BLOBS_SQL = """
|
|
INSERT INTO checkpoint_blobs (thread_id, checkpoint_ns, channel, version, type, blob)
|
|
VALUES (%s, %s, %s, %s, %s, %s)
|
|
ON CONFLICT (thread_id, checkpoint_ns, channel, version) DO NOTHING
|
|
"""
|
|
|
|
UPSERT_CHECKPOINTS_SQL = """
|
|
INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, parent_checkpoint_id, checkpoint, metadata)
|
|
VALUES (%s, %s, %s, %s, %s, %s)
|
|
ON CONFLICT (thread_id, checkpoint_ns, checkpoint_id)
|
|
DO UPDATE SET
|
|
checkpoint = EXCLUDED.checkpoint,
|
|
metadata = EXCLUDED.metadata;
|
|
"""
|
|
|
|
UPSERT_CHECKPOINT_WRITES_SQL = """
|
|
INSERT INTO checkpoint_writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, blob)
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s)
|
|
ON CONFLICT (thread_id, checkpoint_ns, checkpoint_id, task_id, idx) DO UPDATE SET
|
|
channel = EXCLUDED.channel,
|
|
type = EXCLUDED.type,
|
|
blob = EXCLUDED.blob;
|
|
"""
|
|
|
|
INSERT_CHECKPOINT_WRITES_SQL = """
|
|
INSERT INTO checkpoint_writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, blob)
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s)
|
|
ON CONFLICT (thread_id, checkpoint_ns, checkpoint_id, task_id, idx) DO NOTHING
|
|
"""
|
|
|
|
|
|
class _DeltaStage2Row(TypedDict, total=False):
|
|
"""One row from `_build_delta_stage2_sql` (a UNION ALL of writes and blobs)."""
|
|
|
|
_kind: str # "w" or "b"
|
|
checkpoint_id: str | None # "w" rows only
|
|
channel: str | None # set on both "w" and "b" rows
|
|
type: str | None
|
|
blob: bytes | None
|
|
task_id: str | None # "w" rows only
|
|
idx: int | None # "w" rows only
|
|
version: str | None # "b" rows only
|
|
|
|
|
|
# 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:
|
|
"""Build stage 1 SQL with K parallel version lookups + seed probes.
|
|
|
|
For channels=["messages", "files"] (with `paged=True`) the result is::
|
|
|
|
SELECT checkpoint_id, parent_checkpoint_id,
|
|
checkpoint -> 'channel_versions' ->> %s AS ver_0,
|
|
EXISTS (SELECT 1 FROM checkpoint_blobs b0
|
|
WHERE b0.thread_id = checkpoints.thread_id
|
|
AND b0.checkpoint_ns = checkpoints.checkpoint_ns
|
|
AND b0.channel = %s
|
|
AND b0.version = checkpoint -> 'channel_versions' ->> %s
|
|
AND b0.type <> 'empty') AS hb_0,
|
|
checkpoint -> 'channel_values' -> %s AS inline_0,
|
|
checkpoint -> 'channel_versions' ->> %s AS ver_1,
|
|
EXISTS (...) AS hb_1,
|
|
checkpoint -> 'channel_values' -> %s AS inline_1
|
|
FROM checkpoints
|
|
WHERE thread_id = %s AND checkpoint_ns = %s
|
|
AND (%s::text IS NULL OR checkpoint_id < %s)
|
|
ORDER BY checkpoint_id DESC
|
|
LIMIT %s
|
|
|
|
A stored value for a channel lives in one of two places, because `put`
|
|
splits them:
|
|
|
|
* **blob** — non-primitive values (and `_DeltaSnapshot`) are moved to
|
|
`checkpoint_blobs`. `hb_i` ("has blob") probes for one. The probe hits
|
|
that table's primary key `(thread_id, checkpoint_ns, channel, version)`
|
|
exactly, so it is an index lookup per row per channel.
|
|
* **inline** — `None`, `str`, `int`, `float` and `bool` stay in the
|
|
checkpoint's own `channel_values` and get no blob row at all. `inline_i`
|
|
returns that value.
|
|
|
|
Testing only for a key in `channel_values` (the previous approach) missed
|
|
blob-stored plain values, since `put` leaves an inline marker there for
|
|
`_DeltaSnapshot` but not for a plain value — which is what a thread
|
|
migrated from a pre-delta channel type leaves behind. Probing only the
|
|
blobs table would conversely miss inline primitives. Both are needed, and
|
|
the caller treats "either present" as the seed.
|
|
|
|
`hb_i` also disambiguates the two: for a `_DeltaSnapshot`, `inline_i` is the
|
|
literal `true` marker rather than the value, so a blob must win over an
|
|
inline reading whenever one exists. That ordering is what makes a genuine
|
|
inline `true` (a bool channel) distinguishable from the marker.
|
|
|
|
The `type <> 'empty'` predicate mirrors the check stage 2 already applies
|
|
when resolving the seed blob. `put` does not currently produce `empty` rows
|
|
on this path — `blob_versions` is filtered to keys present in
|
|
`channel_values`, so `_dump_blobs`' empty branch is unreachable from it —
|
|
but without the predicate the two stages could disagree: stage 1 would
|
|
terminate the walk on a row stage 2 then discards, yielding no seed *and* a
|
|
truncated write chain, which is the failure this function exists to avoid.
|
|
|
|
Channel names are passed as `%s` parameters (safe from SQL injection).
|
|
Only the column aliases `ver_i` / `hb_i` / `inline_i` and the subquery alias
|
|
`b{i}` are interpolated into the SQL string (i is bounded by len(channels)
|
|
and uses safe identifiers).
|
|
|
|
Caller must extend params with `[ch_0 x4, ch_1 x4, ..., thread_id, ns,
|
|
cursor, cursor, page_size]` when `paged=True` — four per channel: the
|
|
version lookup, the blob's channel, the version the blob must match, and the
|
|
inline lookup.
|
|
|
|
When `paged=False`, the WHERE has no cursor predicate and there's no
|
|
LIMIT/ORDER BY — kept as a non-public helper for tests/diagnostics.
|
|
"""
|
|
cols = []
|
|
for i in range(len(channels)):
|
|
cols.append(
|
|
f"checkpoint -> 'channel_versions' ->> %s AS ver_{i}, "
|
|
f"EXISTS (SELECT 1 FROM checkpoint_blobs b{i} "
|
|
f"WHERE b{i}.thread_id = checkpoints.thread_id "
|
|
f"AND b{i}.checkpoint_ns = checkpoints.checkpoint_ns "
|
|
f"AND b{i}.channel = %s "
|
|
f"AND b{i}.version = checkpoint -> 'channel_versions' ->> %s "
|
|
f"AND b{i}.type <> 'empty') AS hb_{i}, "
|
|
f"checkpoint -> 'channel_values' -> %s AS inline_{i}"
|
|
)
|
|
sql = (
|
|
"SELECT checkpoint_id, parent_checkpoint_id, "
|
|
+ ", ".join(cols)
|
|
+ " FROM checkpoints WHERE thread_id = %s AND checkpoint_ns = %s"
|
|
)
|
|
if paged:
|
|
sql += (
|
|
" AND (%s::text IS NULL OR checkpoint_id < %s)"
|
|
" ORDER BY checkpoint_id DESC LIMIT %s"
|
|
)
|
|
return sql
|
|
|
|
|
|
def _build_delta_stage2_sql(
|
|
*,
|
|
channels_with_chain: Sequence[str],
|
|
channels_with_seed: Sequence[str],
|
|
) -> str:
|
|
"""Build stage 2 SQL as a per-channel UNION ALL.
|
|
|
|
For each channel with a non-empty chain, emit one branch reading
|
|
`checkpoint_writes` for that specific channel + chain_cids. For each
|
|
channel with a seed_version, emit one branch reading `checkpoint_blobs`
|
|
for that channel + version. This avoids the over-fetch of the prior
|
|
`channel = ANY(channels) AND checkpoint_id = ANY(union)` form when
|
|
channels have different chain depths.
|
|
|
|
The caller must pass parameters in matching order:
|
|
|
|
for ch in channels_with_chain:
|
|
params += [thread_id, checkpoint_ns, ch, chain_cids[ch]]
|
|
for ch in channels_with_seed:
|
|
params += [thread_id, checkpoint_ns, ch, seed_version[ch]]
|
|
|
|
Returns an empty SQL string if both channel lists are empty (caller
|
|
must skip executing in that case).
|
|
"""
|
|
branches: list[str] = []
|
|
for _ in channels_with_chain:
|
|
branches.append(
|
|
"SELECT 'w'::text AS _kind, "
|
|
"checkpoint_id, channel, "
|
|
"type, blob, task_id, idx, NULL::text AS version "
|
|
"FROM checkpoint_writes "
|
|
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
|
"AND checkpoint_id = ANY(%s)"
|
|
)
|
|
for _ in channels_with_seed:
|
|
branches.append(
|
|
"SELECT 'b'::text AS _kind, NULL::text AS checkpoint_id, channel, "
|
|
"type, blob, NULL::text AS task_id, NULL::int AS idx, version "
|
|
"FROM checkpoint_blobs "
|
|
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
|
"AND version = %s"
|
|
)
|
|
return " UNION ALL ".join(branches)
|
|
|
|
|
|
# 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]):
|
|
SELECT_SQL = SELECT_SQL
|
|
SELECT_PENDING_SENDS_SQL = SELECT_PENDING_SENDS_SQL
|
|
MIGRATIONS = MIGRATIONS
|
|
UPSERT_CHECKPOINT_BLOBS_SQL = UPSERT_CHECKPOINT_BLOBS_SQL
|
|
UPSERT_CHECKPOINTS_SQL = UPSERT_CHECKPOINTS_SQL
|
|
UPSERT_CHECKPOINT_WRITES_SQL = UPSERT_CHECKPOINT_WRITES_SQL
|
|
INSERT_CHECKPOINT_WRITES_SQL = INSERT_CHECKPOINT_WRITES_SQL
|
|
|
|
supports_pipeline: bool
|
|
|
|
def _migrate_pending_sends(
|
|
self,
|
|
pending_sends: list[tuple[bytes, bytes]],
|
|
checkpoint: dict[str, Any],
|
|
channel_values: list[tuple[bytes, bytes, bytes]],
|
|
) -> None:
|
|
if not pending_sends:
|
|
return
|
|
# add to values
|
|
enc, blob = self.serde.dumps_typed(
|
|
[self.serde.loads_typed((c.decode(), b)) for c, b in pending_sends],
|
|
)
|
|
channel_values.append((TASKS.encode(), enc.encode(), blob))
|
|
# add to versions
|
|
checkpoint["channel_versions"][TASKS] = (
|
|
max(checkpoint["channel_versions"].values())
|
|
if checkpoint["channel_versions"]
|
|
else self.get_next_version(None, None)
|
|
)
|
|
|
|
def _load_blobs(
|
|
self, blob_values: list[tuple[bytes, bytes, bytes]]
|
|
) -> dict[str, Any]:
|
|
if not blob_values:
|
|
return {}
|
|
return {
|
|
k.decode(): self.serde.loads_typed((t.decode(), v))
|
|
for k, t, v in blob_values
|
|
if t.decode() != "empty"
|
|
}
|
|
|
|
@staticmethod
|
|
def _ingest_stage1_page(
|
|
stage1_rows: Sequence[Mapping[str, Any]],
|
|
channels: Sequence[str],
|
|
parent_of: dict[str, str | None],
|
|
ver_by_i_by_cid: list[dict[str, str | None]],
|
|
hb_by_i_by_cid: list[dict[str, bool]],
|
|
inline_by_i_by_cid: list[dict[str, Any]],
|
|
) -> str | None:
|
|
"""Fold one stage-1 page into the running walk-state mappings.
|
|
|
|
Returns the oldest checkpoint_id seen on this page (smallest, since
|
|
pages come back DESC). Caller uses it as the cursor for the next
|
|
page (`AND checkpoint_id < cursor`).
|
|
"""
|
|
oldest: str | None = None
|
|
for r in stage1_rows:
|
|
cid = cast(str, r["checkpoint_id"])
|
|
parent_of[cid] = cast("str | None", r["parent_checkpoint_id"])
|
|
for i in range(len(channels)):
|
|
ver_by_i_by_cid[i][cid] = cast("str | None", r.get(f"ver_{i}"))
|
|
hb_by_i_by_cid[i][cid] = bool(r.get(f"hb_{i}"))
|
|
inline_by_i_by_cid[i][cid] = r.get(f"inline_{i}")
|
|
# Rows are DESC; the last one is the smallest cid in the page.
|
|
oldest = cid
|
|
return oldest
|
|
|
|
@staticmethod
|
|
def _try_advance_walks(
|
|
target_id: str,
|
|
channels: Sequence[str],
|
|
parent_of: Mapping[str, str | None],
|
|
ver_by_i_by_cid: Sequence[Mapping[str, str | None]],
|
|
hb_by_i_by_cid: Sequence[Mapping[str, bool]],
|
|
inline_by_i_by_cid: Sequence[Mapping[str, Any]],
|
|
chain_by_ch: dict[str, list[str]],
|
|
seed_ver_by_ch: dict[str, str | None],
|
|
seed_inline_by_ch: dict[str, Any],
|
|
walk_cursor_by_ch: dict[str, str | None],
|
|
seeded: set[str],
|
|
) -> None:
|
|
"""Advance each not-yet-seeded channel's walk as far as possible.
|
|
|
|
Uses the partial `parent_of` map accumulated so far. A walk stops
|
|
either because:
|
|
(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),
|
|
(c) the next ancestor cid isn't in `parent_of` yet (waiting for
|
|
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.
|
|
"""
|
|
for i, ch in enumerate(channels):
|
|
if ch in seeded:
|
|
continue
|
|
# 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:
|
|
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]
|
|
inline_i = inline_by_i_by_cid[i]
|
|
ver_i = ver_by_i_by_cid[i]
|
|
while cur_cid is not None:
|
|
if cur_cid not in parent_of:
|
|
# Need more pages to continue this walk.
|
|
break
|
|
ch_chain.append(cur_cid)
|
|
has_blob = hb_i.get(cur_cid, False)
|
|
inline = inline_i.get(cur_cid)
|
|
if has_blob or inline is not None:
|
|
# A blob wins: for a `_DeltaSnapshot` the inline reading is
|
|
# the `true` marker, not the value.
|
|
seed_ver_by_ch[ch] = ver_i.get(cur_cid)
|
|
if not has_blob:
|
|
seed_inline_by_ch[ch] = inline
|
|
seeded.add(ch)
|
|
cur_cid = None
|
|
break
|
|
cur_cid = parent_of[cur_cid]
|
|
walk_cursor_by_ch[ch] = cur_cid
|
|
|
|
def _build_delta_channels_writes_history(
|
|
self,
|
|
*,
|
|
channels: Sequence[str],
|
|
chain_by_ch: Mapping[str, list[str]],
|
|
seed_ver_by_ch: Mapping[str, str | None],
|
|
seed_inline_by_ch: Mapping[str, Any],
|
|
stage2_rows: Sequence[_DeltaStage2Row],
|
|
) -> dict[str, DeltaChannelHistory]:
|
|
"""Demux stage 2 rows per channel; produce per-channel histories.
|
|
|
|
stage2_rows carry `channel` on every row. We build per-channel
|
|
`writes_by_cid` and per-channel `seed_blob` dicts, then assemble
|
|
a `DeltaChannelHistory` per requested channel.
|
|
|
|
A seed comes from the blobs table when the walk found one there, and
|
|
otherwise from `seed_inline_by_ch` — `put` keeps `None`, `str`, `int`,
|
|
`float` and `bool` values in the checkpoint's own `channel_values` with
|
|
no blob row, so those never appear in `stage2_rows`.
|
|
|
|
The `seed` key is omitted when the walk reached root without finding a
|
|
stored value, or when the seed blob is sentinel "empty" — in both cases
|
|
the consumer treats absence as "start empty".
|
|
"""
|
|
# writes_by_ch_by_cid[channel][cid] = list of (type, blob, task_id, idx)
|
|
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
|
|
ch: {} for ch in channels
|
|
}
|
|
# seed_blob_by_ver[(channel, version)] = (type, blob)
|
|
seed_blob_by_ver: dict[tuple[str, str], tuple[str, bytes]] = {}
|
|
|
|
for r in stage2_rows:
|
|
ch = cast(str, r["channel"])
|
|
kind = r["_kind"]
|
|
if kind == "w":
|
|
cid = cast(str, r["checkpoint_id"])
|
|
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
|
cast(
|
|
"tuple[str, bytes, str, int]",
|
|
(r["type"], r["blob"], r["task_id"], r["idx"]),
|
|
)
|
|
)
|
|
else: # kind == "b"
|
|
ver = cast(str, r["version"])
|
|
seed_blob_by_ver[(ch, ver)] = cast(
|
|
"tuple[str, bytes]", (r["type"], r["blob"])
|
|
)
|
|
|
|
# Sort writes per (channel, cid) newest-first by (task_id, idx)
|
|
for cid_map in writes_by_ch_by_cid.values():
|
|
for ws in cid_map.values():
|
|
ws.sort(key=lambda w: (w[2], w[3]), reverse=True)
|
|
|
|
result: dict[str, DeltaChannelHistory] = {}
|
|
for ch in channels:
|
|
chain_cids = chain_by_ch.get(ch, [])
|
|
seed_version = seed_ver_by_ch.get(ch)
|
|
|
|
collected: list[PendingWrite] = []
|
|
cid_writes = writes_by_ch_by_cid.get(ch, {})
|
|
for cid in chain_cids:
|
|
for type_tag, write_blob, task_id, _idx in cid_writes.get(cid, []):
|
|
val = self.serde.loads_typed((type_tag, write_blob))
|
|
collected.append((task_id, ch, val))
|
|
collected.reverse()
|
|
|
|
entry: DeltaChannelHistory = {"writes": collected}
|
|
if seed_version is not None:
|
|
blob = seed_blob_by_ver.get((ch, seed_version))
|
|
if blob is not None and blob[0] != "empty":
|
|
entry["seed"] = self.serde.loads_typed(blob)
|
|
elif ch in seed_inline_by_ch:
|
|
# Inline primitive: stored in the checkpoint, not the blobs
|
|
# table, so stage 2 never returned a row for it.
|
|
entry["seed"] = seed_inline_by_ch[ch]
|
|
result[ch] = entry
|
|
return result
|
|
|
|
def _dump_blobs(
|
|
self,
|
|
thread_id: str,
|
|
checkpoint_ns: str,
|
|
values: dict[str, Any],
|
|
versions: ChannelVersions,
|
|
) -> list[tuple[str, str, str, str, str, bytes | None]]:
|
|
if not versions:
|
|
return []
|
|
|
|
return [
|
|
(
|
|
thread_id,
|
|
checkpoint_ns,
|
|
k,
|
|
cast(str, ver),
|
|
*(
|
|
self.serde.dumps_typed(values[k])
|
|
if k in values
|
|
else ("empty", None)
|
|
),
|
|
)
|
|
for k, ver in versions.items()
|
|
]
|
|
|
|
def _load_writes(
|
|
self, writes: list[tuple[bytes, bytes, bytes, bytes]]
|
|
) -> list[tuple[str, str, Any]]:
|
|
return (
|
|
[
|
|
(
|
|
tid.decode(),
|
|
channel.decode(),
|
|
self.serde.loads_typed((t.decode(), v)),
|
|
)
|
|
for tid, channel, t, v in writes
|
|
]
|
|
if writes
|
|
else []
|
|
)
|
|
|
|
def _dump_writes(
|
|
self,
|
|
thread_id: str,
|
|
checkpoint_ns: str,
|
|
checkpoint_id: str,
|
|
task_id: str,
|
|
task_path: str,
|
|
writes: Sequence[tuple[str, Any]],
|
|
) -> list[tuple[str, str, str, str, str, int, str, str, bytes]]:
|
|
return [
|
|
(
|
|
thread_id,
|
|
checkpoint_ns,
|
|
checkpoint_id,
|
|
task_id,
|
|
task_path,
|
|
WRITES_IDX_MAP.get(channel, idx),
|
|
channel,
|
|
*self.serde.dumps_typed(value),
|
|
)
|
|
for idx, (channel, value) in enumerate(writes)
|
|
]
|
|
|
|
def get_next_version(self, current: str | None, channel: None) -> str:
|
|
if current is None:
|
|
current_v = 0
|
|
elif isinstance(current, int):
|
|
current_v = current
|
|
else:
|
|
current_v = int(current.split(".")[0])
|
|
next_v = current_v + 1
|
|
next_h = random.random()
|
|
return f"{next_v:032}.{next_h:016}"
|
|
|
|
def _search_where(
|
|
self,
|
|
config: RunnableConfig | None,
|
|
filter: MetadataInput,
|
|
before: RunnableConfig | None = None,
|
|
) -> tuple[str, list[Any]]:
|
|
"""Return WHERE clause predicates for alist() given config, filter, before.
|
|
|
|
This method returns a tuple of a string and a tuple of values. The string
|
|
is the parametered WHERE clause predicate (including the WHERE keyword):
|
|
"WHERE column1 = $1 AND column2 IS $2". The list of values contains the
|
|
values for each of the corresponding parameters.
|
|
"""
|
|
wheres = []
|
|
param_values = []
|
|
|
|
# construct predicate for config filter
|
|
if config:
|
|
wheres.append("thread_id = %s ")
|
|
param_values.append(config["configurable"]["thread_id"])
|
|
checkpoint_ns = config["configurable"].get("checkpoint_ns")
|
|
if checkpoint_ns is not None:
|
|
wheres.append("checkpoint_ns = %s")
|
|
param_values.append(checkpoint_ns)
|
|
|
|
if checkpoint_id := get_checkpoint_id(config):
|
|
wheres.append("checkpoint_id = %s ")
|
|
param_values.append(checkpoint_id)
|
|
|
|
# construct predicate for metadata filter
|
|
if filter:
|
|
wheres.append("metadata @> %s ")
|
|
param_values.append(Jsonb(filter))
|
|
|
|
# construct predicate for `before`
|
|
if before is not None:
|
|
wheres.append("checkpoint_id < %s ")
|
|
param_values.append(get_checkpoint_id(before))
|
|
|
|
return (
|
|
"WHERE " + " AND ".join(wheres) if wheres else "",
|
|
param_values,
|
|
)
|