Compare commits

..
Author SHA1 Message Date
Elior Nataf Lackritz a8e732c879 fix(langgraph): give each bulk_update_state update its own task id
An update whose node has no pending task to reuse was stored under
uuid5(checkpoint_id, INTERRUPT), so every such update in one superstep
shared a task id. Savers keep one write per (task_id, idx), so all but the
first update's writes were dropped. Plain channels were unaffected, since
their value is stored in the new checkpoint, but a DeltaChannel replays
those writes and lost every update after the first.

The ith update now gets uuid5(checkpoint_id, f"{INTERRUPT}:{i}"). The first
keeps the old id, so a single update stores exactly what it did before.
2026-09-30 12:47:55 -04:00
eb69f67b65 fix(checkpoint-sqlite): walk delta ancestors by parent pointer (#8557)
## Summary

The sqlite delta history silently drops a parent checkpoint whose id sorts above its child's,
losing that parent's stored value and its pending writes. The channel hydrates short with no error.

Fixes #8550

## Problem

Stage 1 walked ancestors with:

```sql
WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id <= ?
ORDER BY checkpoint_id DESC
```

Ancestry is defined by the `parent_checkpoint_id` column. These two predicates add a second
requirement: that every child's id sorts above its parent's. The contract promises monotonic ids,
but that only holds within one process, so ids from processes with different clocks can break it.

When the requirement is violated the parent is excluded from the stream and its seed and writes go
with it. Dropping the range filter alone does not fix it: in `checkpoint_id DESC` order that parent
arrives *before* the target, so the walk streams past it before it has started.

## Fix

A recursive CTE anchored at the target, following `parent_checkpoint_id`:

```sql
WITH RECURSIVE ancestors(checkpoint_id, parent_checkpoint_id, type, checkpoint) AS (
    SELECT ... FROM checkpoints
    WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ?
    UNION ALL
    SELECT c.... 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
```

Rows now arrive in walk order (target, parent, grandparent, ...), so `step_walk_with_row` no longer
needs its off-path skip or its `parent_cid` tracking; both are removed. The query reads only true
ancestors, where the old one read every row at or below the target including sibling branches.

`CROSS JOIN` pins the join order. The saver never runs `ANALYZE`, and with a plain `JOIN` sqlite put
`checkpoints` as the outer loop, scanning the whole thread on every recursion step. With `ancestors`
outside, each step is one primary key lookup. Through `get_delta_channel_history`:

| chain length | plain `JOIN` | `CROSS JOIN` |
| -- | -- | -- |
| 1000 | 0.032s | 0.001s |
| 2000 | 0.124s | 0.003s |
| 4000 | 0.475s | 0.006s |

## Cycle guard

Following pointers can loop where a bounded id scan could not, and a loop is reachable through
`put` alone: `put` writes with `INSERT OR REPLACE`, so re-putting an existing checkpoint id under a
descendant's config repoints that checkpoint at its own descendant. The walk stops on a repeated
`checkpoint_id` (one set insert per row, no depth ceiling that could truncate a long migrated
thread). sqlite yields recursive rows lazily, so abandoning the cursor ends the recursion.

`test_walk_terminates_when_put_makes_the_parent_chain_cycle` fails by hanging, not by asserting, if
the guard regresses (confirmed by deleting the guard locally). The package has no `pytest-timeout`,
so the CI job timeout is the backstop.

## Postgres

No equivalent change needed. It pages the whole thread with no id bound and follows parent pointers
in Python, and its upsert never rewrites `parent_checkpoint_id`, so it can neither miss this parent
nor form the loop. `BaseCheckpointSaver` and `InMemorySaver` also walk parent pointers.

## Test plan

New `libs/checkpoint-sqlite/tests/test_delta_parent_walk.py`:

- [x] Sync and async, parametrised over both id orders; the sync case also asserts equality with
      `BaseCheckpointSaver` on the same rows. `parent_id_sorts_above_child` is the bug,
      `parent_id_sorts_below_child` the control.
- [x] `test_walk_reaches_root_of_long_chain_with_descending_ids`: 40 checkpoints, only stored value
      at the root.
- [x] `test_walk_terminates_when_put_makes_the_parent_chain_cycle`.
- [x] `test_walk_step_looks_up_the_parent_by_primary_key`: asserts the recursive step's
      `EXPLAIN QUERY PLAN` is a key lookup, so a plain `JOIN` can't come back. Fails with it.
- [x] On `main`: 3 of the 6 walk tests fail (both `parent_id_sorts_above_child` cases and the long
      chain). The cycle test passes on `main` too, since the old bounded scan could not loop; it
      guards the new path.
- [x] #8550's repro returns `{'writes': [('task', 'ch', 'write-root')], 'seed': 'seed'}` sync and
      async (was `{'writes': []}` on `main`).
- [x] `libs/checkpoint-sqlite`: `make format`, `make lint` clean; full suite 125 passed, 2 skipped.
- [x] `libs/langgraph`: `-k "delta or sqlite"` 739 passed, 1 skipped.

Thanks to @lylelllll for the report, the minimal repro, the base-saver comparison that isolated it
to the fast path, and for suggesting the recursive CTE.




Co-authored-by: lylelllll <59271327+lylelllll@users.noreply.github.com>
2026-09-30 12:17:02 -04:00
c0279f0910 fix(checkpoint-postgres): derive the delta walk cursor once the target loads (#8556)
## Summary

`get_delta_channel_history` on Postgres returns an empty history for any `DeltaChannel` on a
target checkpoint that is not within the first stage-1 pagination page (1024 rows) of the thread.
No exception, no warning: the channel just hydrates empty.

Fixes #8448

## Problem

Stage 1 pages `checkpoints` newest-first from the head of the thread, and after each page
`_try_advance_walks` tries to move every not-yet-seeded channel's walk along the partial
`parent_of` map accumulated so far. The walk starts at the target's parent:

```python
if ch not in walk_cursor_by_ch:
    walk_cursor_by_ch[ch] = parent_of.get(target_id)
```

The target can be any checkpoint in the thread, not just the head, so on the first page
`parent_of` frequently has no row for it yet. `.get` then returns `None`, which is also what a
target with no parent returns, and the two are stored identically. Because the initialisation is
guarded by `ch not in walk_cursor_by_ch`, it never runs again: once the walk is parked at `None`
it stays there even after the target's real row and real parent load on a later page.

The result is an empty chain and no seed. Downstream `channels_from_checkpoint` does

```python
replay_ch = delta_spec.from_checkpoint(history.get("seed", MISSING))
replay_ch.replay_writes(history["writes"])
```

so `get_state`, `get_state_history` and `update_state` against an older checkpoint reconstruct a
`messages` channel as `[]` on a thread with hundreds of real messages.

## Fix

Start the walk only once `target_id` is actually present in `parent_of`, so "the target has not
loaded yet" stops sharing a representation with "the target is a root":

```python
if ch not in walk_cursor_by_ch:
    if target_id not in parent_of:
        continue
    walk_cursor_by_ch[ch] = parent_of[target_id]
```

`_try_advance_walks` is a static method on `BasePostgresSaver`, so `PostgresSaver` and
`AsyncPostgresSaver` are both covered by the one change.

## Why it's safe

`continue` leaves the channel exactly as it was, so a later page retries. The three existing
stop conditions are untouched: a channel that finds its seed still seeds, one that reaches a real
root still parks at `None`, and one waiting on an ancestor still keeps its cursor. Paging still
terminates on a short page, which is what ends the run for a target that really is a root.

## Long-term

The sibling sqlite implementation avoids this class of bug differently, by starting its stage-1
scan at the target (`checkpoint_id <= ?`) instead of at the head. Postgres could adopt the same
bound and would then never fetch a checkpoint newer than the target at all, which looks like the
bigger win on a long thread. It makes the read path depend on ancestors always sorting below their
descendants, though, which sqlite already assumes but the Postgres fast path currently does not.
#8550 now reports that assumption as a bug in sqlite, on the grounds that ancestry is defined by
`parent_checkpoint_id` and the contract does not require ids to be monotonic, so the bound is the
wrong direction to move Postgres in. Paging the full thread and following parent pointers is what
keeps this path correct when ids are not monotonic, and with this fix Postgres returns the right
history for #8550's scenario at every page size.

## Test plan

New `libs/checkpoint-postgres/tests/test_delta_pagination.py`. Page size is monkeypatched rather
than writing 1024+ real checkpoints per case, since the only thing that decides the behaviour is
which page the target lands on.

- [x] `test_async_target_older_than_the_first_page` and its sync twin, parametrised over page
      sizes `[_DELTA_PAGE_SIZE, 3, 2, 1]`. The thread has 8 checkpoints with a snapshot at step 1
      and the target at step 4, so every size at or below 3 leaves the target off the first page.
      The real page size is the control.
- [x] `test_root_target_has_no_history_and_still_terminates` covers the case where a `None` cursor
      is the correct answer, at page size 1 so the paging loop runs the length of the thread.
- [x] 6 of the 9 fail on `main` (`expected a snapshot seed, got '<missing>'`); the 3 that pass are
      the two controls and the root case.
- [x] `make format`, `make lint_package`, `make lint_tests` clean.
- [x] Full `libs/checkpoint-postgres` suite, rebased on current `main`: 279 passed, 3 skipped on Postgres 16.
- [x] Graph-level repro with `_DELTA_PAGE_SIZE = 5`: 10 invocations, then `get_state` on the 8th-newest
      checkpoint returns `[]` on `main` and the full history on this branch.

Thanks to @Navneet-Scaler for the report, the mechanism write-up, and the fix in #8453, which this
matches.




Co-authored-by: Navneet-Scaler <147032454+Navneet-Scaler@users.noreply.github.com>
2026-09-30 12:16:55 -04:00
ccurmeandGitHub f5804a5bf5 fix(ci): test locally-built wheel and publish to test pypi after pre-release checks (#9124) 2026-09-30 09:30:34 -04:00
14 changed files with 494 additions and 1025 deletions
+23 -36
View File
@@ -139,23 +139,10 @@ jobs:
echo EOF
} >> "$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:
needs:
- build
- release-notes
- test-pypi-publish
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
@@ -180,31 +167,20 @@ jobs:
enable-cache: false
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
working-directory: ${{ inputs.working-directory }}
env:
PKG_NAME: ${{ needs.build.outputs.pkg-name }}
VERSION: ${{ needs.build.outputs.version }}
# Here we use:
# - 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
# Install directly from the locally-built wheel (no index resolution needed).
run: |
uv run pip install \
--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" \
)
uv run pip install dist/*.whl
if [[ "$PKG_NAME" == *prebuilt* ]]; then
uv run pip install langgraph
@@ -226,7 +202,7 @@ jobs:
run: uv sync --group test
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)
working-directory: ${{ inputs.working-directory }}
shell: bash
@@ -234,14 +210,25 @@ jobs:
PKG_NAME: ${{ needs.build.outputs.pkg-name }}
VERSION: ${{ needs.build.outputs.version }}
run: |
uv run pip install \
--extra-index-url https://test.pypi.org/simple/ \
"$PKG_NAME==$VERSION"
uv run pip install dist/*.whl
- name: Run unit tests
run: make test
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:
needs:
- build
@@ -448,11 +448,12 @@ class PostgresSaver(BasePostgresSaver):
Two-stage query, both stages cover ALL requested channels:
* Stage 1 (paged): dynamic SELECT over `checkpoints` with K parallel
JSONB key lookups (one column pair per channel) — no subquery, no
aggregation. Pages newest-first by `checkpoint_id` with a cursor;
page size is `_DELTA_PAGE_SIZE`. Stops paging when every channel
has found its seed or the chain is exhausted.
* Stage 1 (paged): dynamic SELECT over `checkpoints` with three
columns per channel: its version, an `EXISTS` probe for a stored
blob at that version, and its inline value. Pages newest-first by
`checkpoint_id` with a cursor; page size is `_DELTA_PAGE_SIZE`.
Stops paging when every channel has found its seed or a page comes
back short.
* Stage 2 (per-channel UNION ALL): one branch per channel reading
`checkpoint_writes` filtered to that channel's specific
@@ -172,30 +172,8 @@ class _DeltaStage2Row(TypedDict, total=False):
version: str | None # "b" rows only
# Multi-channel two-stage DeltaChannel reconstruction.
#
# Stage 1 scans checkpoint metadata (no blob bytes) and emits one row per
# checkpoint with K parallel JSONB key lookups (one column pair per
# requested delta channel: ver_i / hs_i). No subqueries, no aggregation.
# Python walks the parent chain once across all channels.
#
# Stage 2 fetches all writes and the seed blobs for ALL channels in a
# single roundtrip via `channel = ANY(%s)` and chain/seed-version
# filtering.
#
# Empirical comparison vs an alternative "ship full channel_versions /
# channel_values JSONB and let Python pick" form (1000 checkpoints,
# 8 total channels in graph, 3 delta channels requested):
#
# Postgres execution: A=0.24ms vs B=0.38ms (both negligible)
# End-to-end latency: A=6.83ms vs B=2.28ms (B is 3.0x faster)
# Wire payload: A=836KB vs B=330KB (61% smaller)
# Buffer hits: identical (167 blocks)
#
# B (this dynamic-columns design) wins because it avoids JSONB
# serialization on the wire and JSONB-to-dict deserialization in
# psycopg. Even at K=8 (8 delta channels = 16 dynamic columns), B
# still beats A end-to-end (4.2ms vs 6.8ms).
# Delta history is rebuilt in two queries; `_build_delta_stage1_sql` and
# `_build_delta_stage2_sql` document their shapes.
def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
@@ -335,10 +313,8 @@ def _build_delta_stage2_sql(
return " UNION ALL ".join(branches)
# Stage 1 rows are dynamic-shape dicts: {checkpoint_id, parent_checkpoint_id,
# ver_0, hs_0, ver_1, hs_1, ...}. Walking is parameterized by the channel
# list to map indices back to channel names — no static TypedDict here.
# `dict[str, Any]` is the practical signature.
# Stage 1 rows are dicts keyed by the per-channel aliases
# `_build_delta_stage1_sql` emits, so there is no static TypedDict.
class BasePostgresSaver(BaseCheckpointSaver[str]):
@@ -431,9 +407,11 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
(a) it found a stored value for its channel — a blob or an inline
primitive (channel becomes seeded),
(b) it reached a real root (parent_of[cid] is None — fully
materialized at this point), or
materialized at this point),
(c) the next ancestor cid isn't in `parent_of` yet (waiting for
a later page; the cursor stays put).
a later page; the cursor stays put), or
(d) the target's own row isn't in `parent_of` yet (the walk has
not started; no cursor is set, so a later page retries).
Mutates `chain_by_ch`, `seed_ver_by_ch`, `seed_inline_by_ch`,
`walk_cursor_by_ch`, and `seeded` in place.
@@ -441,9 +419,12 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
for i, ch in enumerate(channels):
if ch in seeded:
continue
# First-time entry: cursor starts at the target's parent.
# Pages start at the thread head, so the target may not have
# loaded yet; a `None` cursor would read as "target is a root".
if ch not in walk_cursor_by_ch:
walk_cursor_by_ch[ch] = parent_of.get(target_id)
if target_id not in parent_of:
continue
walk_cursor_by_ch[ch] = parent_of[target_id]
cur_cid = walk_cursor_by_ch[ch]
ch_chain = chain_by_ch[ch]
hb_i = hb_by_i_by_cid[i]
@@ -0,0 +1,132 @@
from __future__ import annotations
from typing import Any
from uuid import uuid4
import pytest
from langgraph.checkpoint.base import (
Checkpoint,
DeltaChannelHistory,
empty_checkpoint,
)
from langgraph.checkpoint.base.id import uuid6
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from langgraph.checkpoint.postgres import PostgresSaver
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
from langgraph.checkpoint.postgres.base import _DELTA_PAGE_SIZE
from tests.conftest import DEFAULT_URI
CHANNEL = "items"
STEPS = 8
SEED_STEP = 1
SEED_VALUE = [10, 20]
TARGET_STEP = 4
# The real page size is the control; the rest leave the target off the first
# page (three checkpoints are newer than it).
PAGE_SIZES = [_DELTA_PAGE_SIZE, 3, 2, 1]
def _step_args(
thread_id: str, step: int, parent: dict | None
) -> tuple[dict, Checkpoint, dict[str, Any]]:
config: dict = {"configurable": {"thread_id": thread_id, "checkpoint_ns": ""}}
if parent is not None:
config["configurable"]["checkpoint_id"] = parent["configurable"][
"checkpoint_id"
]
checkpoint: Checkpoint = empty_checkpoint()
checkpoint["id"] = str(uuid6(clock_seq=step))
checkpoint["channel_versions"][CHANNEL] = f"v{step}"
if step == SEED_STEP:
checkpoint["channel_values"][CHANNEL] = _DeltaSnapshot(list(SEED_VALUE))
return config, checkpoint, {CHANNEL: f"v{step}"}
return config, checkpoint, {}
async def _abuild_chain(saver: AsyncPostgresSaver) -> list[dict]:
thread_id = str(uuid4())
parent: dict | None = None
configs: list[dict] = []
for step in range(STEPS):
config, checkpoint, new_versions = _step_args(thread_id, step, parent)
parent = await saver.aput(
config,
checkpoint,
{"source": "loop", "step": step, "parents": {}},
new_versions,
)
await saver.aput_writes(parent, [(CHANNEL, f"w{step}")], str(uuid4()))
configs.append(parent)
return configs
def _build_chain(saver: PostgresSaver) -> list[dict]:
thread_id = str(uuid4())
parent: dict | None = None
configs: list[dict] = []
for step in range(STEPS):
config, checkpoint, new_versions = _step_args(thread_id, step, parent)
parent = saver.put(
config,
checkpoint,
{"source": "loop", "step": step, "parents": {}},
new_versions,
)
saver.put_writes(parent, [(CHANNEL, f"w{step}")], str(uuid4()))
configs.append(parent)
return configs
def _assert_history(entry: DeltaChannelHistory, page_size: int) -> None:
seed = entry.get("seed")
assert isinstance(seed, _DeltaSnapshot), (
f"page_size={page_size}: expected a snapshot seed, "
f"got {entry.get('seed', '<missing>')!r}"
)
assert seed.value == SEED_VALUE
assert [w[2] for w in entry["writes"]] == ["w1", "w2", "w3"], (
f"page_size={page_size}: got {[w[2] for w in entry['writes']]}"
)
@pytest.mark.parametrize("page_size", PAGE_SIZES)
async def test_async_target_older_than_the_first_page(
page_size: int, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr("langgraph.checkpoint.postgres.aio._DELTA_PAGE_SIZE", page_size)
async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver:
await saver.setup()
configs = await _abuild_chain(saver)
result = await saver.aget_delta_channel_history(
config=configs[TARGET_STEP], channels=[CHANNEL]
)
_assert_history(result[CHANNEL], page_size)
@pytest.mark.parametrize("page_size", PAGE_SIZES)
def test_sync_target_older_than_the_first_page(
page_size: int, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr("langgraph.checkpoint.postgres._DELTA_PAGE_SIZE", page_size)
with PostgresSaver.from_conn_string(DEFAULT_URI) as saver:
saver.setup()
configs = _build_chain(saver)
result = saver.get_delta_channel_history(
config=configs[TARGET_STEP], channels=[CHANNEL]
)
_assert_history(result[CHANNEL], page_size)
async def test_root_target_has_no_history_and_still_terminates(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr("langgraph.checkpoint.postgres.aio._DELTA_PAGE_SIZE", 1)
async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver:
await saver.setup()
configs = await _abuild_chain(saver)
result = await saver.aget_delta_channel_history(
config=configs[0], channels=[CHANNEL]
)
assert result[CHANNEL] == {"writes": []}
@@ -507,13 +507,12 @@ class SqliteSaver(BaseCheckpointSaver[str]):
Two-stage query:
* Stage 1 (paged): newest-first slice of `checkpoints` returning
`(checkpoint_id, parent_checkpoint_id, type, checkpoint)` per
ancestor. Sqlite has no JSONB, so we ship the full serialized
checkpoint blob and inspect `channel_values` in Python. Pages
newest-first by `checkpoint_id` with a `< cursor` predicate;
page size is `DELTA_PAGE_SIZE`. Stops paging when every channel
has found its seed or the chain is exhausted.
* Stage 1 (streamed): recursive CTE over `checkpoints` following
`parent_checkpoint_id` from the target, returning
`(checkpoint_id, type, checkpoint)` per ancestor. Sqlite has no
JSONB, so we ship the full serialized checkpoint blob and inspect
`channel_values` in Python. Stops reading when every channel has
found its seed or the chain is exhausted.
* Stage 2 (per-channel UNION ALL): one branch per channel reading
`writes` filtered to that channel's specific `chain_cids`. No
@@ -538,12 +537,14 @@ class SqliteSaver(BaseCheckpointSaver[str]):
seeded: set[str] = set()
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:
cid, parent_cid, type_tag, blob = row
cid, type_tag, blob = row
if step_walk_with_row(
cid=cid,
parent_cid=parent_cid,
type_tag=type_tag,
blob=blob,
target_id=checkpoint_id,
@@ -26,16 +26,33 @@ from typing import Any
from langgraph.checkpoint.base import DeltaChannelHistory, PendingWrite
# Stage 1 streams ancestors of `target_cid` newest-first. The `<=`
# predicate keeps target itself in the stream so we can read its
# `parent_checkpoint_id` from the first row without a separate lookup;
# the caller skips target's own writes/seed (matches the
# `BaseCheckpointSaver` contract).
# Stage 1 streams target, then its ancestors nearest-first, by following
# `parent_checkpoint_id` rather than id order: ids are only monotonic within
# one process, so a range scan by id can miss a parent whose id sorts above
# its child's. Target is the anchor row; its own writes/seed are skipped
# (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 = (
"WITH RECURSIVE ancestors(checkpoint_id, parent_checkpoint_id, type, "
"checkpoint) AS ("
"SELECT checkpoint_id, parent_checkpoint_id, type, checkpoint "
"FROM checkpoints "
"WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id <= ? "
"ORDER BY checkpoint_id DESC"
"WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? "
"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"
)
@@ -68,7 +85,6 @@ def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
def step_walk_with_row(
*,
cid: str,
parent_cid: str | None,
type_tag: str,
blob: bytes,
target_id: str,
@@ -81,36 +97,32 @@ def step_walk_with_row(
) -> bool:
"""Process one streamed stage-1 row in the merged ancestor walk.
The cursor returns (cid, parent_cid, type, blob) rows in
`checkpoint_id` DESC order starting at target. The first row is
target itself; we read its parent_cid to seed the walk and otherwise
skip it (target's own writes/seed are not part of the contract).
The cursor returns (cid, type, blob) rows in walk order starting at
target. The first row is target itself and is skipped (target's own
writes/seed are not part of the contract).
For each subsequent row, if `cid` matches the walk's current
position, we deserialize the blob, append the cid to every
not-yet-seeded channel's chain, and check `channel_values` for
For each subsequent row we deserialize the blob, append the cid to
every not-yet-seeded channel's chain, and check `channel_values` for
seeds. The deserialized checkpoint is dropped before advancing — no
cross-row cache, so peak in-flight is one deserialized checkpoint.
Off-path rows (different branch on the same thread) advance the
cursor without doing any work.
Returns True when every requested channel is seeded — the caller
can stop iterating and close the cursor.
Returns True when the caller can stop iterating and close the cursor:
every requested channel is seeded, or the chain revisited a checkpoint.
"""
if "started" not in walk_state:
if cid == target_id:
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["walked"] = {cid}
# Not target yet (or target not present): keep streaming.
return False
active: set[str] = walk_state["active"]
if not active:
return True
if cid != walk_state["cur_cid"]:
# Off-path row from a sibling branch — skip without deserializing.
return False
walked: set[str] = walk_state["walked"]
if cid in walked:
return True
walked.add(cid)
for ch in active:
chain_by_ch[ch].append(cid)
ckpt = serde.loads_typed((type_tag, blob))
@@ -120,7 +132,6 @@ def step_walk_with_row(
seeded.add(ch)
active.discard(ch)
del ckpt, channel_values
walk_state["cur_cid"] = parent_cid
return not active
@@ -625,8 +625,8 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
"""Fast-path override of `BaseCheckpointSaver.aget_delta_channel_history`.
See `SqliteSaver.get_delta_channel_history` for design notes; this
is the async equivalent using `aiosqlite` cursors. Stage 1 pages
the parent chain newest-first and Python-deserializes each
is the async equivalent using `aiosqlite` cursors. Stage 1 streams
the parent chain from the target and Python-deserializes each
checkpoint blob to find per-channel snapshots; stage 2 fetches
only the relevant writes via per-channel UNION ALL.
"""
@@ -650,13 +650,13 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
async with self.lock, self.conn.cursor() as cur:
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:
cid, parent_cid, type_tag, blob = row
cid, type_tag, blob = row
if step_walk_with_row(
cid=cid,
parent_cid=parent_cid,
type_tag=type_tag,
blob=blob,
target_id=checkpoint_id,
@@ -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}"
+21 -62
View File
@@ -8,15 +8,13 @@ from typing import Any, cast
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
ChannelVersions,
Checkpoint,
PendingWrite,
)
from langgraph.checkpoint.base.id import uuid6
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from langgraph._internal._config import DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT
from langgraph._internal._constants import INTERRUPT, PUSH
from langgraph._internal._constants import PUSH
from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel
from langgraph.channels.delta import DeltaChannel
@@ -91,23 +89,6 @@ def get_delta_channels_from_all_channels(
}
def delta_channels_with_pending_writes(
specs: Mapping[str, Any],
pending_writes: Iterable[PendingWrite] | None,
) -> set[str]:
"""DeltaChannels a branch starting from this checkpoint must snapshot.
A checkpoint's pending writes belong to the child that consumed them, and
nothing records which child that was. A new branch snapshots every delta
channel they touch, so its ancestor walk never replays them.
"""
return {
ch
for _, ch, _ in pending_writes or ()
if isinstance(specs.get(ch), DeltaChannel)
}
def create_metadata_for_update_state_api(
channels: Mapping[str, BaseChannel],
updated_channels: set[str],
@@ -141,7 +122,6 @@ def create_checkpoint_plan_for_update_state_api(
parents: dict[str, Any],
saved_metadata: Mapping[str, Any] | None,
is_fresh_thread: bool,
fork_channels: set[str],
) -> tuple[set[str], dict[str, Any]]:
"""Return ``(channels_to_snapshot, metadata)`` for an update_state head."""
metadata: dict[str, Any] = {
@@ -157,9 +137,7 @@ def create_checkpoint_plan_for_update_state_api(
updated_channels,
prev_metadata=saved_metadata,
)
channels_to_snapshot = (
delta_channels_to_snapshot(channels, new_counters) | fork_channels
)
channels_to_snapshot = delta_channels_to_snapshot(channels, new_counters)
for k in channels_to_snapshot:
new_counters[k] = (0, 0)
non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}
@@ -189,7 +167,6 @@ def create_checkpoint(
"""
ts = datetime.now(timezone.utc).isoformat()
channels_to_snapshot = channels_to_snapshot or set()
bumped: dict[str, tuple[Any, Any]] = {}
if channels is None:
values = checkpoint["channel_values"]
channel_versions = checkpoint["channel_versions"]
@@ -197,29 +174,30 @@ def create_checkpoint(
values = {}
channel_versions = dict(checkpoint["channel_versions"])
for k in channels:
ch = channels[k]
if k not in channel_versions:
# A forced snapshot of a never-written channel still has to
# land to stop the ancestor walk, and `put` only stores blobs
# for versioned channels.
if k in channels_to_snapshot and get_next_version is not None:
channel_versions[k] = get_next_version(None, None)
bumped[k] = (None, channel_versions[k])
values[k] = _DeltaSnapshot(
ch.get() if ch.is_available() else ch.typ()
)
continue
ch = channels[k]
if k in channels_to_snapshot:
# `put` only stores a blob for a channel whose version moved,
# so snapshotting a channel this step did not write needs a
# bump: exit mode reaching the cadence on a superstep that
# skipped the channel, and a fork's first checkpoint.
# Callers force a full snapshot blob here: exit mode when a
# delta channel reaches its snapshot cadence, and update_state
# on a fresh thread (no ancestor to replay writes from). The
# manual version-bump below only applies to the exit-mode case.
#
# In exit mode, the snapshot decision is deferred to exit
# time (intermediate steps have do_checkpoint=False). The
# channel's count may have reached snapshot_frequency over
# several supersteps, but the LAST superstep may not have
# written to this channel. In that case apply_writes()
# (in _algo.py) didn't bump this channel's version, so
# saver.put() wouldn't include it in new_versions and
# the snapshot blob would be silently dropped. The manual
# bump below closes the gap. In sync/async durability this
# branch is effectively dead code (the step that pushes
# the count to freq always writes the channel).
if get_next_version is not None and (
updated_channels is None or k not in updated_channels
):
old = channel_versions[k]
channel_versions[k] = get_next_version(old, None)
bumped[k] = (old, channel_versions[k])
channel_versions[k] = get_next_version(channel_versions[k], None)
values[k] = _DeltaSnapshot(ch.get())
else:
v = ch.checkpoint()
@@ -231,30 +209,11 @@ def create_checkpoint(
id=id or str(uuid6(clock_seq=step)),
channel_values=values,
channel_versions=channel_versions,
versions_seen=_mark_bumps_seen(checkpoint["versions_seen"], bumped),
versions_seen=checkpoint["versions_seen"],
updated_channels=None if updated_channels is None else sorted(updated_channels),
)
def _mark_bumps_seen(
versions_seen: dict[str, ChannelVersions],
bumped: Mapping[str, tuple[Any, Any]],
) -> dict[str, ChannelVersions]:
"""Advance whoever had seen a bumped channel's old version to the new one.
A bump that only stores a snapshot is not a write. Left unseen, it would
re-fire `interrupt_before` and rerun the channel's subscribers.
"""
if not bumped:
return versions_seen
out: dict[str, ChannelVersions] = {}
for node, seen in {INTERRUPT: {}, **versions_seen}.items():
advanced = {k: new for k, (old, new) in bumped.items() if seen.get(k) == old}
if advanced or node in versions_seen:
out[node] = {**seen, **advanced}
return out
def _needs_replay(spec: BaseChannel, stored: object) -> bool:
"""True if `spec` is a `DeltaChannel` and no value is stored at this
checkpoint, requiring an ancestor walk to reconstruct.
+11 -22
View File
@@ -102,7 +102,6 @@ from langgraph.pregel._checkpoint import (
copy_checkpoint,
create_checkpoint,
delta_channels_to_snapshot,
delta_channels_with_pending_writes,
empty_checkpoint,
exit_delta_task_id,
)
@@ -223,13 +222,10 @@ class PregelLoop:
# under the saver's `ORDER BY task_id, idx` sorting.
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
# Delta channels that must snapshot at the next checkpoint, whatever their
# cadence counters say:
# * an Overwrite arrived since the last checkpoint, so sparse replay has to
# start from the post-overwrite value;
# * the checkpoint this run starts from has pending writes to them; see
# `delta_channels_with_pending_writes`.
_delta_channels_forced_snapshot: set[str]
# Delta channels that saw an Overwrite since the last checkpoint. These
# channels must snapshot after live update applies overwrite semantics so
# sparse replay starts from the same post-overwrite value.
_delta_channels_with_overwrite: set[str]
# The checkpoint_config that points at the parent loaded at `__enter__`
# (or the synthetic-empty checkpoint, on first run). We capture it
@@ -687,7 +683,7 @@ class PregelLoop:
def after_tick(self) -> None:
# finish superstep
writes = [w for t in self.tasks.values() for w in t.writes]
self._delta_channels_forced_snapshot.update(
self._delta_channels_with_overwrite.update(
ch
for ch, v in writes
if isinstance(self.specs.get(ch), DeltaChannel) and _get_overwrite(v)[0]
@@ -902,15 +898,6 @@ class PregelLoop:
self.checkpoint_pending_writes = [
w for w in self.checkpoint_pending_writes if w[1] != RESUME
]
# A resume that is not replaying reuses the head's pending writes
# instead of rerunning their tasks, so none of them can leak.
self._delta_channels_forced_snapshot = (
set()
if is_resuming and not self.is_replaying
else delta_channels_with_pending_writes(
self.specs, self.checkpoint_pending_writes
)
)
# map command to writes
if input_is_command:
@@ -1004,7 +991,7 @@ class PregelLoop:
manager=None,
updated_channels=updated_channels,
)
self._delta_channels_forced_snapshot.update(
self._delta_channels_with_overwrite.update(
c
for c, v in input_writes
if isinstance(self.specs.get(c), DeltaChannel) and _get_overwrite(v)[0]
@@ -1149,7 +1136,7 @@ class PregelLoop:
# create new checkpoint
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, new_counters)
| self._delta_channels_forced_snapshot
| self._delta_channels_with_overwrite
if do_checkpoint
else set()
)
@@ -1167,7 +1154,7 @@ class PregelLoop:
for k in channels_to_snapshot:
new_counters[k] = (0, 0)
if do_checkpoint:
self._delta_channels_forced_snapshot.difference_update(channels_to_snapshot)
self._delta_channels_with_overwrite.difference_update(channels_to_snapshot)
non_zero = {k: v for k, v in new_counters.items() if v != (0, 0)}
if non_zero:
self.checkpoint_metadata["counters_since_delta_snapshot"] = non_zero
@@ -1252,7 +1239,7 @@ class PregelLoop:
)
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, counters)
| self._delta_channels_forced_snapshot
| self._delta_channels_with_overwrite
)
pending = [
@@ -1697,6 +1684,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
)
self._delta_write_futs = []
self._error_handler_write_futs = []
self._delta_channels_with_overwrite = set()
self._exit_delta_writes = (
[] if self.durability == "exit" and self.checkpointer is not None else None
)
@@ -1954,6 +1942,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
)
self._delta_write_futs = []
self._error_handler_write_futs = []
self._delta_channels_with_overwrite = set()
self._exit_delta_writes = (
[] if self.durability == "exit" and self.checkpointer is not None else None
)
+38 -145
View File
@@ -108,7 +108,6 @@ from langgraph.callbacks import (
get_sync_graph_callback_manager_for_config,
)
from langgraph.channels.base import BaseChannel
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.topic import Topic
from langgraph.config import get_config
from langgraph.constants import END
@@ -134,7 +133,6 @@ from langgraph.pregel._checkpoint import (
copy_checkpoint,
create_checkpoint,
create_checkpoint_plan_for_update_state_api,
delta_channels_with_pending_writes,
empty_checkpoint,
get_updated_channels_from_tasks,
)
@@ -1639,22 +1637,12 @@ class Pregel(
else:
raise ValueError(f"Subgraph {recast} not found")
# Taken from the first superstep's base, and cleared by the first
# checkpoint that carries the snapshots, which `__copy__` does not write.
fork_pending: set[str] | None = None
def perform_superstep(
input_config: RunnableConfig, updates: Sequence[StateUpdate]
) -> RunnableConfig:
nonlocal fork_pending
# get last checkpoint
config = ensure_config(self.config, input_config)
saved = checkpointer.get_tuple(config)
first_superstep = fork_pending is None
if fork_pending is None:
fork_pending = delta_channels_with_pending_writes(
self.channels, saved.pending_writes if saved else None
)
if saved is not None:
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
@@ -1738,17 +1726,9 @@ class Pregel(
self.trigger_to_nodes,
)
# save checkpoint
next_checkpoint = create_checkpoint(
checkpoint,
channels,
step,
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=fork_pending,
)
fork_pending.difference_update(next_checkpoint["channel_values"])
next_config = checkpointer.put(
checkpoint_config,
next_checkpoint,
create_checkpoint(checkpoint, channels, step),
{
"source": "update",
"step": step + 1,
@@ -1756,7 +1736,7 @@ class Pregel(
},
get_new_channel_versions(
checkpoint_previous_versions,
next_checkpoint["channel_versions"],
checkpoint["channel_versions"],
),
)
return patch_checkpoint_map(
@@ -1785,17 +1765,9 @@ class Pregel(
if saved and saved.metadata.get("step") is not None
else -1
)
next_checkpoint = create_checkpoint(
checkpoint,
channels,
next_step,
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=fork_pending,
)
fork_pending.difference_update(next_checkpoint["channel_values"])
next_config = checkpointer.put(
checkpoint_config,
next_checkpoint,
create_checkpoint(checkpoint, channels, next_step),
{
"source": "input",
"step": next_step,
@@ -1805,7 +1777,7 @@ class Pregel(
},
get_new_channel_versions(
checkpoint_previous_versions,
next_checkpoint["channel_versions"],
checkpoint["channel_versions"],
),
)
@@ -1978,7 +1950,7 @@ class Pregel(
run_tasks: list[PregelTaskWrites] = []
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
writers = self.nodes[as_node].flat_writers
if not writers:
@@ -1992,7 +1964,7 @@ class Pregel(
task_id = provided_task_id or (
prepared_task_ids.popleft()
if prepared_task_ids
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
else _update_task_id(checkpoint["id"], i)
)
run_tasks.append(task)
run_task_ids.append(task_id)
@@ -2026,21 +1998,13 @@ class Pregel(
),
)
updated_channels = get_updated_channels_from_tasks(run_tasks)
edited_delta_channels = {
ch
for ch in updated_channels
if isinstance(self.channels.get(ch), DeltaChannel)
}
# The base's other children replay whatever is stored on it, so an
# edit of an older checkpoint snapshots its delta channels here
# instead. Later supersteps address the checkpoint just written.
if (
first_superstep
and saved is not None
and edited_delta_channels
and _is_older_checkpoint(checkpointer, config, saved)
):
fork_pending.update(edited_delta_channels)
if saved is not None:
for task_id, task in zip(run_task_ids, run_tasks):
channel_writes = [w for w in task.writes if w[0] != PUSH]
if channel_writes:
checkpointer.put_writes(
checkpoint_config, channel_writes, task_id
)
apply_writes(
checkpoint,
channels,
@@ -2056,29 +2020,18 @@ class Pregel(
parents=saved.metadata.get("parents", {}) if saved else {},
saved_metadata=saved.metadata if saved else None,
is_fresh_thread=saved is None,
fork_channels=fork_pending,
)
)
checkpoint = create_checkpoint(
checkpoint,
channels,
step + 1,
updated_channels=updated_channels if channels_to_snapshot else None,
get_next_version=checkpointer.get_next_version
if channels_to_snapshot
else None,
channels_to_snapshot=channels_to_snapshot,
)
sealed = fork_pending.intersection(checkpoint["channel_values"])
fork_pending.difference_update(checkpoint["channel_values"])
if saved is not None:
for task_id, task in zip(run_task_ids, run_tasks):
channel_writes = [
w for w in task.writes if w[0] != PUSH and w[0] not in sealed
]
if channel_writes:
checkpointer.put_writes(
checkpoint_config, channel_writes, task_id
)
next_config = checkpointer.put(
checkpoint_config,
checkpoint,
@@ -2150,22 +2103,12 @@ class Pregel(
else:
raise ValueError(f"Subgraph {recast} not found")
# Taken from the first superstep's base, and cleared by the first
# checkpoint that carries the snapshots, which `__copy__` does not write.
fork_pending: set[str] | None = None
async def aperform_superstep(
input_config: RunnableConfig, updates: Sequence[StateUpdate]
) -> RunnableConfig:
nonlocal fork_pending
# get last checkpoint
config = ensure_config(self.config, input_config)
saved = await checkpointer.aget_tuple(config)
first_superstep = fork_pending is None
if fork_pending is None:
fork_pending = delta_channels_with_pending_writes(
self.channels, saved.pending_writes if saved else None
)
if saved is not None:
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
@@ -2247,25 +2190,16 @@ class Pregel(
self.trigger_to_nodes,
)
# save checkpoint
next_checkpoint = create_checkpoint(
checkpoint,
channels,
step,
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=fork_pending,
)
fork_pending.difference_update(next_checkpoint["channel_values"])
next_config = await checkpointer.aput(
checkpoint_config,
next_checkpoint,
create_checkpoint(checkpoint, channels, step),
{
"source": "update",
"step": step + 1,
"parents": saved.metadata.get("parents", {}) if saved else {},
},
get_new_channel_versions(
checkpoint_previous_versions,
next_checkpoint["channel_versions"],
checkpoint_previous_versions, checkpoint["channel_versions"]
),
)
return patch_checkpoint_map(
@@ -2294,17 +2228,9 @@ class Pregel(
if saved and saved.metadata.get("step") is not None
else -1
)
next_checkpoint = create_checkpoint(
checkpoint,
channels,
next_step,
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=fork_pending,
)
fork_pending.difference_update(next_checkpoint["channel_values"])
next_config = await checkpointer.aput(
checkpoint_config,
next_checkpoint,
create_checkpoint(checkpoint, channels, next_step),
{
"source": "input",
"step": next_step,
@@ -2314,7 +2240,7 @@ class Pregel(
},
get_new_channel_versions(
checkpoint_previous_versions,
next_checkpoint["channel_versions"],
checkpoint["channel_versions"],
),
)
@@ -2484,7 +2410,7 @@ class Pregel(
run_tasks: list[PregelTaskWrites] = []
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
writers = self.nodes[as_node].flat_writers
if not writers:
@@ -2498,7 +2424,7 @@ class Pregel(
task_id = provided_task_id or (
prepared_task_ids.popleft()
if prepared_task_ids
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
else _update_task_id(checkpoint["id"], i)
)
run_tasks.append(task)
run_task_ids.append(task_id)
@@ -2532,21 +2458,13 @@ class Pregel(
),
)
updated_channels = get_updated_channels_from_tasks(run_tasks)
edited_delta_channels = {
ch
for ch in updated_channels
if isinstance(self.channels.get(ch), DeltaChannel)
}
# The base's other children replay whatever is stored on it, so an
# edit of an older checkpoint snapshots its delta channels here
# instead. Later supersteps address the checkpoint just written.
if (
first_superstep
and saved is not None
and edited_delta_channels
and await _ais_older_checkpoint(checkpointer, config, saved)
):
fork_pending.update(edited_delta_channels)
if saved is not None:
for task_id, task in zip(run_task_ids, run_tasks):
channel_writes = [w for w in task.writes if w[0] != PUSH]
if channel_writes:
await checkpointer.aput_writes(
checkpoint_config, channel_writes, task_id
)
apply_writes(
checkpoint,
channels,
@@ -2562,29 +2480,18 @@ class Pregel(
parents=saved.metadata.get("parents", {}) if saved else {},
saved_metadata=saved.metadata if saved else None,
is_fresh_thread=saved is None,
fork_channels=fork_pending,
)
)
checkpoint = create_checkpoint(
checkpoint,
channels,
step + 1,
updated_channels=updated_channels if channels_to_snapshot else None,
get_next_version=checkpointer.get_next_version
if channels_to_snapshot
else None,
channels_to_snapshot=channels_to_snapshot,
)
sealed = fork_pending.intersection(checkpoint["channel_values"])
fork_pending.difference_update(checkpoint["channel_values"])
if saved is not None:
for task_id, task in zip(run_task_ids, run_tasks):
channel_writes = [
w for w in task.writes if w[0] != PUSH and w[0] not in sealed
]
if channel_writes:
await checkpointer.aput_writes(
checkpoint_config, channel_writes, task_id
)
next_config = await checkpointer.aput(
checkpoint_config,
checkpoint,
@@ -4265,6 +4172,16 @@ class Pregel(
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]]:
"""Index from a trigger to nodes that depend on it."""
trigger_to_nodes: defaultdict[str, list[str]] = defaultdict(list)
@@ -4274,30 +4191,6 @@ def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str
return dict(trigger_to_nodes)
def _is_older_checkpoint(
checkpointer: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple
) -> bool:
"""Whether `config` addressed a checkpoint the thread has moved past."""
if not config[CONF].get(CONFIG_KEY_CHECKPOINT_ID):
return False
latest = checkpointer.get_tuple(
patch_configurable(config, {CONFIG_KEY_CHECKPOINT_ID: None})
)
return latest is not None and latest.checkpoint["id"] != saved.checkpoint["id"]
async def _ais_older_checkpoint(
checkpointer: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple
) -> bool:
"""Whether `config` addressed a checkpoint the thread has moved past."""
if not config[CONF].get(CONFIG_KEY_CHECKPOINT_ID):
return False
latest = await checkpointer.aget_tuple(
patch_configurable(config, {CONFIG_KEY_CHECKPOINT_ID: None})
)
return latest is not None and latest.checkpoint["id"] != saved.checkpoint["id"]
def _output(
stream_mode: StreamMode | Sequence[StreamMode],
print_mode: StreamMode | Sequence[StreamMode],
+3 -5
View File
@@ -85,13 +85,11 @@ class MemorySaverAssertImmutable(InMemorySaver):
)
== saved
), config["configurable"]["checkpoint_ns"]
next_config = super().put(config, checkpoint, metadata, new_versions)
# Read back, not the object handed in: a DeltaChannel a step did not
# write is refilled on read from the blob its inherited version points at.
self.storage_for_copies[thread_id][checkpoint_ns][checkpoint["id"]] = (
self.serde.dumps_typed(super().get(next_config))
self.serde.dumps_typed(checkpoint)
)
return next_config
# call super to write checkpoint
return super().put(config, checkpoint, metadata, new_versions)
class MemorySaverNoPending(InMemorySaver):
@@ -1,645 +0,0 @@
"""Forking a thread must not replay the abandoned branch into the fork.
Every graph carries a `DeltaChannel` and a plain reducer channel fed the same
values; the plain channel needs no replay, so it is the oracle.
"""
from collections.abc import Sequence
from operator import add
from typing import Annotated, Any
import pytest
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from typing_extensions import TypedDict
from langgraph._internal._constants import INPUT
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import END, START, StateGraph
from langgraph.types import Command, Durability, StateSnapshot, StateUpdate, interrupt
pytestmark = pytest.mark.anyio
def _append(current: list | None, writes: Sequence[Any]) -> list:
out = list(current or [])
for write in writes:
out.extend(write if isinstance(write, list) else [write])
return out
class _State(TypedDict):
log: Annotated[list, DeltaChannel(_append, snapshot_frequency=1000)]
plain: Annotated[list, add]
other: Annotated[list, add]
def _build(checkpointer: BaseCheckpointSaver, tag: str) -> Any:
def node(state: _State) -> dict:
return {"log": [f"{tag}-out"], "plain": [f"{tag}-out"]}
builder = StateGraph(_State)
builder.add_node("n", node)
builder.set_entry_point("n")
builder.set_finish_point("n")
return builder.compile(checkpointer=checkpointer)
def _build_without_delta_writes(checkpointer: BaseCheckpointSaver, tag: str) -> Any:
def node(state: _State) -> dict:
return {"other": [f"{tag}-other"]}
builder = StateGraph(_State)
builder.add_node("n", node)
builder.set_entry_point("n")
builder.set_finish_point("n")
return builder.compile(checkpointer=checkpointer)
def _thread(thread_id: str) -> RunnableConfig:
return {"configurable": {"thread_id": thread_id}}
def _at(config: RunnableConfig, snapshot: StateSnapshot) -> RunnableConfig:
return {
"configurable": {
**config["configurable"],
"checkpoint_ns": "",
"checkpoint_id": snapshot.config["configurable"]["checkpoint_id"],
}
}
def _both(marker: str) -> dict:
return {"log": [marker], "plain": [marker]}
def _snapshotted_checkpoints(
checkpointer: BaseCheckpointSaver, config: RunnableConfig
) -> list[str]:
return [
tuple_.config["configurable"]["checkpoint_id"]
for tuple_ in checkpointer.list(config)
if isinstance(tuple_.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
]
def _assert_fork_is_clean(state: StateSnapshot, abandoned: str) -> None:
assert state.values["log"] == state.values["plain"], (
f"delta channel diverged from the plain channel: "
f"{state.values['log']} != {state.values['plain']}"
)
assert abandoned not in state.values["log"], (
f"{abandoned!r} belongs to the branch the fork replaced, "
f"but was replayed into {state.values['log']}"
)
def test_fork_by_invoke(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
_build(sync_checkpointer, "first").invoke(
_both("in-1"), config, durability=durability
)
graph = _build(sync_checkpointer, "second")
graph.invoke(_both("in-2"), config, durability=durability)
abandoned_head = graph.get_state(config)
base = next(
snapshot
for snapshot in graph.get_state_history(config)
if "in-2" not in snapshot.values["log"]
)
_build(sync_checkpointer, "third").invoke(
_both("in-3"), _at(config, base), durability=durability
)
state = graph.get_state(config)
_assert_fork_is_clean(state, "in-2")
assert state.values["log"] == [*base.values["log"], "in-3", "third-out"]
abandoned = graph.get_state(abandoned_head.config).values
assert abandoned["log"] == abandoned["plain"] == abandoned_head.values["log"]
async def test_afork_by_invoke(
async_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
await _build(async_checkpointer, "first").ainvoke(
_both("in-1"), config, durability=durability
)
graph = _build(async_checkpointer, "second")
await graph.ainvoke(_both("in-2"), config, durability=durability)
abandoned_head = await graph.aget_state(config)
base = await anext(
snapshot
async for snapshot in graph.aget_state_history(config)
if "in-2" not in snapshot.values["log"]
)
await _build(async_checkpointer, "third").ainvoke(
_both("in-3"), _at(config, base), durability=durability
)
state = await graph.aget_state(config)
_assert_fork_is_clean(state, "in-2")
assert state.values["log"] == [*base.values["log"], "in-3", "third-out"]
abandoned = (await graph.aget_state(abandoned_head.config)).values
assert abandoned["log"] == abandoned["plain"] == abandoned_head.values["log"]
def test_fork_off_checkpoint_before_first_input(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build(sync_checkpointer, "first")
graph.invoke(_both("in-1"), config, durability=durability)
root = list(graph.get_state_history(config))[-1]
assert root.values["log"] == []
_build(sync_checkpointer, "third").invoke(
_both("in-9"), _at(config, root), durability=durability
)
state = graph.get_state(config)
_assert_fork_is_clean(state, "in-1")
assert state.values["log"] == ["in-9", "third-out"]
async def test_afork_off_checkpoint_before_first_input(
async_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build(async_checkpointer, "first")
await graph.ainvoke(_both("in-1"), config, durability=durability)
root = [snapshot async for snapshot in graph.aget_state_history(config)][-1]
assert root.values["log"] == []
await _build(async_checkpointer, "third").ainvoke(
_both("in-9"), _at(config, root), durability=durability
)
state = await graph.aget_state(config)
_assert_fork_is_clean(state, "in-1")
assert state.values["log"] == ["in-9", "third-out"]
def test_fork_by_update_state(sync_checkpointer: BaseCheckpointSaver) -> None:
config = _thread("t")
_build(sync_checkpointer, "first").invoke(_both("in-1"), config)
graph = _build(sync_checkpointer, "second")
graph.invoke(_both("in-2"), config)
base = next(
snapshot
for snapshot in graph.get_state_history(config)
if "in-2" not in snapshot.values["log"]
)
forked = graph.update_state(_at(config, base), _both("patched"))
state = graph.get_state(forked)
_assert_fork_is_clean(state, "in-2")
assert state.values["log"] == [*base.values["log"], "patched"]
async def test_afork_by_update_state(
async_checkpointer: BaseCheckpointSaver,
) -> None:
config = _thread("t")
await _build(async_checkpointer, "first").ainvoke(_both("in-1"), config)
graph = _build(async_checkpointer, "second")
await graph.ainvoke(_both("in-2"), config)
base = await anext(
snapshot
async for snapshot in graph.aget_state_history(config)
if "in-2" not in snapshot.values["log"]
)
forked = await graph.aupdate_state(_at(config, base), _both("patched"))
state = await graph.aget_state(forked)
_assert_fork_is_clean(state, "in-2")
assert state.values["log"] == [*base.values["log"], "patched"]
def _assert_branch_unchanged(state: StateSnapshot, expected: list, edit: str) -> None:
assert state.values["log"] == state.values["plain"] == expected, (
f"{edit!r} was written by an update_state on this branch's base, "
f"but this branch now reads {state.values['log']}"
)
# The old checkpoint is either a finished turn, which saved no writes, or one
# whose next node already ran there, so the edit reuses that task's id.
@pytest.mark.parametrize("next_node_ran", [False, True])
def test_update_state_on_an_old_checkpoint_leaves_its_other_branch_alone(
sync_checkpointer: BaseCheckpointSaver, next_node_ran: bool
) -> None:
config = _thread("t")
graph = _build(sync_checkpointer, "first")
graph.invoke(_both("in-1"), config)
_build(sync_checkpointer, "second").invoke(_both("in-2"), config)
branch = graph.get_state(config)
base = next(
snapshot
for snapshot in graph.get_state_history(config)
if "in-2" not in snapshot.values["log"]
and snapshot.next == (("n",) if next_node_ran else ())
)
edited = graph.update_state(_at(config, base), _both("edit"), as_node="n")
_assert_branch_unchanged(
graph.get_state(branch.config), branch.values["log"], "edit"
)
assert graph.get_state(edited).values["log"] == [*base.values["log"], "edit"]
_build(sync_checkpointer, "third").invoke(_both("in-3"), branch.config)
_assert_branch_unchanged(
graph.get_state(config),
[*branch.values["log"], "in-3", "third-out"],
"edit",
)
@pytest.mark.parametrize("next_node_ran", [False, True])
async def test_aupdate_state_on_an_old_checkpoint_leaves_its_other_branch_alone(
async_checkpointer: BaseCheckpointSaver, next_node_ran: bool
) -> None:
config = _thread("t")
graph = _build(async_checkpointer, "first")
await graph.ainvoke(_both("in-1"), config)
await _build(async_checkpointer, "second").ainvoke(_both("in-2"), config)
branch = await graph.aget_state(config)
base = await anext(
snapshot
async for snapshot in graph.aget_state_history(config)
if "in-2" not in snapshot.values["log"]
and snapshot.next == (("n",) if next_node_ran else ())
)
edited = await graph.aupdate_state(_at(config, base), _both("edit"), as_node="n")
_assert_branch_unchanged(
await graph.aget_state(branch.config), branch.values["log"], "edit"
)
assert (await graph.aget_state(edited)).values["log"] == [
*base.values["log"],
"edit",
]
await _build(async_checkpointer, "third").ainvoke(_both("in-3"), branch.config)
_assert_branch_unchanged(
await graph.aget_state(config),
[*branch.values["log"], "in-3", "third-out"],
"edit",
)
def test_bulk_update_on_an_old_checkpoint_leaves_its_other_branch_alone(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
config = _thread("t")
graph = _build(sync_checkpointer, "first")
graph.invoke(_both("in-1"), config)
base = graph.get_state(config)
_build(sync_checkpointer, "second").invoke(_both("in-2"), config)
branch = graph.get_state(config)
edited = graph.bulk_update_state(
_at(config, base),
[[StateUpdate(_both("s1"), "n")], [StateUpdate(_both("s2"), "n")]],
)
_assert_branch_unchanged(graph.get_state(branch.config), branch.values["log"], "s1")
assert graph.get_state(edited).values["log"] == [*base.values["log"], "s1", "s2"]
def test_update_state_with_the_head_checkpoint_id_stores_no_snapshot(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
config = _thread("t")
graph = _build(sync_checkpointer, "first")
graph.invoke(_both("in-1"), config)
for i in range(3):
graph.update_state(graph.get_state(config).config, _both(f"u{i}"))
assert not _snapshotted_checkpoints(sync_checkpointer, config)
assert graph.get_state(config).values["log"] == [
"in-1",
"first-out",
"u0",
"u1",
"u2",
]
def test_unaddressed_run_keeps_snapshot_cadence(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build(sync_checkpointer, "first")
graph.invoke(_both("in-1"), config, durability=durability)
graph.invoke(_both("in-2"), config, durability=durability)
assert not _snapshotted_checkpoints(sync_checkpointer, config)
def test_fork_before_first_value_when_fork_never_writes_the_channel(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build(sync_checkpointer, "first")
graph.invoke(_both("in-1"), config, durability=durability)
root = list(graph.get_state_history(config))[-1]
assert root.values["log"] == []
_build_without_delta_writes(sync_checkpointer, "third").invoke(
{"other": ["in-9"]}, _at(config, root), durability=durability
)
state = graph.get_state(config)
_assert_fork_is_clean(state, "in-1")
assert state.values["log"] == []
async def test_afork_before_first_value_when_fork_never_writes_the_channel(
async_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build(async_checkpointer, "first")
await graph.ainvoke(_both("in-1"), config, durability=durability)
root = [snapshot async for snapshot in graph.aget_state_history(config)][-1]
assert root.values["log"] == []
await _build_without_delta_writes(async_checkpointer, "third").ainvoke(
{"other": ["in-9"]}, _at(config, root), durability=durability
)
state = await graph.aget_state(config)
_assert_fork_is_clean(state, "in-1")
assert state.values["log"] == []
def test_fork_before_first_value_by_bulk_update(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
config = _thread("t")
graph = _build(sync_checkpointer, "first")
graph.invoke(_both("in-1"), config)
root = list(graph.get_state_history(config))[-1]
assert root.values["log"] == []
forked = graph.bulk_update_state(
_at(config, root),
[
[StateUpdate({"other": ["s1"]}, "n")],
[StateUpdate(_both("s2"), "n")],
],
)
state = graph.get_state(forked)
_assert_fork_is_clean(state, "in-1")
assert state.values["log"] == ["s2"]
@pytest.mark.parametrize("first_as_node", [INPUT, END, "__copy__"])
def test_fork_by_bulk_update_whose_first_superstep_skips_the_plan(
sync_checkpointer: BaseCheckpointSaver, first_as_node: str
) -> None:
config = _thread("t")
_build(sync_checkpointer, "first").invoke(_both("in-1"), config)
graph = _build(sync_checkpointer, "second")
graph.invoke(_both("in-2"), config)
base = next(
snapshot
for snapshot in graph.get_state_history(config)
if "in-2" not in snapshot.values["log"]
)
first = (
StateUpdate(_both("first-step"), first_as_node)
if first_as_node == INPUT
else StateUpdate(None, first_as_node)
)
forked = graph.bulk_update_state(
_at(config, base),
[[first], [StateUpdate(_both("second-step"), "n")]],
)
state = graph.get_state(forked)
assert state.values["log"] == state.values["plain"], (
f"delta channel diverged from the plain channel: "
f"{state.values['log']} != {state.values['plain']}"
)
def test_unaddressed_bulk_update_keeps_snapshot_cadence(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
config = _thread("t")
graph = _build(sync_checkpointer, "first")
graph.invoke(_both("in-1"), config)
graph.bulk_update_state(
config,
[[StateUpdate(_both(f"u{i}"), "n")] for i in range(4)],
)
assert not _snapshotted_checkpoints(sync_checkpointer, config)
def _build_paused_before_b(checkpointer: BaseCheckpointSaver) -> Any:
builder = StateGraph(_State)
builder.add_node("a", lambda state: _both("a"))
builder.add_node("b", lambda state: _both("b"))
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("b", END)
return builder.compile(checkpointer=checkpointer, interrupt_before=["b"])
def _build_parallel_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
def ask(state: _State) -> dict:
interrupt("approve?")
return {"other": ["q"]}
builder = StateGraph(_State)
builder.add_node("p", lambda state: _both("p"))
builder.add_node("q", ask)
builder.add_edge(START, "p")
builder.add_edge(START, "q")
return builder.compile(checkpointer=checkpointer)
def test_resume_at_interrupt_before_with_the_head_checkpoint_id_runs_the_node(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build_paused_before_b(sync_checkpointer)
graph.invoke(_both("in"), config, durability=durability)
graph.invoke(None, graph.get_state(config).config, durability=durability)
state = graph.get_state(config)
assert state.next == (), f"resume paused again before {state.next}"
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
async def test_aresume_at_interrupt_before_with_the_head_checkpoint_id_runs_the_node(
async_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build_paused_before_b(async_checkpointer)
await graph.ainvoke(_both("in"), config, durability=durability)
await graph.ainvoke(
None, (await graph.aget_state(config)).config, durability=durability
)
state = await graph.aget_state(config)
assert state.next == (), f"resume paused again before {state.next}"
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
def test_replay_from_a_paused_checkpoint_runs_the_node_once(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
config = _thread("t")
graph = _build_paused_before_b(sync_checkpointer)
graph.invoke(_both("in"), config)
paused = graph.get_state(config).config
graph.invoke(None, config)
graph.invoke(None, paused)
state = graph.get_state(config)
assert state.next == (), f"replay paused again before {state.next}"
assert state.values["log"] == state.values["plain"] == ["in", "a", "b"]
@pytest.mark.parametrize("addressed", [False, True])
def test_new_input_on_an_interrupted_head_does_not_replay_its_pending_writes(
sync_checkpointer: BaseCheckpointSaver, durability: Durability, addressed: bool
) -> None:
config = _thread("t")
graph = _build_parallel_interrupt(sync_checkpointer)
graph.invoke(_both("in-1"), config, durability=durability)
head = graph.get_state(config).config
graph.invoke(_both("in-2"), head if addressed else config, durability=durability)
state = graph.get_state(config)
assert state.values["log"] == state.values["plain"] == ["in-1", "in-2", "p"]
def _build_deferred_after_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
builder = StateGraph(_State)
builder.add_node("a", lambda state: _both("a"))
builder.add_node("b", lambda state: _both("b"), defer=True)
builder.add_node("c", lambda state: {})
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("a", "c")
return builder.compile(checkpointer=checkpointer, interrupt_after=["a"])
@pytest.mark.parametrize(
"durability",
[
"sync",
"async",
pytest.param(
"exit",
marks=pytest.mark.xfail(
reason="exit durability stores a resumed run's loaded writes twice",
strict=True,
),
),
],
)
def test_resume_on_an_interrupted_head_consumes_its_writes_without_a_snapshot(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build_parallel_interrupt(sync_checkpointer)
graph.invoke(_both("in-1"), config, durability=durability)
graph.invoke(Command(resume="yes"), config, durability=durability)
state = graph.get_state(config)
assert state.next == ()
assert state.values["log"] == state.values["plain"] == ["in-1", "p"]
assert not _snapshotted_checkpoints(sync_checkpointer, config)
def test_resume_addressed_at_an_interrupted_head_reruns_its_tasks_once(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
config = _thread("t")
graph = _build_parallel_interrupt(sync_checkpointer)
graph.invoke(_both("in-1"), config, durability=durability)
graph.invoke(
Command(resume="yes"), graph.get_state(config).config, durability=durability
)
state = graph.get_state(config)
assert state.next == ()
assert state.values["log"] == state.values["plain"] == ["in-1", "p"]
def test_update_state_with_the_head_checkpoint_id_keeps_a_deferred_node(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_deferred_after_interrupt(sync_checkpointer)
config = _thread("t")
graph.invoke(_both("in"), config)
graph.update_state(graph.get_state(config).config, _both("u"), as_node="c")
graph.invoke(None, config)
state = graph.get_state(config)
assert state.next == (), f"deferred node never ran, still pending: {state.next}"
assert state.values["log"] == state.values["plain"] == ["in", "a", "u", "b"]
async def test_aupdate_state_with_the_head_checkpoint_id_keeps_a_deferred_node(
async_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_deferred_after_interrupt(async_checkpointer)
config = _thread("t")
await graph.ainvoke(_both("in"), config)
await graph.aupdate_state(
(await graph.aget_state(config)).config, _both("u"), as_node="c"
)
await graph.ainvoke(None, config)
state = await graph.aget_state(config)
assert state.next == (), f"deferred node never ran, still pending: {state.next}"
assert state.values["log"] == state.values["plain"] == ["in", "a", "u", "b"]
def test_turns_addressed_at_the_head_store_no_snapshot(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
config = _thread("t")
graph = _build(sync_checkpointer, "turn")
graph.invoke(_both("in-1"), config)
for turn in range(2, 5):
graph.invoke(_both(f"in-{turn}"), graph.get_state(config).config)
assert not _snapshotted_checkpoints(sync_checkpointer, config)
assert (
graph.get_state(config).values["log"] == graph.get_state(config).values["plain"]
)
@@ -20,6 +20,7 @@ from typing import Annotated, Any
import pytest
from langchain_core.messages import HumanMessage
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from typing_extensions import TypedDict
@@ -27,16 +28,17 @@ from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import _messages_delta_reducer
from langgraph.types import StateUpdate
from langgraph.types import StateSnapshot, StateUpdate
pytestmark = pytest.mark.anyio
def _build_graph(
checkpointer: InMemorySaver,
checkpointer: BaseCheckpointSaver,
*,
two_nodes: bool = False,
snapshot_frequency: int = 1000,
interrupt_before: list[str] | None = None,
) -> Any:
"""Compile a minimal DeltaChannel-backed `messages` graph.
@@ -63,7 +65,7 @@ def _build_graph(
builder.set_finish_point("assistant")
else:
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
`put_writes` outside the per-task loop would persist only the last
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()
@@ -310,6 +308,92 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
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
# ---------------------------------------------------------------------------
@@ -338,29 +422,3 @@ def test_state_history_chain_after_fresh_update_state_delta_channel() -> None:
assert update_snapshot.metadata["step"] == 0
assert update_snapshot.parent_config is None
assert [m.content for m in update_snapshot.values["messages"]] == ["hello"]
def test_update_state_that_snapshots_keeps_a_deferred_node_pending() -> None:
channel = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1)
class State(TypedDict):
messages: Annotated[list, channel]
builder = StateGraph(State)
builder.add_node("a", lambda state: {"messages": [HumanMessage("a", id="a")]})
builder.add_node(
"b", lambda state: {"messages": [HumanMessage("b", id="b")]}, defer=True
)
builder.add_node("c", lambda state: {})
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("a", "c")
graph = builder.compile(checkpointer=InMemorySaver(), interrupt_after=["a"])
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"messages": [HumanMessage("s", id="s")]}, config)
graph.update_state(config, {"messages": [HumanMessage("u", id="u")]}, as_node="c")
final = graph.invoke(None, config)
assert [m.content for m in final["messages"]] == ["s", "a", "u", "b"]
assert graph.get_state(config).next == ()