mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 05:55:14 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a8e732c879 | ||
|
|
eb69f67b65 | ||
|
|
c0279f0910 | ||
|
|
f5804a5bf5 |
@@ -139,23 +139,10 @@ jobs:
|
|||||||
echo EOF
|
echo EOF
|
||||||
} >> "$GITHUB_OUTPUT"
|
} >> "$GITHUB_OUTPUT"
|
||||||
|
|
||||||
test-pypi-publish:
|
|
||||||
needs:
|
|
||||||
- build
|
|
||||||
- release-notes
|
|
||||||
permissions:
|
|
||||||
contents: read
|
|
||||||
id-token: write
|
|
||||||
uses: ./.github/workflows/_test_release.yml
|
|
||||||
with:
|
|
||||||
working-directory: ${{ inputs.working-directory }}
|
|
||||||
secrets: inherit
|
|
||||||
|
|
||||||
pre-release-checks:
|
pre-release-checks:
|
||||||
needs:
|
needs:
|
||||||
- build
|
- build
|
||||||
- release-notes
|
- release-notes
|
||||||
- test-pypi-publish
|
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
@@ -180,31 +167,20 @@ jobs:
|
|||||||
enable-cache: false
|
enable-cache: false
|
||||||
working-directory: ${{ inputs.working-directory }}
|
working-directory: ${{ inputs.working-directory }}
|
||||||
|
|
||||||
- name: Import published package
|
- uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8.0.1
|
||||||
|
with:
|
||||||
|
name: dist
|
||||||
|
path: ${{ inputs.working-directory }}/dist/
|
||||||
|
|
||||||
|
- name: Import dist package
|
||||||
shell: bash
|
shell: bash
|
||||||
working-directory: ${{ inputs.working-directory }}
|
working-directory: ${{ inputs.working-directory }}
|
||||||
env:
|
env:
|
||||||
PKG_NAME: ${{ needs.build.outputs.pkg-name }}
|
PKG_NAME: ${{ needs.build.outputs.pkg-name }}
|
||||||
VERSION: ${{ needs.build.outputs.version }}
|
VERSION: ${{ needs.build.outputs.version }}
|
||||||
# Here we use:
|
# Install directly from the locally-built wheel (no index resolution needed).
|
||||||
# - The default regular PyPI index as the *primary* index, meaning
|
|
||||||
# that it takes priority (https://pypi.org/simple)
|
|
||||||
# - The test PyPI index as an extra index, so that any dependencies that
|
|
||||||
# are not found on test PyPI can be resolved and installed anyway.
|
|
||||||
# (https://test.pypi.org/simple). This will include the PKG_NAME==VERSION
|
|
||||||
# package because VERSION will not have been uploaded to regular PyPI yet.
|
|
||||||
# - attempt install again after 5 seconds if it fails because there is
|
|
||||||
# sometimes a delay in availability on test pypi
|
|
||||||
run: |
|
run: |
|
||||||
uv run pip install \
|
uv run pip install dist/*.whl
|
||||||
--extra-index-url https://test.pypi.org/simple/ \
|
|
||||||
"$PKG_NAME==$VERSION" || \
|
|
||||||
( \
|
|
||||||
sleep 5 && \
|
|
||||||
uv run pip install \
|
|
||||||
--extra-index-url https://test.pypi.org/simple/ \
|
|
||||||
"$PKG_NAME==$VERSION" \
|
|
||||||
)
|
|
||||||
|
|
||||||
if [[ "$PKG_NAME" == *prebuilt* ]]; then
|
if [[ "$PKG_NAME" == *prebuilt* ]]; then
|
||||||
uv run pip install langgraph
|
uv run pip install langgraph
|
||||||
@@ -226,7 +202,7 @@ jobs:
|
|||||||
run: uv sync --group test
|
run: uv sync --group test
|
||||||
working-directory: ${{ inputs.working-directory }}
|
working-directory: ${{ inputs.working-directory }}
|
||||||
|
|
||||||
# Overwrite the local version of the package with the test PyPI version.
|
# Overwrite the local version of the package with the built version
|
||||||
- name: Import published package (again)
|
- name: Import published package (again)
|
||||||
working-directory: ${{ inputs.working-directory }}
|
working-directory: ${{ inputs.working-directory }}
|
||||||
shell: bash
|
shell: bash
|
||||||
@@ -234,14 +210,25 @@ jobs:
|
|||||||
PKG_NAME: ${{ needs.build.outputs.pkg-name }}
|
PKG_NAME: ${{ needs.build.outputs.pkg-name }}
|
||||||
VERSION: ${{ needs.build.outputs.version }}
|
VERSION: ${{ needs.build.outputs.version }}
|
||||||
run: |
|
run: |
|
||||||
uv run pip install \
|
uv run pip install dist/*.whl
|
||||||
--extra-index-url https://test.pypi.org/simple/ \
|
|
||||||
"$PKG_NAME==$VERSION"
|
|
||||||
|
|
||||||
- name: Run unit tests
|
- name: Run unit tests
|
||||||
run: make test
|
run: make test
|
||||||
working-directory: ${{ inputs.working-directory }}
|
working-directory: ${{ inputs.working-directory }}
|
||||||
|
|
||||||
|
test-pypi-publish:
|
||||||
|
needs:
|
||||||
|
- build
|
||||||
|
- release-notes
|
||||||
|
- pre-release-checks
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
id-token: write
|
||||||
|
uses: ./.github/workflows/_test_release.yml
|
||||||
|
with:
|
||||||
|
working-directory: ${{ inputs.working-directory }}
|
||||||
|
secrets: inherit
|
||||||
|
|
||||||
publish:
|
publish:
|
||||||
needs:
|
needs:
|
||||||
- build
|
- build
|
||||||
|
|||||||
-57
@@ -267,61 +267,6 @@ async def test_history_seed_ancestor_own_writes_are_replayed(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# Every uuid4 `build_delta_chain` tags its own writes with sorts between these
|
|
||||||
# two, so task_id order is fixed and always disagrees with task_path order.
|
|
||||||
TASK_ID_SORTS_FIRST = "00000000-0000-0000-0000-000000000000"
|
|
||||||
TASK_ID_SORTS_LAST = "ffffffff-ffff-ffff-ffff-ffffffffffff"
|
|
||||||
|
|
||||||
|
|
||||||
async def test_history_orders_parallel_writes_by_task_path(
|
|
||||||
saver: BaseCheckpointSaver,
|
|
||||||
) -> None:
|
|
||||||
"""Writes from parallel tasks replay in task_path order, not task_id order."""
|
|
||||||
configs = await build_delta_chain(
|
|
||||||
saver,
|
|
||||||
thread_id=str(uuid4()),
|
|
||||||
channel="ch",
|
|
||||||
snapshots_at_steps=[0],
|
|
||||||
total_steps=3,
|
|
||||||
)
|
|
||||||
step_1, head = configs[1], configs[2]
|
|
||||||
await saver.aput_writes(
|
|
||||||
step_1, [("ch", "second")], TASK_ID_SORTS_FIRST, "~pull, 02"
|
|
||||||
)
|
|
||||||
await saver.aput_writes(step_1, [("ch", "first")], TASK_ID_SORTS_LAST, "~pull, 01")
|
|
||||||
|
|
||||||
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
|
|
||||||
values = [w[2] for w in result["ch"]["writes"]]
|
|
||||||
assert values == [1, "first", "second"], (
|
|
||||||
f"Expected task_path order [1, 'first', 'second'], got {values}. "
|
|
||||||
"Ordering by (task_id, idx) alone yields [1, 'second', 'first']."
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
async def test_history_orders_pathless_writes_first(
|
|
||||||
saver: BaseCheckpointSaver,
|
|
||||||
) -> None:
|
|
||||||
"""Writes stored without a task_path (graph input) replay before task writes."""
|
|
||||||
configs = await build_delta_chain(
|
|
||||||
saver,
|
|
||||||
thread_id=str(uuid4()),
|
|
||||||
channel="ch",
|
|
||||||
snapshots_at_steps=[0],
|
|
||||||
total_steps=3,
|
|
||||||
)
|
|
||||||
step_1, head = configs[1], configs[2]
|
|
||||||
await saver.aput_writes(
|
|
||||||
step_1, [("ch", "from_node")], TASK_ID_SORTS_FIRST, "~pull, a"
|
|
||||||
)
|
|
||||||
await saver.aput_writes(step_1, [("ch", "from_input")], TASK_ID_SORTS_LAST)
|
|
||||||
|
|
||||||
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
|
|
||||||
values = [w[2] for w in result["ch"]["writes"]]
|
|
||||||
assert values == [1, "from_input", "from_node"], (
|
|
||||||
f"Expected pathless writes first, got {values}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
||||||
test_history_returns_writes_oldest_first,
|
test_history_returns_writes_oldest_first,
|
||||||
test_history_seed_is_nearest_snapshot,
|
test_history_seed_is_nearest_snapshot,
|
||||||
@@ -331,8 +276,6 @@ ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
|||||||
test_history_walk_to_root_no_seed,
|
test_history_walk_to_root_no_seed,
|
||||||
test_history_migration_plain_value_as_seed,
|
test_history_migration_plain_value_as_seed,
|
||||||
test_history_seed_ancestor_own_writes_are_replayed,
|
test_history_seed_ancestor_own_writes_are_replayed,
|
||||||
test_history_orders_parallel_writes_by_task_path,
|
|
||||||
test_history_orders_pathless_writes_first,
|
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -448,11 +448,12 @@ class PostgresSaver(BasePostgresSaver):
|
|||||||
|
|
||||||
Two-stage query, both stages cover ALL requested channels:
|
Two-stage query, both stages cover ALL requested channels:
|
||||||
|
|
||||||
* Stage 1 (paged): dynamic SELECT over `checkpoints` with K parallel
|
* Stage 1 (paged): dynamic SELECT over `checkpoints` with three
|
||||||
JSONB key lookups (one column pair per channel) — no subquery, no
|
columns per channel: its version, an `EXISTS` probe for a stored
|
||||||
aggregation. Pages newest-first by `checkpoint_id` with a cursor;
|
blob at that version, and its inline value. Pages newest-first by
|
||||||
page size is `_DELTA_PAGE_SIZE`. Stops paging when every channel
|
`checkpoint_id` with a cursor; page size is `_DELTA_PAGE_SIZE`.
|
||||||
has found its seed or the chain is exhausted.
|
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
|
* Stage 2 (per-channel UNION ALL): one branch per channel reading
|
||||||
`checkpoint_writes` filtered to that channel's specific
|
`checkpoint_writes` filtered to that channel's specific
|
||||||
|
|||||||
@@ -168,35 +168,12 @@ class _DeltaStage2Row(TypedDict, total=False):
|
|||||||
type: str | None
|
type: str | None
|
||||||
blob: bytes | None
|
blob: bytes | None
|
||||||
task_id: str | None # "w" rows only
|
task_id: str | None # "w" rows only
|
||||||
task_path: str | None # "w" rows only
|
|
||||||
idx: int | None # "w" rows only
|
idx: int | None # "w" rows only
|
||||||
version: str | None # "b" rows only
|
version: str | None # "b" rows only
|
||||||
|
|
||||||
|
|
||||||
# Multi-channel two-stage DeltaChannel reconstruction.
|
# Delta history is rebuilt in two queries; `_build_delta_stage1_sql` and
|
||||||
#
|
# `_build_delta_stage2_sql` document their shapes.
|
||||||
# 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).
|
|
||||||
|
|
||||||
|
|
||||||
def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
|
def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
|
||||||
@@ -320,7 +297,7 @@ def _build_delta_stage2_sql(
|
|||||||
branches.append(
|
branches.append(
|
||||||
"SELECT 'w'::text AS _kind, "
|
"SELECT 'w'::text AS _kind, "
|
||||||
"checkpoint_id, channel, "
|
"checkpoint_id, channel, "
|
||||||
"type, blob, task_id, task_path, idx, NULL::text AS version "
|
"type, blob, task_id, idx, NULL::text AS version "
|
||||||
"FROM checkpoint_writes "
|
"FROM checkpoint_writes "
|
||||||
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
||||||
"AND checkpoint_id = ANY(%s)"
|
"AND checkpoint_id = ANY(%s)"
|
||||||
@@ -328,8 +305,7 @@ def _build_delta_stage2_sql(
|
|||||||
for _ in channels_with_seed:
|
for _ in channels_with_seed:
|
||||||
branches.append(
|
branches.append(
|
||||||
"SELECT 'b'::text AS _kind, NULL::text AS checkpoint_id, channel, "
|
"SELECT 'b'::text AS _kind, NULL::text AS checkpoint_id, channel, "
|
||||||
"type, blob, NULL::text AS task_id, NULL::text AS task_path, "
|
"type, blob, NULL::text AS task_id, NULL::int AS idx, version "
|
||||||
"NULL::int AS idx, version "
|
|
||||||
"FROM checkpoint_blobs "
|
"FROM checkpoint_blobs "
|
||||||
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
||||||
"AND version = %s"
|
"AND version = %s"
|
||||||
@@ -337,10 +313,8 @@ def _build_delta_stage2_sql(
|
|||||||
return " UNION ALL ".join(branches)
|
return " UNION ALL ".join(branches)
|
||||||
|
|
||||||
|
|
||||||
# Stage 1 rows are dynamic-shape dicts: {checkpoint_id, parent_checkpoint_id,
|
# Stage 1 rows are dicts keyed by the per-channel aliases
|
||||||
# ver_0, hs_0, ver_1, hs_1, ...}. Walking is parameterized by the channel
|
# `_build_delta_stage1_sql` emits, so there is no static TypedDict.
|
||||||
# list to map indices back to channel names — no static TypedDict here.
|
|
||||||
# `dict[str, Any]` is the practical signature.
|
|
||||||
|
|
||||||
|
|
||||||
class BasePostgresSaver(BaseCheckpointSaver[str]):
|
class BasePostgresSaver(BaseCheckpointSaver[str]):
|
||||||
@@ -433,9 +407,11 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
(a) it found a stored value for its channel — a blob or an inline
|
(a) it found a stored value for its channel — a blob or an inline
|
||||||
primitive (channel becomes seeded),
|
primitive (channel becomes seeded),
|
||||||
(b) it reached a real root (parent_of[cid] is None — fully
|
(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
|
(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`,
|
Mutates `chain_by_ch`, `seed_ver_by_ch`, `seed_inline_by_ch`,
|
||||||
`walk_cursor_by_ch`, and `seeded` in place.
|
`walk_cursor_by_ch`, and `seeded` in place.
|
||||||
@@ -443,9 +419,12 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
for i, ch in enumerate(channels):
|
for i, ch in enumerate(channels):
|
||||||
if ch in seeded:
|
if ch in seeded:
|
||||||
continue
|
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:
|
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]
|
cur_cid = walk_cursor_by_ch[ch]
|
||||||
ch_chain = chain_by_ch[ch]
|
ch_chain = chain_by_ch[ch]
|
||||||
hb_i = hb_by_i_by_cid[i]
|
hb_i = hb_by_i_by_cid[i]
|
||||||
@@ -494,11 +473,10 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
stored value, or when the seed blob is sentinel "empty" — in both cases
|
stored value, or when the seed blob is sentinel "empty" — in both cases
|
||||||
the consumer treats absence as "start empty".
|
the consumer treats absence as "start empty".
|
||||||
"""
|
"""
|
||||||
# writes_by_ch_by_cid[channel][cid] = list of
|
# writes_by_ch_by_cid[channel][cid] = list of (type, blob, task_id, idx)
|
||||||
# (type, blob, task_id, idx, task_path)
|
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
|
||||||
writes_by_ch_by_cid: dict[
|
ch: {} for ch in channels
|
||||||
str, dict[str, list[tuple[str, bytes, str, int, str]]]
|
}
|
||||||
] = {ch: {} for ch in channels}
|
|
||||||
# seed_blob_by_ver[(channel, version)] = (type, blob)
|
# seed_blob_by_ver[(channel, version)] = (type, blob)
|
||||||
seed_blob_by_ver: dict[tuple[str, str], tuple[str, bytes]] = {}
|
seed_blob_by_ver: dict[tuple[str, str], tuple[str, bytes]] = {}
|
||||||
|
|
||||||
@@ -509,14 +487,8 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
cid = cast(str, r["checkpoint_id"])
|
cid = cast(str, r["checkpoint_id"])
|
||||||
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
||||||
cast(
|
cast(
|
||||||
"tuple[str, bytes, str, int, str]",
|
"tuple[str, bytes, str, int]",
|
||||||
(
|
(r["type"], r["blob"], r["task_id"], r["idx"]),
|
||||||
r["type"],
|
|
||||||
r["blob"],
|
|
||||||
r["task_id"],
|
|
||||||
r["idx"],
|
|
||||||
r["task_path"],
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
else: # kind == "b"
|
else: # kind == "b"
|
||||||
@@ -525,10 +497,10 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
"tuple[str, bytes]", (r["type"], r["blob"])
|
"tuple[str, bytes]", (r["type"], r["blob"])
|
||||||
)
|
)
|
||||||
|
|
||||||
# Sort writes per (channel, cid) newest-first by (task_path, task_id, idx)
|
# Sort writes per (channel, cid) newest-first by (task_id, idx)
|
||||||
for cid_map in writes_by_ch_by_cid.values():
|
for cid_map in writes_by_ch_by_cid.values():
|
||||||
for ws in cid_map.values():
|
for ws in cid_map.values():
|
||||||
ws.sort(key=lambda w: (w[4], w[2], w[3]), reverse=True)
|
ws.sort(key=lambda w: (w[2], w[3]), reverse=True)
|
||||||
|
|
||||||
result: dict[str, DeltaChannelHistory] = {}
|
result: dict[str, DeltaChannelHistory] = {}
|
||||||
for ch in channels:
|
for ch in channels:
|
||||||
@@ -538,9 +510,7 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|||||||
collected: list[PendingWrite] = []
|
collected: list[PendingWrite] = []
|
||||||
cid_writes = writes_by_ch_by_cid.get(ch, {})
|
cid_writes = writes_by_ch_by_cid.get(ch, {})
|
||||||
for cid in chain_cids:
|
for cid in chain_cids:
|
||||||
for type_tag, write_blob, task_id, _idx, _path in cid_writes.get(
|
for type_tag, write_blob, task_id, _idx in cid_writes.get(cid, []):
|
||||||
cid, []
|
|
||||||
):
|
|
||||||
val = self.serde.loads_typed((type_tag, write_blob))
|
val = self.serde.loads_typed((type_tag, write_blob))
|
||||||
collected.append((task_id, ch, val))
|
collected.append((task_id, ch, val))
|
||||||
collected.reverse()
|
collected.reverse()
|
||||||
|
|||||||
@@ -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": []}
|
||||||
@@ -81,7 +81,6 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
|
|
||||||
conn: sqlite3.Connection
|
conn: sqlite3.Connection
|
||||||
is_setup: bool
|
is_setup: bool
|
||||||
_has_task_path: bool = True
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -155,7 +154,6 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||||
checkpoint_id TEXT NOT NULL,
|
checkpoint_id TEXT NOT NULL,
|
||||||
task_id TEXT NOT NULL,
|
task_id TEXT NOT NULL,
|
||||||
task_path TEXT NOT NULL DEFAULT '',
|
|
||||||
idx INTEGER NOT NULL,
|
idx INTEGER NOT NULL,
|
||||||
channel TEXT NOT NULL,
|
channel TEXT NOT NULL,
|
||||||
type TEXT,
|
type TEXT,
|
||||||
@@ -164,19 +162,6 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
);
|
);
|
||||||
"""
|
"""
|
||||||
)
|
)
|
||||||
# sqlite has no ADD COLUMN IF NOT EXISTS; this migrates databases
|
|
||||||
# created before `task_path` existed and is a no-op on the rest.
|
|
||||||
try:
|
|
||||||
self.conn.execute(
|
|
||||||
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
|
|
||||||
)
|
|
||||||
except sqlite3.OperationalError as e:
|
|
||||||
# A read-only database from before the column can still be read;
|
|
||||||
# its rows would all read back as '' anyway.
|
|
||||||
if "readonly database" in str(e):
|
|
||||||
self._has_task_path = False
|
|
||||||
elif "duplicate column name" not in str(e):
|
|
||||||
raise
|
|
||||||
|
|
||||||
self.is_setup = True
|
self.is_setup = True
|
||||||
|
|
||||||
@@ -475,9 +460,9 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
task_path: Path of the task creating the writes.
|
task_path: Path of the task creating the writes.
|
||||||
"""
|
"""
|
||||||
query = (
|
query = (
|
||||||
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||||
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
||||||
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||||
)
|
)
|
||||||
with self.cursor() as cur:
|
with self.cursor() as cur:
|
||||||
cur.executemany(
|
cur.executemany(
|
||||||
@@ -488,7 +473,6 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
str(config["configurable"]["checkpoint_ns"]),
|
str(config["configurable"]["checkpoint_ns"]),
|
||||||
str(config["configurable"]["checkpoint_id"]),
|
str(config["configurable"]["checkpoint_id"]),
|
||||||
task_id,
|
task_id,
|
||||||
task_path,
|
|
||||||
WRITES_IDX_MAP.get(channel, idx),
|
WRITES_IDX_MAP.get(channel, idx),
|
||||||
channel,
|
channel,
|
||||||
*self.serde.dumps_typed(value),
|
*self.serde.dumps_typed(value),
|
||||||
@@ -523,13 +507,12 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
|
|
||||||
Two-stage query:
|
Two-stage query:
|
||||||
|
|
||||||
* Stage 1 (paged): newest-first slice of `checkpoints` returning
|
* Stage 1 (streamed): recursive CTE over `checkpoints` following
|
||||||
`(checkpoint_id, parent_checkpoint_id, type, checkpoint)` per
|
`parent_checkpoint_id` from the target, returning
|
||||||
ancestor. Sqlite has no JSONB, so we ship the full serialized
|
`(checkpoint_id, type, checkpoint)` per ancestor. Sqlite has no
|
||||||
checkpoint blob and inspect `channel_values` in Python. Pages
|
JSONB, so we ship the full serialized checkpoint blob and inspect
|
||||||
newest-first by `checkpoint_id` with a `< cursor` predicate;
|
`channel_values` in Python. Stops reading when every channel has
|
||||||
page size is `DELTA_PAGE_SIZE`. Stops paging when every channel
|
found its seed or the chain is exhausted.
|
||||||
has found its seed or the chain is exhausted.
|
|
||||||
|
|
||||||
* Stage 2 (per-channel UNION ALL): one branch per channel reading
|
* Stage 2 (per-channel UNION ALL): one branch per channel reading
|
||||||
`writes` filtered to that channel's specific `chain_cids`. No
|
`writes` filtered to that channel's specific `chain_cids`. No
|
||||||
@@ -554,12 +537,14 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
seeded: set[str] = set()
|
seeded: set[str] = set()
|
||||||
|
|
||||||
with self.cursor(transaction=False) as cur:
|
with self.cursor(transaction=False) as cur:
|
||||||
cur.execute(DELTA_STAGE1_SQL, (thread_id, checkpoint_ns, checkpoint_id))
|
cur.execute(
|
||||||
|
DELTA_STAGE1_SQL,
|
||||||
|
(thread_id, checkpoint_ns, checkpoint_id, thread_id, checkpoint_ns),
|
||||||
|
)
|
||||||
for row in cur:
|
for row in cur:
|
||||||
cid, parent_cid, type_tag, blob = row
|
cid, type_tag, blob = row
|
||||||
if step_walk_with_row(
|
if step_walk_with_row(
|
||||||
cid=cid,
|
cid=cid,
|
||||||
parent_cid=parent_cid,
|
|
||||||
type_tag=type_tag,
|
type_tag=type_tag,
|
||||||
blob=blob,
|
blob=blob,
|
||||||
target_id=checkpoint_id,
|
target_id=checkpoint_id,
|
||||||
@@ -574,7 +559,6 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
|
|
||||||
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
||||||
stage2_sql = build_delta_stage2_sql(
|
stage2_sql = build_delta_stage2_sql(
|
||||||
has_task_path=self._has_task_path,
|
|
||||||
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
|
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
|
||||||
)
|
)
|
||||||
if stage2_sql:
|
if stage2_sql:
|
||||||
@@ -585,7 +569,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
)
|
)
|
||||||
cur.execute(stage2_sql, stage2_params)
|
cur.execute(stage2_sql, stage2_params)
|
||||||
stage2_rows = cast(
|
stage2_rows = cast(
|
||||||
"list[tuple[str, str, str, int, str, bytes, str]]", cur.fetchall()
|
"list[tuple[str, str, str, int, str, bytes]]", cur.fetchall()
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
stage2_rows = []
|
stage2_rows = []
|
||||||
|
|||||||
@@ -26,22 +26,37 @@ from typing import Any
|
|||||||
|
|
||||||
from langgraph.checkpoint.base import DeltaChannelHistory, PendingWrite
|
from langgraph.checkpoint.base import DeltaChannelHistory, PendingWrite
|
||||||
|
|
||||||
# Stage 1 streams ancestors of `target_cid` newest-first. The `<=`
|
# Stage 1 streams target, then its ancestors nearest-first, by following
|
||||||
# predicate keeps target itself in the stream so we can read its
|
# `parent_checkpoint_id` rather than id order: ids are only monotonic within
|
||||||
# `parent_checkpoint_id` from the first row without a separate lookup;
|
# one process, so a range scan by id can miss a parent whose id sorts above
|
||||||
# the caller skips target's own writes/seed (matches the
|
# its child's. Target is the anchor row; its own writes/seed are skipped
|
||||||
# `BaseCheckpointSaver` contract).
|
# (matches the `BaseCheckpointSaver` contract).
|
||||||
|
#
|
||||||
|
# `put` is `INSERT OR REPLACE`, so re-putting an existing id under a
|
||||||
|
# descendant's config makes the chain a loop. `step_walk_with_row` stops on a
|
||||||
|
# repeated id; sqlite yields recursive rows lazily, so abandoning the cursor
|
||||||
|
# ends the recursion.
|
||||||
|
#
|
||||||
|
# `CROSS JOIN` pins `ancestors` as the outer loop, so each step is one primary
|
||||||
|
# key lookup. With a plain `JOIN` and no `ANALYZE` stats, sqlite can put
|
||||||
|
# `checkpoints` outside and scan the whole thread per step.
|
||||||
DELTA_STAGE1_SQL = (
|
DELTA_STAGE1_SQL = (
|
||||||
|
"WITH RECURSIVE ancestors(checkpoint_id, parent_checkpoint_id, type, "
|
||||||
|
"checkpoint) AS ("
|
||||||
"SELECT checkpoint_id, parent_checkpoint_id, type, checkpoint "
|
"SELECT checkpoint_id, parent_checkpoint_id, type, checkpoint "
|
||||||
"FROM checkpoints "
|
"FROM checkpoints "
|
||||||
"WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id <= ? "
|
"WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? "
|
||||||
"ORDER BY checkpoint_id DESC"
|
"UNION ALL "
|
||||||
|
"SELECT c.checkpoint_id, c.parent_checkpoint_id, c.type, c.checkpoint "
|
||||||
|
"FROM ancestors a CROSS JOIN checkpoints c "
|
||||||
|
"ON c.checkpoint_id = a.parent_checkpoint_id "
|
||||||
|
"WHERE c.thread_id = ? AND c.checkpoint_ns = ?"
|
||||||
|
") "
|
||||||
|
"SELECT checkpoint_id, type, checkpoint FROM ancestors"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_delta_stage2_sql(
|
def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
|
||||||
*, chain_lens: Sequence[int], has_task_path: bool = True
|
|
||||||
) -> str:
|
|
||||||
"""Stage-2 per-channel UNION ALL fetching writes from `writes`.
|
"""Stage-2 per-channel UNION ALL fetching writes from `writes`.
|
||||||
|
|
||||||
One branch per channel with a non-empty chain. Each branch inlines its
|
One branch per channel with a non-empty chain. Each branch inlines its
|
||||||
@@ -55,12 +70,11 @@ def build_delta_stage2_sql(
|
|||||||
of a single `channel = ANY(channels)` filter when channels have
|
of a single `channel = ANY(channels)` filter when channels have
|
||||||
different chain depths — same rationale as postgres.
|
different chain depths — same rationale as postgres.
|
||||||
"""
|
"""
|
||||||
task_path = "task_path" if has_task_path else "''"
|
|
||||||
branches: list[str] = []
|
branches: list[str] = []
|
||||||
for n in chain_lens:
|
for n in chain_lens:
|
||||||
cid_placeholders = ",".join("?" * n)
|
cid_placeholders = ",".join("?" * n)
|
||||||
branches.append(
|
branches.append(
|
||||||
f"SELECT checkpoint_id, channel, task_id, idx, type, value, {task_path} "
|
"SELECT checkpoint_id, channel, task_id, idx, type, value "
|
||||||
"FROM writes "
|
"FROM writes "
|
||||||
"WHERE thread_id = ? AND checkpoint_ns = ? AND channel = ? "
|
"WHERE thread_id = ? AND checkpoint_ns = ? AND channel = ? "
|
||||||
f"AND checkpoint_id IN ({cid_placeholders})"
|
f"AND checkpoint_id IN ({cid_placeholders})"
|
||||||
@@ -71,7 +85,6 @@ def build_delta_stage2_sql(
|
|||||||
def step_walk_with_row(
|
def step_walk_with_row(
|
||||||
*,
|
*,
|
||||||
cid: str,
|
cid: str,
|
||||||
parent_cid: str | None,
|
|
||||||
type_tag: str,
|
type_tag: str,
|
||||||
blob: bytes,
|
blob: bytes,
|
||||||
target_id: str,
|
target_id: str,
|
||||||
@@ -84,36 +97,32 @@ def step_walk_with_row(
|
|||||||
) -> bool:
|
) -> bool:
|
||||||
"""Process one streamed stage-1 row in the merged ancestor walk.
|
"""Process one streamed stage-1 row in the merged ancestor walk.
|
||||||
|
|
||||||
The cursor returns (cid, parent_cid, type, blob) rows in
|
The cursor returns (cid, type, blob) rows in walk order starting at
|
||||||
`checkpoint_id` DESC order starting at target. The first row is
|
target. The first row is target itself and is skipped (target's own
|
||||||
target itself; we read its parent_cid to seed the walk and otherwise
|
writes/seed are not part of the contract).
|
||||||
skip it (target's own writes/seed are not part of the contract).
|
|
||||||
|
|
||||||
For each subsequent row, if `cid` matches the walk's current
|
For each subsequent row we deserialize the blob, append the cid to
|
||||||
position, we deserialize the blob, append the cid to every
|
every not-yet-seeded channel's chain, and check `channel_values` for
|
||||||
not-yet-seeded channel's chain, and check `channel_values` for
|
|
||||||
seeds. The deserialized checkpoint is dropped before advancing — no
|
seeds. The deserialized checkpoint is dropped before advancing — no
|
||||||
cross-row cache, so peak in-flight is one deserialized checkpoint.
|
cross-row cache, so peak in-flight is one deserialized checkpoint.
|
||||||
|
|
||||||
Off-path rows (different branch on the same thread) advance the
|
Returns True when the caller can stop iterating and close the cursor:
|
||||||
cursor without doing any work.
|
every requested channel is seeded, or the chain revisited a checkpoint.
|
||||||
|
|
||||||
Returns True when every requested channel is seeded — the caller
|
|
||||||
can stop iterating and close the cursor.
|
|
||||||
"""
|
"""
|
||||||
if "started" not in walk_state:
|
if "started" not in walk_state:
|
||||||
if cid == target_id:
|
if cid == target_id:
|
||||||
walk_state["started"] = True
|
walk_state["started"] = True
|
||||||
walk_state["cur_cid"] = parent_cid
|
|
||||||
walk_state["active"] = {ch for ch in channels if ch not in seeded}
|
walk_state["active"] = {ch for ch in channels if ch not in seeded}
|
||||||
|
walk_state["walked"] = {cid}
|
||||||
# Not target yet (or target not present): keep streaming.
|
# Not target yet (or target not present): keep streaming.
|
||||||
return False
|
return False
|
||||||
active: set[str] = walk_state["active"]
|
active: set[str] = walk_state["active"]
|
||||||
if not active:
|
if not active:
|
||||||
return True
|
return True
|
||||||
if cid != walk_state["cur_cid"]:
|
walked: set[str] = walk_state["walked"]
|
||||||
# Off-path row from a sibling branch — skip without deserializing.
|
if cid in walked:
|
||||||
return False
|
return True
|
||||||
|
walked.add(cid)
|
||||||
for ch in active:
|
for ch in active:
|
||||||
chain_by_ch[ch].append(cid)
|
chain_by_ch[ch].append(cid)
|
||||||
ckpt = serde.loads_typed((type_tag, blob))
|
ckpt = serde.loads_typed((type_tag, blob))
|
||||||
@@ -123,7 +132,6 @@ def step_walk_with_row(
|
|||||||
seeded.add(ch)
|
seeded.add(ch)
|
||||||
active.discard(ch)
|
active.discard(ch)
|
||||||
del ckpt, channel_values
|
del ckpt, channel_values
|
||||||
walk_state["cur_cid"] = parent_cid
|
|
||||||
return not active
|
return not active
|
||||||
|
|
||||||
|
|
||||||
@@ -133,31 +141,29 @@ def build_delta_channels_writes_history(
|
|||||||
chain_by_ch: Mapping[str, list[str]],
|
chain_by_ch: Mapping[str, list[str]],
|
||||||
seed_val_by_ch: Mapping[str, Any],
|
seed_val_by_ch: Mapping[str, Any],
|
||||||
seeded: set[str],
|
seeded: set[str],
|
||||||
stage2_rows: Sequence[tuple[str, str, str, int, str, bytes, str]],
|
stage2_rows: Sequence[tuple[str, str, str, int, str, bytes]],
|
||||||
serde: Any,
|
serde: Any,
|
||||||
) -> dict[str, DeltaChannelHistory]:
|
) -> dict[str, DeltaChannelHistory]:
|
||||||
"""Demux stage-2 rows per channel; produce per-channel histories.
|
"""Demux stage-2 rows per channel; produce per-channel histories.
|
||||||
|
|
||||||
Stage-2 rows are
|
Stage-2 rows are `(checkpoint_id, channel, task_id, idx, type, value)`.
|
||||||
`(checkpoint_id, channel, task_id, idx, type, value, task_path)`.
|
Final write order is oldest→newest globally and `(task_id, idx)` within
|
||||||
Final write order is oldest→newest globally and
|
a checkpoint, matching the contract on `DeltaChannelHistory.writes`.
|
||||||
`(task_path, task_id, idx)` within a checkpoint, matching the contract
|
|
||||||
on `DeltaChannelHistory.writes`.
|
|
||||||
|
|
||||||
`seed` is omitted when the walk reached a true root with no snapshot
|
`seed` is omitted when the walk reached a true root with no snapshot
|
||||||
found (channel never entered `seeded`); consumers treat absence as
|
found (channel never entered `seeded`); consumers treat absence as
|
||||||
"start empty".
|
"start empty".
|
||||||
"""
|
"""
|
||||||
writes_by_ch_by_cid: dict[
|
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
|
||||||
str, dict[str, list[tuple[str, bytes, str, int, str]]]
|
ch: {} for ch in channels
|
||||||
] = {ch: {} for ch in channels}
|
}
|
||||||
for cid, ch, task_id, idx, type_tag, value_blob, task_path in stage2_rows:
|
for cid, ch, task_id, idx, type_tag, value_blob in stage2_rows:
|
||||||
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
||||||
(type_tag, value_blob, task_id, idx, task_path)
|
(type_tag, value_blob, task_id, idx)
|
||||||
)
|
)
|
||||||
for cid_map in writes_by_ch_by_cid.values():
|
for cid_map in writes_by_ch_by_cid.values():
|
||||||
for ws in cid_map.values():
|
for ws in cid_map.values():
|
||||||
ws.sort(key=lambda w: (w[4], w[2], w[3]))
|
ws.sort(key=lambda w: (w[2], w[3]))
|
||||||
|
|
||||||
result: dict[str, DeltaChannelHistory] = {}
|
result: dict[str, DeltaChannelHistory] = {}
|
||||||
for ch in channels:
|
for ch in channels:
|
||||||
@@ -166,7 +172,7 @@ def build_delta_channels_writes_history(
|
|||||||
collected: list[PendingWrite] = []
|
collected: list[PendingWrite] = []
|
||||||
# Chain is newest-first; iterate oldest-first for the public order.
|
# Chain is newest-first; iterate oldest-first for the public order.
|
||||||
for cid in reversed(chain_cids):
|
for cid in reversed(chain_cids):
|
||||||
for type_tag, value_blob, task_id, _idx, _path in cid_writes.get(cid, []):
|
for type_tag, value_blob, task_id, _idx in cid_writes.get(cid, []):
|
||||||
collected.append(
|
collected.append(
|
||||||
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
|
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -114,7 +114,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
|
|
||||||
lock: asyncio.Lock
|
lock: asyncio.Lock
|
||||||
is_setup: bool
|
is_setup: bool
|
||||||
_has_task_path: bool = True
|
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -332,7 +331,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||||
checkpoint_id TEXT NOT NULL,
|
checkpoint_id TEXT NOT NULL,
|
||||||
task_id TEXT NOT NULL,
|
task_id TEXT NOT NULL,
|
||||||
task_path TEXT NOT NULL DEFAULT '',
|
|
||||||
idx INTEGER NOT NULL,
|
idx INTEGER NOT NULL,
|
||||||
channel TEXT NOT NULL,
|
channel TEXT NOT NULL,
|
||||||
type TEXT,
|
type TEXT,
|
||||||
@@ -343,21 +341,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
):
|
):
|
||||||
await self.conn.commit()
|
await self.conn.commit()
|
||||||
|
|
||||||
# sqlite has no ADD COLUMN IF NOT EXISTS; this migrates databases
|
|
||||||
# created before `task_path` existed and is a no-op on the rest.
|
|
||||||
try:
|
|
||||||
await self.conn.execute(
|
|
||||||
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
|
|
||||||
)
|
|
||||||
await self.conn.commit()
|
|
||||||
except aiosqlite.OperationalError as e:
|
|
||||||
# A read-only database from before the column can still be read;
|
|
||||||
# its rows would all read back as '' anyway.
|
|
||||||
if "readonly database" in str(e):
|
|
||||||
self._has_task_path = False
|
|
||||||
elif "duplicate column name" not in str(e):
|
|
||||||
raise
|
|
||||||
|
|
||||||
self.is_setup = True
|
self.is_setup = True
|
||||||
|
|
||||||
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
||||||
@@ -593,9 +576,9 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
task_path: Path of the task creating the writes.
|
task_path: Path of the task creating the writes.
|
||||||
"""
|
"""
|
||||||
query = (
|
query = (
|
||||||
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||||
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
||||||
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||||
)
|
)
|
||||||
await self.setup()
|
await self.setup()
|
||||||
async with self.lock, self.conn.cursor() as cur:
|
async with self.lock, self.conn.cursor() as cur:
|
||||||
@@ -607,7 +590,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
str(config["configurable"]["checkpoint_ns"]),
|
str(config["configurable"]["checkpoint_ns"]),
|
||||||
str(config["configurable"]["checkpoint_id"]),
|
str(config["configurable"]["checkpoint_id"]),
|
||||||
task_id,
|
task_id,
|
||||||
task_path,
|
|
||||||
WRITES_IDX_MAP.get(channel, idx),
|
WRITES_IDX_MAP.get(channel, idx),
|
||||||
channel,
|
channel,
|
||||||
*self.serde.dumps_typed(value),
|
*self.serde.dumps_typed(value),
|
||||||
@@ -643,8 +625,8 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
"""Fast-path override of `BaseCheckpointSaver.aget_delta_channel_history`.
|
"""Fast-path override of `BaseCheckpointSaver.aget_delta_channel_history`.
|
||||||
|
|
||||||
See `SqliteSaver.get_delta_channel_history` for design notes; this
|
See `SqliteSaver.get_delta_channel_history` for design notes; this
|
||||||
is the async equivalent using `aiosqlite` cursors. Stage 1 pages
|
is the async equivalent using `aiosqlite` cursors. Stage 1 streams
|
||||||
the parent chain newest-first and Python-deserializes each
|
the parent chain from the target and Python-deserializes each
|
||||||
checkpoint blob to find per-channel snapshots; stage 2 fetches
|
checkpoint blob to find per-channel snapshots; stage 2 fetches
|
||||||
only the relevant writes via per-channel UNION ALL.
|
only the relevant writes via per-channel UNION ALL.
|
||||||
"""
|
"""
|
||||||
@@ -668,13 +650,13 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
|
|
||||||
async with self.lock, self.conn.cursor() as cur:
|
async with self.lock, self.conn.cursor() as cur:
|
||||||
await cur.execute(
|
await cur.execute(
|
||||||
DELTA_STAGE1_SQL, (thread_id, checkpoint_ns, checkpoint_id)
|
DELTA_STAGE1_SQL,
|
||||||
|
(thread_id, checkpoint_ns, checkpoint_id, thread_id, checkpoint_ns),
|
||||||
)
|
)
|
||||||
async for row in cur:
|
async for row in cur:
|
||||||
cid, parent_cid, type_tag, blob = row
|
cid, type_tag, blob = row
|
||||||
if step_walk_with_row(
|
if step_walk_with_row(
|
||||||
cid=cid,
|
cid=cid,
|
||||||
parent_cid=parent_cid,
|
|
||||||
type_tag=type_tag,
|
type_tag=type_tag,
|
||||||
blob=blob,
|
blob=blob,
|
||||||
target_id=checkpoint_id,
|
target_id=checkpoint_id,
|
||||||
@@ -689,7 +671,6 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
|
|
||||||
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
||||||
stage2_sql = build_delta_stage2_sql(
|
stage2_sql = build_delta_stage2_sql(
|
||||||
has_task_path=self._has_task_path,
|
|
||||||
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
|
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
|
||||||
)
|
)
|
||||||
if stage2_sql:
|
if stage2_sql:
|
||||||
@@ -700,7 +681,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
)
|
)
|
||||||
await cur.execute(stage2_sql, stage2_params)
|
await cur.execute(stage2_sql, stage2_params)
|
||||||
stage2_rows = cast(
|
stage2_rows = cast(
|
||||||
"list[tuple[str, str, str, int, str, bytes, str]]",
|
"list[tuple[str, str, str, int, str, bytes]]",
|
||||||
await cur.fetchall(),
|
await cur.fetchall(),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -0,0 +1,104 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from langgraph.checkpoint.base import (
|
||||||
|
BaseCheckpointSaver,
|
||||||
|
Checkpoint,
|
||||||
|
DeltaChannelHistory,
|
||||||
|
empty_checkpoint,
|
||||||
|
)
|
||||||
|
|
||||||
|
from langgraph.checkpoint.sqlite import SqliteSaver
|
||||||
|
from langgraph.checkpoint.sqlite._delta import DELTA_STAGE1_SQL
|
||||||
|
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
||||||
|
|
||||||
|
CHANNEL = "ch"
|
||||||
|
CONFIG: dict[str, Any] = {"configurable": {"thread_id": "t", "checkpoint_ns": ""}}
|
||||||
|
EXPECTED: DeltaChannelHistory = {
|
||||||
|
"writes": [("task", CHANNEL, "write-root")],
|
||||||
|
"seed": "seed",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _checkpoint(checkpoint_id: str, values: dict[str, Any]) -> Checkpoint:
|
||||||
|
value = empty_checkpoint()
|
||||||
|
value["id"] = checkpoint_id
|
||||||
|
value["channel_values"] = values
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
PARENT_ID_ORDERS = [
|
||||||
|
pytest.param("z-older", "a-newer", id="parent_id_sorts_above_child"),
|
||||||
|
pytest.param("a-older", "z-newer", id="parent_id_sorts_below_child"),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(("root_id", "child_id"), PARENT_ID_ORDERS)
|
||||||
|
def test_sync_walk_reaches_parent_whatever_the_id_order(
|
||||||
|
root_id: str, child_id: str
|
||||||
|
) -> None:
|
||||||
|
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||||
|
root = saver.put(CONFIG, _checkpoint(root_id, {CHANNEL: "seed"}), {}, {})
|
||||||
|
saver.put_writes(root, [(CHANNEL, "write-root")], "task")
|
||||||
|
child = saver.put(root, _checkpoint(child_id, {}), {}, {})
|
||||||
|
|
||||||
|
got = saver.get_delta_channel_history(config=child, channels=[CHANNEL])
|
||||||
|
reference = BaseCheckpointSaver.get_delta_channel_history(
|
||||||
|
saver, config=child, channels=[CHANNEL]
|
||||||
|
)
|
||||||
|
assert got[CHANNEL] == EXPECTED
|
||||||
|
assert got[CHANNEL] == reference[CHANNEL], "fast path disagrees with base"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(("root_id", "child_id"), PARENT_ID_ORDERS)
|
||||||
|
async def test_async_walk_reaches_parent_whatever_the_id_order(
|
||||||
|
root_id: str, child_id: str
|
||||||
|
) -> None:
|
||||||
|
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
|
||||||
|
root = await saver.aput(CONFIG, _checkpoint(root_id, {CHANNEL: "seed"}), {}, {})
|
||||||
|
await saver.aput_writes(root, [(CHANNEL, "write-root")], "task")
|
||||||
|
child = await saver.aput(root, _checkpoint(child_id, {}), {}, {})
|
||||||
|
|
||||||
|
got = await saver.aget_delta_channel_history(config=child, channels=[CHANNEL])
|
||||||
|
assert got[CHANNEL] == EXPECTED
|
||||||
|
|
||||||
|
|
||||||
|
def test_walk_reaches_root_of_long_chain_with_descending_ids() -> None:
|
||||||
|
steps = 40
|
||||||
|
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||||
|
parent = saver.put(
|
||||||
|
CONFIG, _checkpoint(f"id-{steps:03d}", {CHANNEL: "seed"}), {}, {}
|
||||||
|
)
|
||||||
|
saver.put_writes(parent, [(CHANNEL, "write-root")], "task")
|
||||||
|
for step in range(steps - 1, 0, -1):
|
||||||
|
parent = saver.put(parent, _checkpoint(f"id-{step:03d}", {}), {}, {})
|
||||||
|
|
||||||
|
got = saver.get_delta_channel_history(config=parent, channels=[CHANNEL])
|
||||||
|
assert got[CHANNEL] == EXPECTED
|
||||||
|
|
||||||
|
|
||||||
|
def test_walk_terminates_when_put_makes_the_parent_chain_cycle() -> None:
|
||||||
|
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||||
|
a = saver.put(CONFIG, _checkpoint("cid-a", {}), {}, {})
|
||||||
|
b = saver.put(a, _checkpoint("cid-b", {}), {}, {})
|
||||||
|
repoint_a_under_b = _checkpoint("cid-a", {})
|
||||||
|
saver.put(b, repoint_a_under_b, {}, {})
|
||||||
|
|
||||||
|
got = saver.get_delta_channel_history(config=b, channels=[CHANNEL])
|
||||||
|
assert got[CHANNEL] == {"writes": []}
|
||||||
|
|
||||||
|
|
||||||
|
def test_walk_step_looks_up_the_parent_by_primary_key() -> None:
|
||||||
|
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||||
|
saver.setup()
|
||||||
|
plan = [
|
||||||
|
row[3]
|
||||||
|
for row in saver.conn.execute(
|
||||||
|
f"EXPLAIN QUERY PLAN {DELTA_STAGE1_SQL}", ("t", "", "id", "t", "")
|
||||||
|
)
|
||||||
|
]
|
||||||
|
assert any(
|
||||||
|
step.startswith("SEARCH c ") and "checkpoint_id=?" in step for step in plan
|
||||||
|
), f"recursive step should look up the parent by key, got {plan}"
|
||||||
@@ -1,128 +0,0 @@
|
|||||||
import sqlite3
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
import aiosqlite
|
|
||||||
import pytest
|
|
||||||
from langgraph.checkpoint.base import empty_checkpoint
|
|
||||||
|
|
||||||
from langgraph.checkpoint.sqlite import SqliteSaver
|
|
||||||
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
|
||||||
|
|
||||||
WRITES_BEFORE_TASK_PATH = """
|
|
||||||
CREATE TABLE 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,
|
|
||||||
value BLOB,
|
|
||||||
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
|
|
||||||
);
|
|
||||||
INSERT INTO writes VALUES ('t', '', 'c', 'old-task', 0, 'ch', 'null', X'');
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def legacy_db(tmp_path: Path) -> Path:
|
|
||||||
db = tmp_path / "legacy.sqlite"
|
|
||||||
with sqlite3.connect(db) as conn:
|
|
||||||
conn.executescript(WRITES_BEFORE_TASK_PATH)
|
|
||||||
return db
|
|
||||||
|
|
||||||
|
|
||||||
def test_setup_migrates_legacy_writes_table_repeatably(legacy_db: Path) -> None:
|
|
||||||
for _ in range(2):
|
|
||||||
with SqliteSaver.from_conn_string(str(legacy_db)) as saver:
|
|
||||||
saver.setup()
|
|
||||||
rows = saver.conn.execute(
|
|
||||||
"SELECT task_id, task_path FROM writes"
|
|
||||||
).fetchall()
|
|
||||||
assert rows == [("old-task", "")]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("fresh", [True, False], ids=["fresh", "legacy"])
|
|
||||||
def test_put_writes_persists_task_path(
|
|
||||||
tmp_path: Path, legacy_db: Path, fresh: bool
|
|
||||||
) -> None:
|
|
||||||
db = tmp_path / "fresh.sqlite" if fresh else legacy_db
|
|
||||||
with SqliteSaver.from_conn_string(str(db)) as saver:
|
|
||||||
config = saver.put(
|
|
||||||
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
|
|
||||||
empty_checkpoint(),
|
|
||||||
{},
|
|
||||||
{},
|
|
||||||
)
|
|
||||||
saver.put_writes(config, [("ch", "v")], "task-1", "~__pregel_pull, node")
|
|
||||||
stored = saver.conn.execute(
|
|
||||||
"SELECT task_path FROM writes WHERE task_id = 'task-1'"
|
|
||||||
).fetchall()
|
|
||||||
assert stored == [("~__pregel_pull, node",)]
|
|
||||||
|
|
||||||
|
|
||||||
async def test_async_setup_migrates_legacy_writes_table_repeatably(
|
|
||||||
legacy_db: Path,
|
|
||||||
) -> None:
|
|
||||||
for _ in range(2):
|
|
||||||
async with AsyncSqliteSaver.from_conn_string(str(legacy_db)) as saver:
|
|
||||||
await saver.setup()
|
|
||||||
config = await saver.aput(
|
|
||||||
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
|
|
||||||
empty_checkpoint(),
|
|
||||||
{},
|
|
||||||
{},
|
|
||||||
)
|
|
||||||
await saver.aput_writes(
|
|
||||||
config, [("ch", "v")], "task-1", "~__pregel_pull, node"
|
|
||||||
)
|
|
||||||
|
|
||||||
async with aiosqlite.connect(legacy_db) as conn:
|
|
||||||
async with conn.execute(
|
|
||||||
"SELECT DISTINCT task_id, task_path FROM writes ORDER BY task_id"
|
|
||||||
) as cur:
|
|
||||||
assert await cur.fetchall() == [
|
|
||||||
("old-task", ""),
|
|
||||||
("task-1", "~__pregel_pull, node"),
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
def _legacy_database_with_history(db: Path) -> dict:
|
|
||||||
root = empty_checkpoint()
|
|
||||||
root["channel_values"] = {"ch": "seed"}
|
|
||||||
root["channel_versions"] = {"ch": 1}
|
|
||||||
with SqliteSaver.from_conn_string(str(db)) as saver:
|
|
||||||
root_config = saver.put(
|
|
||||||
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
|
|
||||||
root,
|
|
||||||
{},
|
|
||||||
{"ch": 1},
|
|
||||||
)
|
|
||||||
saver.put_writes(root_config, [("ch", "write")], "task", "~__pregel_pull, n")
|
|
||||||
child = saver.put(root_config, empty_checkpoint(), {}, {})
|
|
||||||
saver.conn.execute("ALTER TABLE writes DROP COLUMN task_path")
|
|
||||||
saver.conn.commit()
|
|
||||||
return child
|
|
||||||
|
|
||||||
|
|
||||||
def test_read_only_legacy_database_still_reads_delta_history(tmp_path: Path) -> None:
|
|
||||||
db = tmp_path / "legacy.sqlite"
|
|
||||||
child = _legacy_database_with_history(db)
|
|
||||||
|
|
||||||
saver = SqliteSaver(sqlite3.connect(f"file:{db}?mode=ro", uri=True))
|
|
||||||
got = saver.get_delta_channel_history(config=child, channels=["ch"])
|
|
||||||
|
|
||||||
assert got["ch"] == {"seed": "seed", "writes": [("task", "ch", "write")]}
|
|
||||||
|
|
||||||
|
|
||||||
async def test_async_read_only_legacy_database_still_reads_delta_history(
|
|
||||||
tmp_path: Path,
|
|
||||||
) -> None:
|
|
||||||
db = tmp_path / "legacy.sqlite"
|
|
||||||
child = _legacy_database_with_history(db)
|
|
||||||
|
|
||||||
async with aiosqlite.connect(f"file:{db}?mode=ro", uri=True) as conn:
|
|
||||||
saver = AsyncSqliteSaver(conn)
|
|
||||||
got = await saver.aget_delta_channel_history(config=child, channels=["ch"])
|
|
||||||
|
|
||||||
assert got["ch"] == {"seed": "seed", "writes": [("task", "ch", "write")]}
|
|
||||||
@@ -162,14 +162,6 @@ class DeltaChannelHistory(TypedDict):
|
|||||||
Always present; possibly empty. Already filtered to one channel.
|
Always present; possibly empty. Already filtered to one channel.
|
||||||
Writes stored at the target checkpoint itself are pending for the
|
Writes stored at the target checkpoint itself are pending for the
|
||||||
next super-step and are excluded.
|
next super-step and are excluded.
|
||||||
|
|
||||||
Within a single checkpoint, writes are ordered by
|
|
||||||
`(task_path, task_id, idx)`, which is the order live execution applies
|
|
||||||
a super-step's task writes in. `task_id` is a hash of the path, so
|
|
||||||
ordering by it permutes parallel tasks writing one channel, and
|
|
||||||
reducers need not be order-invariant. Writes stored without a
|
|
||||||
`task_path` (graph input, `update_state` updates, exit-durability runs,
|
|
||||||
rows predating the column) sort first, by `task_id`.
|
|
||||||
* `seed` — the stored value at the nearest ancestor whose
|
* `seed` — the stored value at the nearest ancestor whose
|
||||||
`channel_values[ch]` is populated. Omitted if the walk reached the
|
`channel_values[ch]` is populated. Omitted if the walk reached the
|
||||||
root without finding any stored value (consumer treats absence as
|
root without finding any stored value (consumer treats absence as
|
||||||
@@ -619,11 +611,6 @@ class BaseCheckpointSaver(Generic[V]):
|
|||||||
`PostgresSaver`) override for performance; the return contract is
|
`PostgresSaver`) override for performance; the return contract is
|
||||||
fixed here.
|
fixed here.
|
||||||
|
|
||||||
`PendingWrite` carries no `task_path`, so this default replays each
|
|
||||||
checkpoint's writes in `get_tuple`'s `pending_writes` order. Savers
|
|
||||||
that do not return `pending_writes` ordered by
|
|
||||||
`(task_path, task_id, idx)` must override it.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
config: Configuration identifying the target checkpoint.
|
config: Configuration identifying the target checkpoint.
|
||||||
channels: Channel names to walk for. Empty → empty mapping.
|
channels: Channel names to walk for. Empty → empty mapping.
|
||||||
|
|||||||
@@ -199,8 +199,8 @@ class InMemorySaver(
|
|||||||
terminated_here.add(ch)
|
terminated_here.add(ch)
|
||||||
|
|
||||||
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
|
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
|
||||||
for _, (tid, ch, serialized, _) in sorted(
|
for (_task_id, _idx), (tid, ch, serialized, _) in sorted(
|
||||||
step_writes.items(), key=lambda kv: (kv[1][3], kv[0]), reverse=True
|
step_writes.items(), reverse=True
|
||||||
):
|
):
|
||||||
if ch not in remaining:
|
if ch not in remaining:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -1950,7 +1950,7 @@ class Pregel(
|
|||||||
run_tasks: list[PregelTaskWrites] = []
|
run_tasks: list[PregelTaskWrites] = []
|
||||||
run_task_ids: list[str] = []
|
run_task_ids: list[str] = []
|
||||||
|
|
||||||
for as_node, values, provided_task_id in valid_updates:
|
for i, (as_node, values, provided_task_id) in enumerate(valid_updates):
|
||||||
# create task to run all writers of the chosen node
|
# create task to run all writers of the chosen node
|
||||||
writers = self.nodes[as_node].flat_writers
|
writers = self.nodes[as_node].flat_writers
|
||||||
if not writers:
|
if not writers:
|
||||||
@@ -1964,7 +1964,7 @@ class Pregel(
|
|||||||
task_id = provided_task_id or (
|
task_id = provided_task_id or (
|
||||||
prepared_task_ids.popleft()
|
prepared_task_ids.popleft()
|
||||||
if prepared_task_ids
|
if prepared_task_ids
|
||||||
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
|
else _update_task_id(checkpoint["id"], i)
|
||||||
)
|
)
|
||||||
run_tasks.append(task)
|
run_tasks.append(task)
|
||||||
run_task_ids.append(task_id)
|
run_task_ids.append(task_id)
|
||||||
@@ -2410,7 +2410,7 @@ class Pregel(
|
|||||||
run_tasks: list[PregelTaskWrites] = []
|
run_tasks: list[PregelTaskWrites] = []
|
||||||
run_task_ids: list[str] = []
|
run_task_ids: list[str] = []
|
||||||
|
|
||||||
for as_node, values, provided_task_id in valid_updates:
|
for i, (as_node, values, provided_task_id) in enumerate(valid_updates):
|
||||||
# create task to run all writers of the chosen node
|
# create task to run all writers of the chosen node
|
||||||
writers = self.nodes[as_node].flat_writers
|
writers = self.nodes[as_node].flat_writers
|
||||||
if not writers:
|
if not writers:
|
||||||
@@ -2424,7 +2424,7 @@ class Pregel(
|
|||||||
task_id = provided_task_id or (
|
task_id = provided_task_id or (
|
||||||
prepared_task_ids.popleft()
|
prepared_task_ids.popleft()
|
||||||
if prepared_task_ids
|
if prepared_task_ids
|
||||||
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
|
else _update_task_id(checkpoint["id"], i)
|
||||||
)
|
)
|
||||||
run_tasks.append(task)
|
run_tasks.append(task)
|
||||||
run_task_ids.append(task_id)
|
run_task_ids.append(task_id)
|
||||||
@@ -4172,6 +4172,16 @@ class Pregel(
|
|||||||
await self.cache.aclear(namespaces)
|
await self.cache.aclear(namespaces)
|
||||||
|
|
||||||
|
|
||||||
|
def _update_task_id(checkpoint_id: str, i: int) -> str:
|
||||||
|
"""Task id for the `i`th update of a superstep that has no task to reuse.
|
||||||
|
|
||||||
|
Savers keep one write per `(task_id, idx)`, so updates sharing an id lose
|
||||||
|
all but the first one's writes, which a `DeltaChannel` replays from. The
|
||||||
|
first update keeps the id a lone update has always had.
|
||||||
|
"""
|
||||||
|
return str(uuid5(UUID(checkpoint_id), INTERRUPT if i == 0 else f"{INTERRUPT}:{i}"))
|
||||||
|
|
||||||
|
|
||||||
def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str]]:
|
def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str]]:
|
||||||
"""Index from a trigger to nodes that depend on it."""
|
"""Index from a trigger to nodes that depend on it."""
|
||||||
trigger_to_nodes: defaultdict[str, list[str]] = defaultdict(list)
|
trigger_to_nodes: defaultdict[str, list[str]] = defaultdict(list)
|
||||||
|
|||||||
@@ -1,113 +0,0 @@
|
|||||||
"""`DeltaChannel` replay must apply parallel writes in the order `invoke` did."""
|
|
||||||
|
|
||||||
from typing import Annotated, Any
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
|
||||||
from typing_extensions import TypedDict
|
|
||||||
|
|
||||||
from langgraph.channels.delta import DeltaChannel
|
|
||||||
from langgraph.graph import END, START, StateGraph
|
|
||||||
from langgraph.types import Send
|
|
||||||
|
|
||||||
pytestmark = pytest.mark.anyio
|
|
||||||
|
|
||||||
# Sorted, because live execution applies PULL tasks in node-name order.
|
|
||||||
FAN_OUT_NAMES = ["a", "b", "c", "d", "e", "f", "g", "h"]
|
|
||||||
SEND_ARGS = [f"send-{i:02d}" for i in range(12)]
|
|
||||||
|
|
||||||
|
|
||||||
def _append_reducer(current: list, updates: list) -> list:
|
|
||||||
return [*current, *(x for u in updates for x in u)]
|
|
||||||
|
|
||||||
|
|
||||||
def _build_fan_out_graph(checkpointer: BaseCheckpointSaver) -> Any:
|
|
||||||
class State(TypedDict):
|
|
||||||
items: Annotated[
|
|
||||||
list, DeltaChannel(_append_reducer, list, snapshot_frequency=10_000)
|
|
||||||
]
|
|
||||||
|
|
||||||
def make_node(label: str) -> Any:
|
|
||||||
return lambda state: {"items": [label]}
|
|
||||||
|
|
||||||
builder = StateGraph(State)
|
|
||||||
for name in FAN_OUT_NAMES:
|
|
||||||
builder.add_node(name, make_node(name))
|
|
||||||
builder.add_edge(START, name)
|
|
||||||
builder.add_edge(name, END)
|
|
||||||
return builder.compile(checkpointer=checkpointer)
|
|
||||||
|
|
||||||
|
|
||||||
def _build_send_fan_out_graph(checkpointer: BaseCheckpointSaver) -> Any:
|
|
||||||
class State(TypedDict):
|
|
||||||
items: Annotated[
|
|
||||||
list, DeltaChannel(_append_reducer, list, snapshot_frequency=10_000)
|
|
||||||
]
|
|
||||||
|
|
||||||
builder = StateGraph(State)
|
|
||||||
builder.add_node("worker", lambda arg: {"items": [arg]})
|
|
||||||
builder.add_conditional_edges(
|
|
||||||
START, lambda state: [Send("worker", n) for n in SEND_ARGS]
|
|
||||||
)
|
|
||||||
builder.add_edge("worker", END)
|
|
||||||
return builder.compile(checkpointer=checkpointer)
|
|
||||||
|
|
||||||
|
|
||||||
async def test_get_state_matches_live_send_order(
|
|
||||||
async_checkpointer: BaseCheckpointSaver,
|
|
||||||
) -> None:
|
|
||||||
graph = _build_send_fan_out_graph(async_checkpointer)
|
|
||||||
config = {"configurable": {"thread_id": "1"}}
|
|
||||||
|
|
||||||
live = (await graph.ainvoke({"items": []}, config))["items"]
|
|
||||||
replayed = (await graph.aget_state(config)).values["items"]
|
|
||||||
|
|
||||||
assert live == SEND_ARGS
|
|
||||||
assert replayed == live
|
|
||||||
|
|
||||||
|
|
||||||
async def test_get_state_matches_live_invoke_order(
|
|
||||||
async_checkpointer: BaseCheckpointSaver,
|
|
||||||
) -> None:
|
|
||||||
graph = _build_fan_out_graph(async_checkpointer)
|
|
||||||
config = {"configurable": {"thread_id": "1"}}
|
|
||||||
|
|
||||||
live = (await graph.ainvoke({"items": []}, config))["items"]
|
|
||||||
replayed = (await graph.aget_state(config)).values["items"]
|
|
||||||
|
|
||||||
assert live == FAN_OUT_NAMES
|
|
||||||
assert replayed == live
|
|
||||||
|
|
||||||
|
|
||||||
async def test_continuing_thread_preserves_committed_prefix(
|
|
||||||
async_checkpointer: BaseCheckpointSaver,
|
|
||||||
) -> None:
|
|
||||||
graph = _build_fan_out_graph(async_checkpointer)
|
|
||||||
config = {"configurable": {"thread_id": "1"}}
|
|
||||||
|
|
||||||
first = (await graph.ainvoke({"items": []}, config))["items"]
|
|
||||||
second = (await graph.ainvoke({"items": []}, config))["items"]
|
|
||||||
|
|
||||||
assert second == first + first
|
|
||||||
assert (await graph.aget_state(config)).values["items"] == second
|
|
||||||
|
|
||||||
|
|
||||||
async def test_state_history_reports_live_order_at_every_step(
|
|
||||||
async_checkpointer: BaseCheckpointSaver,
|
|
||||||
) -> None:
|
|
||||||
runs = 3
|
|
||||||
graph = _build_fan_out_graph(async_checkpointer)
|
|
||||||
config = {"configurable": {"thread_id": "1"}}
|
|
||||||
for _ in range(runs):
|
|
||||||
await graph.ainvoke({"items": []}, config)
|
|
||||||
live = FAN_OUT_NAMES * runs
|
|
||||||
|
|
||||||
seen = [
|
|
||||||
s.values["items"]
|
|
||||||
async for s in graph.aget_state_history(config)
|
|
||||||
if "items" in s.values
|
|
||||||
]
|
|
||||||
|
|
||||||
assert max(map(len, seen)) == len(live)
|
|
||||||
for values in seen:
|
|
||||||
assert values == live[: len(values)], f"{values} is not a prefix of {live}"
|
|
||||||
@@ -20,6 +20,7 @@ from typing import Annotated, Any
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from langchain_core.messages import HumanMessage
|
from langchain_core.messages import HumanMessage
|
||||||
|
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||||
from langgraph.checkpoint.memory import InMemorySaver
|
from langgraph.checkpoint.memory import InMemorySaver
|
||||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||||
from typing_extensions import TypedDict
|
from typing_extensions import TypedDict
|
||||||
@@ -27,16 +28,17 @@ from typing_extensions import TypedDict
|
|||||||
from langgraph.channels.delta import DeltaChannel
|
from langgraph.channels.delta import DeltaChannel
|
||||||
from langgraph.graph import START, StateGraph
|
from langgraph.graph import START, StateGraph
|
||||||
from langgraph.graph.message import _messages_delta_reducer
|
from langgraph.graph.message import _messages_delta_reducer
|
||||||
from langgraph.types import StateUpdate
|
from langgraph.types import StateSnapshot, StateUpdate
|
||||||
|
|
||||||
pytestmark = pytest.mark.anyio
|
pytestmark = pytest.mark.anyio
|
||||||
|
|
||||||
|
|
||||||
def _build_graph(
|
def _build_graph(
|
||||||
checkpointer: InMemorySaver,
|
checkpointer: BaseCheckpointSaver,
|
||||||
*,
|
*,
|
||||||
two_nodes: bool = False,
|
two_nodes: bool = False,
|
||||||
snapshot_frequency: int = 1000,
|
snapshot_frequency: int = 1000,
|
||||||
|
interrupt_before: list[str] | None = None,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""Compile a minimal DeltaChannel-backed `messages` graph.
|
"""Compile a minimal DeltaChannel-backed `messages` graph.
|
||||||
|
|
||||||
@@ -63,7 +65,7 @@ def _build_graph(
|
|||||||
builder.set_finish_point("assistant")
|
builder.set_finish_point("assistant")
|
||||||
else:
|
else:
|
||||||
builder.set_finish_point("model")
|
builder.set_finish_point("model")
|
||||||
return builder.compile(checkpointer=checkpointer)
|
return builder.compile(checkpointer=checkpointer, interrupt_before=interrupt_before)
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -273,10 +275,6 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
|
|||||||
that each call `put_writes`. Guards the regression where moving
|
that each call `put_writes`. Guards the regression where moving
|
||||||
`put_writes` outside the per-task loop would persist only the last
|
`put_writes` outside the per-task loop would persist only the last
|
||||||
task's writes.
|
task's writes.
|
||||||
|
|
||||||
Explicit `task_id`s are required to disambiguate writes belonging to
|
|
||||||
different `StateUpdate`s targeting the same node — otherwise both share
|
|
||||||
the deterministic interrupt-derived id and collide in the saver.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
saver = InMemorySaver()
|
saver = InMemorySaver()
|
||||||
@@ -310,6 +308,92 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
|
|||||||
assert sorted(ids) == ["m1", "m2"]
|
assert sorted(ids) == ["m1", "m2"]
|
||||||
|
|
||||||
|
|
||||||
|
def _update(content: str, as_node: str) -> StateUpdate:
|
||||||
|
return StateUpdate(
|
||||||
|
values={"messages": [HumanMessage(content=content, id=content)]},
|
||||||
|
as_node=as_node,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _contents(state: StateSnapshot) -> list[str]:
|
||||||
|
return [m.content for m in state.values["messages"]]
|
||||||
|
|
||||||
|
|
||||||
|
def test_bulk_update_state_keeps_every_update_without_task_ids(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
graph = _build_graph(sync_checkpointer, two_nodes=True)
|
||||||
|
config = {"configurable": {"thread_id": "bulk-no-task-ids"}}
|
||||||
|
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||||
|
|
||||||
|
graph.bulk_update_state(
|
||||||
|
config,
|
||||||
|
[
|
||||||
|
[
|
||||||
|
_update("first", "model"),
|
||||||
|
_update("second", "model"),
|
||||||
|
_update("third", "assistant"),
|
||||||
|
]
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
contents = _contents(graph.get_state(config))
|
||||||
|
assert sorted(contents) == ["first", "hi", "second", "third"], (
|
||||||
|
f"every update's writes must persist; got {contents}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_abulk_update_state_keeps_every_update_without_task_ids(
|
||||||
|
async_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
graph = _build_graph(async_checkpointer, two_nodes=True)
|
||||||
|
config = {"configurable": {"thread_id": "bulk-no-task-ids"}}
|
||||||
|
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||||
|
|
||||||
|
await graph.abulk_update_state(
|
||||||
|
config,
|
||||||
|
[
|
||||||
|
[
|
||||||
|
_update("first", "model"),
|
||||||
|
_update("second", "model"),
|
||||||
|
_update("third", "assistant"),
|
||||||
|
]
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
contents = _contents(await graph.aget_state(config))
|
||||||
|
assert sorted(contents) == ["first", "hi", "second", "third"], (
|
||||||
|
f"every update's writes must persist; got {contents}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_bulk_update_state_keeps_every_update_next_to_a_pending_task(
|
||||||
|
sync_checkpointer: BaseCheckpointSaver,
|
||||||
|
) -> None:
|
||||||
|
graph = _build_graph(
|
||||||
|
sync_checkpointer, two_nodes=True, interrupt_before=["assistant"]
|
||||||
|
)
|
||||||
|
config = {"configurable": {"thread_id": "bulk-pending-task"}}
|
||||||
|
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||||
|
assert graph.get_state(config).next == ("assistant",)
|
||||||
|
|
||||||
|
graph.bulk_update_state(
|
||||||
|
config,
|
||||||
|
[
|
||||||
|
[
|
||||||
|
_update("first", "assistant"),
|
||||||
|
_update("second", "model"),
|
||||||
|
_update("third", "model"),
|
||||||
|
]
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
contents = _contents(graph.get_state(config))
|
||||||
|
assert sorted(contents) == ["first", "hi", "second", "third"], (
|
||||||
|
f"every update's writes must persist; got {contents}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Public-API observation of fresh-thread checkpoint shape
|
# Public-API observation of fresh-thread checkpoint shape
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user