mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 05:25:05 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a9a10dedaf | ||
|
|
89ff2d33de | ||
|
|
02c4bc992b | ||
|
|
35c3609b65 |
+57
@@ -267,6 +267,61 @@ async def test_history_seed_ancestor_own_writes_are_replayed(
|
||||
)
|
||||
|
||||
|
||||
# Every uuid4 `build_delta_chain` tags its own writes with sorts between these
|
||||
# two, so task_id order is fixed and always disagrees with task_path order.
|
||||
TASK_ID_SORTS_FIRST = "00000000-0000-0000-0000-000000000000"
|
||||
TASK_ID_SORTS_LAST = "ffffffff-ffff-ffff-ffff-ffffffffffff"
|
||||
|
||||
|
||||
async def test_history_orders_parallel_writes_by_task_path(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Writes from parallel tasks replay in task_path order, not task_id order."""
|
||||
configs = await build_delta_chain(
|
||||
saver,
|
||||
thread_id=str(uuid4()),
|
||||
channel="ch",
|
||||
snapshots_at_steps=[0],
|
||||
total_steps=3,
|
||||
)
|
||||
step_1, head = configs[1], configs[2]
|
||||
await saver.aput_writes(
|
||||
step_1, [("ch", "second")], TASK_ID_SORTS_FIRST, "~pull, 02"
|
||||
)
|
||||
await saver.aput_writes(step_1, [("ch", "first")], TASK_ID_SORTS_LAST, "~pull, 01")
|
||||
|
||||
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
|
||||
values = [w[2] for w in result["ch"]["writes"]]
|
||||
assert values == [1, "first", "second"], (
|
||||
f"Expected task_path order [1, 'first', 'second'], got {values}. "
|
||||
"Ordering by (task_id, idx) alone yields [1, 'second', 'first']."
|
||||
)
|
||||
|
||||
|
||||
async def test_history_orders_pathless_writes_first(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Writes stored without a task_path (graph input) replay before task writes."""
|
||||
configs = await build_delta_chain(
|
||||
saver,
|
||||
thread_id=str(uuid4()),
|
||||
channel="ch",
|
||||
snapshots_at_steps=[0],
|
||||
total_steps=3,
|
||||
)
|
||||
step_1, head = configs[1], configs[2]
|
||||
await saver.aput_writes(
|
||||
step_1, [("ch", "from_node")], TASK_ID_SORTS_FIRST, "~pull, a"
|
||||
)
|
||||
await saver.aput_writes(step_1, [("ch", "from_input")], TASK_ID_SORTS_LAST)
|
||||
|
||||
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
|
||||
values = [w[2] for w in result["ch"]["writes"]]
|
||||
assert values == [1, "from_input", "from_node"], (
|
||||
f"Expected pathless writes first, got {values}"
|
||||
)
|
||||
|
||||
|
||||
ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
||||
test_history_returns_writes_oldest_first,
|
||||
test_history_seed_is_nearest_snapshot,
|
||||
@@ -276,6 +331,8 @@ ALL_DELTA_CHANNEL_HISTORY_TESTS = [
|
||||
test_history_walk_to_root_no_seed,
|
||||
test_history_migration_plain_value_as_seed,
|
||||
test_history_seed_ancestor_own_writes_are_replayed,
|
||||
test_history_orders_parallel_writes_by_task_path,
|
||||
test_history_orders_pathless_writes_first,
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -168,6 +168,7 @@ class _DeltaStage2Row(TypedDict, total=False):
|
||||
type: str | None
|
||||
blob: bytes | None
|
||||
task_id: str | None # "w" rows only
|
||||
task_path: str | None # "w" rows only
|
||||
idx: int | None # "w" rows only
|
||||
version: str | None # "b" rows only
|
||||
|
||||
@@ -319,7 +320,7 @@ def _build_delta_stage2_sql(
|
||||
branches.append(
|
||||
"SELECT 'w'::text AS _kind, "
|
||||
"checkpoint_id, channel, "
|
||||
"type, blob, task_id, idx, NULL::text AS version "
|
||||
"type, blob, task_id, task_path, idx, NULL::text AS version "
|
||||
"FROM checkpoint_writes "
|
||||
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
||||
"AND checkpoint_id = ANY(%s)"
|
||||
@@ -327,7 +328,8 @@ def _build_delta_stage2_sql(
|
||||
for _ in channels_with_seed:
|
||||
branches.append(
|
||||
"SELECT 'b'::text AS _kind, NULL::text AS checkpoint_id, channel, "
|
||||
"type, blob, NULL::text AS task_id, NULL::int AS idx, version "
|
||||
"type, blob, NULL::text AS task_id, NULL::text AS task_path, "
|
||||
"NULL::int AS idx, version "
|
||||
"FROM checkpoint_blobs "
|
||||
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
|
||||
"AND version = %s"
|
||||
@@ -492,10 +494,11 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
||||
stored value, or when the seed blob is sentinel "empty" — in both cases
|
||||
the consumer treats absence as "start empty".
|
||||
"""
|
||||
# writes_by_ch_by_cid[channel][cid] = list of (type, blob, task_id, idx)
|
||||
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
|
||||
ch: {} for ch in channels
|
||||
}
|
||||
# writes_by_ch_by_cid[channel][cid] = list of
|
||||
# (type, blob, task_id, idx, task_path)
|
||||
writes_by_ch_by_cid: dict[
|
||||
str, dict[str, list[tuple[str, bytes, str, int, str]]]
|
||||
] = {ch: {} for ch in channels}
|
||||
# seed_blob_by_ver[(channel, version)] = (type, blob)
|
||||
seed_blob_by_ver: dict[tuple[str, str], tuple[str, bytes]] = {}
|
||||
|
||||
@@ -506,8 +509,14 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
||||
cid = cast(str, r["checkpoint_id"])
|
||||
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
||||
cast(
|
||||
"tuple[str, bytes, str, int]",
|
||||
(r["type"], r["blob"], r["task_id"], r["idx"]),
|
||||
"tuple[str, bytes, str, int, str]",
|
||||
(
|
||||
r["type"],
|
||||
r["blob"],
|
||||
r["task_id"],
|
||||
r["idx"],
|
||||
r["task_path"],
|
||||
),
|
||||
)
|
||||
)
|
||||
else: # kind == "b"
|
||||
@@ -516,10 +525,10 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
||||
"tuple[str, bytes]", (r["type"], r["blob"])
|
||||
)
|
||||
|
||||
# Sort writes per (channel, cid) newest-first by (task_id, idx)
|
||||
# Sort writes per (channel, cid) newest-first by (task_path, task_id, idx)
|
||||
for cid_map in writes_by_ch_by_cid.values():
|
||||
for ws in cid_map.values():
|
||||
ws.sort(key=lambda w: (w[2], w[3]), reverse=True)
|
||||
ws.sort(key=lambda w: (w[4], w[2], w[3]), reverse=True)
|
||||
|
||||
result: dict[str, DeltaChannelHistory] = {}
|
||||
for ch in channels:
|
||||
@@ -529,7 +538,9 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
||||
collected: list[PendingWrite] = []
|
||||
cid_writes = writes_by_ch_by_cid.get(ch, {})
|
||||
for cid in chain_cids:
|
||||
for type_tag, write_blob, task_id, _idx in cid_writes.get(cid, []):
|
||||
for type_tag, write_blob, task_id, _idx, _path in cid_writes.get(
|
||||
cid, []
|
||||
):
|
||||
val = self.serde.loads_typed((type_tag, write_blob))
|
||||
collected.append((task_id, ch, val))
|
||||
collected.reverse()
|
||||
|
||||
@@ -81,6 +81,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
|
||||
conn: sqlite3.Connection
|
||||
is_setup: bool
|
||||
_has_task_path: bool = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -154,6 +155,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||
checkpoint_id TEXT NOT NULL,
|
||||
task_id TEXT NOT NULL,
|
||||
task_path TEXT NOT NULL DEFAULT '',
|
||||
idx INTEGER NOT NULL,
|
||||
channel TEXT NOT NULL,
|
||||
type TEXT,
|
||||
@@ -162,6 +164,19 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
);
|
||||
"""
|
||||
)
|
||||
# sqlite has no ADD COLUMN IF NOT EXISTS; this migrates databases
|
||||
# created before `task_path` existed and is a no-op on the rest.
|
||||
try:
|
||||
self.conn.execute(
|
||||
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
|
||||
)
|
||||
except sqlite3.OperationalError as e:
|
||||
# A read-only database from before the column can still be read;
|
||||
# its rows would all read back as '' anyway.
|
||||
if "readonly database" in str(e):
|
||||
self._has_task_path = False
|
||||
elif "duplicate column name" not in str(e):
|
||||
raise
|
||||
|
||||
self.is_setup = True
|
||||
|
||||
@@ -460,9 +475,9 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
task_path: Path of the task creating the writes.
|
||||
"""
|
||||
query = (
|
||||
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
||||
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
||||
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
||||
)
|
||||
with self.cursor() as cur:
|
||||
cur.executemany(
|
||||
@@ -473,6 +488,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
str(config["configurable"]["checkpoint_ns"]),
|
||||
str(config["configurable"]["checkpoint_id"]),
|
||||
task_id,
|
||||
task_path,
|
||||
WRITES_IDX_MAP.get(channel, idx),
|
||||
channel,
|
||||
*self.serde.dumps_typed(value),
|
||||
@@ -558,6 +574,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
|
||||
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
||||
stage2_sql = build_delta_stage2_sql(
|
||||
has_task_path=self._has_task_path,
|
||||
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
|
||||
)
|
||||
if stage2_sql:
|
||||
@@ -568,7 +585,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
)
|
||||
cur.execute(stage2_sql, stage2_params)
|
||||
stage2_rows = cast(
|
||||
"list[tuple[str, str, str, int, str, bytes]]", cur.fetchall()
|
||||
"list[tuple[str, str, str, int, str, bytes, str]]", cur.fetchall()
|
||||
)
|
||||
else:
|
||||
stage2_rows = []
|
||||
|
||||
@@ -39,7 +39,9 @@ DELTA_STAGE1_SQL = (
|
||||
)
|
||||
|
||||
|
||||
def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
|
||||
def build_delta_stage2_sql(
|
||||
*, chain_lens: Sequence[int], has_task_path: bool = True
|
||||
) -> str:
|
||||
"""Stage-2 per-channel UNION ALL fetching writes from `writes`.
|
||||
|
||||
One branch per channel with a non-empty chain. Each branch inlines its
|
||||
@@ -53,11 +55,12 @@ def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
|
||||
of a single `channel = ANY(channels)` filter when channels have
|
||||
different chain depths — same rationale as postgres.
|
||||
"""
|
||||
task_path = "task_path" if has_task_path else "''"
|
||||
branches: list[str] = []
|
||||
for n in chain_lens:
|
||||
cid_placeholders = ",".join("?" * n)
|
||||
branches.append(
|
||||
"SELECT checkpoint_id, channel, task_id, idx, type, value "
|
||||
f"SELECT checkpoint_id, channel, task_id, idx, type, value, {task_path} "
|
||||
"FROM writes "
|
||||
"WHERE thread_id = ? AND checkpoint_ns = ? AND channel = ? "
|
||||
f"AND checkpoint_id IN ({cid_placeholders})"
|
||||
@@ -130,29 +133,31 @@ def build_delta_channels_writes_history(
|
||||
chain_by_ch: Mapping[str, list[str]],
|
||||
seed_val_by_ch: Mapping[str, Any],
|
||||
seeded: set[str],
|
||||
stage2_rows: Sequence[tuple[str, str, str, int, str, bytes]],
|
||||
stage2_rows: Sequence[tuple[str, str, str, int, str, bytes, str]],
|
||||
serde: Any,
|
||||
) -> dict[str, DeltaChannelHistory]:
|
||||
"""Demux stage-2 rows per channel; produce per-channel histories.
|
||||
|
||||
Stage-2 rows are `(checkpoint_id, channel, task_id, idx, type, value)`.
|
||||
Final write order is oldest→newest globally and `(task_id, idx)` within
|
||||
a checkpoint, matching the contract on `DeltaChannelHistory.writes`.
|
||||
Stage-2 rows are
|
||||
`(checkpoint_id, channel, task_id, idx, type, value, task_path)`.
|
||||
Final write order is oldest→newest globally and
|
||||
`(task_path, task_id, idx)` within a checkpoint, matching the contract
|
||||
on `DeltaChannelHistory.writes`.
|
||||
|
||||
`seed` is omitted when the walk reached a true root with no snapshot
|
||||
found (channel never entered `seeded`); consumers treat absence as
|
||||
"start empty".
|
||||
"""
|
||||
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
|
||||
ch: {} for ch in channels
|
||||
}
|
||||
for cid, ch, task_id, idx, type_tag, value_blob in stage2_rows:
|
||||
writes_by_ch_by_cid: dict[
|
||||
str, dict[str, list[tuple[str, bytes, str, int, str]]]
|
||||
] = {ch: {} for ch in channels}
|
||||
for cid, ch, task_id, idx, type_tag, value_blob, task_path in stage2_rows:
|
||||
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
|
||||
(type_tag, value_blob, task_id, idx)
|
||||
(type_tag, value_blob, task_id, idx, task_path)
|
||||
)
|
||||
for cid_map in writes_by_ch_by_cid.values():
|
||||
for ws in cid_map.values():
|
||||
ws.sort(key=lambda w: (w[2], w[3]))
|
||||
ws.sort(key=lambda w: (w[4], w[2], w[3]))
|
||||
|
||||
result: dict[str, DeltaChannelHistory] = {}
|
||||
for ch in channels:
|
||||
@@ -161,7 +166,7 @@ def build_delta_channels_writes_history(
|
||||
collected: list[PendingWrite] = []
|
||||
# Chain is newest-first; iterate oldest-first for the public order.
|
||||
for cid in reversed(chain_cids):
|
||||
for type_tag, value_blob, task_id, _idx in cid_writes.get(cid, []):
|
||||
for type_tag, value_blob, task_id, _idx, _path in cid_writes.get(cid, []):
|
||||
collected.append(
|
||||
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
|
||||
)
|
||||
|
||||
@@ -114,6 +114,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
|
||||
lock: asyncio.Lock
|
||||
is_setup: bool
|
||||
_has_task_path: bool = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -331,6 +332,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||
checkpoint_id TEXT NOT NULL,
|
||||
task_id TEXT NOT NULL,
|
||||
task_path TEXT NOT NULL DEFAULT '',
|
||||
idx INTEGER NOT NULL,
|
||||
channel TEXT NOT NULL,
|
||||
type TEXT,
|
||||
@@ -341,6 +343,21 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
):
|
||||
await self.conn.commit()
|
||||
|
||||
# sqlite has no ADD COLUMN IF NOT EXISTS; this migrates databases
|
||||
# created before `task_path` existed and is a no-op on the rest.
|
||||
try:
|
||||
await self.conn.execute(
|
||||
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
|
||||
)
|
||||
await self.conn.commit()
|
||||
except aiosqlite.OperationalError as e:
|
||||
# A read-only database from before the column can still be read;
|
||||
# its rows would all read back as '' anyway.
|
||||
if "readonly database" in str(e):
|
||||
self._has_task_path = False
|
||||
elif "duplicate column name" not in str(e):
|
||||
raise
|
||||
|
||||
self.is_setup = True
|
||||
|
||||
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
|
||||
@@ -576,9 +593,9 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
task_path: Path of the task creating the writes.
|
||||
"""
|
||||
query = (
|
||||
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
||||
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
||||
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
|
||||
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
|
||||
)
|
||||
await self.setup()
|
||||
async with self.lock, self.conn.cursor() as cur:
|
||||
@@ -590,6 +607,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
str(config["configurable"]["checkpoint_ns"]),
|
||||
str(config["configurable"]["checkpoint_id"]),
|
||||
task_id,
|
||||
task_path,
|
||||
WRITES_IDX_MAP.get(channel, idx),
|
||||
channel,
|
||||
*self.serde.dumps_typed(value),
|
||||
@@ -671,6 +689,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
|
||||
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
||||
stage2_sql = build_delta_stage2_sql(
|
||||
has_task_path=self._has_task_path,
|
||||
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
|
||||
)
|
||||
if stage2_sql:
|
||||
@@ -681,7 +700,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
)
|
||||
await cur.execute(stage2_sql, stage2_params)
|
||||
stage2_rows = cast(
|
||||
"list[tuple[str, str, str, int, str, bytes]]",
|
||||
"list[tuple[str, str, str, int, str, bytes, str]]",
|
||||
await cur.fetchall(),
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
import sqlite3
|
||||
from pathlib import Path
|
||||
|
||||
import aiosqlite
|
||||
import pytest
|
||||
from langgraph.checkpoint.base import empty_checkpoint
|
||||
|
||||
from langgraph.checkpoint.sqlite import SqliteSaver
|
||||
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
||||
|
||||
WRITES_BEFORE_TASK_PATH = """
|
||||
CREATE TABLE writes (
|
||||
thread_id TEXT NOT NULL,
|
||||
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
||||
checkpoint_id TEXT NOT NULL,
|
||||
task_id TEXT NOT NULL,
|
||||
idx INTEGER NOT NULL,
|
||||
channel TEXT NOT NULL,
|
||||
type TEXT,
|
||||
value BLOB,
|
||||
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
|
||||
);
|
||||
INSERT INTO writes VALUES ('t', '', 'c', 'old-task', 0, 'ch', 'null', X'');
|
||||
"""
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def legacy_db(tmp_path: Path) -> Path:
|
||||
db = tmp_path / "legacy.sqlite"
|
||||
with sqlite3.connect(db) as conn:
|
||||
conn.executescript(WRITES_BEFORE_TASK_PATH)
|
||||
return db
|
||||
|
||||
|
||||
def test_setup_migrates_legacy_writes_table_repeatably(legacy_db: Path) -> None:
|
||||
for _ in range(2):
|
||||
with SqliteSaver.from_conn_string(str(legacy_db)) as saver:
|
||||
saver.setup()
|
||||
rows = saver.conn.execute(
|
||||
"SELECT task_id, task_path FROM writes"
|
||||
).fetchall()
|
||||
assert rows == [("old-task", "")]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("fresh", [True, False], ids=["fresh", "legacy"])
|
||||
def test_put_writes_persists_task_path(
|
||||
tmp_path: Path, legacy_db: Path, fresh: bool
|
||||
) -> None:
|
||||
db = tmp_path / "fresh.sqlite" if fresh else legacy_db
|
||||
with SqliteSaver.from_conn_string(str(db)) as saver:
|
||||
config = saver.put(
|
||||
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
|
||||
empty_checkpoint(),
|
||||
{},
|
||||
{},
|
||||
)
|
||||
saver.put_writes(config, [("ch", "v")], "task-1", "~__pregel_pull, node")
|
||||
stored = saver.conn.execute(
|
||||
"SELECT task_path FROM writes WHERE task_id = 'task-1'"
|
||||
).fetchall()
|
||||
assert stored == [("~__pregel_pull, node",)]
|
||||
|
||||
|
||||
async def test_async_setup_migrates_legacy_writes_table_repeatably(
|
||||
legacy_db: Path,
|
||||
) -> None:
|
||||
for _ in range(2):
|
||||
async with AsyncSqliteSaver.from_conn_string(str(legacy_db)) as saver:
|
||||
await saver.setup()
|
||||
config = await saver.aput(
|
||||
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
|
||||
empty_checkpoint(),
|
||||
{},
|
||||
{},
|
||||
)
|
||||
await saver.aput_writes(
|
||||
config, [("ch", "v")], "task-1", "~__pregel_pull, node"
|
||||
)
|
||||
|
||||
async with aiosqlite.connect(legacy_db) as conn:
|
||||
async with conn.execute(
|
||||
"SELECT DISTINCT task_id, task_path FROM writes ORDER BY task_id"
|
||||
) as cur:
|
||||
assert await cur.fetchall() == [
|
||||
("old-task", ""),
|
||||
("task-1", "~__pregel_pull, node"),
|
||||
]
|
||||
|
||||
|
||||
def _legacy_database_with_history(db: Path) -> dict:
|
||||
root = empty_checkpoint()
|
||||
root["channel_values"] = {"ch": "seed"}
|
||||
root["channel_versions"] = {"ch": 1}
|
||||
with SqliteSaver.from_conn_string(str(db)) as saver:
|
||||
root_config = saver.put(
|
||||
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
|
||||
root,
|
||||
{},
|
||||
{"ch": 1},
|
||||
)
|
||||
saver.put_writes(root_config, [("ch", "write")], "task", "~__pregel_pull, n")
|
||||
child = saver.put(root_config, empty_checkpoint(), {}, {})
|
||||
saver.conn.execute("ALTER TABLE writes DROP COLUMN task_path")
|
||||
saver.conn.commit()
|
||||
return child
|
||||
|
||||
|
||||
def test_read_only_legacy_database_still_reads_delta_history(tmp_path: Path) -> None:
|
||||
db = tmp_path / "legacy.sqlite"
|
||||
child = _legacy_database_with_history(db)
|
||||
|
||||
saver = SqliteSaver(sqlite3.connect(f"file:{db}?mode=ro", uri=True))
|
||||
got = saver.get_delta_channel_history(config=child, channels=["ch"])
|
||||
|
||||
assert got["ch"] == {"seed": "seed", "writes": [("task", "ch", "write")]}
|
||||
|
||||
|
||||
async def test_async_read_only_legacy_database_still_reads_delta_history(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
db = tmp_path / "legacy.sqlite"
|
||||
child = _legacy_database_with_history(db)
|
||||
|
||||
async with aiosqlite.connect(f"file:{db}?mode=ro", uri=True) as conn:
|
||||
saver = AsyncSqliteSaver(conn)
|
||||
got = await saver.aget_delta_channel_history(config=child, channels=["ch"])
|
||||
|
||||
assert got["ch"] == {"seed": "seed", "writes": [("task", "ch", "write")]}
|
||||
@@ -162,6 +162,14 @@ class DeltaChannelHistory(TypedDict):
|
||||
Always present; possibly empty. Already filtered to one channel.
|
||||
Writes stored at the target checkpoint itself are pending for the
|
||||
next super-step and are excluded.
|
||||
|
||||
Within a single checkpoint, writes are ordered by
|
||||
`(task_path, task_id, idx)`, which is the order live execution applies
|
||||
a super-step's task writes in. `task_id` is a hash of the path, so
|
||||
ordering by it permutes parallel tasks writing one channel, and
|
||||
reducers need not be order-invariant. Writes stored without a
|
||||
`task_path` (graph input, `update_state` updates, exit-durability runs,
|
||||
rows predating the column) sort first, by `task_id`.
|
||||
* `seed` — the stored value at the nearest ancestor whose
|
||||
`channel_values[ch]` is populated. Omitted if the walk reached the
|
||||
root without finding any stored value (consumer treats absence as
|
||||
@@ -611,6 +619,11 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
`PostgresSaver`) override for performance; the return contract is
|
||||
fixed here.
|
||||
|
||||
`PendingWrite` carries no `task_path`, so this default replays each
|
||||
checkpoint's writes in `get_tuple`'s `pending_writes` order. Savers
|
||||
that do not return `pending_writes` ordered by
|
||||
`(task_path, task_id, idx)` must override it.
|
||||
|
||||
Args:
|
||||
config: Configuration identifying the target checkpoint.
|
||||
channels: Channel names to walk for. Empty → empty mapping.
|
||||
|
||||
@@ -199,8 +199,8 @@ class InMemorySaver(
|
||||
terminated_here.add(ch)
|
||||
|
||||
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
|
||||
for (_task_id, _idx), (tid, ch, serialized, _) in sorted(
|
||||
step_writes.items(), reverse=True
|
||||
for _, (tid, ch, serialized, _) in sorted(
|
||||
step_writes.items(), key=lambda kv: (kv[1][3], kv[0]), reverse=True
|
||||
):
|
||||
if ch not in remaining:
|
||||
continue
|
||||
|
||||
@@ -47,6 +47,16 @@ def exit_delta_task_id(step: int, task_id: str) -> str:
|
||||
return f"{step:08d}-{parts[1]}-{parts[2]}-{parts[3]}-{parts[4]}"
|
||||
|
||||
|
||||
def exit_delta_late_task_id(step: int, task_id: str) -> str:
|
||||
"""Synthetic task id for exit-mode writes of a superstep after the anchor's own.
|
||||
|
||||
Sorts after every real task id, in step order, so replay keeps them after
|
||||
the anchor's own superstep whether a saver orders by task path or task id.
|
||||
"""
|
||||
parts = str(uuid.UUID(task_id)).split("-")
|
||||
return f"ffffffff-{step >> 16:04x}-{step & 0xFFFF:04x}-{parts[3]}-{parts[4]}"
|
||||
|
||||
|
||||
def delta_channels_to_snapshot(
|
||||
channels: Mapping[str, BaseChannel],
|
||||
counters_since_delta_snapshot: Mapping[str, tuple[int, int]],
|
||||
|
||||
@@ -103,6 +103,7 @@ from langgraph.pregel._checkpoint import (
|
||||
create_checkpoint,
|
||||
delta_channels_to_snapshot,
|
||||
empty_checkpoint,
|
||||
exit_delta_late_task_id,
|
||||
exit_delta_task_id,
|
||||
)
|
||||
from langgraph.pregel._executor import (
|
||||
@@ -119,7 +120,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,
|
||||
@@ -218,10 +218,15 @@ class PregelLoop:
|
||||
# `after_tick`). At exit, `_put_exit_delta_writes` filters out channels
|
||||
# that will snapshot, then persists the rest under an anchor parent.
|
||||
# `None` when not in exit mode (so the capture sites are no-ops).
|
||||
# Each tuple is `(step, task_id, channel, value)` — `step` drives the
|
||||
# synthetic step-prefixed task_id used to preserve chronological order
|
||||
# under the saver's `ORDER BY task_id, idx` sorting.
|
||||
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
|
||||
# Each tuple is `(step, task_id, task_path, channel, value)`; see
|
||||
# `_put_exit_delta_writes` for how they are ordered.
|
||||
_exit_delta_writes: list[tuple[int, str, str, str, Any]] | None = None
|
||||
|
||||
# The pending writes loaded with the checkpoint, already stored on it, kept
|
||||
# alive so their ids stay unique; and the checkpoint's own superstep, the
|
||||
# first one this run ticks.
|
||||
_loaded_write_ids: dict[int, tuple[str, str, Any]]
|
||||
_exit_first_step: int | None = None
|
||||
|
||||
# Delta channels that saw an Overwrite since the last checkpoint. These
|
||||
# channels must snapshot after live update applies overwrite semantics so
|
||||
@@ -708,9 +713,18 @@ class PregelLoop:
|
||||
)
|
||||
# capture delta-channel writes for exit-mode accumulator before clearing
|
||||
if self._exit_delta_writes is not None:
|
||||
for tid, ch, v in self.checkpoint_pending_writes:
|
||||
if isinstance(self.specs.get(ch), DeltaChannel):
|
||||
self._exit_delta_writes.append((self.step, tid, ch, v))
|
||||
if self._exit_first_step is None:
|
||||
self._exit_first_step = self.step
|
||||
for w in self.checkpoint_pending_writes:
|
||||
tid, ch, v = w
|
||||
if not isinstance(self.specs.get(ch), DeltaChannel):
|
||||
continue
|
||||
if id(w) in self._loaded_write_ids:
|
||||
continue
|
||||
task = self.tasks.get(tid)
|
||||
path = task_path_str(task.path) if task else ""
|
||||
self._exit_delta_writes.append((self.step, tid, path, ch, v))
|
||||
self._loaded_write_ids = {}
|
||||
# clear pending writes
|
||||
self.checkpoint_pending_writes.clear()
|
||||
# only replay (re-execute) done tasks on the first tick
|
||||
@@ -737,14 +751,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,16 +831,39 @@ 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:
|
||||
self._loaded_write_ids = {id(w): w for w in self.checkpoint_pending_writes}
|
||||
# Resuming from a previous checkpoint requires two things:
|
||||
# 1. A prior checkpoint exists (channel_versions is non-empty)
|
||||
# 2. The input signals continuation (not a fresh run with new input)
|
||||
@@ -993,7 +1033,9 @@ class PregelLoop:
|
||||
if self._exit_delta_writes is not None:
|
||||
for c, v in input_writes:
|
||||
if isinstance(self.specs.get(c), DeltaChannel):
|
||||
self._exit_delta_writes.append((self.step, NULL_TASK_ID, c, v))
|
||||
self._exit_delta_writes.append(
|
||||
(self.step, NULL_TASK_ID, "", c, v)
|
||||
)
|
||||
# Persist delta-channel input writes so sub-freq inputs are
|
||||
# recoverable via ancestor walk (mirrors the Command input path).
|
||||
if self.durability != "exit":
|
||||
@@ -1219,9 +1261,7 @@ class PregelLoop:
|
||||
)
|
||||
|
||||
pending = [
|
||||
(step, tid, ch, v)
|
||||
for (step, tid, ch, v) in self._exit_delta_writes
|
||||
if ch not in channels_to_snapshot
|
||||
w for w in self._exit_delta_writes if w[3] not in channels_to_snapshot
|
||||
]
|
||||
if not pending:
|
||||
return
|
||||
@@ -1256,11 +1296,19 @@ class PregelLoop:
|
||||
# sees the stub as its parent.
|
||||
self.checkpoint_config = anchor_config
|
||||
|
||||
# Step-prefixed synthetic task_id preserves chronological superstep
|
||||
# order under the saver's ORDER BY task_id, idx sorting.
|
||||
grouped: dict[tuple[int, str], list[tuple[str, Any]]] = {}
|
||||
for step, tid, ch, v in pending:
|
||||
grouped.setdefault((step, tid), []).append((ch, v))
|
||||
# The checkpoint's own superstep is stored as sync durability stores
|
||||
# it, so it interleaves with the writes a resume loaded from it. Later
|
||||
# supersteps sort after every real task path and task id, in step
|
||||
# order, so this holds whether a saver orders by path or by id.
|
||||
grouped: dict[tuple[str, str], list[tuple[str, Any]]] = {}
|
||||
for step, tid, path, ch, v in pending:
|
||||
if tid == NULL_TASK_ID:
|
||||
key = (exit_delta_task_id(step, tid), "")
|
||||
elif step == self._exit_first_step:
|
||||
key = (tid, path)
|
||||
else:
|
||||
key = (exit_delta_late_task_id(step, tid), f"~~{step:010d}{path}")
|
||||
grouped.setdefault(key, []).append((ch, v))
|
||||
anchor_write_config = patch_configurable(
|
||||
anchor_config,
|
||||
{
|
||||
@@ -1270,22 +1318,21 @@ class PregelLoop:
|
||||
CONFIG_KEY_CHECKPOINT_ID: anchor_config[CONF][CONFIG_KEY_CHECKPOINT_ID],
|
||||
},
|
||||
)
|
||||
for (step, tid), entries in grouped.items():
|
||||
synth_tid = exit_delta_task_id(step, tid)
|
||||
for (tid, path), entries in grouped.items():
|
||||
if self.checkpointer_put_writes_accepts_task_path:
|
||||
fut = self.submit(
|
||||
self.checkpointer_put_writes,
|
||||
anchor_write_config,
|
||||
entries,
|
||||
synth_tid,
|
||||
"",
|
||||
tid,
|
||||
path,
|
||||
)
|
||||
else:
|
||||
fut = self.submit(
|
||||
self.checkpointer_put_writes,
|
||||
anchor_write_config,
|
||||
entries,
|
||||
synth_tid,
|
||||
tid,
|
||||
)
|
||||
if self._delta_write_futs is not None:
|
||||
self._delta_write_futs.append(fut)
|
||||
|
||||
@@ -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,
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -6,11 +6,13 @@ channel), lazy stub creation when no parent exists, and proper read-path
|
||||
reconstruction via ancestor walks.
|
||||
"""
|
||||
|
||||
import operator
|
||||
import uuid
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import AIMessage, 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
|
||||
@@ -19,6 +21,7 @@ from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.graph.message import _messages_delta_reducer
|
||||
from langgraph.pregel._checkpoint import exit_delta_task_id
|
||||
from langgraph.types import Command, Durability, interrupt
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
@@ -389,3 +392,112 @@ async def test_exit_snapshot_then_tail_deltas() -> None:
|
||||
assert "seed-msg" in contents
|
||||
assert "tail-msg" in contents
|
||||
assert contents.index("seed-msg") < contents.index("tail-msg")
|
||||
|
||||
|
||||
def _append(current: list, writes: list) -> list:
|
||||
out = list(current)
|
||||
for write in writes:
|
||||
out.extend(write)
|
||||
return out
|
||||
|
||||
|
||||
class _ResumeState(TypedDict):
|
||||
log: Annotated[list, DeltaChannel(_append)]
|
||||
plain: Annotated[list, operator.add]
|
||||
|
||||
|
||||
def _both(marker: str) -> dict:
|
||||
return {"log": [marker], "plain": [marker]}
|
||||
|
||||
|
||||
def _ask(marker: str) -> Any:
|
||||
def ask(state: _ResumeState) -> dict:
|
||||
interrupt("approve?")
|
||||
return _both(marker)
|
||||
|
||||
return ask
|
||||
|
||||
|
||||
@pytest.mark.parametrize("addressed", [False, True])
|
||||
def test_resume_after_a_parallel_interrupt_replays_in_live_order(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability, addressed: bool
|
||||
) -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("done", lambda state: _both("done"))
|
||||
builder.add_node("ask", _ask("ask"))
|
||||
builder.add_node("after", lambda state: _both("after"))
|
||||
builder.add_edge(START, "done")
|
||||
builder.add_edge(START, "ask")
|
||||
builder.add_edge("ask", "after")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability=durability)
|
||||
head = graph.get_state(config).config
|
||||
|
||||
graph.invoke(
|
||||
Command(resume="yes"), head if addressed else config, durability=durability
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.values["log"] == state.values["plain"]
|
||||
assert sorted(state.values["log"]) == ["after", "ask", "done", "in"]
|
||||
|
||||
|
||||
def test_resume_interleaves_the_resumed_superstep_by_task_path(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("z_done", lambda state: _both("z"))
|
||||
builder.add_node("a_asks", _ask("a"))
|
||||
builder.add_edge(START, "z_done")
|
||||
builder.add_edge(START, "a_asks")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability=durability)
|
||||
|
||||
graph.invoke(Command(resume="yes"), config, durability=durability)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.values["log"] == state.values["plain"] == ["in", "a", "z"]
|
||||
|
||||
|
||||
class _TaskIdOrderSaver(InMemorySaver):
|
||||
"""Replays each checkpoint's writes by task id, as savers without task path
|
||||
ordering do."""
|
||||
|
||||
def get_tuple(self, config: Any) -> Any:
|
||||
tup = super().get_tuple(config)
|
||||
if tup and tup.pending_writes:
|
||||
tup = tup._replace(pending_writes=sorted(tup.pending_writes))
|
||||
return tup
|
||||
|
||||
get_delta_channel_history = BaseCheckpointSaver.get_delta_channel_history
|
||||
|
||||
|
||||
def test_exit_run_replays_supersteps_in_order_on_a_task_id_ordered_saver() -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
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")
|
||||
graph = builder.compile(checkpointer=_TaskIdOrderSaver())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
graph.invoke(_both("in"), config, durability="exit")
|
||||
|
||||
assert graph.get_state(config).values["log"] == ["in", "a", "b"]
|
||||
|
||||
|
||||
def test_exit_resume_replays_supersteps_in_order_on_a_task_id_ordered_saver() -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("ask", _ask("ask"))
|
||||
builder.add_node("after", lambda state: _both("after"))
|
||||
builder.add_edge(START, "ask")
|
||||
builder.add_edge("ask", "after")
|
||||
graph = builder.compile(checkpointer=_TaskIdOrderSaver())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability="exit")
|
||||
|
||||
graph.invoke(Command(resume="yes"), config, durability="exit")
|
||||
|
||||
assert graph.get_state(config).values["log"] == ["in", "ask", "after"]
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
"""`DeltaChannel` replay must apply parallel writes in the order `invoke` did."""
|
||||
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
from langgraph.types import Send
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
# Sorted, because live execution applies PULL tasks in node-name order.
|
||||
FAN_OUT_NAMES = ["a", "b", "c", "d", "e", "f", "g", "h"]
|
||||
SEND_ARGS = [f"send-{i:02d}" for i in range(12)]
|
||||
|
||||
|
||||
def _append_reducer(current: list, updates: list) -> list:
|
||||
return [*current, *(x for u in updates for x in u)]
|
||||
|
||||
|
||||
def _build_fan_out_graph(checkpointer: BaseCheckpointSaver) -> Any:
|
||||
class State(TypedDict):
|
||||
items: Annotated[
|
||||
list, DeltaChannel(_append_reducer, list, snapshot_frequency=10_000)
|
||||
]
|
||||
|
||||
def make_node(label: str) -> Any:
|
||||
return lambda state: {"items": [label]}
|
||||
|
||||
builder = StateGraph(State)
|
||||
for name in FAN_OUT_NAMES:
|
||||
builder.add_node(name, make_node(name))
|
||||
builder.add_edge(START, name)
|
||||
builder.add_edge(name, END)
|
||||
return builder.compile(checkpointer=checkpointer)
|
||||
|
||||
|
||||
def _build_send_fan_out_graph(checkpointer: BaseCheckpointSaver) -> Any:
|
||||
class State(TypedDict):
|
||||
items: Annotated[
|
||||
list, DeltaChannel(_append_reducer, list, snapshot_frequency=10_000)
|
||||
]
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("worker", lambda arg: {"items": [arg]})
|
||||
builder.add_conditional_edges(
|
||||
START, lambda state: [Send("worker", n) for n in SEND_ARGS]
|
||||
)
|
||||
builder.add_edge("worker", END)
|
||||
return builder.compile(checkpointer=checkpointer)
|
||||
|
||||
|
||||
async def test_get_state_matches_live_send_order(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_send_fan_out_graph(async_checkpointer)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
live = (await graph.ainvoke({"items": []}, config))["items"]
|
||||
replayed = (await graph.aget_state(config)).values["items"]
|
||||
|
||||
assert live == SEND_ARGS
|
||||
assert replayed == live
|
||||
|
||||
|
||||
async def test_get_state_matches_live_invoke_order(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_fan_out_graph(async_checkpointer)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
live = (await graph.ainvoke({"items": []}, config))["items"]
|
||||
replayed = (await graph.aget_state(config)).values["items"]
|
||||
|
||||
assert live == FAN_OUT_NAMES
|
||||
assert replayed == live
|
||||
|
||||
|
||||
async def test_continuing_thread_preserves_committed_prefix(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
graph = _build_fan_out_graph(async_checkpointer)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
first = (await graph.ainvoke({"items": []}, config))["items"]
|
||||
second = (await graph.ainvoke({"items": []}, config))["items"]
|
||||
|
||||
assert second == first + first
|
||||
assert (await graph.aget_state(config)).values["items"] == second
|
||||
|
||||
|
||||
async def test_state_history_reports_live_order_at_every_step(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
runs = 3
|
||||
graph = _build_fan_out_graph(async_checkpointer)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
for _ in range(runs):
|
||||
await graph.ainvoke({"items": []}, config)
|
||||
live = FAN_OUT_NAMES * runs
|
||||
|
||||
seen = [
|
||||
s.values["items"]
|
||||
async for s in graph.aget_state_history(config)
|
||||
if "items" in s.values
|
||||
]
|
||||
|
||||
assert max(map(len, seen)) == len(live)
|
||||
for values in seen:
|
||||
assert values == live[: len(values)], f"{values} is not a prefix of {live}"
|
||||
@@ -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