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
|
||||
} >> "$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}"
|
||||
@@ -119,7 +119,6 @@ from langgraph.pregel._io import (
|
||||
)
|
||||
from langgraph.pregel._messages import ensure_message_ids
|
||||
from langgraph.pregel._read import PregelNode
|
||||
from langgraph.pregel._task_status import read_task_statuses
|
||||
from langgraph.pregel._utils import get_new_channel_versions, is_xxh3_128_hexdigest
|
||||
from langgraph.pregel.debug import (
|
||||
map_debug_checkpoint,
|
||||
@@ -737,14 +736,17 @@ class PregelLoop:
|
||||
def _reapply_writes_to_succeeded_nodes(
|
||||
self, tasks: Mapping[str, PregelExecutableTask]
|
||||
) -> None:
|
||||
"""Restore the output of finished tasks from checkpoint to in-memory tasks.
|
||||
"""Restore successful channel writes from checkpoint to in-memory tasks.
|
||||
|
||||
Unfinished (failed or interrupted) tasks keep empty writes, so the
|
||||
runner re-executes them or routes them to error handlers.
|
||||
Skips control signals (ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME)
|
||||
so that failed/interrupted tasks remain with empty writes and will be
|
||||
re-executed (or routed to error handlers) by the runner.
|
||||
"""
|
||||
for tid, status in read_task_statuses(self.checkpoint_pending_writes).items():
|
||||
for tid, k, v in self.checkpoint_pending_writes:
|
||||
if k in (ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME):
|
||||
continue
|
||||
if task := tasks.get(tid):
|
||||
task.writes.extend(status.output)
|
||||
task.writes.append((k, v))
|
||||
|
||||
def _resume_error_handlers_if_applicable(self) -> None:
|
||||
"""On resume, schedule error handlers for tasks that failed in a prior run.
|
||||
@@ -814,13 +816,35 @@ class PregelLoop:
|
||||
self.tasks[handler_task.id] = handler_task
|
||||
|
||||
def _pending_interrupts(self) -> set[str]:
|
||||
"""Return the ids of interrupts that are still waiting for an answer."""
|
||||
return {
|
||||
interrupt.id
|
||||
for status in read_task_statuses(self.checkpoint_pending_writes).values()
|
||||
for interrupt in status.pending_interrupts
|
||||
"""Return the set of interrupt ids that are pending without corresponding resume values."""
|
||||
# mapping of task ids to interrupt ids
|
||||
pending_interrupts: dict[str, str] = {}
|
||||
|
||||
# set of resume task ids
|
||||
pending_resumes: set[str] = set()
|
||||
|
||||
for task_id, write_type, value in self.checkpoint_pending_writes:
|
||||
if write_type == INTERRUPT:
|
||||
# interrupts is always a list, but there should only be one element
|
||||
pending_interrupts[task_id] = value[0].id
|
||||
elif write_type == RESUME:
|
||||
pending_resumes.add(task_id)
|
||||
|
||||
resumed_interrupt_ids = {
|
||||
pending_interrupts[task_id]
|
||||
for task_id in pending_resumes
|
||||
if task_id in pending_interrupts
|
||||
}
|
||||
|
||||
# Keep only interrupts whose interrupt_id is not resumed
|
||||
hanging_interrupts: set[str] = {
|
||||
interrupt_id
|
||||
for interrupt_id in pending_interrupts.values()
|
||||
if interrupt_id not in resumed_interrupt_ids
|
||||
}
|
||||
|
||||
return hanging_interrupts
|
||||
|
||||
def _first(
|
||||
self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None
|
||||
) -> set[str] | None:
|
||||
|
||||
@@ -45,7 +45,6 @@ from langgraph.errors import GraphBubbleUp, GraphInterrupt
|
||||
from langgraph.pregel._algo import Call
|
||||
from langgraph.pregel._executor import Submit
|
||||
from langgraph.pregel._retry import arun_with_retry, run_with_retry
|
||||
from langgraph.pregel._task_status import CONTROL_WRITES
|
||||
from langgraph.types import (
|
||||
CachePolicy,
|
||||
PregelExecutableTask,
|
||||
@@ -607,9 +606,8 @@ class PregelRunner:
|
||||
task.config is None or TAG_HIDDEN not in task.config.get("tags", [])
|
||||
):
|
||||
self.node_finished(task.name)
|
||||
if all(chan in CONTROL_WRITES for chan, _ in task.writes):
|
||||
# record that the task finished, even if it produced no output
|
||||
# (see `langgraph.pregel._task_status`)
|
||||
if not task.writes:
|
||||
# add no writes marker
|
||||
task.writes.append((NO_WRITES, None))
|
||||
# save task writes to checkpointer
|
||||
self.put_writes()(task.id, task.writes) # type: ignore[misc]
|
||||
|
||||
@@ -1,127 +0,0 @@
|
||||
"""Read the status of each task from the writes recorded for a superstep.
|
||||
|
||||
While a superstep is open, the checkpointer keeps a log of writes for each
|
||||
task in that step. Entries are added as tasks run and are only discarded when
|
||||
the whole superstep finishes and a new checkpoint is saved. When a task runs
|
||||
again, for example after being resumed, its earlier entries stay in the log.
|
||||
|
||||
This module is the single place that turns that log into task status. Code that
|
||||
needs to know whether a task finished, which interrupts it raised, which of them
|
||||
are still waiting for an answer, or which output it produced must use
|
||||
`read_task_statuses` instead of inspecting the writes directly.
|
||||
|
||||
The log uses two kinds of writes:
|
||||
|
||||
- Control writes describe what happened to a task: `INTERRUPT` (the task asked
|
||||
a question), `RESUME` (answers the task has received), `ERROR`, and
|
||||
`ERROR_SOURCE_NODE`. `INTERRUPT`, `RESUME` and `ERROR` each have a fixed slot
|
||||
per task (`WRITES_IDX_MAP`), so a newer write of the same kind can replace an
|
||||
older one.
|
||||
- Every other write is output: channel writes, `RETURN` for functional tasks,
|
||||
and the `NO_WRITES` marker.
|
||||
|
||||
The rules are:
|
||||
|
||||
1. When a task that ran finishes successfully, `PregelRunner.commit` records at
|
||||
least one output write, adding `NO_WRITES` if the task produced no other
|
||||
output.
|
||||
2. A task that pauses at an interrupt records only control writes.
|
||||
3. A task is therefore treated as finished if and only if it has an output
|
||||
write.
|
||||
4. Because `INTERRUPT` is stored in a fixed slot, its recorded value is the most
|
||||
recent question the task asked. That question is waiting for an answer only
|
||||
while the task is unfinished.
|
||||
|
||||
A `RESUME` write never means a task is finished: it can hold the answer to an
|
||||
earlier question while the task waits on a later one.
|
||||
|
||||
What these rules cannot see:
|
||||
|
||||
- A task whose result came from the cache does not go through
|
||||
`PregelRunner.commit`, so nothing is recorded for it. It reads as not
|
||||
finished.
|
||||
- A task that fails can record partial output writes along with its error. It
|
||||
reads as finished, which is how the executor has always treated it.
|
||||
- Writes recorded before rule 1 existed may describe a finished task with no
|
||||
output using only control writes. Those tasks read as unfinished, which
|
||||
matches how they were treated before.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from langgraph.checkpoint.base import PendingWrite
|
||||
|
||||
from langgraph._internal._constants import (
|
||||
ERROR,
|
||||
ERROR_SOURCE_NODE,
|
||||
INTERRUPT,
|
||||
NULL_TASK_ID,
|
||||
RESUME,
|
||||
)
|
||||
from langgraph.types import Interrupt
|
||||
|
||||
__all__ = ("CONTROL_WRITES", "TaskStatus", "read_task_statuses")
|
||||
|
||||
CONTROL_WRITES = frozenset((ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME))
|
||||
"""Channels that describe what happened to a task rather than what it produced."""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TaskStatus:
|
||||
"""The status of one task, read from the writes recorded for its superstep."""
|
||||
|
||||
output: tuple[tuple[str, Any], ...] = ()
|
||||
"""Output writes in recorded order. Empty if the task has not finished."""
|
||||
|
||||
interrupts: tuple[Interrupt, ...] = ()
|
||||
"""The most recent interrupts the task raised, whether or not they were answered."""
|
||||
|
||||
error: BaseException | None = None
|
||||
"""The recorded error, if any."""
|
||||
|
||||
@property
|
||||
def finished(self) -> bool:
|
||||
"""Whether the task ran to completion."""
|
||||
return bool(self.output)
|
||||
|
||||
@property
|
||||
def pending_interrupts(self) -> tuple[Interrupt, ...]:
|
||||
"""Interrupts waiting for an answer. Always empty for a finished task."""
|
||||
return () if self.finished else self.interrupts
|
||||
|
||||
|
||||
def read_task_statuses(
|
||||
pending_writes: Iterable[PendingWrite],
|
||||
) -> dict[str, TaskStatus]:
|
||||
"""Return the status of every task that has recorded writes, keyed by task id.
|
||||
|
||||
Writes from `NULL_TASK_ID` are input to the superstep, not task activity, so
|
||||
they are not included.
|
||||
"""
|
||||
output: dict[str, list[tuple[str, Any]]] = {}
|
||||
interrupts: dict[str, list[Interrupt]] = {}
|
||||
errors: dict[str, BaseException] = {}
|
||||
for task_id, channel, value in pending_writes:
|
||||
if task_id == NULL_TASK_ID:
|
||||
continue
|
||||
output.setdefault(task_id, [])
|
||||
if channel == INTERRUPT:
|
||||
interrupts.setdefault(task_id, []).extend(
|
||||
value if isinstance(value, Sequence) else [value]
|
||||
)
|
||||
elif channel == ERROR:
|
||||
errors.setdefault(task_id, value)
|
||||
elif channel not in CONTROL_WRITES:
|
||||
output[task_id].append((channel, value))
|
||||
return {
|
||||
task_id: TaskStatus(
|
||||
output=tuple(task_output),
|
||||
interrupts=tuple(interrupts.get(task_id, ())),
|
||||
error=errors.get(task_id),
|
||||
)
|
||||
for task_id, task_output in output.items()
|
||||
}
|
||||
@@ -26,7 +26,6 @@ from langgraph._internal._typing import MISSING
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.constants import TAG_HIDDEN
|
||||
from langgraph.pregel._io import read_channels
|
||||
from langgraph.pregel._task_status import TaskStatus, read_task_statuses
|
||||
from langgraph.types import (
|
||||
CheckpointPayload,
|
||||
PregelExecutableTask,
|
||||
@@ -38,8 +37,6 @@ from langgraph.types import (
|
||||
|
||||
TASK_NAMESPACE = UUID("6ba7b831-9dad-11d1-80b4-00c04fd430c8")
|
||||
|
||||
_NOT_STARTED = TaskStatus()
|
||||
|
||||
|
||||
def map_debug_tasks(tasks: Iterable[PregelExecutableTask]) -> Iterator[TaskPayload]:
|
||||
"""Produce "task" events for stream_mode=debug."""
|
||||
@@ -214,21 +211,35 @@ def tasks_w_writes(
|
||||
pending_writes: list[PendingWrite] | None,
|
||||
states: dict[str, RunnableConfig | StateSnapshot] | None,
|
||||
output_keys: str | Sequence[str],
|
||||
*,
|
||||
live: bool = False,
|
||||
) -> tuple[PregelTask, ...]:
|
||||
"""Apply writes / subgraph states to tasks to be returned in a StateSnapshot.
|
||||
|
||||
With `live=True`, tasks report only the interrupts still waiting for an
|
||||
answer, as of the most recent writes. Otherwise tasks report the interrupts
|
||||
they raised in the step, including answered ones, as a record of the step.
|
||||
"""
|
||||
statuses = read_task_statuses(pending_writes or [])
|
||||
"""Apply writes / subgraph states to tasks to be returned in a StateSnapshot."""
|
||||
pending_writes = pending_writes or []
|
||||
out: list[PregelTask] = []
|
||||
for task in tasks:
|
||||
status = statuses.get(task.id, _NOT_STARTED)
|
||||
rtn = next((val for chan, val in status.output if chan == RETURN), MISSING)
|
||||
task_writes = [(chan, val) for chan, val in status.output if chan != RETURN]
|
||||
rtn = next(
|
||||
(
|
||||
val
|
||||
for tid, chan, val in pending_writes
|
||||
if tid == task.id and chan == RETURN
|
||||
),
|
||||
MISSING,
|
||||
)
|
||||
task_error = next(
|
||||
(exc for tid, n, exc in pending_writes if tid == task.id and n == ERROR),
|
||||
None,
|
||||
)
|
||||
task_interrupts = tuple(
|
||||
v
|
||||
for tid, n, vv in pending_writes
|
||||
if tid == task.id and n == INTERRUPT
|
||||
for v in (vv if isinstance(vv, Sequence) else [vv])
|
||||
)
|
||||
|
||||
task_writes = [
|
||||
(chan, val)
|
||||
for tid, chan, val in pending_writes
|
||||
if tid == task.id and chan not in (ERROR, INTERRUPT, RETURN)
|
||||
]
|
||||
|
||||
if rtn is not MISSING:
|
||||
task_result = rtn
|
||||
@@ -250,15 +261,19 @@ def tasks_w_writes(
|
||||
mapped_writes = map_task_result_writes(filtered_writes)
|
||||
task_result = mapped_writes if filtered_writes else {}
|
||||
|
||||
has_writes = rtn is not MISSING or any(
|
||||
w[0] == task.id and w[1] not in (ERROR, INTERRUPT) for w in pending_writes
|
||||
)
|
||||
|
||||
out.append(
|
||||
PregelTask(
|
||||
task.id,
|
||||
task.name,
|
||||
task.path,
|
||||
status.error,
|
||||
status.pending_interrupts if live else status.interrupts,
|
||||
task_error,
|
||||
task_interrupts,
|
||||
states.get(task.id) if states else None,
|
||||
task_result if status.finished else None,
|
||||
task_result if has_writes else None,
|
||||
)
|
||||
)
|
||||
return tuple(out)
|
||||
|
||||
@@ -79,6 +79,7 @@ from langgraph._internal._constants import (
|
||||
CONFIG_KEY_STREAM_MESSAGES_V2,
|
||||
CONFIG_KEY_TASK_ID,
|
||||
CONFIG_KEY_THREAD_ID,
|
||||
ERROR,
|
||||
INPUT,
|
||||
INTERRUPT,
|
||||
NS_END,
|
||||
@@ -148,7 +149,6 @@ from langgraph.pregel._messages import (
|
||||
from langgraph.pregel._read import DEFAULT_BOUND, PregelNode
|
||||
from langgraph.pregel._retry import RetryPolicy
|
||||
from langgraph.pregel._runner import PregelRunner
|
||||
from langgraph.pregel._task_status import read_task_statuses
|
||||
from langgraph.pregel._tools import StreamToolCallHandler
|
||||
from langgraph.pregel._utils import (
|
||||
get_new_channel_versions,
|
||||
@@ -1147,16 +1147,8 @@ class Pregel(
|
||||
config: RunnableConfig,
|
||||
saved: CheckpointTuple | None,
|
||||
recurse: BaseCheckpointSaver | None = None,
|
||||
live: bool = False,
|
||||
apply_pending_writes: bool = False,
|
||||
) -> StateSnapshot:
|
||||
"""Build a `StateSnapshot` from a saved checkpoint and its pending writes.
|
||||
|
||||
With `live=True` the snapshot shows current status: values include the
|
||||
output of tasks that already finished, `next` lists only tasks that still
|
||||
need to run, and `interrupts` lists only questions still waiting for an
|
||||
answer. Otherwise the snapshot is a record of the step: values as of the
|
||||
start of the step, every task in the step, and the interrupts they raised.
|
||||
"""
|
||||
if not saved:
|
||||
return StateSnapshot(
|
||||
values={},
|
||||
@@ -1244,10 +1236,13 @@ class Pregel(
|
||||
None,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
if live and saved.pending_writes:
|
||||
for tid, status in read_task_statuses(saved.pending_writes).items():
|
||||
if tid in next_tasks:
|
||||
next_tasks[tid].writes.extend(status.output)
|
||||
if apply_pending_writes and saved.pending_writes:
|
||||
for tid, k, v in saved.pending_writes:
|
||||
if k in (ERROR, INTERRUPT):
|
||||
continue
|
||||
if tid not in next_tasks:
|
||||
continue
|
||||
next_tasks[tid].writes.append((k, v))
|
||||
if tasks := [t for t in next_tasks.values() if t.writes]:
|
||||
apply_writes(
|
||||
saved.checkpoint, channels, tasks, None, self.trigger_to_nodes
|
||||
@@ -1257,7 +1252,6 @@ class Pregel(
|
||||
saved.pending_writes,
|
||||
task_states,
|
||||
self.stream_channels_asis,
|
||||
live=live,
|
||||
)
|
||||
# assemble the state snapshot
|
||||
return StateSnapshot(
|
||||
@@ -1276,16 +1270,8 @@ class Pregel(
|
||||
config: RunnableConfig,
|
||||
saved: CheckpointTuple | None,
|
||||
recurse: BaseCheckpointSaver | None = None,
|
||||
live: bool = False,
|
||||
apply_pending_writes: bool = False,
|
||||
) -> StateSnapshot:
|
||||
"""Build a `StateSnapshot` from a saved checkpoint and its pending writes.
|
||||
|
||||
With `live=True` the snapshot shows current status: values include the
|
||||
output of tasks that already finished, `next` lists only tasks that still
|
||||
need to run, and `interrupts` lists only questions still waiting for an
|
||||
answer. Otherwise the snapshot is a record of the step: values as of the
|
||||
start of the step, every task in the step, and the interrupts they raised.
|
||||
"""
|
||||
if not saved:
|
||||
return StateSnapshot(
|
||||
values={},
|
||||
@@ -1373,10 +1359,13 @@ class Pregel(
|
||||
None,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
if live and saved.pending_writes:
|
||||
for tid, status in read_task_statuses(saved.pending_writes).items():
|
||||
if tid in next_tasks:
|
||||
next_tasks[tid].writes.extend(status.output)
|
||||
if apply_pending_writes and saved.pending_writes:
|
||||
for tid, k, v in saved.pending_writes:
|
||||
if k in (ERROR, INTERRUPT):
|
||||
continue
|
||||
if tid not in next_tasks:
|
||||
continue
|
||||
next_tasks[tid].writes.append((k, v))
|
||||
if tasks := [t for t in next_tasks.values() if t.writes]:
|
||||
apply_writes(
|
||||
saved.checkpoint, channels, tasks, None, self.trigger_to_nodes
|
||||
@@ -1387,7 +1376,6 @@ class Pregel(
|
||||
saved.pending_writes,
|
||||
task_states,
|
||||
self.stream_channels_asis,
|
||||
live=live,
|
||||
)
|
||||
# assemble the state snapshot
|
||||
return StateSnapshot(
|
||||
@@ -1442,7 +1430,7 @@ class Pregel(
|
||||
config,
|
||||
saved,
|
||||
recurse=checkpointer if subgraphs else None,
|
||||
live=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
|
||||
apply_pending_writes=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
|
||||
)
|
||||
|
||||
async def aget_state(
|
||||
@@ -1486,7 +1474,7 @@ class Pregel(
|
||||
config,
|
||||
saved,
|
||||
recurse=checkpointer if subgraphs else None,
|
||||
live=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
|
||||
apply_pending_writes=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
|
||||
)
|
||||
|
||||
def get_state_history(
|
||||
@@ -1722,12 +1710,13 @@ class Pregel(
|
||||
checkpointer.get_next_version,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
# apply writes from tasks that already finished
|
||||
for tid, status in read_task_statuses(
|
||||
saved.pending_writes or []
|
||||
).items():
|
||||
if tid in next_tasks:
|
||||
next_tasks[tid].writes.extend(status.output)
|
||||
# apply writes from tasks that already ran
|
||||
for tid, k, v in saved.pending_writes or []:
|
||||
if k in (ERROR, INTERRUPT):
|
||||
continue
|
||||
if tid not in next_tasks:
|
||||
continue
|
||||
next_tasks[tid].writes.append((k, v))
|
||||
# clear all current tasks
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
@@ -1961,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:
|
||||
@@ -1975,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)
|
||||
@@ -2185,12 +2174,13 @@ class Pregel(
|
||||
checkpointer.get_next_version,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
# apply writes from tasks that already finished
|
||||
for tid, status in read_task_statuses(
|
||||
saved.pending_writes or []
|
||||
).items():
|
||||
if tid in next_tasks:
|
||||
next_tasks[tid].writes.extend(status.output)
|
||||
# apply writes from tasks that already ran
|
||||
for tid, k, v in saved.pending_writes or []:
|
||||
if k in (ERROR, INTERRUPT):
|
||||
continue
|
||||
if tid not in next_tasks:
|
||||
continue
|
||||
next_tasks[tid].writes.append((k, v))
|
||||
# clear all current tasks
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
@@ -2420,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:
|
||||
@@ -2434,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)
|
||||
@@ -4182,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)
|
||||
|
||||
@@ -726,13 +726,7 @@ class StateSnapshot(NamedTuple):
|
||||
tasks: tuple[PregelTask, ...]
|
||||
"""Tasks to execute in this step. If already attempted, may contain an error."""
|
||||
interrupts: tuple[Interrupt, ...]
|
||||
"""Interrupts that occurred in this step.
|
||||
|
||||
When reading the latest state (`get_state` without a `checkpoint_id`), this
|
||||
contains only interrupts still waiting for an answer. When reading a specific
|
||||
checkpoint or state history, it contains the most recent interrupt each task
|
||||
raised in that step, including ones answered later in the same step.
|
||||
"""
|
||||
"""Interrupts that occurred in this step that are pending resolution."""
|
||||
|
||||
|
||||
class Send:
|
||||
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -1,534 +0,0 @@
|
||||
"""State reads while some tasks of a superstep are finished and others are paused.
|
||||
|
||||
When parallel tasks each call `interrupt()` and only some of them are resumed,
|
||||
the superstep stays open. Its recorded writes then contain the old interrupt of
|
||||
each finished task next to that task's output. These tests check that state
|
||||
reads, which are rebuilt from the checkpointer, report only the interrupts that
|
||||
still need an answer.
|
||||
"""
|
||||
|
||||
import operator
|
||||
import sys
|
||||
import uuid
|
||||
from collections import Counter
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph._internal._constants import (
|
||||
ERROR,
|
||||
INTERRUPT,
|
||||
NO_WRITES,
|
||||
NULL_TASK_ID,
|
||||
RESUME,
|
||||
RETURN,
|
||||
)
|
||||
from langgraph.func import entrypoint, task
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
from langgraph.pregel._task_status import read_task_statuses
|
||||
from langgraph.types import Command, Durability, Interrupt, Send, interrupt
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
NEEDS_CONTEXTVARS = pytest.mark.skipif(
|
||||
sys.version_info < (3, 11),
|
||||
reason="Python 3.11+ is required for async contextvars support",
|
||||
)
|
||||
|
||||
|
||||
class State(TypedDict, total=False):
|
||||
log: Annotated[list[str], operator.add]
|
||||
count: int
|
||||
|
||||
|
||||
def _config() -> dict[str, Any]:
|
||||
return {"configurable": {"thread_id": str(uuid.uuid4())}}
|
||||
|
||||
|
||||
def _build_parallel(
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
calls: Counter[str],
|
||||
*,
|
||||
a_questions: int = 1,
|
||||
a_returns: Any = "log",
|
||||
):
|
||||
"""Build a graph where nodes `a` and `b` start in parallel and both ask questions.
|
||||
|
||||
`a` asks `a_questions` questions in a row. `a_returns` controls what `a`
|
||||
returns after its last answer. The default `"log"` returns the answers in
|
||||
`log`. Any other value is returned as-is.
|
||||
"""
|
||||
|
||||
def a(state: State) -> Any:
|
||||
calls["a"] += 1
|
||||
answers = [interrupt(f"A{i + 1}") for i in range(a_questions)]
|
||||
if a_returns == "log":
|
||||
return {"log": [f"a:{answer}" for answer in answers]}
|
||||
return a_returns
|
||||
|
||||
def b(state: State) -> State:
|
||||
calls["b"] += 1
|
||||
return {"log": [f"b:{interrupt('B')}"]}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("a", a)
|
||||
builder.add_node("b", b)
|
||||
builder.add_edge(START, "a")
|
||||
builder.add_edge(START, "b")
|
||||
builder.add_edge("a", END)
|
||||
builder.add_edge("b", END)
|
||||
return builder.compile(checkpointer=checkpointer)
|
||||
|
||||
|
||||
def _interrupt_by_value(snapshot: Any, value: str) -> Interrupt:
|
||||
return next(i for i in snapshot.interrupts if i.value == value)
|
||||
|
||||
|
||||
def _task(snapshot: Any, name: str) -> Any:
|
||||
return next(t for t in snapshot.tasks if t.name == name)
|
||||
|
||||
|
||||
def _interrupt_values(interrupts: Any) -> list[str]:
|
||||
return sorted(i.value for i in interrupts)
|
||||
|
||||
|
||||
# --- Task A answered and finished, task B still paused ---
|
||||
|
||||
|
||||
def test_finished_task_does_not_report_answered_interrupt(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
calls: Counter[str] = Counter()
|
||||
graph = _build_parallel(sync_checkpointer, calls)
|
||||
config = _config()
|
||||
|
||||
graph.invoke({"log": []}, config, durability=durability)
|
||||
snapshot = graph.get_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["A1", "B"]
|
||||
|
||||
graph.invoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "A1").id: "yes"}),
|
||||
config,
|
||||
durability=durability,
|
||||
)
|
||||
|
||||
snapshot = graph.get_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["B"]
|
||||
assert snapshot.next == ("b",)
|
||||
assert _task(snapshot, "a").interrupts == ()
|
||||
assert _task(snapshot, "a").result == {"log": ["a:yes"]}
|
||||
assert _interrupt_values(_task(snapshot, "b").interrupts) == ["B"]
|
||||
assert _task(snapshot, "b").result is None
|
||||
|
||||
# Reading the same checkpoint by id gives the record of the step: every task
|
||||
# in it, and every question asked, including the one A already answered.
|
||||
record = graph.get_state(snapshot.config)
|
||||
assert sorted(record.next) == ["a", "b"]
|
||||
assert _interrupt_values(record.interrupts) == ["A1", "B"]
|
||||
assert _interrupt_values(_task(record, "a").interrupts) == ["A1"]
|
||||
assert _task(record, "a").result == {"log": ["a:yes"]}
|
||||
|
||||
# B can still be answered, and the graph finishes normally.
|
||||
result = graph.invoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "B").id: "ok"}),
|
||||
config,
|
||||
durability=durability,
|
||||
)
|
||||
assert sorted(result["log"]) == ["a:yes", "b:ok"]
|
||||
assert calls == {"a": 2, "b": 3}
|
||||
snapshot = graph.get_state(config)
|
||||
assert snapshot.next == ()
|
||||
assert snapshot.interrupts == ()
|
||||
|
||||
# History still shows where each question was asked.
|
||||
asked = [
|
||||
_interrupt_values(s.interrupts)
|
||||
for s in graph.get_state_history(config)
|
||||
if s.interrupts
|
||||
]
|
||||
if durability != "exit":
|
||||
assert asked == [["A1", "B"]]
|
||||
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_finished_task_does_not_report_answered_interrupt_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
calls: Counter[str] = Counter()
|
||||
graph = _build_parallel(async_checkpointer, calls)
|
||||
config = _config()
|
||||
|
||||
await graph.ainvoke({"log": []}, config)
|
||||
snapshot = await graph.aget_state(config)
|
||||
await graph.ainvoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "A1").id: "yes"}), config
|
||||
)
|
||||
|
||||
snapshot = await graph.aget_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["B"]
|
||||
assert snapshot.next == ("b",)
|
||||
assert _task(snapshot, "a").interrupts == ()
|
||||
assert _task(snapshot, "a").result == {"log": ["a:yes"]}
|
||||
assert _interrupt_values(_task(snapshot, "b").interrupts) == ["B"]
|
||||
|
||||
record = await graph.aget_state(snapshot.config)
|
||||
assert _interrupt_values(record.interrupts) == ["A1", "B"]
|
||||
assert _interrupt_values(_task(record, "a").interrupts) == ["A1"]
|
||||
|
||||
result = await graph.ainvoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "B").id: "ok"}), config
|
||||
)
|
||||
assert sorted(result["log"]) == ["a:yes", "b:ok"]
|
||||
assert calls == {"a": 2, "b": 3}
|
||||
|
||||
|
||||
# --- Task A answered its first question and asked a second one ---
|
||||
|
||||
|
||||
def test_task_paused_at_second_question_stays_pending(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
calls: Counter[str] = Counter()
|
||||
graph = _build_parallel(sync_checkpointer, calls, a_questions=2)
|
||||
config = _config()
|
||||
|
||||
graph.invoke({"log": []}, config)
|
||||
snapshot = graph.get_state(config)
|
||||
graph.invoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "A1").id: "one"}), config
|
||||
)
|
||||
|
||||
snapshot = graph.get_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["A2", "B"]
|
||||
# A is not finished: it has a saved answer, but no output.
|
||||
assert sorted(snapshot.next) == ["a", "b"]
|
||||
assert _interrupt_values(_task(snapshot, "a").interrupts) == ["A2"]
|
||||
assert _task(snapshot, "a").result is None
|
||||
assert _interrupt_values(_task(snapshot, "b").interrupts) == ["B"]
|
||||
|
||||
# Both remaining questions can be answered together.
|
||||
result = graph.invoke(
|
||||
Command(
|
||||
resume={
|
||||
_interrupt_by_value(snapshot, "A2").id: "two",
|
||||
_interrupt_by_value(snapshot, "B").id: "ok",
|
||||
}
|
||||
),
|
||||
config,
|
||||
)
|
||||
assert sorted(result["log"]) == ["a:one", "a:two", "b:ok"]
|
||||
snapshot = graph.get_state(config)
|
||||
assert snapshot.next == ()
|
||||
assert snapshot.interrupts == ()
|
||||
|
||||
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_task_paused_at_second_question_stays_pending_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
calls: Counter[str] = Counter()
|
||||
graph = _build_parallel(async_checkpointer, calls, a_questions=2)
|
||||
config = _config()
|
||||
|
||||
await graph.ainvoke({"log": []}, config)
|
||||
snapshot = await graph.aget_state(config)
|
||||
await graph.ainvoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "A1").id: "one"}), config
|
||||
)
|
||||
|
||||
snapshot = await graph.aget_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["A2", "B"]
|
||||
assert sorted(snapshot.next) == ["a", "b"]
|
||||
assert _interrupt_values(_task(snapshot, "a").interrupts) == ["A2"]
|
||||
assert _task(snapshot, "a").result is None
|
||||
|
||||
|
||||
def test_task_paused_at_second_question_then_other_task_finishes(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
calls: Counter[str] = Counter()
|
||||
graph = _build_parallel(sync_checkpointer, calls, a_questions=2)
|
||||
config = _config()
|
||||
|
||||
graph.invoke({"log": []}, config)
|
||||
snapshot = graph.get_state(config)
|
||||
graph.invoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "A1").id: "one"}), config
|
||||
)
|
||||
snapshot = graph.get_state(config)
|
||||
graph.invoke(Command(resume={_interrupt_by_value(snapshot, "B").id: "ok"}), config)
|
||||
|
||||
snapshot = graph.get_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["A2"]
|
||||
assert snapshot.next == ("a",)
|
||||
assert _task(snapshot, "b").interrupts == ()
|
||||
assert _task(snapshot, "b").result == {"log": ["b:ok"]}
|
||||
|
||||
result = graph.invoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "A2").id: "two"}), config
|
||||
)
|
||||
assert sorted(result["log"]) == ["a:one", "a:two", "b:ok"]
|
||||
|
||||
|
||||
def test_resume_without_id_rejected_when_second_question_and_other_task_pending(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
calls: Counter[str] = Counter()
|
||||
graph = _build_parallel(sync_checkpointer, calls, a_questions=2)
|
||||
config = _config()
|
||||
|
||||
graph.invoke({"log": []}, config)
|
||||
snapshot = graph.get_state(config)
|
||||
graph.invoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "A1").id: "one"}), config
|
||||
)
|
||||
|
||||
# A2 and B are both waiting, so a resume value without an id is ambiguous.
|
||||
with pytest.raises(RuntimeError, match="multiple pending interrupts"):
|
||||
graph.invoke(Command(resume="ambiguous"), config)
|
||||
|
||||
|
||||
def test_resume_without_id_rejected_when_subgraph_has_parallel_interrupts(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
# A subgraph node whose child graph pauses in two parallel nodes records
|
||||
# both interrupts under one parent task. Both count as pending, so a resume
|
||||
# value without an id is ambiguous. (Before, only the first was counted and
|
||||
# the value went to whichever interrupt consumed it first.)
|
||||
child_builder = StateGraph(State)
|
||||
child_builder.add_node("a", lambda s: {"log": [f"a:{interrupt('A')}"]})
|
||||
child_builder.add_node("b", lambda s: {"log": [f"b:{interrupt('B')}"]})
|
||||
child_builder.add_edge(START, "a")
|
||||
child_builder.add_edge(START, "b")
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("child", child_builder.compile())
|
||||
builder.add_edge(START, "child")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = _config()
|
||||
|
||||
graph.invoke({"log": []}, config)
|
||||
snapshot = graph.get_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["A", "B"]
|
||||
|
||||
with pytest.raises(RuntimeError, match="multiple pending interrupts"):
|
||||
graph.invoke(Command(resume="ambiguous"), config)
|
||||
|
||||
result = graph.invoke(
|
||||
Command(
|
||||
resume={
|
||||
_interrupt_by_value(snapshot, "A").id: "x",
|
||||
_interrupt_by_value(snapshot, "B").id: "y",
|
||||
}
|
||||
),
|
||||
config,
|
||||
)
|
||||
assert sorted(result["log"]) == ["a:x", "b:y"]
|
||||
|
||||
|
||||
# --- Task A finished with an empty or falsy result ---
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"a_returns",
|
||||
[None, {}, {"count": 0}, {"log": []}],
|
||||
ids=["none", "empty_dict", "zero", "empty_list"],
|
||||
)
|
||||
def test_task_finished_with_falsy_result(
|
||||
sync_checkpointer: BaseCheckpointSaver, a_returns: Any
|
||||
) -> None:
|
||||
calls: Counter[str] = Counter()
|
||||
graph = _build_parallel(sync_checkpointer, calls, a_returns=a_returns)
|
||||
config = _config()
|
||||
|
||||
graph.invoke({"log": []}, config)
|
||||
snapshot = graph.get_state(config)
|
||||
graph.invoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "A1").id: "yes"}), config
|
||||
)
|
||||
|
||||
snapshot = graph.get_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["B"]
|
||||
assert snapshot.next == ("b",)
|
||||
assert _task(snapshot, "a").interrupts == ()
|
||||
|
||||
graph.invoke(Command(resume={_interrupt_by_value(snapshot, "B").id: "ok"}), config)
|
||||
# A already finished, so resuming B must not run A again.
|
||||
assert calls == {"a": 2, "b": 3}
|
||||
snapshot = graph.get_state(config)
|
||||
assert snapshot.next == ()
|
||||
assert snapshot.interrupts == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("a_returns", [None, {"count": 0}], ids=["none", "zero"])
|
||||
@NEEDS_CONTEXTVARS
|
||||
async def test_task_finished_with_falsy_result_async(
|
||||
async_checkpointer: BaseCheckpointSaver, a_returns: Any
|
||||
) -> None:
|
||||
calls: Counter[str] = Counter()
|
||||
graph = _build_parallel(async_checkpointer, calls, a_returns=a_returns)
|
||||
config = _config()
|
||||
|
||||
await graph.ainvoke({"log": []}, config)
|
||||
snapshot = await graph.aget_state(config)
|
||||
await graph.ainvoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "A1").id: "yes"}), config
|
||||
)
|
||||
|
||||
snapshot = await graph.aget_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["B"]
|
||||
assert snapshot.next == ("b",)
|
||||
assert _task(snapshot, "a").interrupts == ()
|
||||
|
||||
await graph.ainvoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "B").id: "ok"}), config
|
||||
)
|
||||
assert calls == {"a": 2, "b": 3}
|
||||
|
||||
|
||||
# --- Subgraphs and the functional API ---
|
||||
|
||||
|
||||
def test_parallel_subgraphs_report_only_pending_interrupts(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
class ChildState(TypedDict):
|
||||
prompt: str
|
||||
answers: Annotated[list[str], operator.add]
|
||||
|
||||
def ask(state: ChildState) -> dict[str, Any]:
|
||||
return {"answers": [interrupt(state["prompt"])]}
|
||||
|
||||
child_builder = StateGraph(ChildState)
|
||||
child_builder.add_node("ask", ask)
|
||||
child_builder.add_edge(START, "ask")
|
||||
child = child_builder.compile()
|
||||
|
||||
class ParentState(TypedDict):
|
||||
answers: Annotated[list[str], operator.add]
|
||||
|
||||
builder = StateGraph(ParentState)
|
||||
builder.add_node("child", child)
|
||||
builder.add_conditional_edges(
|
||||
START,
|
||||
lambda _: [Send("child", {"prompt": p, "answers": []}) for p in ("a", "b")],
|
||||
["child"],
|
||||
)
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = _config()
|
||||
|
||||
graph.invoke({"answers": []}, config)
|
||||
snapshot = graph.get_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["a", "b"]
|
||||
graph.invoke(Command(resume={_interrupt_by_value(snapshot, "a").id: "x"}), config)
|
||||
|
||||
snapshot = graph.get_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["b"]
|
||||
assert snapshot.next == ("child",)
|
||||
finished = next(t for t in snapshot.tasks if t.result is not None)
|
||||
assert finished.interrupts == ()
|
||||
assert finished.result == {"answers": ["x"]}
|
||||
|
||||
result = graph.invoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "b").id: "y"}), config
|
||||
)
|
||||
assert sorted(result["answers"]) == ["x", "y"]
|
||||
|
||||
|
||||
def test_functional_task_finished_with_none_is_not_rerun(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
calls: Counter[str] = Counter()
|
||||
|
||||
@task
|
||||
def ask_a() -> None:
|
||||
calls["a"] += 1
|
||||
interrupt("A")
|
||||
|
||||
@task
|
||||
def ask_b() -> str:
|
||||
calls["b"] += 1
|
||||
return interrupt("B")
|
||||
|
||||
@entrypoint(checkpointer=sync_checkpointer)
|
||||
def workflow(_: Any) -> list[Any]:
|
||||
a, b = ask_a(), ask_b()
|
||||
return [a.result(), b.result()]
|
||||
|
||||
config = _config()
|
||||
workflow.invoke(1, config)
|
||||
snapshot = workflow.get_state(config)
|
||||
workflow.invoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "A").id: "x"}), config
|
||||
)
|
||||
|
||||
snapshot = workflow.get_state(config)
|
||||
assert _interrupt_values(snapshot.interrupts) == ["B"]
|
||||
|
||||
result = workflow.invoke(
|
||||
Command(resume={_interrupt_by_value(snapshot, "B").id: "y"}), config
|
||||
)
|
||||
assert result == [None, "y"]
|
||||
assert calls == {"a": 2, "b": 3}
|
||||
|
||||
|
||||
# --- Reading task status from recorded writes ---
|
||||
|
||||
|
||||
def test_read_task_statuses() -> None:
|
||||
a1 = Interrupt(value="A1", id="a")
|
||||
a2 = Interrupt(value="A2", id="a")
|
||||
b = Interrupt(value="B", id="b")
|
||||
error = ValueError("boom")
|
||||
|
||||
statuses = read_task_statuses(
|
||||
[
|
||||
# answered and finished: old interrupt stays recorded
|
||||
("finished", INTERRUPT, (a1,)),
|
||||
("finished", RESUME, ["yes"]),
|
||||
("finished", "log", ["a:yes"]),
|
||||
# answered once, then paused at a second question
|
||||
("paused", INTERRUPT, (a2,)),
|
||||
("paused", RESUME, ["one"]),
|
||||
# finished with no output
|
||||
("no_output", INTERRUPT, (b,)),
|
||||
("no_output", RESUME, ["ok"]),
|
||||
("no_output", NO_WRITES, None),
|
||||
# functional task that returned None
|
||||
("returned_none", RETURN, None),
|
||||
# failed
|
||||
("failed", ERROR, error),
|
||||
# not a task
|
||||
(NULL_TASK_ID, RESUME, "global"),
|
||||
]
|
||||
)
|
||||
|
||||
assert set(statuses) == {
|
||||
"finished",
|
||||
"paused",
|
||||
"no_output",
|
||||
"returned_none",
|
||||
"failed",
|
||||
}
|
||||
|
||||
assert statuses["finished"].finished
|
||||
assert statuses["finished"].interrupts == (a1,)
|
||||
assert statuses["finished"].pending_interrupts == ()
|
||||
assert statuses["finished"].output == (("log", ["a:yes"]),)
|
||||
|
||||
assert not statuses["paused"].finished
|
||||
assert statuses["paused"].interrupts == (a2,)
|
||||
assert statuses["paused"].pending_interrupts == (a2,)
|
||||
assert statuses["paused"].output == ()
|
||||
|
||||
assert statuses["no_output"].finished
|
||||
assert statuses["no_output"].interrupts == (b,)
|
||||
assert statuses["no_output"].pending_interrupts == ()
|
||||
|
||||
assert statuses["returned_none"].finished
|
||||
assert statuses["returned_none"].output == ((RETURN, None),)
|
||||
|
||||
assert not statuses["failed"].finished
|
||||
assert statuses["failed"].error is error
|
||||
Reference in New Issue
Block a user