mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-11 18:55:17 +02:00
Compare commits
23
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
536059f99b | ||
|
|
b4991f1ba3 | ||
|
|
6aa0afba68 | ||
|
|
12aeb0fddb | ||
|
|
d05236f805 | ||
|
|
9d92f33cca | ||
|
|
26356227c4 | ||
|
|
a5dbacae0d | ||
|
|
cba111d8d6 | ||
|
|
93a5a28008 | ||
|
|
bfcfea554e | ||
|
|
93f5eaff21 | ||
|
|
5965d72ff7 | ||
|
|
40a2e6d845 | ||
|
|
a0053bb616 | ||
|
|
87f1c8eb9a | ||
|
|
736a18eaab | ||
|
|
ec4f8535f8 | ||
|
|
ce14d33687 | ||
|
|
708aaa4ebc | ||
|
|
39c523eb0a | ||
|
|
3a1796ecf0 | ||
|
|
2b93e8b79f |
+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,
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -176,6 +176,35 @@ async def test_get_tuple_pending_writes(saver: BaseCheckpointSaver) -> None:
|
||||
)
|
||||
|
||||
|
||||
async def test_get_tuple_pending_writes_in_writes_sort_key_order(
|
||||
saver: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""pending_writes come back in writes_sort_key order, not put or task_id order."""
|
||||
config = generate_config(str(uuid4()))
|
||||
stored = await saver.aput(config, generate_checkpoint(), generate_metadata(), {})
|
||||
await saver.aput_writes(
|
||||
stored, [("ch", "b")], "00000000-0000-0000-0000-000000000000", "~pull, 02"
|
||||
)
|
||||
await saver.aput_writes(
|
||||
stored,
|
||||
[("ch", "a1"), ("ch", "a2")],
|
||||
"ffffffff-ffff-ffff-ffff-ffffffffffff",
|
||||
"~pull, 01",
|
||||
)
|
||||
await saver.aput_writes(
|
||||
stored, [("ch", "input")], "88888888-8888-8888-8888-888888888888"
|
||||
)
|
||||
|
||||
tup = await saver.aget_tuple(stored)
|
||||
assert tup is not None
|
||||
values = [w[2] for w in tup.pending_writes or []]
|
||||
assert values == ["input", "a1", "a2", "b"], (
|
||||
f"Expected writes_sort_key order ['input', 'a1', 'a2', 'b'], got {values}. "
|
||||
"Put order gives ['b', 'a1', 'a2', 'input'], task_id order gives "
|
||||
"['b', 'input', 'a1', 'a2']."
|
||||
)
|
||||
|
||||
|
||||
async def test_get_tuple_respects_namespace(saver: BaseCheckpointSaver) -> None:
|
||||
"""checkpoint_ns filtering."""
|
||||
tid = str(uuid4())
|
||||
@@ -223,6 +252,7 @@ ALL_GET_TUPLE_TESTS = [
|
||||
test_get_tuple_metadata,
|
||||
test_get_tuple_parent_config,
|
||||
test_get_tuple_pending_writes,
|
||||
test_get_tuple_pending_writes_in_writes_sort_key_order,
|
||||
test_get_tuple_respects_namespace,
|
||||
test_get_tuple_nonexistent_checkpoint_id,
|
||||
]
|
||||
|
||||
@@ -14,8 +14,8 @@ async def memory_checkpointer():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_memory_base():
|
||||
"""InMemorySaver passes all base capability tests."""
|
||||
async def test_validate_memory():
|
||||
"""InMemorySaver passes the tests of every capability it implements."""
|
||||
report = await validate(memory_checkpointer)
|
||||
report.print_report()
|
||||
assert report.passed_all_base(), f"Base tests failed: {report.to_dict()}"
|
||||
assert report.passed_all(), f"Capability tests failed: {report.to_dict()}"
|
||||
|
||||
Generated
+1
-1
@@ -279,7 +279,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -451,7 +451,8 @@ class PostgresSaver(BasePostgresSaver):
|
||||
* 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`.
|
||||
`checkpoint_id`, starting at the target; page size is
|
||||
`_DELTA_PAGE_SIZE`.
|
||||
Stops paging when every channel has found its seed or a page comes
|
||||
back short.
|
||||
|
||||
@@ -468,7 +469,7 @@ class PostgresSaver(BasePostgresSaver):
|
||||
thread_id = config["configurable"]["thread_id"]
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
checkpoint_id = get_checkpoint_id(config)
|
||||
if checkpoint_id is None:
|
||||
if not checkpoint_id:
|
||||
target = self.get_tuple(config)
|
||||
if target is None:
|
||||
return {ch: {"writes": []} for ch in channels}
|
||||
@@ -476,7 +477,7 @@ class PostgresSaver(BasePostgresSaver):
|
||||
|
||||
# Stage 1: paged K-JSONB-lookup scan, walking the parent chain in
|
||||
# Python after each page. Stops as soon as every channel has its seed.
|
||||
stage1_sql = _build_delta_stage1_sql(channels, paged=True)
|
||||
stage1_sql = _build_delta_stage1_sql(channels, paged=True, include_cursor=True)
|
||||
parent_of: dict[str, str | None] = {}
|
||||
ver_by_i_by_cid: list[dict[str, str | None]] = [{} for _ in channels]
|
||||
hb_by_i_by_cid: list[dict[str, bool]] = [{} for _ in channels]
|
||||
@@ -486,7 +487,7 @@ class PostgresSaver(BasePostgresSaver):
|
||||
seed_inline_by_ch: dict[str, Any] = {}
|
||||
walk_cursor_by_ch: dict[str, str | None] = {}
|
||||
seeded: set[str] = set()
|
||||
cursor: str | None = None
|
||||
cursor: str | None = checkpoint_id
|
||||
|
||||
with self._cursor() as cur:
|
||||
while True:
|
||||
@@ -495,7 +496,7 @@ class PostgresSaver(BasePostgresSaver):
|
||||
# ver_i, blob channel, blob version, inline_i
|
||||
stage1_params.extend([ch, ch, ch, ch])
|
||||
stage1_params.extend(
|
||||
[thread_id, checkpoint_ns, cursor, cursor, _DELTA_PAGE_SIZE]
|
||||
[thread_id, checkpoint_ns, cursor, _DELTA_PAGE_SIZE]
|
||||
)
|
||||
cur.execute(stage1_sql, stage1_params)
|
||||
page = cur.fetchall()
|
||||
@@ -527,6 +528,7 @@ class PostgresSaver(BasePostgresSaver):
|
||||
if len(seeded) == len(channels) or len(page) < _DELTA_PAGE_SIZE:
|
||||
break
|
||||
cursor = oldest
|
||||
stage1_sql = _build_delta_stage1_sql(channels, paged=True)
|
||||
|
||||
# Stage 2: per-channel UNION ALL — one writes branch per channel
|
||||
# with non-empty chain, plus one blob branch per seeded channel.
|
||||
|
||||
@@ -417,13 +417,13 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
thread_id = config["configurable"]["thread_id"]
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
checkpoint_id = get_checkpoint_id(config)
|
||||
if checkpoint_id is None:
|
||||
if not checkpoint_id:
|
||||
target = await self.aget_tuple(config)
|
||||
if target is None:
|
||||
return {ch: {"writes": []} for ch in channels}
|
||||
checkpoint_id = target.config["configurable"]["checkpoint_id"]
|
||||
|
||||
stage1_sql = _build_delta_stage1_sql(channels, paged=True)
|
||||
stage1_sql = _build_delta_stage1_sql(channels, paged=True, include_cursor=True)
|
||||
parent_of: dict[str, str | None] = {}
|
||||
ver_by_i_by_cid: list[dict[str, str | None]] = [{} for _ in channels]
|
||||
hb_by_i_by_cid: list[dict[str, bool]] = [{} for _ in channels]
|
||||
@@ -433,7 +433,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
seed_inline_by_ch: dict[str, Any] = {}
|
||||
walk_cursor_by_ch: dict[str, str | None] = {}
|
||||
seeded: set[str] = set()
|
||||
cursor: str | None = None
|
||||
cursor: str | None = checkpoint_id
|
||||
|
||||
async with self._cursor() as cur:
|
||||
while True:
|
||||
@@ -442,7 +442,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
# ver_i, blob channel, blob version, inline_i
|
||||
stage1_params.extend([ch, ch, ch, ch])
|
||||
stage1_params.extend(
|
||||
[thread_id, checkpoint_ns, cursor, cursor, _DELTA_PAGE_SIZE]
|
||||
[thread_id, checkpoint_ns, cursor, _DELTA_PAGE_SIZE]
|
||||
)
|
||||
await cur.execute(stage1_sql, stage1_params)
|
||||
page = await cur.fetchall()
|
||||
@@ -472,6 +472,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
if len(seeded) == len(channels) or len(page) < _DELTA_PAGE_SIZE:
|
||||
break
|
||||
cursor = oldest
|
||||
stage1_sql = _build_delta_stage1_sql(channels, paged=True)
|
||||
|
||||
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
|
||||
channels_with_seed = [ch for ch in channels if seed_ver_by_ch[ch] is not None]
|
||||
@@ -687,5 +688,23 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
self.adelete_thread(thread_id), self.loop
|
||||
).result()
|
||||
|
||||
def get_delta_channel_history(
|
||||
self, *, config: RunnableConfig, channels: Sequence[str]
|
||||
) -> Mapping[str, DeltaChannelHistory]:
|
||||
"""Sync bridge to `aget_delta_channel_history`, guarded like `get_tuple`."""
|
||||
try:
|
||||
if asyncio.get_running_loop() is self.loop:
|
||||
raise asyncio.InvalidStateError(
|
||||
"Synchronous calls to AsyncPostgresSaver are only allowed from a "
|
||||
"different thread. From the main thread, use the async interface. "
|
||||
"For example, use `await checkpointer.aget_delta_channel_history(...)`."
|
||||
)
|
||||
except RuntimeError:
|
||||
pass
|
||||
return asyncio.run_coroutine_threadsafe(
|
||||
self.aget_delta_channel_history(config=config, channels=channels),
|
||||
self.loop,
|
||||
).result()
|
||||
|
||||
|
||||
__all__ = ["AsyncPostgresSaver", "AsyncShallowPostgresSaver", "Conn"]
|
||||
|
||||
@@ -14,6 +14,7 @@ from langgraph.checkpoint.base import (
|
||||
DeltaChannelHistory,
|
||||
PendingWrite,
|
||||
get_checkpoint_id,
|
||||
writes_sort_key,
|
||||
)
|
||||
from langgraph.checkpoint.serde.types import TASKS
|
||||
from psycopg.types.json import Jsonb
|
||||
@@ -109,7 +110,7 @@ select
|
||||
) as channel_values,
|
||||
(
|
||||
select
|
||||
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
|
||||
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob, convert_to(cw.task_path, 'UTF8'), cw.idx::text::bytea])
|
||||
from checkpoint_writes cw
|
||||
where cw.thread_id = checkpoints.thread_id
|
||||
and cw.checkpoint_ns = checkpoints.checkpoint_ns
|
||||
@@ -168,6 +169,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
|
||||
|
||||
@@ -176,10 +178,13 @@ class _DeltaStage2Row(TypedDict, total=False):
|
||||
# `_build_delta_stage2_sql` document their shapes.
|
||||
|
||||
|
||||
def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
|
||||
def _build_delta_stage1_sql(
|
||||
channels: Sequence[str], *, paged: bool, include_cursor: bool = False
|
||||
) -> str:
|
||||
"""Build stage 1 SQL with K parallel version lookups + seed probes.
|
||||
|
||||
For channels=["messages", "files"] (with `paged=True`) the result is::
|
||||
For channels=["messages", "files"] (with `paged=True, include_cursor=True`)
|
||||
the result is::
|
||||
|
||||
SELECT checkpoint_id, parent_checkpoint_id,
|
||||
checkpoint -> 'channel_versions' ->> %s AS ver_0,
|
||||
@@ -195,7 +200,7 @@ def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
|
||||
checkpoint -> 'channel_values' -> %s AS inline_1
|
||||
FROM checkpoints
|
||||
WHERE thread_id = %s AND checkpoint_ns = %s
|
||||
AND (%s::text IS NULL OR checkpoint_id < %s)
|
||||
AND checkpoint_id <= %s
|
||||
ORDER BY checkpoint_id DESC
|
||||
LIMIT %s
|
||||
|
||||
@@ -236,9 +241,16 @@ def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
|
||||
and uses safe identifiers).
|
||||
|
||||
Caller must extend params with `[ch_0 x4, ch_1 x4, ..., thread_id, ns,
|
||||
cursor, cursor, page_size]` when `paged=True` — four per channel: the
|
||||
version lookup, the blob's channel, the version the blob must match, and the
|
||||
inline lookup.
|
||||
cursor, page_size]` when `paged=True` — four per channel: the version
|
||||
lookup, the blob's channel, the version the blob must match, and the inline
|
||||
lookup.
|
||||
|
||||
Pages run newest-first from the target down. The first page passes the
|
||||
target as the cursor with `include_cursor=True`, so it opens with the
|
||||
target's own row, whose parent starts the walk; each later page continues
|
||||
below the oldest row read. A checkpoint's ancestors have smaller ids (uuid6
|
||||
is time-ordered, which `get_tuple` also relies on to find the latest
|
||||
checkpoint), so no row newer than the target is part of its chain.
|
||||
|
||||
When `paged=False`, the WHERE has no cursor predicate and there's no
|
||||
LIMIT/ORDER BY — kept as a non-public helper for tests/diagnostics.
|
||||
@@ -262,7 +274,7 @@ def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
|
||||
)
|
||||
if paged:
|
||||
sql += (
|
||||
" AND (%s::text IS NULL OR checkpoint_id < %s)"
|
||||
f" AND checkpoint_id {'<=' if include_cursor else '<'} %s"
|
||||
" ORDER BY checkpoint_id DESC LIMIT %s"
|
||||
)
|
||||
return sql
|
||||
@@ -297,7 +309,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)"
|
||||
@@ -305,7 +317,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"
|
||||
@@ -410,8 +423,8 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
||||
materialized at this point),
|
||||
(c) the next ancestor cid isn't in `parent_of` yet (waiting for
|
||||
a later page; the cursor stays put), or
|
||||
(d) the target's own row isn't in `parent_of` yet (the walk has
|
||||
not started; no cursor is set, so a later page retries).
|
||||
(d) the target's own row isn't in `parent_of` (the target doesn't
|
||||
exist, so the walk never starts).
|
||||
|
||||
Mutates `chain_by_ch`, `seed_ver_by_ch`, `seed_inline_by_ch`,
|
||||
`walk_cursor_by_ch`, and `seeded` in place.
|
||||
@@ -419,8 +432,9 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
||||
for i, ch in enumerate(channels):
|
||||
if ch in seeded:
|
||||
continue
|
||||
# Pages start at the thread head, so the target may not have
|
||||
# loaded yet; a `None` cursor would read as "target is a root".
|
||||
# The first page opens with the target's row, so it's missing only
|
||||
# when the target doesn't exist; a `None` cursor would read as
|
||||
# "target is a root".
|
||||
if ch not in walk_cursor_by_ch:
|
||||
if target_id not in parent_of:
|
||||
continue
|
||||
@@ -473,10 +487,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]] = {}
|
||||
|
||||
@@ -487,8 +502,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"
|
||||
@@ -497,10 +518,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
|
||||
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: writes_sort_key(w[4], w[2], w[3]), reverse=True)
|
||||
|
||||
result: dict[str, DeltaChannelHistory] = {}
|
||||
for ch in channels:
|
||||
@@ -510,7 +531,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()
|
||||
@@ -553,20 +576,15 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
|
||||
]
|
||||
|
||||
def _load_writes(
|
||||
self, writes: list[tuple[bytes, bytes, bytes, bytes]]
|
||||
self, writes: list[tuple[bytes, bytes, bytes, bytes, bytes, bytes]] | None
|
||||
) -> list[tuple[str, str, Any]]:
|
||||
return (
|
||||
[
|
||||
(
|
||||
tid.decode(),
|
||||
channel.decode(),
|
||||
self.serde.loads_typed((t.decode(), v)),
|
||||
)
|
||||
for tid, channel, t, v in writes
|
||||
]
|
||||
if writes
|
||||
else []
|
||||
)
|
||||
return [
|
||||
(tid.decode(), channel.decode(), self.serde.loads_typed((t.decode(), v)))
|
||||
for tid, channel, t, v, _, _ in sorted(
|
||||
writes or [],
|
||||
key=lambda w: writes_sort_key(w[4].decode(), w[0].decode(), int(w[5])),
|
||||
)
|
||||
]
|
||||
|
||||
def _dump_writes(
|
||||
self,
|
||||
|
||||
@@ -97,7 +97,7 @@ select
|
||||
) as channel_values,
|
||||
(
|
||||
select
|
||||
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
|
||||
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob, convert_to(cw.task_path, 'UTF8'), cw.idx::text::bytea])
|
||||
from checkpoint_writes cw
|
||||
where cw.thread_id = checkpoints.thread_id
|
||||
and cw.checkpoint_ns = checkpoints.checkpoint_ns
|
||||
|
||||
@@ -36,6 +36,7 @@ from langgraph.store.base import (
|
||||
ensure_embeddings,
|
||||
get_text_at_path,
|
||||
tokenize_path,
|
||||
validate_op_namespace,
|
||||
)
|
||||
from psycopg import Capabilities, Connection, Cursor, Pipeline
|
||||
from psycopg.rows import DictRow, dict_row
|
||||
@@ -1386,6 +1387,7 @@ def _group_ops(ops: Iterable[Op]) -> tuple[dict[type, list[tuple[int, Op]]], int
|
||||
grouped_ops: dict[type, list[tuple[int, Op]]] = defaultdict(list)
|
||||
tot = 0
|
||||
for idx, op in enumerate(ops):
|
||||
validate_op_namespace(op)
|
||||
grouped_ops[type(op)].append((idx, op))
|
||||
tot += 1
|
||||
return grouped_ops, tot
|
||||
|
||||
@@ -12,7 +12,7 @@ readme = "README.md"
|
||||
license = "MIT"
|
||||
license-files = ['LICENSE']
|
||||
dependencies = [
|
||||
"langgraph-checkpoint>=4.1.0,<5.0.0",
|
||||
"langgraph-checkpoint>=4.3.0,<5.0.0",
|
||||
"orjson>=3.11.5",
|
||||
"psycopg>=3.2.0",
|
||||
"psycopg-pool>=3.2.0",
|
||||
|
||||
@@ -13,8 +13,10 @@ import pytest
|
||||
from langchain_core.embeddings import Embeddings
|
||||
from langgraph.store.base import (
|
||||
GetOp,
|
||||
InvalidNamespaceError,
|
||||
Item,
|
||||
ListNamespacesOp,
|
||||
MatchCondition,
|
||||
PutOp,
|
||||
SearchOp,
|
||||
)
|
||||
@@ -871,3 +873,65 @@ async def test_omit_expired_search_pagination(store: AsyncPostgresStore) -> None
|
||||
page2 = await store.asearch(ns, limit=2, offset=2)
|
||||
assert [i.key for i in page1] == ["a", "b"]
|
||||
assert [i.key for i in page2] == ["c"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace", [("foo.bar",), ("foo", ""), ("foo", 1)])
|
||||
async def test_abatch_rejects_invalid_namespace_labels(
|
||||
store: AsyncPostgresStore, namespace: tuple
|
||||
) -> None:
|
||||
await store.aput(("foo", "bar"), "key", {"original": True})
|
||||
|
||||
for op in (
|
||||
GetOp(namespace, "key"),
|
||||
GetOp(namespace, "key", refresh_ttl=True),
|
||||
PutOp(namespace, "key", {"changed": True}),
|
||||
PutOp(namespace, "key", None),
|
||||
SearchOp(namespace),
|
||||
ListNamespacesOp((MatchCondition("prefix", namespace),)),
|
||||
ListNamespacesOp((MatchCondition("suffix", namespace),)),
|
||||
):
|
||||
with pytest.raises(InvalidNamespaceError):
|
||||
await store.abatch([op])
|
||||
|
||||
item = await store.aget(("foo", "bar"), "key")
|
||||
assert item is not None and item.value == {"original": True}
|
||||
|
||||
|
||||
async def test_invalid_namespace_only_fails_its_own_call(
|
||||
store: AsyncPostgresStore,
|
||||
) -> None:
|
||||
"""Concurrent calls share one `abatch`, which fails every op if it raises.
|
||||
|
||||
Labels are checked before an op is queued, so one caller's bad label cannot
|
||||
fail another caller's request.
|
||||
"""
|
||||
await store.aput(("foo", "bar"), "key", {"original": True})
|
||||
|
||||
valid, invalid = await asyncio.gather(
|
||||
store.aget(("foo", "bar"), "key"),
|
||||
store.aget(("foo.bar",), "key"),
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
assert isinstance(valid, Item) and valid.value == {"original": True}
|
||||
assert isinstance(invalid, InvalidNamespaceError)
|
||||
|
||||
|
||||
async def test_sync_methods_reject_invalid_namespace_labels(
|
||||
store: AsyncPostgresStore,
|
||||
) -> None:
|
||||
"""The sync wrappers run off the event loop thread and must validate too."""
|
||||
await store.aput(("foo", "bar"), "key", {"original": True})
|
||||
|
||||
for call in (
|
||||
lambda: store.get(("foo.bar",), "key"),
|
||||
lambda: store.search(("foo.bar",)),
|
||||
lambda: store.delete(("foo.bar",), "key"),
|
||||
lambda: store.list_namespaces(prefix=("foo.bar",)),
|
||||
lambda: store.batch([GetOp(("foo.bar",), "key")]),
|
||||
):
|
||||
with pytest.raises(InvalidNamespaceError):
|
||||
await asyncio.to_thread(call)
|
||||
|
||||
item = await store.aget(("foo", "bar"), "key")
|
||||
assert item is not None and item.value == {"original": True}
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
@@ -14,7 +15,7 @@ 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 langgraph.checkpoint.postgres.base import _DELTA_PAGE_SIZE, BasePostgresSaver
|
||||
from tests.conftest import DEFAULT_URI
|
||||
|
||||
CHANNEL = "items"
|
||||
@@ -23,8 +24,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).
|
||||
# The real page size is the control; the rest split the walk from the target
|
||||
# to its seed across pages.
|
||||
PAGE_SIZES = [_DELTA_PAGE_SIZE, 3, 2, 1]
|
||||
|
||||
|
||||
@@ -92,7 +93,7 @@ def _assert_history(entry: DeltaChannelHistory, page_size: int) -> None:
|
||||
|
||||
|
||||
@pytest.mark.parametrize("page_size", PAGE_SIZES)
|
||||
async def test_async_target_older_than_the_first_page(
|
||||
async def test_async_walk_continues_across_pages(
|
||||
page_size: int, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setattr("langgraph.checkpoint.postgres.aio._DELTA_PAGE_SIZE", page_size)
|
||||
@@ -106,7 +107,7 @@ async def test_async_target_older_than_the_first_page(
|
||||
|
||||
|
||||
@pytest.mark.parametrize("page_size", PAGE_SIZES)
|
||||
def test_sync_target_older_than_the_first_page(
|
||||
def test_sync_walk_continues_across_pages(
|
||||
page_size: int, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setattr("langgraph.checkpoint.postgres._DELTA_PAGE_SIZE", page_size)
|
||||
@@ -119,6 +120,84 @@ def test_sync_target_older_than_the_first_page(
|
||||
_assert_history(result[CHANNEL], page_size)
|
||||
|
||||
|
||||
def _record_rows_read(monkeypatch: pytest.MonkeyPatch) -> list[str]:
|
||||
read: list[str] = []
|
||||
ingest = BasePostgresSaver._ingest_stage1_page
|
||||
|
||||
def record(rows: Sequence[Mapping[str, Any]], *args: Any) -> str | None:
|
||||
read.extend(row["checkpoint_id"] for row in rows)
|
||||
return ingest(rows, *args)
|
||||
|
||||
monkeypatch.setattr(BasePostgresSaver, "_ingest_stage1_page", staticmethod(record))
|
||||
return read
|
||||
|
||||
|
||||
def _ids_from_target_down(configs: list[dict]) -> list[str]:
|
||||
return [c["configurable"]["checkpoint_id"] for c in configs[TARGET_STEP::-1]]
|
||||
|
||||
|
||||
async def test_async_walk_reads_nothing_newer_than_the_target(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
read = _record_rows_read(monkeypatch)
|
||||
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], _DELTA_PAGE_SIZE)
|
||||
assert read == _ids_from_target_down(configs)
|
||||
|
||||
|
||||
def test_sync_walk_reads_nothing_newer_than_the_target(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
read = _record_rows_read(monkeypatch)
|
||||
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], _DELTA_PAGE_SIZE)
|
||||
assert read == _ids_from_target_down(configs)
|
||||
|
||||
|
||||
def _empty_checkpoint_id(config: dict) -> dict:
|
||||
return {"configurable": {**config["configurable"], "checkpoint_id": ""}}
|
||||
|
||||
|
||||
async def test_async_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
|
||||
async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver:
|
||||
await saver.setup()
|
||||
configs = await _abuild_chain(saver)
|
||||
latest = await saver.aget_delta_channel_history(
|
||||
config=configs[-1], channels=[CHANNEL]
|
||||
)
|
||||
|
||||
result = await saver.aget_delta_channel_history(
|
||||
config=_empty_checkpoint_id(configs[-1]), channels=[CHANNEL]
|
||||
)
|
||||
|
||||
assert latest[CHANNEL]["writes"]
|
||||
assert result == latest
|
||||
|
||||
|
||||
def test_sync_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
|
||||
with PostgresSaver.from_conn_string(DEFAULT_URI) as saver:
|
||||
saver.setup()
|
||||
configs = _build_chain(saver)
|
||||
latest = saver.get_delta_channel_history(config=configs[-1], channels=[CHANNEL])
|
||||
|
||||
result = saver.get_delta_channel_history(
|
||||
config=_empty_checkpoint_id(configs[-1]), channels=[CHANNEL]
|
||||
)
|
||||
|
||||
assert latest[CHANNEL]["writes"]
|
||||
assert result == latest
|
||||
|
||||
|
||||
async def test_root_target_has_no_history_and_still_terminates(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
|
||||
@@ -11,6 +11,7 @@ import pytest
|
||||
from langchain_core.embeddings import Embeddings
|
||||
from langgraph.store.base import (
|
||||
GetOp,
|
||||
InvalidNamespaceError,
|
||||
Item,
|
||||
ListNamespacesOp,
|
||||
MatchCondition,
|
||||
@@ -1164,3 +1165,51 @@ def test_namespace_labels_with_trailing_newline(store) -> None:
|
||||
assert set(store.list_namespaces(prefix=["users", "alice"], limit=100)) == {
|
||||
("users", "alice"),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace", [("foo.bar",), ("foo", ""), ("foo", 1)])
|
||||
@pytest.mark.parametrize(
|
||||
"kind", ["get", "put", "delete", "search", "list_prefix", "list_suffix"]
|
||||
)
|
||||
def test_batch_rejects_invalid_namespace_labels(
|
||||
store, namespace: tuple, kind: str
|
||||
) -> None:
|
||||
"""Ops passed straight to `batch` must not reach another namespace.
|
||||
|
||||
Namespaces are stored dot-joined, so `("foo.bar",)` flattens to the same
|
||||
text as `("foo", "bar")`. `BaseStore` methods validate labels themselves,
|
||||
but `batch` takes ops as given.
|
||||
"""
|
||||
op = {
|
||||
"get": GetOp(namespace, "key"),
|
||||
"put": PutOp(namespace, "key", {"changed": True}),
|
||||
"delete": PutOp(namespace, "key", None),
|
||||
"search": SearchOp(namespace),
|
||||
"list_prefix": ListNamespacesOp((MatchCondition("prefix", namespace),)),
|
||||
"list_suffix": ListNamespacesOp((MatchCondition("suffix", namespace),)),
|
||||
}[kind]
|
||||
store.put(("foo", "bar"), "key", {"original": True})
|
||||
|
||||
with pytest.raises(InvalidNamespaceError):
|
||||
store.batch([PutOp(("valid",), "key", {}), op])
|
||||
|
||||
item = store.get(("foo", "bar"), "key")
|
||||
assert item is not None and item.value == {"original": True}
|
||||
# The whole batch is rejected before any SQL runs.
|
||||
assert store.get(("valid",), "key") is None
|
||||
|
||||
|
||||
def test_batch_allows_empty_search_prefix_and_listing_wildcards(
|
||||
store,
|
||||
) -> None:
|
||||
store.put(("foo", "bar"), "key", {"v": 1})
|
||||
|
||||
found, listed = store.batch(
|
||||
[
|
||||
SearchOp(()),
|
||||
ListNamespacesOp((MatchCondition("prefix", ("foo", "*")),)),
|
||||
]
|
||||
)
|
||||
|
||||
assert [item.namespace for item in found] == [("foo", "bar")]
|
||||
assert listed == [("foo", "bar")]
|
||||
|
||||
Generated
+1
-1
@@ -278,7 +278,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -29,7 +29,11 @@ from langgraph.checkpoint.sqlite._delta import (
|
||||
build_delta_stage2_sql,
|
||||
step_walk_with_row,
|
||||
)
|
||||
from langgraph.checkpoint.sqlite.utils import search_where
|
||||
from langgraph.checkpoint.sqlite.utils import (
|
||||
load_pending_writes,
|
||||
pending_writes_sql,
|
||||
search_where,
|
||||
)
|
||||
|
||||
_AIO_ERROR_MSG = (
|
||||
"The SqliteSaver does not support async methods. "
|
||||
@@ -81,6 +85,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
|
||||
conn: sqlite3.Connection
|
||||
is_setup: bool
|
||||
_has_task_path: bool = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -154,6 +159,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 +168,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
|
||||
|
||||
@@ -260,7 +279,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
}
|
||||
# find any pending writes
|
||||
cur.execute(
|
||||
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
|
||||
pending_writes_sql(self._has_task_path),
|
||||
(
|
||||
str(config["configurable"]["thread_id"]),
|
||||
checkpoint_ns,
|
||||
@@ -286,10 +305,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
if parent_checkpoint_id
|
||||
else None
|
||||
),
|
||||
[
|
||||
(task_id, channel, self.serde.loads_typed((type, value)))
|
||||
for task_id, channel, type, value in cur
|
||||
],
|
||||
load_pending_writes(cur, self.serde),
|
||||
)
|
||||
|
||||
def list(
|
||||
@@ -351,7 +367,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
metadata,
|
||||
) in cur:
|
||||
wcur.execute(
|
||||
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
|
||||
pending_writes_sql(self._has_task_path),
|
||||
(thread_id, checkpoint_ns, checkpoint_id),
|
||||
)
|
||||
yield CheckpointTuple(
|
||||
@@ -378,10 +394,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
if parent_checkpoint_id
|
||||
else None
|
||||
),
|
||||
[
|
||||
(task_id, channel, self.serde.loads_typed((type, value)))
|
||||
for task_id, channel, type, value in wcur
|
||||
],
|
||||
load_pending_writes(wcur, self.serde),
|
||||
)
|
||||
|
||||
def put(
|
||||
@@ -460,9 +473,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 +486,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),
|
||||
@@ -525,7 +539,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
||||
thread_id = str(config["configurable"]["thread_id"])
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
checkpoint_id = get_checkpoint_id(config)
|
||||
if checkpoint_id is None:
|
||||
if not checkpoint_id:
|
||||
target = self.get_tuple(config)
|
||||
if target is None:
|
||||
return {ch: {"writes": []} for ch in channels}
|
||||
@@ -559,6 +573,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:
|
||||
@@ -569,7 +584,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 = []
|
||||
|
||||
@@ -24,7 +24,11 @@ from __future__ import annotations
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any
|
||||
|
||||
from langgraph.checkpoint.base import DeltaChannelHistory, PendingWrite
|
||||
from langgraph.checkpoint.base import (
|
||||
DeltaChannelHistory,
|
||||
PendingWrite,
|
||||
writes_sort_key,
|
||||
)
|
||||
|
||||
# Stage 1 streams target, then its ancestors nearest-first, by following
|
||||
# `parent_checkpoint_id` rather than id order: ids are only monotonic within
|
||||
@@ -56,7 +60,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
|
||||
@@ -70,11 +76,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})"
|
||||
@@ -141,29 +148,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 `writes_sort_key`
|
||||
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: writes_sort_key(w[4], w[2], w[3]))
|
||||
|
||||
result: dict[str, DeltaChannelHistory] = {}
|
||||
for ch in channels:
|
||||
@@ -172,7 +181,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)))
|
||||
)
|
||||
|
||||
@@ -30,7 +30,11 @@ from langgraph.checkpoint.sqlite._delta import (
|
||||
build_delta_stage2_sql,
|
||||
step_walk_with_row,
|
||||
)
|
||||
from langgraph.checkpoint.sqlite.utils import search_where
|
||||
from langgraph.checkpoint.sqlite.utils import (
|
||||
load_pending_writes,
|
||||
pending_writes_sql,
|
||||
search_where,
|
||||
)
|
||||
|
||||
T = TypeVar("T", bound=Callable)
|
||||
|
||||
@@ -114,6 +118,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
|
||||
lock: asyncio.Lock
|
||||
is_setup: bool
|
||||
_has_task_path: bool = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -331,6 +336,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 +347,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:
|
||||
@@ -395,7 +416,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
}
|
||||
# find any pending writes
|
||||
await cur.execute(
|
||||
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
|
||||
pending_writes_sql(self._has_task_path),
|
||||
(
|
||||
str(config["configurable"]["thread_id"]),
|
||||
checkpoint_ns,
|
||||
@@ -421,10 +442,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
if parent_checkpoint_id
|
||||
else None
|
||||
),
|
||||
[
|
||||
(task_id, channel, self.serde.loads_typed((type, value)))
|
||||
async for task_id, channel, type, value in cur
|
||||
],
|
||||
load_pending_writes(await cur.fetchall(), self.serde),
|
||||
)
|
||||
|
||||
async def alist(
|
||||
@@ -473,7 +491,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
metadata,
|
||||
) in cur:
|
||||
await wcur.execute(
|
||||
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
|
||||
pending_writes_sql(self._has_task_path),
|
||||
(thread_id, checkpoint_ns, checkpoint_id),
|
||||
)
|
||||
yield CheckpointTuple(
|
||||
@@ -500,10 +518,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
if parent_checkpoint_id
|
||||
else None
|
||||
),
|
||||
[
|
||||
(task_id, channel, self.serde.loads_typed((type, value)))
|
||||
async for task_id, channel, type, value in wcur
|
||||
],
|
||||
load_pending_writes(await wcur.fetchall(), self.serde),
|
||||
)
|
||||
|
||||
async def aput(
|
||||
@@ -576,9 +591,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 +605,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),
|
||||
@@ -637,7 +653,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
|
||||
thread_id = str(config["configurable"]["thread_id"])
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
checkpoint_id = get_checkpoint_id(config)
|
||||
if checkpoint_id is None:
|
||||
if not checkpoint_id:
|
||||
target = await self.aget_tuple(config)
|
||||
if target is None:
|
||||
return {ch: {"writes": []} for ch in channels}
|
||||
@@ -671,6 +687,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 +698,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:
|
||||
|
||||
@@ -2,11 +2,12 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from collections.abc import Sequence
|
||||
from collections.abc import Iterable, Sequence
|
||||
from typing import Any
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langgraph.checkpoint.base import get_checkpoint_id
|
||||
from langgraph.checkpoint.base import PendingWrite, get_checkpoint_id, writes_sort_key
|
||||
from langgraph.checkpoint.serde.base import SerializerProtocol
|
||||
|
||||
_FILTER_PATTERN = re.compile(r"^[a-zA-Z0-9_.-]+$")
|
||||
|
||||
@@ -114,3 +115,23 @@ def search_where(
|
||||
param_values.append(get_checkpoint_id(before))
|
||||
|
||||
return ("WHERE " + " AND ".join(wheres) if wheres else "", param_values)
|
||||
|
||||
|
||||
def pending_writes_sql(has_task_path: bool) -> str:
|
||||
task_path = "task_path" if has_task_path else "''"
|
||||
return (
|
||||
f"SELECT task_id, channel, type, value, {task_path}, idx FROM writes "
|
||||
"WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ?"
|
||||
)
|
||||
|
||||
|
||||
def load_pending_writes(
|
||||
rows: Iterable[Any], serde: SerializerProtocol
|
||||
) -> list[PendingWrite]:
|
||||
"""Deserialize `pending_writes_sql` rows in `writes_sort_key` order."""
|
||||
return [
|
||||
(task_id, channel, serde.loads_typed((type_, value)))
|
||||
for task_id, channel, type_, value, _, _ in sorted(
|
||||
rows, key=lambda r: writes_sort_key(r[4], r[0], r[5])
|
||||
)
|
||||
]
|
||||
|
||||
@@ -28,6 +28,7 @@ from langgraph.store.base import (
|
||||
ensure_embeddings,
|
||||
get_text_at_path,
|
||||
tokenize_path,
|
||||
validate_op_namespace,
|
||||
)
|
||||
|
||||
_AIO_ERROR_MSG = (
|
||||
@@ -257,6 +258,7 @@ def _group_ops(ops: Iterable[Op]) -> tuple[dict[type, list[tuple[int, Op]]], int
|
||||
grouped_ops: dict[type, list[tuple[int, Op]]] = defaultdict(list)
|
||||
tot = 0
|
||||
for idx, op in enumerate(ops):
|
||||
validate_op_namespace(op)
|
||||
grouped_ops[type(op)].append((idx, op))
|
||||
tot += 1
|
||||
return grouped_ops, tot
|
||||
|
||||
@@ -12,7 +12,7 @@ readme = "README.md"
|
||||
license = "MIT"
|
||||
license-files = ['LICENSE']
|
||||
dependencies = [
|
||||
"langgraph-checkpoint>=4.1.0,<5.0.0",
|
||||
"langgraph-checkpoint>=4.3.0,<5.0.0",
|
||||
"aiosqlite>=0.20",
|
||||
"sqlite-vec>=0.1.6",
|
||||
]
|
||||
|
||||
@@ -9,8 +9,10 @@ from typing import cast
|
||||
import pytest
|
||||
from langgraph.store.base import (
|
||||
GetOp,
|
||||
InvalidNamespaceError,
|
||||
Item,
|
||||
ListNamespacesOp,
|
||||
MatchCondition,
|
||||
PutOp,
|
||||
SearchOp,
|
||||
)
|
||||
@@ -745,3 +747,65 @@ async def test_async_namespace_segment_boundary(store: AsyncSqliteStore) -> None
|
||||
assert set(await store.alist_namespaces(suffix=["alice"], limit=100)) == {
|
||||
("uid", "users", "alice"),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace", [("foo.bar",), ("foo", ""), ("foo", 1)])
|
||||
async def test_abatch_rejects_invalid_namespace_labels(
|
||||
store: AsyncSqliteStore, namespace: tuple
|
||||
) -> None:
|
||||
await store.aput(("foo", "bar"), "key", {"original": True})
|
||||
|
||||
for op in (
|
||||
GetOp(namespace, "key"),
|
||||
GetOp(namespace, "key", refresh_ttl=True),
|
||||
PutOp(namespace, "key", {"changed": True}),
|
||||
PutOp(namespace, "key", None),
|
||||
SearchOp(namespace),
|
||||
ListNamespacesOp((MatchCondition("prefix", namespace),)),
|
||||
ListNamespacesOp((MatchCondition("suffix", namespace),)),
|
||||
):
|
||||
with pytest.raises(InvalidNamespaceError):
|
||||
await store.abatch([op])
|
||||
|
||||
item = await store.aget(("foo", "bar"), "key")
|
||||
assert item is not None and item.value == {"original": True}
|
||||
|
||||
|
||||
async def test_invalid_namespace_only_fails_its_own_call(
|
||||
store: AsyncSqliteStore,
|
||||
) -> None:
|
||||
"""Concurrent calls share one `abatch`, which fails every op if it raises.
|
||||
|
||||
Labels are checked before an op is queued, so one caller's bad label cannot
|
||||
fail another caller's request.
|
||||
"""
|
||||
await store.aput(("foo", "bar"), "key", {"original": True})
|
||||
|
||||
valid, invalid = await asyncio.gather(
|
||||
store.aget(("foo", "bar"), "key"),
|
||||
store.aget(("foo.bar",), "key"),
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
assert isinstance(valid, Item) and valid.value == {"original": True}
|
||||
assert isinstance(invalid, InvalidNamespaceError)
|
||||
|
||||
|
||||
async def test_sync_methods_reject_invalid_namespace_labels(
|
||||
store: AsyncSqliteStore,
|
||||
) -> None:
|
||||
"""The sync wrappers run off the event loop thread and must validate too."""
|
||||
await store.aput(("foo", "bar"), "key", {"original": True})
|
||||
|
||||
for call in (
|
||||
lambda: store.get(("foo.bar",), "key"),
|
||||
lambda: store.search(("foo.bar",)),
|
||||
lambda: store.delete(("foo.bar",), "key"),
|
||||
lambda: store.list_namespaces(prefix=("foo.bar",)),
|
||||
lambda: store.batch([GetOp(("foo.bar",), "key")]),
|
||||
):
|
||||
with pytest.raises(InvalidNamespaceError):
|
||||
await asyncio.to_thread(call)
|
||||
|
||||
item = await store.aget(("foo", "bar"), "key")
|
||||
assert item is not None and item.value == {"original": True}
|
||||
|
||||
@@ -65,6 +65,39 @@ async def test_async_walk_reaches_parent_whatever_the_id_order(
|
||||
assert got[CHANNEL] == EXPECTED
|
||||
|
||||
|
||||
EMPTY_CHECKPOINT_ID: dict[str, Any] = {
|
||||
"configurable": {**CONFIG["configurable"], "checkpoint_id": ""}
|
||||
}
|
||||
|
||||
|
||||
def test_sync_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
|
||||
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||
root = saver.put(CONFIG, _checkpoint("a-older", {CHANNEL: "seed"}), {}, {})
|
||||
saver.put_writes(root, [(CHANNEL, "write-root")], "task")
|
||||
saver.put(root, _checkpoint("z-newer", {}), {}, {})
|
||||
|
||||
got = saver.get_delta_channel_history(
|
||||
config=EMPTY_CHECKPOINT_ID, channels=[CHANNEL]
|
||||
)
|
||||
|
||||
assert got[CHANNEL] == EXPECTED
|
||||
|
||||
|
||||
async def test_async_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
|
||||
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
|
||||
root = await saver.aput(
|
||||
CONFIG, _checkpoint("a-older", {CHANNEL: "seed"}), {}, {}
|
||||
)
|
||||
await saver.aput_writes(root, [(CHANNEL, "write-root")], "task")
|
||||
await saver.aput(root, _checkpoint("z-newer", {}), {}, {})
|
||||
|
||||
got = await saver.aget_delta_channel_history(
|
||||
config=EMPTY_CHECKPOINT_ID, 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:
|
||||
|
||||
@@ -14,6 +14,7 @@ import pytest
|
||||
from langchain_core.embeddings import Embeddings
|
||||
from langgraph.store.base import (
|
||||
GetOp,
|
||||
InvalidNamespaceError,
|
||||
Item,
|
||||
ListNamespacesOp,
|
||||
MatchCondition,
|
||||
@@ -1435,3 +1436,51 @@ def test_list_namespaces_metacharacter_labels(store: SqliteStore) -> None:
|
||||
assert set(store.list_namespaces(prefix=[label, "child"], limit=100)) == {
|
||||
(label, "child"),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace", [("foo.bar",), ("foo", ""), ("foo", 1)])
|
||||
@pytest.mark.parametrize(
|
||||
"kind", ["get", "put", "delete", "search", "list_prefix", "list_suffix"]
|
||||
)
|
||||
def test_batch_rejects_invalid_namespace_labels(
|
||||
store: SqliteStore, namespace: tuple, kind: str
|
||||
) -> None:
|
||||
"""Ops passed straight to `batch` must not reach another namespace.
|
||||
|
||||
Namespaces are stored dot-joined, so `("foo.bar",)` flattens to the same
|
||||
text as `("foo", "bar")`. `BaseStore` methods validate labels themselves,
|
||||
but `batch` takes ops as given.
|
||||
"""
|
||||
op = {
|
||||
"get": GetOp(namespace, "key"),
|
||||
"put": PutOp(namespace, "key", {"changed": True}),
|
||||
"delete": PutOp(namespace, "key", None),
|
||||
"search": SearchOp(namespace),
|
||||
"list_prefix": ListNamespacesOp((MatchCondition("prefix", namespace),)),
|
||||
"list_suffix": ListNamespacesOp((MatchCondition("suffix", namespace),)),
|
||||
}[kind]
|
||||
store.put(("foo", "bar"), "key", {"original": True})
|
||||
|
||||
with pytest.raises(InvalidNamespaceError):
|
||||
store.batch([PutOp(("valid",), "key", {}), op])
|
||||
|
||||
item = store.get(("foo", "bar"), "key")
|
||||
assert item is not None and item.value == {"original": True}
|
||||
# The whole batch is rejected before any SQL runs.
|
||||
assert store.get(("valid",), "key") is None
|
||||
|
||||
|
||||
def test_batch_allows_empty_search_prefix_and_listing_wildcards(
|
||||
store: SqliteStore,
|
||||
) -> None:
|
||||
store.put(("foo", "bar"), "key", {"v": 1})
|
||||
|
||||
found, listed = store.batch(
|
||||
[
|
||||
SearchOp(()),
|
||||
ListNamespacesOp((MatchCondition("prefix", ("foo", "*")),)),
|
||||
]
|
||||
)
|
||||
|
||||
assert [item.namespace for item in found] == [("foo", "bar")]
|
||||
assert listed == [("foo", "bar")]
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
import sqlite3
|
||||
from collections.abc import Iterator
|
||||
from contextlib import closing
|
||||
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"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def busy_db(tmp_path: Path) -> Iterator[Path]:
|
||||
db = tmp_path / "busy.sqlite"
|
||||
with SqliteSaver.from_conn_string(str(db)) as saver:
|
||||
saver.setup()
|
||||
with closing(sqlite3.connect(db, isolation_level=None)) as writer:
|
||||
writer.execute("BEGIN IMMEDIATE")
|
||||
yield db
|
||||
writer.execute("ROLLBACK")
|
||||
|
||||
|
||||
def test_setup_does_not_wait_on_another_writer(busy_db: Path) -> None:
|
||||
with closing(sqlite3.connect(busy_db, timeout=0)) as conn:
|
||||
SqliteSaver(conn).setup()
|
||||
|
||||
|
||||
async def test_async_setup_does_not_wait_on_another_writer(busy_db: Path) -> None:
|
||||
async with aiosqlite.connect(busy_db, timeout=0) as conn:
|
||||
await AsyncSqliteSaver(conn).setup()
|
||||
|
||||
|
||||
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")]}
|
||||
Generated
+1
-1
@@ -285,7 +285,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -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 `writes_sort_key`,
|
||||
`(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` and `Command` updates, 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
|
||||
@@ -174,6 +182,20 @@ class DeltaChannelHistory(TypedDict):
|
||||
seed: NotRequired[Any]
|
||||
|
||||
|
||||
def writes_sort_key(
|
||||
task_path: str, task_id: str = "", idx: int = 0
|
||||
) -> tuple[str, str, int]:
|
||||
"""Sort key for the writes of one super-step.
|
||||
|
||||
Live execution applies a super-step's tasks in this order, so a saver
|
||||
that replays stored writes, as `get_delta_channel_history` does, must
|
||||
sort them by it too, or an order-sensitive reducer rebuilds a different
|
||||
value than the run produced. `task_path` is the string passed to
|
||||
`put_writes`.
|
||||
"""
|
||||
return (task_path, task_id, idx)
|
||||
|
||||
|
||||
class BaseCheckpointSaver(Generic[V]):
|
||||
"""Base class for creating a graph checkpointer.
|
||||
|
||||
@@ -244,7 +266,8 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
config: Configuration specifying which checkpoint to retrieve.
|
||||
|
||||
Returns:
|
||||
The requested checkpoint tuple, or `None` if not found.
|
||||
The requested checkpoint tuple, or `None` if not found. Its
|
||||
`pending_writes` must be in `writes_sort_key` order.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Implement this method in your custom checkpoint saver.
|
||||
@@ -434,7 +457,8 @@ class BaseCheckpointSaver(Generic[V]):
|
||||
config: Configuration specifying which checkpoint to retrieve.
|
||||
|
||||
Returns:
|
||||
The requested checkpoint tuple, or `None` if not found.
|
||||
The requested checkpoint tuple, or `None` if not found. Its
|
||||
`pending_writes` must be in `writes_sort_key` order.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: Implement this method in your custom checkpoint saver.
|
||||
@@ -611,6 +635,10 @@ 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, which
|
||||
`get_tuple` must return in `writes_sort_key` order.
|
||||
|
||||
Args:
|
||||
config: Configuration identifying the target checkpoint.
|
||||
channels: Channel names to walk for. Empty → empty mapping.
|
||||
|
||||
@@ -25,6 +25,7 @@ from langgraph.checkpoint.base import (
|
||||
SerializerProtocol,
|
||||
get_checkpoint_id,
|
||||
get_checkpoint_metadata,
|
||||
writes_sort_key,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -139,6 +140,15 @@ class InMemorySaver(
|
||||
result[k] = self.serde.loads_typed(vv)
|
||||
return result
|
||||
|
||||
def _ordered_writes(
|
||||
self, thread_id: str, checkpoint_ns: str, checkpoint_id: str
|
||||
) -> list[tuple[str, str, tuple[str, bytes], str]]:
|
||||
stored = self.writes.get((thread_id, checkpoint_ns, checkpoint_id), {})
|
||||
return [
|
||||
stored[k]
|
||||
for k in sorted(stored, key=lambda k: writes_sort_key(stored[k][3], *k))
|
||||
]
|
||||
|
||||
def get_delta_channel_history(
|
||||
self, *, config: RunnableConfig, channels: Sequence[str]
|
||||
) -> Mapping[str, DeltaChannelHistory]:
|
||||
@@ -160,8 +170,8 @@ class InMemorySaver(
|
||||
return {}
|
||||
thread_id = config["configurable"]["thread_id"]
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
checkpoint_id = config["configurable"].get("checkpoint_id", "")
|
||||
ns_storage = self.storage.get(thread_id, {}).get(checkpoint_ns, {})
|
||||
checkpoint_id = get_checkpoint_id(config) or max(ns_storage, default="")
|
||||
|
||||
chain: list[str] = []
|
||||
target_entry = ns_storage.get(checkpoint_id)
|
||||
@@ -198,9 +208,8 @@ class InMemorySaver(
|
||||
blob_value_by_ch[ch] = self.serde.loads_typed(blob_entry)
|
||||
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 reversed(
|
||||
self._ordered_writes(thread_id, checkpoint_ns, cp_id)
|
||||
):
|
||||
if ch not in remaining:
|
||||
continue
|
||||
@@ -246,7 +255,7 @@ class InMemorySaver(
|
||||
if checkpoint_id := get_checkpoint_id(config):
|
||||
if saved := self.storage[thread_id][checkpoint_ns].get(checkpoint_id):
|
||||
checkpoint, metadata, parent_checkpoint_id = saved
|
||||
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
|
||||
writes = self._ordered_writes(thread_id, checkpoint_ns, checkpoint_id)
|
||||
checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)
|
||||
return CheckpointTuple(
|
||||
config=config,
|
||||
@@ -276,7 +285,7 @@ class InMemorySaver(
|
||||
if checkpoints := self.storage[thread_id][checkpoint_ns]:
|
||||
checkpoint_id = max(checkpoints.keys())
|
||||
checkpoint, metadata, parent_checkpoint_id = checkpoints[checkpoint_id]
|
||||
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
|
||||
writes = self._ordered_writes(thread_id, checkpoint_ns, checkpoint_id)
|
||||
checkpoint_ = self.serde.loads_typed(checkpoint)
|
||||
return CheckpointTuple(
|
||||
config={
|
||||
@@ -379,9 +388,9 @@ class InMemorySaver(
|
||||
elif limit is not None:
|
||||
limit -= 1
|
||||
|
||||
writes = self.writes[
|
||||
(thread_id, checkpoint_ns, checkpoint_id)
|
||||
].values()
|
||||
writes = self._ordered_writes(
|
||||
thread_id, checkpoint_ns, checkpoint_id
|
||||
)
|
||||
|
||||
checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)
|
||||
|
||||
|
||||
@@ -55,6 +55,7 @@ logger = logging.getLogger(__name__)
|
||||
_MAX_WARNED_TYPES = 1000
|
||||
_warned_unregistered_types: set[tuple[str, str]] = set()
|
||||
_warned_blocked_types: set[tuple[str, str]] = set()
|
||||
_warned_unreconstructable_types: set[tuple[str, str]] = set()
|
||||
|
||||
|
||||
def _is_safe_json_type(id_list: list[str]) -> bool:
|
||||
@@ -79,6 +80,27 @@ def _warn_once(
|
||||
logger.warning(msg, *args)
|
||||
|
||||
|
||||
def _reconstruction_fallback(tup: Any, exc: Exception) -> Any:
|
||||
"""Return the serialized payload of an object that could not be rebuilt.
|
||||
|
||||
Returning `None` here would silently erase the value from restored state.
|
||||
"""
|
||||
try:
|
||||
module, name, payload = tup[0], tup[1], tup[2]
|
||||
except Exception:
|
||||
return None
|
||||
_warn_once(
|
||||
_warned_unreconstructable_types,
|
||||
(str(module), str(name)),
|
||||
"Could not reconstruct %s.%s from checkpoint (%s); "
|
||||
"returning its serialized data instead.",
|
||||
module,
|
||||
name,
|
||||
type(exc).__name__,
|
||||
)
|
||||
return payload
|
||||
|
||||
|
||||
class JsonPlusSerializer(SerializerProtocol):
|
||||
"""Serializer that uses ormsgpack, with optional fallbacks.
|
||||
|
||||
@@ -638,6 +660,7 @@ def _create_msgpack_ext_hook(
|
||||
)
|
||||
)
|
||||
elif code == EXT_CONSTRUCTOR_SINGLE_ARG:
|
||||
tup = None
|
||||
try:
|
||||
tup = ormsgpack.unpackb(
|
||||
data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
|
||||
@@ -649,9 +672,10 @@ def _create_msgpack_ext_hook(
|
||||
return tup[2]
|
||||
# module, name, arg
|
||||
return getattr(importlib.import_module(tup[0]), tup[1])(tup[2])
|
||||
except Exception:
|
||||
return None
|
||||
except Exception as exc:
|
||||
return _reconstruction_fallback(tup, exc)
|
||||
elif code == EXT_CONSTRUCTOR_POS_ARGS:
|
||||
tup = None
|
||||
try:
|
||||
tup = ormsgpack.unpackb(
|
||||
data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
|
||||
@@ -662,9 +686,10 @@ def _create_msgpack_ext_hook(
|
||||
return _send_from_args(tup[2])
|
||||
# module, name, args
|
||||
return getattr(importlib.import_module(tup[0]), tup[1])(*tup[2])
|
||||
except Exception:
|
||||
return None
|
||||
except Exception as exc:
|
||||
return _reconstruction_fallback(tup, exc)
|
||||
elif code == EXT_CONSTRUCTOR_KW_ARGS:
|
||||
tup = None
|
||||
try:
|
||||
tup = ormsgpack.unpackb(
|
||||
data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
|
||||
@@ -673,9 +698,10 @@ def _create_msgpack_ext_hook(
|
||||
return tup[2]
|
||||
# module, name, kwargs
|
||||
return getattr(importlib.import_module(tup[0]), tup[1])(**tup[2])
|
||||
except Exception:
|
||||
return None
|
||||
except Exception as exc:
|
||||
return _reconstruction_fallback(tup, exc)
|
||||
elif code == EXT_METHOD_SINGLE_ARG:
|
||||
tup = None
|
||||
try:
|
||||
tup = ormsgpack.unpackb(
|
||||
data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
|
||||
@@ -686,8 +712,8 @@ def _create_msgpack_ext_hook(
|
||||
return getattr(
|
||||
getattr(importlib.import_module(tup[0]), tup[1]), tup[3]
|
||||
)(tup[2])
|
||||
except Exception:
|
||||
return None
|
||||
except Exception as exc:
|
||||
return _reconstruction_fallback(tup, exc)
|
||||
elif code == EXT_PYDANTIC_V1:
|
||||
try:
|
||||
tup = ormsgpack.unpackb(
|
||||
|
||||
@@ -771,7 +771,12 @@ class BaseStore(ABC):
|
||||
|
||||
Returns:
|
||||
The retrieved item or `None` if not found.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If a namespace label is empty, is not a string,
|
||||
or contains a period (`.`).
|
||||
"""
|
||||
_validate_namespace_labels(namespace)
|
||||
return self.batch(
|
||||
[GetOp(namespace, str(key), _ensure_refresh(self.ttl_config, refresh_ttl))]
|
||||
)[0]
|
||||
@@ -801,6 +806,10 @@ class BaseStore(ABC):
|
||||
Returns:
|
||||
List of items matching the search criteria.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If a `namespace_prefix` label is empty, is not a
|
||||
string, or contains a period (`.`).
|
||||
|
||||
???+ example "Examples"
|
||||
|
||||
Basic filtering:
|
||||
@@ -840,6 +849,7 @@ class BaseStore(ABC):
|
||||
Natural language search support depends on your store implementation
|
||||
and requires proper embedding configuration.
|
||||
"""
|
||||
_validate_namespace_labels(namespace_prefix)
|
||||
return self.batch(
|
||||
[
|
||||
SearchOp(
|
||||
@@ -887,6 +897,11 @@ class BaseStore(ABC):
|
||||
By default, the expiration timer refreshes on both read operations (get/search)
|
||||
and write operations (put/update), whenever the item is included in the operation.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If the namespace is empty, its root label is
|
||||
`"langgraph"`, or a label is empty, is not a string, or contains a
|
||||
period (`.`).
|
||||
|
||||
Note:
|
||||
Indexing support depends on your store implementation.
|
||||
If you do not initialize the store with indexing capabilities,
|
||||
@@ -940,7 +955,12 @@ class BaseStore(ABC):
|
||||
Args:
|
||||
namespace: Hierarchical path for the item.
|
||||
key: Unique identifier within the namespace.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If a namespace label is empty, is not a string,
|
||||
or contains a period (`.`).
|
||||
"""
|
||||
_validate_namespace_labels(namespace)
|
||||
self.batch([PutOp(namespace, str(key), None, ttl=None)])
|
||||
|
||||
def list_namespaces(
|
||||
@@ -969,6 +989,10 @@ class BaseStore(ABC):
|
||||
A list of namespace tuples that match the criteria. Each tuple represents a
|
||||
full namespace path up to `max_depth`.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If a `prefix` or `suffix` label is empty, is not a
|
||||
string, or contains a period (`.`).
|
||||
|
||||
???+ example "Examples":
|
||||
|
||||
Setting `max_depth=3`. Given the namespaces:
|
||||
@@ -984,6 +1008,8 @@ class BaseStore(ABC):
|
||||
# [("a", "b", "c"), ("a", "b", "d"), ("a", "b", "f")]
|
||||
```
|
||||
"""
|
||||
_validate_namespace_labels(prefix or ())
|
||||
_validate_namespace_labels(suffix or ())
|
||||
match_conditions = []
|
||||
if prefix:
|
||||
match_conditions.append(MatchCondition(match_type="prefix", path=prefix))
|
||||
@@ -1013,7 +1039,12 @@ class BaseStore(ABC):
|
||||
|
||||
Returns:
|
||||
The retrieved item or `None` if not found.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If a namespace label is empty, is not a string,
|
||||
or contains a period (`.`).
|
||||
"""
|
||||
_validate_namespace_labels(namespace)
|
||||
return (
|
||||
await self.abatch(
|
||||
[
|
||||
@@ -1052,6 +1083,10 @@ class BaseStore(ABC):
|
||||
Returns:
|
||||
List of items matching the search criteria.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If a `namespace_prefix` label is empty, is not a
|
||||
string, or contains a period (`.`).
|
||||
|
||||
???+ example "Examples"
|
||||
|
||||
Basic filtering:
|
||||
@@ -1091,6 +1126,7 @@ class BaseStore(ABC):
|
||||
Natural language search support depends on your store implementation
|
||||
and requires proper embedding configuration.
|
||||
"""
|
||||
_validate_namespace_labels(namespace_prefix)
|
||||
return (
|
||||
await self.abatch(
|
||||
[
|
||||
@@ -1140,6 +1176,11 @@ class BaseStore(ABC):
|
||||
By default, the expiration timer refreshes on both read operations (get/search)
|
||||
and write operations (put/update), whenever the item is included in the operation.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If the namespace is empty, its root label is
|
||||
`"langgraph"`, or a label is empty, is not a string, or contains a
|
||||
period (`.`).
|
||||
|
||||
Note:
|
||||
Indexing support depends on your store implementation.
|
||||
If you do not initialize the store with indexing capabilities,
|
||||
@@ -1201,7 +1242,12 @@ class BaseStore(ABC):
|
||||
Args:
|
||||
namespace: Hierarchical path for the item.
|
||||
key: Unique identifier within the namespace.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If a namespace label is empty, is not a string,
|
||||
or contains a period (`.`).
|
||||
"""
|
||||
_validate_namespace_labels(namespace)
|
||||
await self.abatch([PutOp(namespace, str(key), None)])
|
||||
|
||||
async def alist_namespaces(
|
||||
@@ -1230,6 +1276,10 @@ class BaseStore(ABC):
|
||||
A list of namespace tuples that match the criteria. Each tuple represents a
|
||||
full namespace path up to `max_depth`.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If a `prefix` or `suffix` label is empty, is not a
|
||||
string, or contains a period (`.`).
|
||||
|
||||
???+ example "Examples"
|
||||
|
||||
Setting `max_depth=3` with existing namespaces:
|
||||
@@ -1245,6 +1295,8 @@ class BaseStore(ABC):
|
||||
# Returns: [("a", "b", "c"), ("a", "b", "d"), ("a", "b", "f")]
|
||||
```
|
||||
"""
|
||||
_validate_namespace_labels(prefix or ())
|
||||
_validate_namespace_labels(suffix or ())
|
||||
match_conditions = []
|
||||
if prefix:
|
||||
match_conditions.append(MatchCondition(match_type="prefix", path=prefix))
|
||||
@@ -1263,6 +1315,14 @@ class BaseStore(ABC):
|
||||
def _validate_namespace(namespace: tuple[str, ...]) -> None:
|
||||
if not namespace:
|
||||
raise InvalidNamespaceError("Namespace cannot be empty.")
|
||||
_validate_namespace_labels(namespace)
|
||||
if namespace[0] == "langgraph":
|
||||
raise InvalidNamespaceError(
|
||||
f'Root label for namespace cannot be "langgraph". Got: {namespace}'
|
||||
)
|
||||
|
||||
|
||||
def _validate_namespace_labels(namespace: tuple[str, ...]) -> None:
|
||||
for label in namespace:
|
||||
if not isinstance(label, str):
|
||||
raise InvalidNamespaceError(
|
||||
@@ -1277,10 +1337,27 @@ def _validate_namespace(namespace: tuple[str, ...]) -> None:
|
||||
raise InvalidNamespaceError(
|
||||
f"Namespace labels cannot be empty strings. Got {label} in {namespace}"
|
||||
)
|
||||
if namespace[0] == "langgraph":
|
||||
raise InvalidNamespaceError(
|
||||
f'Root label for namespace cannot be "langgraph". Got: {namespace}'
|
||||
)
|
||||
|
||||
|
||||
def validate_op_namespace(op: Op) -> None:
|
||||
"""Validate the namespace labels an op carries before a store executes it.
|
||||
|
||||
`BaseStore` methods check labels before batching, but ops passed directly to
|
||||
`batch`/`abatch` skip those methods. Stores that serialize namespaces as
|
||||
delimited text should call this for every op they execute, so a label such
|
||||
as `"foo.bar"` cannot address the namespace `("foo", "bar")`.
|
||||
|
||||
Raises:
|
||||
InvalidNamespaceError: If a label is empty, is not a string, or contains
|
||||
a period (`.`).
|
||||
"""
|
||||
if isinstance(op, (GetOp, PutOp)):
|
||||
_validate_namespace_labels(op.namespace)
|
||||
elif isinstance(op, SearchOp):
|
||||
_validate_namespace_labels(op.namespace_prefix)
|
||||
elif isinstance(op, ListNamespacesOp):
|
||||
for condition in op.match_conditions or ():
|
||||
_validate_namespace_labels(condition.path)
|
||||
|
||||
|
||||
def _ensure_refresh(
|
||||
@@ -1319,4 +1396,5 @@ __all__ = [
|
||||
"ensure_embeddings",
|
||||
"tokenize_path",
|
||||
"get_text_at_path",
|
||||
"validate_op_namespace",
|
||||
]
|
||||
|
||||
@@ -25,6 +25,7 @@ from langgraph.store.base import (
|
||||
_ensure_refresh,
|
||||
_ensure_ttl,
|
||||
_validate_namespace,
|
||||
_validate_namespace_labels,
|
||||
)
|
||||
|
||||
F = TypeVar("F", bound=Callable)
|
||||
@@ -86,6 +87,7 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
*,
|
||||
refresh_ttl: bool | None = None,
|
||||
) -> Item | None:
|
||||
_validate_namespace_labels(namespace)
|
||||
self._ensure_task()
|
||||
fut = self._loop.create_future()
|
||||
self._aqueue.put_nowait(
|
||||
@@ -111,6 +113,7 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
offset: int = 0,
|
||||
refresh_ttl: bool | None = None,
|
||||
) -> list[SearchItem]:
|
||||
_validate_namespace_labels(namespace_prefix)
|
||||
self._ensure_task()
|
||||
fut = self._loop.create_future()
|
||||
self._aqueue.put_nowait(
|
||||
@@ -155,6 +158,7 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
namespace: tuple[str, ...],
|
||||
key: str,
|
||||
) -> None:
|
||||
_validate_namespace_labels(namespace)
|
||||
self._ensure_task()
|
||||
fut = self._loop.create_future()
|
||||
self._aqueue.put_nowait((fut, PutOp(namespace, key, None)))
|
||||
@@ -169,6 +173,8 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
) -> list[tuple[str, ...]]:
|
||||
_validate_namespace_labels(prefix or ())
|
||||
_validate_namespace_labels(suffix or ())
|
||||
self._ensure_task()
|
||||
fut = self._loop.create_future()
|
||||
match_conditions = []
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
description = "Library with base interfaces for LangGraph checkpoint savers."
|
||||
authors = []
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -1,37 +0,0 @@
|
||||
"""Run delta-channel conformance capabilities against InMemorySaver."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
conformance = pytest.importorskip(
|
||||
"langgraph.checkpoint.conformance",
|
||||
reason="langgraph-checkpoint-conformance not installed",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delta_channel_conformance():
|
||||
# Imported inside the test: the module-level importorskip above is what
|
||||
# makes these safe, so they cannot move to the top of the file.
|
||||
from langgraph.checkpoint.conformance import validate # noqa: PLC0415
|
||||
from langgraph.checkpoint.conformance.initializer import ( # noqa: PLC0415
|
||||
checkpointer_test,
|
||||
)
|
||||
|
||||
from langgraph.checkpoint.memory import InMemorySaver # noqa: PLC0415
|
||||
|
||||
@checkpointer_test(name="InMemorySaver")
|
||||
async def mem_saver():
|
||||
yield InMemorySaver()
|
||||
|
||||
report = await validate(
|
||||
mem_saver,
|
||||
capabilities={
|
||||
"delta_channel_history",
|
||||
},
|
||||
)
|
||||
for cap, result in report.results.items():
|
||||
if result.passed is False:
|
||||
details = "\n".join(result.failures or [])
|
||||
pytest.fail(f"Capability {cap} failed:\n{details}")
|
||||
@@ -7,6 +7,7 @@ import pickle
|
||||
import re
|
||||
import sys
|
||||
import tempfile
|
||||
import types
|
||||
import uuid
|
||||
from collections import deque
|
||||
from datetime import date, datetime, time, timezone
|
||||
@@ -33,12 +34,16 @@ from langgraph.checkpoint.serde.event_hooks import (
|
||||
register_serde_event_listener,
|
||||
)
|
||||
from langgraph.checkpoint.serde.jsonplus import (
|
||||
EXT_CONSTRUCTOR_KW_ARGS,
|
||||
EXT_CONSTRUCTOR_POS_ARGS,
|
||||
EXT_CONSTRUCTOR_SINGLE_ARG,
|
||||
EXT_METHOD_SINGLE_ARG,
|
||||
InvalidModuleError,
|
||||
JsonPlusSerializer,
|
||||
_msgpack_enc,
|
||||
_msgpack_ext_hook_to_json,
|
||||
_warned_blocked_types,
|
||||
_warned_unreconstructable_types,
|
||||
_warned_unregistered_types,
|
||||
)
|
||||
from langgraph.store.base import Item
|
||||
@@ -821,6 +826,7 @@ def _reset_warned_types() -> None:
|
||||
# a fresh slate and assertions about warning emission are stable.
|
||||
_warned_unregistered_types.clear()
|
||||
_warned_blocked_types.clear()
|
||||
_warned_unreconstructable_types.clear()
|
||||
|
||||
|
||||
def test_msgpack_pydantic_warns_by_default(caplog: pytest.LogCaptureFixture) -> None:
|
||||
@@ -1230,3 +1236,58 @@ def test_msgpack_nested_pydantic_serializes_as_dict(
|
||||
# No blocking should occur - inner is serialized as dict, not ext
|
||||
assert "blocked" not in caplog.text.lower()
|
||||
assert result == obj
|
||||
|
||||
|
||||
def test_msgpack_dataclass_from_removed_module_restores_payload(
|
||||
monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture
|
||||
) -> None:
|
||||
@dataclasses.dataclass
|
||||
class SavedObject:
|
||||
value: int
|
||||
|
||||
SavedObject.__module__ = "removed_module"
|
||||
module = types.ModuleType("removed_module")
|
||||
module.SavedObject = SavedObject
|
||||
monkeypatch.setitem(sys.modules, "removed_module", module)
|
||||
serde = JsonPlusSerializer(
|
||||
allowed_msgpack_modules=[("removed_module", "SavedObject")]
|
||||
)
|
||||
dumped = serde.dumps_typed({"state": SavedObject(123)})
|
||||
assert serde.loads_typed(dumped) == {"state": SavedObject(123)}
|
||||
|
||||
monkeypatch.delitem(sys.modules, "removed_module")
|
||||
caplog.set_level(logging.WARNING, logger="langgraph.checkpoint.serde.jsonplus")
|
||||
|
||||
assert serde.loads_typed(dumped) == {"state": {"value": 123}}
|
||||
assert (
|
||||
"could not reconstruct removed_module.savedobject from checkpoint "
|
||||
"(modulenotfounderror)" in caplog.text.lower()
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("code", "tup", "expected"),
|
||||
[
|
||||
(EXT_CONSTRUCTOR_SINGLE_ARG, ("missing_module", "Thing", "x"), "x"),
|
||||
(EXT_CONSTRUCTOR_POS_ARGS, ("missing_module", "Thing", [1, 2]), [1, 2]),
|
||||
(
|
||||
EXT_CONSTRUCTOR_KW_ARGS,
|
||||
("missing_module", "Thing", {"value": 123}),
|
||||
{"value": 123},
|
||||
),
|
||||
(
|
||||
EXT_METHOD_SINGLE_ARG,
|
||||
("datetime", "datetime", "not-a-date", "fromisoformat"),
|
||||
"not-a-date",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_msgpack_failed_reconstruction_returns_payload(
|
||||
code: int, tup: tuple, expected: object
|
||||
) -> None:
|
||||
serde = JsonPlusSerializer(allowed_msgpack_modules=True)
|
||||
payload = ormsgpack.packb(
|
||||
ormsgpack.Ext(code, _msgpack_enc(tup)), option=ormsgpack.OPT_NON_STR_KEYS
|
||||
)
|
||||
|
||||
assert serde.loads_typed(("msgpack", payload)) == expected
|
||||
|
||||
@@ -420,6 +420,19 @@ class TestInMemorySaverDeltaChannel:
|
||||
assert "seed" not in result
|
||||
assert result["writes"] == []
|
||||
|
||||
def test_get_channel_writes_without_checkpoint_id_reads_the_latest(self) -> None:
|
||||
saver = InMemorySaver()
|
||||
thread: RunnableConfig = {
|
||||
"configurable": {"thread_id": "t1", "checkpoint_ns": ""}
|
||||
}
|
||||
parent = saver.put(thread, empty_checkpoint(), {}, {})
|
||||
saver.put_writes(parent, [("messages", "hi")], "task1")
|
||||
saver.put(parent, empty_checkpoint(), {}, {})
|
||||
|
||||
result = saver.get_delta_channel_history(config=thread, channels=["messages"])
|
||||
|
||||
assert result == {"messages": {"writes": [("task1", "messages", "hi")]}}
|
||||
|
||||
|
||||
class TestBaseFallbackGetChannelWrites:
|
||||
"""Exercises the `BaseCheckpointSaver.get_delta_channel_history` default
|
||||
|
||||
@@ -13,10 +13,14 @@ from langgraph.store.base import (
|
||||
GetOp,
|
||||
InvalidNamespaceError,
|
||||
Item,
|
||||
ListNamespacesOp,
|
||||
MatchCondition,
|
||||
Op,
|
||||
PutOp,
|
||||
Result,
|
||||
SearchOp,
|
||||
get_text_at_path,
|
||||
validate_op_namespace,
|
||||
)
|
||||
from langgraph.store.base.batch import AsyncBatchedBaseStore
|
||||
from langgraph.store.memory import InMemoryStore
|
||||
@@ -528,6 +532,127 @@ async def test_cannot_put_empty_namespace() -> None:
|
||||
assert (await async_store.aget(("valid", "namespace"), "key")) is None
|
||||
|
||||
|
||||
INVALID_NAMESPACES = [("foo.bar",), ("foo", ""), (123,)]
|
||||
NAMESPACE_METHODS = ["get", "delete", "search", "prefix", "suffix"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace", INVALID_NAMESPACES)
|
||||
@pytest.mark.parametrize("method", NAMESPACE_METHODS)
|
||||
def test_rejects_invalid_namespace_labels(
|
||||
mocker: MockerFixture, namespace: tuple, method: str
|
||||
) -> None:
|
||||
store = InMemoryStore()
|
||||
batch = mocker.spy(InMemoryStore, "batch")
|
||||
call = {
|
||||
"get": lambda: store.get(namespace, "key"),
|
||||
"delete": lambda: store.delete(namespace, "key"),
|
||||
"search": lambda: store.search(namespace),
|
||||
"prefix": lambda: store.list_namespaces(prefix=namespace),
|
||||
"suffix": lambda: store.list_namespaces(suffix=namespace),
|
||||
}[method]
|
||||
|
||||
with pytest.raises(InvalidNamespaceError):
|
||||
call()
|
||||
|
||||
batch.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batched", [False, True])
|
||||
@pytest.mark.parametrize("namespace", INVALID_NAMESPACES)
|
||||
@pytest.mark.parametrize("method", NAMESPACE_METHODS)
|
||||
async def test_async_rejects_invalid_namespace_labels(
|
||||
mocker: MockerFixture, batched: bool, namespace: tuple, method: str
|
||||
) -> None:
|
||||
# The batched store must reject before queueing: a failure inside the
|
||||
# shared `abatch` would fail every op queued alongside this one.
|
||||
store = MockAsyncBatchedStore() if batched else InMemoryStore()
|
||||
# `MockAsyncBatchedStore` dispatches through `InMemoryStore.batch`.
|
||||
batch = mocker.spy(InMemoryStore, "batch")
|
||||
abatch = mocker.spy(InMemoryStore, "abatch")
|
||||
call = {
|
||||
"get": lambda: store.aget(namespace, "key"),
|
||||
"delete": lambda: store.adelete(namespace, "key"),
|
||||
"search": lambda: store.asearch(namespace),
|
||||
"prefix": lambda: store.alist_namespaces(prefix=namespace),
|
||||
"suffix": lambda: store.alist_namespaces(suffix=namespace),
|
||||
}[method]
|
||||
|
||||
with pytest.raises(InvalidNamespaceError):
|
||||
await call()
|
||||
|
||||
batch.assert_not_called()
|
||||
abatch.assert_not_called()
|
||||
|
||||
|
||||
def test_search_and_listing_keep_empty_prefixes_and_wildcards() -> None:
|
||||
store = InMemoryStore()
|
||||
store.put(("tenant", "a_%"), "key", {"v": 1})
|
||||
store.put(("tenant", "b", "child"), "key", {"v": 1})
|
||||
|
||||
assert len(store.search(())) == 2
|
||||
assert [item.namespace for item in store.search(("tenant", "a_%"))] == [
|
||||
("tenant", "a_%")
|
||||
]
|
||||
assert sorted(store.list_namespaces(prefix=("tenant", "*"), suffix=("*",))) == [
|
||||
("tenant", "a_%"),
|
||||
("tenant", "b", "child"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batched", [False, True])
|
||||
async def test_async_search_and_listing_keep_empty_prefixes_and_wildcards(
|
||||
batched: bool,
|
||||
) -> None:
|
||||
store = MockAsyncBatchedStore() if batched else InMemoryStore()
|
||||
await store.aput(("tenant", "a_%"), "key", {"v": 1})
|
||||
await store.aput(("tenant", "b", "child"), "key", {"v": 1})
|
||||
|
||||
assert len(await store.asearch(())) == 2
|
||||
assert [item.namespace for item in await store.asearch(("tenant", "a_%"))] == [
|
||||
("tenant", "a_%")
|
||||
]
|
||||
assert sorted(
|
||||
await store.alist_namespaces(prefix=("tenant", "*"), suffix=("*",))
|
||||
) == [("tenant", "a_%"), ("tenant", "b", "child")]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("namespace", INVALID_NAMESPACES)
|
||||
@pytest.mark.parametrize(
|
||||
"kind", ["get", "put", "delete", "search", "list_prefix", "list_suffix"]
|
||||
)
|
||||
def test_validate_op_namespace_rejects_invalid_labels(
|
||||
namespace: tuple, kind: str
|
||||
) -> None:
|
||||
op = {
|
||||
"get": GetOp(namespace, "key"),
|
||||
"put": PutOp(namespace, "key", {"v": 1}),
|
||||
"delete": PutOp(namespace, "key", None),
|
||||
"search": SearchOp(namespace),
|
||||
"list_prefix": ListNamespacesOp((MatchCondition("prefix", namespace),)),
|
||||
"list_suffix": ListNamespacesOp((MatchCondition("suffix", namespace),)),
|
||||
}[kind]
|
||||
|
||||
with pytest.raises(InvalidNamespaceError):
|
||||
validate_op_namespace(op)
|
||||
|
||||
|
||||
def test_validate_op_namespace_allows_empty_prefix_and_wildcards() -> None:
|
||||
for op in (
|
||||
SearchOp(()),
|
||||
ListNamespacesOp(),
|
||||
ListNamespacesOp(
|
||||
(
|
||||
MatchCondition("prefix", ("tenant", "*")),
|
||||
MatchCondition("suffix", ("*",)),
|
||||
)
|
||||
),
|
||||
GetOp(("tenant", "a_%"), "key"),
|
||||
# Write-only rules belong to `put`, not to op validation.
|
||||
PutOp(("langgraph", "x"), "key", {"v": 1}),
|
||||
):
|
||||
validate_op_namespace(op)
|
||||
|
||||
|
||||
async def test_async_batch_store_deduplication(mocker: MockerFixture) -> None:
|
||||
abatch = mocker.spy(InMemoryStore, "batch")
|
||||
store = MockAsyncBatchedStore()
|
||||
|
||||
Generated
+1
-1
@@ -301,7 +301,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -1 +1 @@
|
||||
__version__ = "0.4.32"
|
||||
__version__ = "0.4.33"
|
||||
|
||||
@@ -32,7 +32,7 @@ from langgraph_cli.host_backend import (
|
||||
HostBackendError,
|
||||
SourceName,
|
||||
)
|
||||
from langgraph_cli.image_reference import ImageReference
|
||||
from langgraph_cli.image_reference import DIGEST_SEPARATOR, ImageReference
|
||||
from langgraph_cli.progress import Progress
|
||||
from langgraph_cli.util import warn_non_wolfi_distro
|
||||
|
||||
@@ -624,6 +624,41 @@ def format_deployments_table(deployments: Sequence[dict[str, object]]) -> str:
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _extract_listener_namespaces(listener: dict[str, object]) -> str:
|
||||
compute_config = listener.get("compute_config")
|
||||
namespaces = (
|
||||
compute_config.get("k8s_namespaces")
|
||||
if isinstance(compute_config, dict)
|
||||
else None
|
||||
)
|
||||
if isinstance(namespaces, list) and namespaces:
|
||||
return ", ".join(str(namespace) for namespace in namespaces)
|
||||
return "-"
|
||||
|
||||
|
||||
def format_listeners_table(listeners: Sequence[dict[str, object]]) -> str:
|
||||
headers = ("Listener ID", "Compute ID", "Namespaces")
|
||||
rows = [
|
||||
(
|
||||
str(listener.get("id", "-") or "-"),
|
||||
str(listener.get("compute_id", "-") or "-"),
|
||||
_extract_listener_namespaces(listener),
|
||||
)
|
||||
for listener in listeners
|
||||
]
|
||||
widths = [
|
||||
max(len(headers[index]), *(len(row[index]) for row in rows))
|
||||
for index in range(len(headers))
|
||||
]
|
||||
|
||||
def format_row(row: Sequence[str]) -> str:
|
||||
return " ".join(value.ljust(widths[index]) for index, value in enumerate(row))
|
||||
|
||||
lines = [format_row(headers), format_row(tuple("-" * width for width in widths))]
|
||||
lines.extend(format_row(row) for row in rows)
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def format_revisions_table(revisions: Sequence[dict[str, object]]) -> str:
|
||||
headers = ("Revision ID", "Status", "Created At")
|
||||
latest_deployed_seen = False
|
||||
@@ -1544,9 +1579,9 @@ def _ensure_customer_registry_source(existing: ExistingDeployment) -> None:
|
||||
if existing.source != _CUSTOMER_REGISTRY_SOURCE:
|
||||
raise click.UsageError(
|
||||
f"Deployment {existing.id} was not created from an external image "
|
||||
"and cannot be updated with --push-to. Run without --push-to to keep "
|
||||
"its current build mode, or use a different --name to create a new "
|
||||
"deployment."
|
||||
"and cannot be updated with --push-to or --image-uri. Run without "
|
||||
"either flag to keep its current build mode, or use a different "
|
||||
"--name to create a new deployment."
|
||||
)
|
||||
|
||||
|
||||
@@ -1599,9 +1634,22 @@ class RemoteBuildSource:
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CustomerRegistrySource:
|
||||
class BuildAndPush:
|
||||
reference: ImageReference
|
||||
prebuilt_image: str | None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PublishedImage:
|
||||
image_uri: str
|
||||
|
||||
|
||||
ImageSource = BuildAndPush | PublishedImage
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CustomerRegistrySource:
|
||||
image: ImageSource
|
||||
requested_placement: RequestedPlacement
|
||||
|
||||
def run(self, ctx: DeployContext) -> DeployOutcome:
|
||||
@@ -1686,16 +1734,22 @@ class CustomerRegistrySource:
|
||||
)
|
||||
|
||||
def _publish(self, ctx: DeployContext, step: int) -> tuple[str, int]:
|
||||
image = str(self.reference)
|
||||
if isinstance(self.image, PublishedImage):
|
||||
return self.image.image_uri, step
|
||||
image = str(self.image.reference)
|
||||
with Runner() as runner:
|
||||
if self.prebuilt_image:
|
||||
_log_deploy_step(step, f"Validating image {self.prebuilt_image}")
|
||||
if self.image.prebuilt_image:
|
||||
_log_deploy_step(step, f"Validating image {self.image.prebuilt_image}")
|
||||
_validate_prebuilt_image(
|
||||
runner, self.prebuilt_image, verbose=ctx.verbose
|
||||
runner, self.image.prebuilt_image, verbose=ctx.verbose
|
||||
)
|
||||
runner.run(
|
||||
subp_exec(
|
||||
"docker", "tag", self.prebuilt_image, image, verbose=ctx.verbose
|
||||
"docker",
|
||||
"tag",
|
||||
self.image.prebuilt_image,
|
||||
image,
|
||||
verbose=ctx.verbose,
|
||||
)
|
||||
)
|
||||
else:
|
||||
@@ -1733,20 +1787,35 @@ def _push_reference(push_to: str, tag: str | None) -> ImageReference:
|
||||
return reference.with_tag(normalize_image_tag(tag or _DEFAULT_IMAGE_TAG))
|
||||
|
||||
|
||||
def _validate_image_uri(image_uri: str) -> str:
|
||||
value = image_uri.strip()
|
||||
if not value:
|
||||
raise click.UsageError("--image-uri must not be empty.")
|
||||
if DIGEST_SEPARATOR not in value:
|
||||
raise click.UsageError(
|
||||
"--image-uri must pin a digest, e.g. "
|
||||
f"repository{DIGEST_SEPARATOR}<sha256 hex>. Kubernetes can cache "
|
||||
"images by tag, so redeploying a mutable tag may silently keep "
|
||||
"running the previous image."
|
||||
)
|
||||
return value
|
||||
|
||||
|
||||
def _select_source(
|
||||
*,
|
||||
push_to: str | None,
|
||||
image: str | None,
|
||||
image_uri: str | None,
|
||||
image_name: str | None,
|
||||
tag: str | None,
|
||||
remote_build_flag: bool | None,
|
||||
placement: RequestedPlacement,
|
||||
selector: DeploymentSelector,
|
||||
) -> DeploymentSource:
|
||||
if push_to is None and placement.requested:
|
||||
if push_to is None and image_uri is None and placement.requested:
|
||||
raise click.UsageError(
|
||||
"--listener-id and --k8s-namespace only apply when creating a "
|
||||
"deployment with --push-to."
|
||||
"deployment with --push-to or --image-uri."
|
||||
)
|
||||
if placement.requested and isinstance(selector, ById):
|
||||
raise click.UsageError(
|
||||
@@ -1754,6 +1823,19 @@ def _select_source(
|
||||
"they cannot be set for an existing --deployment-id. Drop them, or "
|
||||
"create a new deployment with --name."
|
||||
)
|
||||
if image_uri is not None:
|
||||
if push_to is not None:
|
||||
raise click.UsageError("--image-uri cannot be combined with --push-to.")
|
||||
if image is not None:
|
||||
raise click.UsageError("--image-uri cannot be combined with --image.")
|
||||
if tag is not None:
|
||||
raise click.UsageError("--image-uri cannot be combined with --tag.")
|
||||
if remote_build_flag is not None:
|
||||
raise click.UsageError("--image-uri cannot be combined with --remote.")
|
||||
return CustomerRegistrySource(
|
||||
image=PublishedImage(_validate_image_uri(image_uri)),
|
||||
requested_placement=placement,
|
||||
)
|
||||
if push_to is not None:
|
||||
if remote_build_flag is True:
|
||||
raise click.UsageError("--push-to cannot be combined with --remote.")
|
||||
@@ -1761,8 +1843,7 @@ def _select_source(
|
||||
if image is None:
|
||||
_require_local_docker()
|
||||
return CustomerRegistrySource(
|
||||
reference=reference,
|
||||
prebuilt_image=image,
|
||||
image=BuildAndPush(reference=reference, prebuilt_image=image),
|
||||
requested_placement=placement,
|
||||
)
|
||||
if image and remote_build_flag is True:
|
||||
@@ -2075,19 +2156,30 @@ def _deploy_base_options(
|
||||
"Give the tag here or with --tag (default: latest)."
|
||||
),
|
||||
),
|
||||
click.option(
|
||||
"--image-uri",
|
||||
help=(
|
||||
"Deploy an image that's already in a registry you manage, "
|
||||
"without building, retagging, or pushing anything. For "
|
||||
"self-hosted and hybrid LangSmith. Give the full reference, "
|
||||
"e.g. 123456789.dkr.ecr.us-east-1.amazonaws.com/agents/"
|
||||
"my-agent:v1.2.3 or ...@sha256:<digest>. Cannot be combined "
|
||||
"with --push-to, --image, --tag, or --remote."
|
||||
),
|
||||
),
|
||||
click.option(
|
||||
"--listener-id",
|
||||
help=(
|
||||
"Listener that will run the deployment, for workspaces that "
|
||||
"deploy through a listener in your own cluster. Only used when "
|
||||
"creating a deployment with --push-to."
|
||||
"creating a deployment with --push-to or --image-uri."
|
||||
),
|
||||
),
|
||||
click.option(
|
||||
"--k8s-namespace",
|
||||
help=(
|
||||
"Kubernetes namespace the listener deploys into. Only used when "
|
||||
"creating a deployment with --push-to."
|
||||
"creating a deployment with --push-to or --image-uri."
|
||||
),
|
||||
),
|
||||
click.option(
|
||||
@@ -2200,6 +2292,7 @@ def _deploy_cmd(
|
||||
image_name: str | None,
|
||||
image: str | None,
|
||||
push_to: str | None,
|
||||
image_uri: str | None,
|
||||
listener_id: str | None,
|
||||
k8s_namespace: str | None,
|
||||
tag: str | None,
|
||||
@@ -2273,6 +2366,7 @@ def _deploy_cmd(
|
||||
source = _select_source(
|
||||
push_to=push_to,
|
||||
image=image,
|
||||
image_uri=image_uri,
|
||||
image_name=image_name,
|
||||
tag=tag,
|
||||
remote_build_flag=remote_build_flag,
|
||||
@@ -2408,6 +2502,41 @@ def deploy_list(
|
||||
click.echo(format_deployments_table(deployments))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# deploy listeners
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@deploy.group(
|
||||
"listeners",
|
||||
cls=NestedHelpGroup,
|
||||
help="[Beta] Inspect listeners available to this workspace.",
|
||||
)
|
||||
def deploy_listeners() -> None:
|
||||
pass
|
||||
|
||||
|
||||
@OPT_HOST_API_KEY
|
||||
@OPT_HOST_URL
|
||||
@deploy_listeners.command(
|
||||
"list",
|
||||
help=(
|
||||
"[Beta] List listeners available to this workspace.\n\n"
|
||||
"Pass a listener's id to `langgraph deploy --push-to ... "
|
||||
"--listener-id <id>` to deploy through it."
|
||||
),
|
||||
)
|
||||
def deploy_listeners_list(api_key: str | None, host_url: str | None) -> None:
|
||||
client = _create_host_backend_client(host_url, api_key)
|
||||
listeners = _call_host_backend_with_optional_tenant(
|
||||
client, lambda c: c.list_listeners()
|
||||
)
|
||||
if not listeners:
|
||||
click.echo("No listeners found for this workspace.")
|
||||
return
|
||||
click.echo(format_listeners_table(listeners))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# deploy revisions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -453,6 +453,88 @@ def test_deploy_list_command_no_results(monkeypatch) -> None:
|
||||
assert result.output.strip() == "No deployments found."
|
||||
|
||||
|
||||
def test_deploy_listeners_list_command(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
|
||||
captured["host_url"] = host_url
|
||||
captured["api_key"] = api_key
|
||||
captured["tenant_id"] = tenant_id or ""
|
||||
|
||||
def list_listeners(self):
|
||||
return [
|
||||
{
|
||||
"id": "listener-1",
|
||||
"compute_id": "prod-cluster",
|
||||
"compute_config": {"k8s_namespaces": ["agents"]},
|
||||
},
|
||||
{
|
||||
"id": "listener-2",
|
||||
"compute_id": "multi-cluster",
|
||||
"compute_config": {"k8s_namespaces": ["agents", "agents-staging"]},
|
||||
},
|
||||
]
|
||||
|
||||
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
|
||||
|
||||
result = runner.invoke(
|
||||
cli,
|
||||
[
|
||||
"deploy",
|
||||
"listeners",
|
||||
"list",
|
||||
"--api-key",
|
||||
"test-key",
|
||||
"--host-url",
|
||||
"https://api.example.com",
|
||||
],
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert captured == {
|
||||
"host_url": "https://api.example.com",
|
||||
"api_key": "test-key",
|
||||
"tenant_id": "",
|
||||
}
|
||||
assert "Listener ID" in result.output
|
||||
assert "Compute ID" in result.output
|
||||
assert "Namespaces" in result.output
|
||||
assert "listener-1" in result.output
|
||||
assert "prod-cluster" in result.output
|
||||
assert "agents, agents-staging" in result.output
|
||||
|
||||
|
||||
def test_deploy_listeners_list_command_no_results(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
|
||||
pass
|
||||
|
||||
def list_listeners(self):
|
||||
return []
|
||||
|
||||
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
|
||||
|
||||
result = runner.invoke(
|
||||
cli,
|
||||
[
|
||||
"deploy",
|
||||
"listeners",
|
||||
"list",
|
||||
"--api-key",
|
||||
"test-key",
|
||||
"--host-url",
|
||||
"https://api.example.com",
|
||||
],
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert result.output.strip() == "No listeners found for this workspace."
|
||||
|
||||
|
||||
def test_deploy_revisions_list_command(monkeypatch) -> None:
|
||||
runner = CliRunner()
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
@@ -695,6 +695,106 @@ def test_push_to_with_deployment_id_fetches_the_deployment_once(
|
||||
]
|
||||
|
||||
|
||||
def test_image_uri_creates_an_external_deployment_without_any_docker_work(
|
||||
deploy_project: DeployProject,
|
||||
) -> None:
|
||||
result = deploy_project.run("--image-uri", EXTERNAL_DIGEST)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert deploy_project.docker.verbs() == []
|
||||
assert deploy_project.timeline == [LIST_DEPLOYMENTS, CREATE_DEPLOYMENT]
|
||||
assert deploy_project.control_plane.bodies[CREATE_DEPLOYMENT][
|
||||
"source_revision_config"
|
||||
] == {"image_uri": EXTERNAL_DIGEST}
|
||||
|
||||
|
||||
def test_image_uri_updates_an_existing_external_deployment_without_any_docker_work(
|
||||
deploy_project: DeployProject,
|
||||
) -> None:
|
||||
deploy_project.control_plane.existing_deployments = [
|
||||
{"id": "dep-ext", "name": "my-app", "source": "external_docker"}
|
||||
]
|
||||
|
||||
result = deploy_project.run("--image-uri", EXTERNAL_DIGEST)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert deploy_project.docker.verbs() == []
|
||||
assert deploy_project.timeline == [LIST_DEPLOYMENTS, _patch("dep-ext")]
|
||||
assert deploy_project.control_plane.bodies[_patch("dep-ext")] == {
|
||||
"source_revision_config": {"image_uri": EXTERNAL_DIGEST},
|
||||
"secrets": [],
|
||||
"tracked_packages": TRACKED_PACKAGES,
|
||||
}
|
||||
|
||||
|
||||
def test_image_uri_places_a_new_deployment_on_the_only_listener(
|
||||
deploy_project: DeployProject,
|
||||
) -> None:
|
||||
deploy_project.control_plane.listeners = [LISTENER]
|
||||
|
||||
result = deploy_project.run(
|
||||
"--image-uri", EXTERNAL_DIGEST, host_url=CLOUD_CONTROL_PLANE_URL
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert deploy_project.docker.verbs() == []
|
||||
assert deploy_project.control_plane.bodies[CREATE_DEPLOYMENT]["source_config"] == {
|
||||
"resource_spec": {},
|
||||
"listener_id": LISTENER_ID,
|
||||
"listener_config": {"k8s_namespace": "agents"},
|
||||
}
|
||||
|
||||
|
||||
def test_image_uri_rejects_a_non_external_deployment_before_any_docker_work(
|
||||
deploy_project: DeployProject,
|
||||
) -> None:
|
||||
deploy_project.control_plane.existing_deployments = [
|
||||
{"id": "dep-cli", "name": "my-app", "source": "internal_docker"}
|
||||
]
|
||||
|
||||
result = deploy_project.run("--image-uri", EXTERNAL_DIGEST)
|
||||
|
||||
assert result.exit_code != 0
|
||||
assert "cannot be updated with --push-to or --image-uri" in result.output
|
||||
assert deploy_project.docker.verbs() == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("args", "message"),
|
||||
[
|
||||
pytest.param(
|
||||
("--image-uri", EXTERNAL_DIGEST, "--push-to", PUSH_REPOSITORY),
|
||||
"--image-uri cannot be combined with --push-to.",
|
||||
id="with_push_to",
|
||||
),
|
||||
pytest.param(
|
||||
("--image-uri", EXTERNAL_DIGEST, "--image", "local/app:dev"),
|
||||
"--image-uri cannot be combined with --image.",
|
||||
id="with_image",
|
||||
),
|
||||
pytest.param(
|
||||
("--image-uri", EXTERNAL_DIGEST, "--tag", "v2"),
|
||||
"--image-uri cannot be combined with --tag.",
|
||||
id="with_tag",
|
||||
),
|
||||
pytest.param(
|
||||
("--image-uri", EXTERNAL_DIGEST, "--remote"),
|
||||
"--image-uri cannot be combined with --remote.",
|
||||
id="with_remote",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_image_uri_conflicting_flags_are_rejected_before_any_docker_work(
|
||||
deploy_project: DeployProject, args: tuple[str, ...], message: str
|
||||
) -> None:
|
||||
result = deploy_project.run(*args)
|
||||
|
||||
assert result.exit_code != 0
|
||||
assert message in result.output
|
||||
assert deploy_project.docker.verbs() == []
|
||||
assert deploy_project.timeline == []
|
||||
|
||||
|
||||
def test_invalid_tag_fails_before_any_control_plane_call(
|
||||
deploy_project: DeployProject,
|
||||
) -> None:
|
||||
|
||||
@@ -13,6 +13,7 @@ import pytest
|
||||
|
||||
import langgraph_cli.deploy as deploy_mod
|
||||
from langgraph_cli.deploy import (
|
||||
BuildAndPush,
|
||||
ById,
|
||||
ByName,
|
||||
CustomerRegistrySource,
|
||||
@@ -21,6 +22,7 @@ from langgraph_cli.deploy import (
|
||||
Listener,
|
||||
ManagedRegistrySource,
|
||||
OnListener,
|
||||
PublishedImage,
|
||||
RemoteBuildSource,
|
||||
RequestedPlacement,
|
||||
Unplaced,
|
||||
@@ -614,6 +616,7 @@ class TestSelectSource:
|
||||
OPTIONS = {
|
||||
"push_to": None,
|
||||
"image": None,
|
||||
"image_uri": None,
|
||||
"image_name": None,
|
||||
"tag": None,
|
||||
"remote_build_flag": None,
|
||||
@@ -629,8 +632,10 @@ class TestSelectSource:
|
||||
{"push_to": REPOSITORY},
|
||||
True,
|
||||
CustomerRegistrySource(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image=None,
|
||||
image=BuildAndPush(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image=None,
|
||||
),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="push_to_selects_the_external_source_with_the_default_tag",
|
||||
@@ -639,8 +644,10 @@ class TestSelectSource:
|
||||
{"push_to": f"{REPOSITORY}:v2"},
|
||||
True,
|
||||
CustomerRegistrySource(
|
||||
reference=ImageReference(REPOSITORY, "v2"),
|
||||
prebuilt_image=None,
|
||||
image=BuildAndPush(
|
||||
reference=ImageReference(REPOSITORY, "v2"),
|
||||
prebuilt_image=None,
|
||||
),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="push_to_keeps_a_tag_given_in_the_reference",
|
||||
@@ -649,8 +656,10 @@ class TestSelectSource:
|
||||
{"push_to": REPOSITORY, "tag": "v3"},
|
||||
True,
|
||||
CustomerRegistrySource(
|
||||
reference=ImageReference(REPOSITORY, "v3"),
|
||||
prebuilt_image=None,
|
||||
image=BuildAndPush(
|
||||
reference=ImageReference(REPOSITORY, "v3"),
|
||||
prebuilt_image=None,
|
||||
),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="tag_flag_composes_with_push_to",
|
||||
@@ -659,8 +668,10 @@ class TestSelectSource:
|
||||
{"push_to": REPOSITORY, "image": "app:dev"},
|
||||
False,
|
||||
CustomerRegistrySource(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image="app:dev",
|
||||
image=BuildAndPush(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image="app:dev",
|
||||
),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="prebuilt_image_is_retagged_for_push_to_without_docker_checks",
|
||||
@@ -672,12 +683,44 @@ class TestSelectSource:
|
||||
},
|
||||
True,
|
||||
CustomerRegistrySource(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image=None,
|
||||
image=BuildAndPush(
|
||||
reference=ImageReference(REPOSITORY, "latest"),
|
||||
prebuilt_image=None,
|
||||
),
|
||||
requested_placement=RequestedPlacement("listener-1", "agents"),
|
||||
),
|
||||
id="push_to_carries_the_requested_placement",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": f"{REPOSITORY}@sha256:abc123"},
|
||||
True,
|
||||
CustomerRegistrySource(
|
||||
image=PublishedImage(f"{REPOSITORY}@sha256:abc123"),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="image_uri_selects_the_published_image_source_by_digest",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": f" {REPOSITORY}@sha256:abc123 "},
|
||||
False,
|
||||
CustomerRegistrySource(
|
||||
image=PublishedImage(f"{REPOSITORY}@sha256:abc123"),
|
||||
requested_placement=RequestedPlacement(),
|
||||
),
|
||||
id="image_uri_needs_no_local_docker_and_is_trimmed",
|
||||
),
|
||||
pytest.param(
|
||||
{
|
||||
"image_uri": f"{REPOSITORY}@sha256:abc123",
|
||||
"placement": RequestedPlacement("listener-1", "agents"),
|
||||
},
|
||||
False,
|
||||
CustomerRegistrySource(
|
||||
image=PublishedImage(f"{REPOSITORY}@sha256:abc123"),
|
||||
requested_placement=RequestedPlacement("listener-1", "agents"),
|
||||
),
|
||||
id="image_uri_carries_the_requested_placement",
|
||||
),
|
||||
pytest.param(
|
||||
{"remote_build_flag": True},
|
||||
True,
|
||||
@@ -763,6 +806,46 @@ class TestSelectSource:
|
||||
"only apply when creating a deployment with --push-to",
|
||||
id="namespace_without_push_to",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": REPOSITORY, "push_to": REPOSITORY},
|
||||
"--image-uri cannot be combined with --push-to.",
|
||||
id="image_uri_with_push_to",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": REPOSITORY, "image": "app:dev"},
|
||||
"--image-uri cannot be combined with --image.",
|
||||
id="image_uri_with_image",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": REPOSITORY, "tag": "v2"},
|
||||
"--image-uri cannot be combined with --tag.",
|
||||
id="image_uri_with_tag",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": REPOSITORY, "remote_build_flag": True},
|
||||
"--image-uri cannot be combined with --remote.",
|
||||
id="image_uri_with_remote",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": ""},
|
||||
"--image-uri must not be empty.",
|
||||
id="image_uri_empty",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": " "},
|
||||
"--image-uri must not be empty.",
|
||||
id="image_uri_blank",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": f"{REPOSITORY}:v1.2.3"},
|
||||
"--image-uri must pin a digest",
|
||||
id="image_uri_with_a_mutable_tag",
|
||||
),
|
||||
pytest.param(
|
||||
{"image_uri": REPOSITORY},
|
||||
"--image-uri must pin a digest",
|
||||
id="image_uri_without_any_tag_or_digest",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_conflicting_flags_are_rejected(self, monkeypatch, flags, message):
|
||||
@@ -1164,6 +1247,7 @@ def test_a_deployment_id_with_listener_flags_is_refused_without_probing_docker(
|
||||
_select_source(
|
||||
push_to="registry.example.com/app",
|
||||
image=None,
|
||||
image_uri=None,
|
||||
image_name=None,
|
||||
tag=None,
|
||||
remote_build_flag=None,
|
||||
|
||||
@@ -3,6 +3,7 @@ from unittest.mock import patch
|
||||
from langgraph_cli.deploy import (
|
||||
_extract_deployment_url,
|
||||
format_deployments_table,
|
||||
format_listeners_table,
|
||||
format_revisions_table,
|
||||
)
|
||||
from langgraph_cli.util import clean_empty_lines, warn_non_wolfi_distro
|
||||
@@ -255,3 +256,34 @@ def test_format_revisions_table():
|
||||
assert "rev-456" in output
|
||||
assert "rev-789" in output
|
||||
assert "REPLACED" in output
|
||||
|
||||
|
||||
def test_format_listeners_table():
|
||||
output = format_listeners_table(
|
||||
[
|
||||
{
|
||||
"id": "listener-1",
|
||||
"compute_id": "prod-cluster",
|
||||
"compute_config": {"k8s_namespaces": ["agents"]},
|
||||
},
|
||||
{
|
||||
"id": "listener-2",
|
||||
"compute_id": "multi-cluster",
|
||||
"compute_config": {"k8s_namespaces": ["agents", "agents-staging"]},
|
||||
},
|
||||
{
|
||||
"id": "listener-3",
|
||||
"compute_id": "broken-cluster",
|
||||
},
|
||||
]
|
||||
)
|
||||
assert "Listener ID" in output
|
||||
assert "Compute ID" in output
|
||||
assert "Namespaces" in output
|
||||
assert "listener-1" in output
|
||||
assert "prod-cluster" in output
|
||||
assert "agents" in output
|
||||
assert "listener-2" in output
|
||||
assert "agents, agents-staging" in output
|
||||
assert "listener-3" in output
|
||||
assert "broken-cluster" in output
|
||||
|
||||
@@ -28,6 +28,7 @@ __all__ = (
|
||||
"ParentCommand",
|
||||
"EmptyInputError",
|
||||
"TaskNotFound",
|
||||
"is_invalid_resume",
|
||||
)
|
||||
|
||||
|
||||
@@ -239,3 +240,21 @@ class NodeTimeoutError(Exception):
|
||||
self.kind = kind
|
||||
self.idle_timeout = idle_timeout
|
||||
self.run_timeout = run_timeout
|
||||
|
||||
|
||||
_INVALID_RESUME = "_langgraph_invalid_resume"
|
||||
|
||||
|
||||
def _mark_invalid_resume(error: BaseException) -> None:
|
||||
setattr(error, _INVALID_RESUME, True)
|
||||
|
||||
|
||||
def is_invalid_resume(error: BaseException) -> bool:
|
||||
"""Whether `error` was raised because a resume value didn't match `response_schema`.
|
||||
|
||||
`interrupt()` raises a `pydantic.ValidationError` in that case. `ToolNode` uses
|
||||
this to tell it apart from invalid tool arguments when a tool calls `interrupt()`
|
||||
or runs a graph that does, so the resume fails and the interrupt can be answered
|
||||
again.
|
||||
"""
|
||||
return getattr(error, _INVALID_RESUME, False) is True
|
||||
|
||||
@@ -27,6 +27,7 @@ from langgraph.checkpoint.base import (
|
||||
Checkpoint,
|
||||
PendingWrite,
|
||||
V,
|
||||
writes_sort_key,
|
||||
)
|
||||
from langgraph.store.base import BaseStore
|
||||
from xxhash import xxh3_128_hexdigest
|
||||
@@ -251,9 +252,7 @@ def apply_writes(
|
||||
Set of channels that were updated in this step.
|
||||
"""
|
||||
# sort tasks on path, to ensure deterministic order for update application
|
||||
# any path parts after the 3rd are ignored for sorting
|
||||
# (we use them for eg. task ids which aren't good for sorting)
|
||||
tasks = sorted(tasks, key=lambda t: task_path_str(t.path[:3]))
|
||||
tasks = sorted(tasks, key=lambda t: writes_sort_key(task_path_str(t.path)))
|
||||
# if no task has triggers this is applying writes from the null task only
|
||||
# so we don't do anything other than update the channels written to
|
||||
bump_step = any(t.triggers for t in tasks)
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterable, Mapping
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from inspect import signature
|
||||
from typing import Any, Literal, cast
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
@@ -25,11 +26,13 @@ from langgraph._internal._constants import (
|
||||
CONFIG_KEY_CHECKPOINT_ID,
|
||||
NS_END,
|
||||
NS_SEP,
|
||||
NULL_TASK_ID,
|
||||
PUSH,
|
||||
SNAPSHOT_BUMPS,
|
||||
)
|
||||
from langgraph._internal._typing import MISSING
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.channels.binop import _get_overwrite
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.managed.base import ManagedValueMapping, ManagedValueSpec
|
||||
|
||||
@@ -49,15 +52,38 @@ def empty_checkpoint() -> Checkpoint:
|
||||
)
|
||||
|
||||
|
||||
def put_writes_accepts_task_path(put_writes: Callable[..., Any]) -> bool:
|
||||
"""Whether a saver's `put_writes` or `aput_writes` takes `task_path`.
|
||||
|
||||
Savers written before the parameter existed don't, so it is passed only
|
||||
when this is true.
|
||||
"""
|
||||
return signature(put_writes).parameters.get("task_path") is not None
|
||||
|
||||
|
||||
def exit_delta_task_id(step: int, task_id: str) -> str:
|
||||
"""Synthetic task id for exit-mode DeltaChannel writes.
|
||||
|
||||
Embeds the superstep in the first UUID group so `ORDER BY task_id, idx`
|
||||
preserves chronological order while remaining a valid RFC UUID (required by
|
||||
Postgres `checkpoint_writes.task_id uuid` columns).
|
||||
Postgres `checkpoint_writes.task_id uuid` columns). Never `NULL_TASK_ID`:
|
||||
readers apply writes under it as the anchor checkpoint's own pending writes.
|
||||
"""
|
||||
parts = str(uuid.UUID(task_id)).split("-")
|
||||
return f"{step:08d}-{parts[1]}-{parts[2]}-{parts[3]}-{parts[4]}"
|
||||
synthetic = f"{step:08d}-{parts[1]}-{parts[2]}-{parts[3]}-{parts[4]}"
|
||||
if synthetic == NULL_TASK_ID:
|
||||
return f"{step:08d}-0000-0000-0000-000000000001"
|
||||
return synthetic
|
||||
|
||||
|
||||
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(
|
||||
@@ -127,6 +153,23 @@ def delta_channels_with_pending_writes(
|
||||
}
|
||||
|
||||
|
||||
def delta_channels_overwritten(
|
||||
specs: Mapping[str, Any], writes: Iterable[tuple[str, Any]]
|
||||
) -> set[str]:
|
||||
"""Return the names of the DeltaChannels that `writes` set with an `Overwrite`.
|
||||
|
||||
`update_state` saves a full snapshot of these channels in the checkpoint it
|
||||
creates, like the loop does when a node returns an `Overwrite`. Otherwise,
|
||||
reading the channel later starts from an older snapshot and replays the
|
||||
writes the `Overwrite` threw away.
|
||||
"""
|
||||
return {
|
||||
ch
|
||||
for ch, value in writes
|
||||
if isinstance(specs.get(ch), DeltaChannel) and _get_overwrite(value)[0]
|
||||
}
|
||||
|
||||
|
||||
def checkpoint_superseded(
|
||||
saver: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple
|
||||
) -> bool:
|
||||
@@ -231,6 +274,7 @@ def create_checkpoint(
|
||||
get_next_version: GetNextVersion | None = None,
|
||||
channels_to_snapshot: set[str] | None = None,
|
||||
stored_versions: ChannelVersions | None = None,
|
||||
trigger_to_nodes: Mapping[str, Sequence[str]] | None = None,
|
||||
) -> Checkpoint:
|
||||
"""Build a new Checkpoint from the previous one and live channel state.
|
||||
|
||||
@@ -289,7 +333,9 @@ def create_checkpoint(
|
||||
id=id or str(uuid6(clock_seq=step)),
|
||||
channel_values=values,
|
||||
channel_versions=channel_versions,
|
||||
versions_seen=_mark_bumps_seen(checkpoint["versions_seen"], bumped),
|
||||
versions_seen=_mark_bumps_seen(
|
||||
checkpoint["versions_seen"], bumped, trigger_to_nodes or {}
|
||||
),
|
||||
updated_channels=None if updated_channels is None else sorted(updated_channels),
|
||||
)
|
||||
|
||||
@@ -297,19 +343,27 @@ def create_checkpoint(
|
||||
def _mark_bumps_seen(
|
||||
versions_seen: dict[str, ChannelVersions],
|
||||
bumped: Mapping[str, tuple[Any, Any]],
|
||||
trigger_to_nodes: Mapping[str, Sequence[str]],
|
||||
) -> dict[str, ChannelVersions]:
|
||||
"""Advance whoever had seen a bumped channel's old version to the new one.
|
||||
|
||||
A bump that only stores a snapshot is not a write. Left unseen, it would
|
||||
re-fire `interrupt_before` and rerun the channel's subscribers. For each
|
||||
entry it advances, `SNAPSHOT_BUMPS` keeps the new version and the one the
|
||||
node really read, so `versions_seen_without_bumps` can put the read back.
|
||||
re-fire `interrupt_before` and rerun the channel's subscribers. A channel
|
||||
bumped from no version was never written, so it also goes to the
|
||||
subscribers that never ran: they have no entry, and would start on the bump.
|
||||
For each entry it advances, `SNAPSHOT_BUMPS` keeps the new version and the
|
||||
one the node really read, so `versions_seen_without_bumps` can put the read
|
||||
back.
|
||||
"""
|
||||
if not bumped:
|
||||
return versions_seen
|
||||
out = dict(versions_seen)
|
||||
for k, (old, _) in bumped.items():
|
||||
if old is None:
|
||||
for node in trigger_to_nodes.get(k, ()):
|
||||
out.setdefault(node, {})
|
||||
marks = dict(versions_seen.get(SNAPSHOT_BUMPS, {}))
|
||||
for node, seen in versions_seen.items():
|
||||
for node, seen in out.items():
|
||||
if node == SNAPSHOT_BUMPS:
|
||||
continue
|
||||
for k, (old, new) in bumped.items():
|
||||
|
||||
@@ -12,7 +12,6 @@ from contextlib import (
|
||||
ExitStack,
|
||||
)
|
||||
from datetime import datetime, timezone
|
||||
from inspect import signature
|
||||
from types import TracebackType
|
||||
from typing import (
|
||||
Any,
|
||||
@@ -20,6 +19,7 @@ from typing import (
|
||||
TypeVar,
|
||||
cast,
|
||||
)
|
||||
from uuid import UUID, uuid5
|
||||
|
||||
from langchain_core.callbacks import AsyncParentRunManager, ParentRunManager
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
@@ -106,7 +106,9 @@ from langgraph.pregel._checkpoint import (
|
||||
delta_channels_to_snapshot,
|
||||
delta_channels_with_pending_writes,
|
||||
empty_checkpoint,
|
||||
exit_delta_late_task_id,
|
||||
exit_delta_task_id,
|
||||
put_writes_accepts_task_path,
|
||||
)
|
||||
from langgraph.pregel._executor import (
|
||||
AsyncBackgroundExecutor,
|
||||
@@ -190,6 +192,7 @@ class PregelLoop:
|
||||
Callable[
|
||||
[
|
||||
concurrent.futures.Future | None,
|
||||
Sequence[Any],
|
||||
RunnableConfig,
|
||||
Checkpoint,
|
||||
str,
|
||||
@@ -203,11 +206,13 @@ class PregelLoop:
|
||||
submit: Submit
|
||||
channels: Mapping[str, BaseChannel]
|
||||
# Futures from `checkpointer.put_writes` calls that produced delta-channel
|
||||
# writes. `_checkpointer_put_after_previous` drains this list (swap to a
|
||||
# local `futs` then reset to `[]` and wait/gather) before putting the
|
||||
# next checkpoint, so a checkpoint never becomes durable before the
|
||||
# writes that produced it. Initialised to `[]` in both sync and async
|
||||
# `__enter__`; stays `None` only when no checkpointer.
|
||||
# writes. `_put_checkpoint` hands this list to the save it submits, which
|
||||
# waits for them first, so a checkpoint never becomes durable before the
|
||||
# writes that produced it. If a write or the previous save failed, the
|
||||
# save fails too: a DeltaChannel is rebuilt from its writes along the
|
||||
# parent chain, so a checkpoint saved past either gap reads back short
|
||||
# for good. Initialised to `[]` in both sync and async `__enter__`;
|
||||
# stays `None` only when no checkpointer.
|
||||
_delta_write_futs: list[Any] | None = None
|
||||
|
||||
# Same pattern as `_delta_write_futs` but for error-handler writes.
|
||||
@@ -221,10 +226,18 @@ 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 (task_id, channel) pairs whose delta writes are stored on the loaded
|
||||
# checkpoint, which a resume addressed by `checkpoint_id` can rerun; this
|
||||
# run's `Command` delta writes, kept apart from the NULL_TASK_ID writes
|
||||
# loaded with the checkpoint; and the checkpoint's own superstep, the
|
||||
# first one this run ticks.
|
||||
_stored_delta_writes: set[tuple[str, str]]
|
||||
_exit_command_writes: list[tuple[str, Any]]
|
||||
_exit_first_step: int | None = None
|
||||
|
||||
# Delta channels that must snapshot at the next checkpoint, whatever their
|
||||
# cadence counters say:
|
||||
@@ -257,6 +270,9 @@ class PregelLoop:
|
||||
# `_put_exit_delta_writes` uses this to decide between anchoring on
|
||||
# the existing parent (True) or creating a lazy stub (False).
|
||||
_has_persisted_parent: bool = False
|
||||
# True iff `__enter__` loaded the thread's latest checkpoint, not one a
|
||||
# `checkpoint_id` addressed, so nothing has been built on it yet.
|
||||
_loaded_latest: bool = False
|
||||
|
||||
managed: ManagedValueMapping
|
||||
checkpoint: Checkpoint
|
||||
@@ -669,10 +685,6 @@ class PregelLoop:
|
||||
self.status = "done"
|
||||
return False
|
||||
|
||||
if self.control is not None and self.control.drain_requested:
|
||||
self.status = "draining"
|
||||
return False
|
||||
|
||||
# if there are pending writes from a previous loop, apply them
|
||||
if self._reapplies_pending_writes and self.checkpoint_pending_writes:
|
||||
self._reapply_writes_to_succeeded_nodes(self.tasks)
|
||||
@@ -685,6 +697,13 @@ class PregelLoop:
|
||||
self.status = "interrupt_before"
|
||||
raise GraphInterrupt()
|
||||
|
||||
# stop before running the next tasks if a drain was requested. after
|
||||
# the breakpoint check: a resume passes the next step's breakpoints,
|
||||
# so draining in front of one would skip it
|
||||
if self.control is not None and self.control.drain_requested:
|
||||
self.status = "draining"
|
||||
return False
|
||||
|
||||
# produce debug output
|
||||
self._emit("tasks", map_debug_tasks, self.tasks.values())
|
||||
|
||||
@@ -737,9 +756,25 @@ class PregelLoop:
|
||||
)
|
||||
# capture delta-channel writes for exit-mode accumulator before clearing
|
||||
if self._exit_delta_writes is not None:
|
||||
# On the first tick the pending writes still hold the ones loaded
|
||||
# with the checkpoint, which are already stored on it.
|
||||
first = self._exit_first_step is None
|
||||
if first:
|
||||
self._exit_first_step = self.step
|
||||
self._exit_delta_writes.extend(
|
||||
(self.step, NULL_TASK_ID, "", ch, v)
|
||||
for ch, v in self._exit_command_writes
|
||||
)
|
||||
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 not isinstance(self.specs.get(ch), DeltaChannel):
|
||||
continue
|
||||
if first and (
|
||||
tid == NULL_TASK_ID or (tid, ch) in self._stored_delta_writes
|
||||
):
|
||||
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))
|
||||
# clear pending writes
|
||||
self.checkpoint_pending_writes.clear()
|
||||
# only replay (re-execute) done tasks on the first tick
|
||||
@@ -860,6 +895,12 @@ class PregelLoop:
|
||||
def _first(
|
||||
self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None
|
||||
) -> set[str] | None:
|
||||
self._stored_delta_writes = {
|
||||
(tid, ch)
|
||||
for tid, ch, _ in self.checkpoint_pending_writes
|
||||
if tid != NULL_TASK_ID and isinstance(self.specs.get(ch), DeltaChannel)
|
||||
}
|
||||
self._exit_command_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)
|
||||
@@ -977,6 +1018,12 @@ class PregelLoop:
|
||||
carried.extend((tid, c, v) for c, v in ws)
|
||||
else:
|
||||
self.put_writes(tid, ws)
|
||||
if self._exit_delta_writes is not None and tid == NULL_TASK_ID:
|
||||
self._exit_command_writes.extend(
|
||||
(c, v)
|
||||
for c, v in ws
|
||||
if isinstance(self.specs.get(c), DeltaChannel)
|
||||
)
|
||||
self._delta_channels_forced_snapshot.update(
|
||||
delta_channels_with_pending_writes(self.specs, carried)
|
||||
)
|
||||
@@ -1074,17 +1121,29 @@ 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))
|
||||
# Persist delta-channel input writes so sub-freq inputs are
|
||||
# recoverable via ancestor walk (mirrors the Command input path).
|
||||
self._exit_delta_writes.append(
|
||||
(self.step, NULL_TASK_ID, "", c, v)
|
||||
)
|
||||
# A DeltaChannel reads its input from the writes stored on the
|
||||
# checkpoint this run starts from, under a task id of their own:
|
||||
# readers apply a checkpoint's NULL_TASK_ID writes as its own state.
|
||||
# A new thread has no checkpoint to store them on, and one a
|
||||
# `checkpoint_id` addressed may have children that would read them,
|
||||
# so then the input checkpoint snapshots the channel instead.
|
||||
if self.durability != "exit":
|
||||
delta_input = [
|
||||
(c, v)
|
||||
for c, v in input_writes
|
||||
if isinstance(self.specs.get(c), DeltaChannel)
|
||||
]
|
||||
if delta_input:
|
||||
self.put_writes(NULL_TASK_ID, delta_input)
|
||||
if delta_input and self._has_persisted_parent and self._loaded_latest:
|
||||
self.put_writes(
|
||||
str(uuid5(UUID(self.checkpoint["id"]), INPUT)), delta_input
|
||||
)
|
||||
else:
|
||||
self._delta_channels_forced_snapshot.update(
|
||||
c for c, _ in delta_input
|
||||
)
|
||||
# save input checkpoint
|
||||
self.updated_channels = updated_channels
|
||||
self._put_checkpoint({"source": "input"})
|
||||
@@ -1214,6 +1273,7 @@ class PregelLoop:
|
||||
else None,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=self.checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
for k in channels_to_snapshot:
|
||||
new_counters[k] = (0, 0)
|
||||
@@ -1260,12 +1320,17 @@ class PregelLoop:
|
||||
)
|
||||
self.checkpoint_previous_versions = channel_versions
|
||||
|
||||
# Take this checkpoint's writes now: saves run in the background
|
||||
# and can start out of order, so a save that took them itself
|
||||
# could get another checkpoint's writes.
|
||||
delta_write_futs, self._delta_write_futs = self._delta_write_futs, []
|
||||
# save it, without blocking
|
||||
# if there's a previous checkpoint save in progress, wait for it
|
||||
# ensuring checkpointers receive checkpoints in order
|
||||
self._put_checkpoint_fut = self.submit(
|
||||
self._checkpointer_put_after_previous,
|
||||
getattr(self, "_put_checkpoint_fut", None),
|
||||
delta_write_futs,
|
||||
self.checkpoint_config,
|
||||
copy_checkpoint(self.checkpoint),
|
||||
self.checkpoint_metadata,
|
||||
@@ -1309,9 +1374,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
|
||||
@@ -1337,6 +1400,7 @@ class PregelLoop:
|
||||
self._put_checkpoint_fut = self.submit(
|
||||
self._checkpointer_put_after_previous,
|
||||
getattr(self, "_put_checkpoint_fut", None),
|
||||
(),
|
||||
stub_put_config,
|
||||
stub_cp,
|
||||
{"step": -2},
|
||||
@@ -1346,11 +1410,21 @@ 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 keeps its real task paths, so it
|
||||
# interleaves with the writes a resume loaded from it. Its task ids stay
|
||||
# synthetic: under the real id, a run whose final checkpoint fails to
|
||||
# save would leave the resumed task looking done to the next resume.
|
||||
# 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 = (exit_delta_task_id(step, 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,
|
||||
{
|
||||
@@ -1360,22 +1434,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)
|
||||
@@ -1584,8 +1657,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
self.checkpointer_get_next_version = checkpointer.get_next_version
|
||||
self.checkpointer_put_writes = checkpointer.put_writes
|
||||
self.checkpointer_put_writes_accepts_task_path = (
|
||||
signature(checkpointer.put_writes).parameters.get("task_path")
|
||||
is not None
|
||||
put_writes_accepts_task_path(checkpointer.put_writes)
|
||||
)
|
||||
else:
|
||||
self.checkpointer_get_next_version = increment
|
||||
@@ -1596,21 +1668,19 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
def _checkpointer_put_after_previous(
|
||||
self,
|
||||
prev: concurrent.futures.Future | None,
|
||||
delta_write_futs: Sequence[concurrent.futures.Future],
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
if self._delta_write_futs:
|
||||
futs, self._delta_write_futs = self._delta_write_futs, []
|
||||
concurrent.futures.wait(futs)
|
||||
try:
|
||||
if prev is not None:
|
||||
prev.result()
|
||||
finally:
|
||||
cast(BaseCheckpointSaver, self.checkpointer).put(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
for fut in delta_write_futs:
|
||||
fut.result()
|
||||
if prev is not None:
|
||||
prev.result()
|
||||
cast(BaseCheckpointSaver, self.checkpointer).put(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
|
||||
def match_cached_writes(self) -> Sequence[PregelExecutableTask]:
|
||||
if self.cache is None:
|
||||
@@ -1722,6 +1792,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
# Normal case: fetch the most recent checkpoint for this
|
||||
# graph/thread. Returns None on first invocation.
|
||||
saved = self.checkpointer.get_tuple(self.checkpoint_config)
|
||||
self._loaded_latest = True
|
||||
|
||||
# Capture before the synthetic-empty fallback below overwrites `saved`.
|
||||
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
|
||||
@@ -1840,8 +1911,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
self.checkpointer_get_next_version = checkpointer.get_next_version
|
||||
self.checkpointer_put_writes = checkpointer.aput_writes
|
||||
self.checkpointer_put_writes_accepts_task_path = (
|
||||
signature(checkpointer.aput_writes).parameters.get("task_path")
|
||||
is not None
|
||||
put_writes_accepts_task_path(checkpointer.aput_writes)
|
||||
)
|
||||
else:
|
||||
self.checkpointer_get_next_version = increment
|
||||
@@ -1852,23 +1922,19 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
async def _checkpointer_put_after_previous(
|
||||
self,
|
||||
prev: asyncio.Task | None,
|
||||
delta_write_futs: Sequence[asyncio.Future],
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
# Drain DeltaChannel write futures before committing the checkpoint so
|
||||
# ancestor walks never see a checkpoint without its backing writes.
|
||||
if self._delta_write_futs:
|
||||
futs, self._delta_write_futs = self._delta_write_futs, []
|
||||
await asyncio.gather(*futs)
|
||||
try:
|
||||
if prev is not None:
|
||||
await prev
|
||||
finally:
|
||||
await cast(BaseCheckpointSaver, self.checkpointer).aput(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
if delta_write_futs:
|
||||
await asyncio.gather(*delta_write_futs)
|
||||
if prev is not None:
|
||||
await prev
|
||||
await cast(BaseCheckpointSaver, self.checkpointer).aput(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
|
||||
async def amatch_cached_writes(self) -> Sequence[PregelExecutableTask]:
|
||||
if self.cache is None:
|
||||
@@ -1983,6 +2049,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
# Normal case: fetch the most recent checkpoint for this
|
||||
# graph/thread. Returns None on first invocation.
|
||||
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
|
||||
self._loaded_latest = True
|
||||
|
||||
# Capture before the synthetic-empty fallback below overwrites `saved`.
|
||||
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
|
||||
|
||||
@@ -126,6 +126,7 @@ from langgraph.pregel._algo import (
|
||||
apply_writes,
|
||||
local_read,
|
||||
prepare_next_tasks,
|
||||
task_path_str,
|
||||
)
|
||||
from langgraph.pregel._call import identifier
|
||||
from langgraph.pregel._checkpoint import (
|
||||
@@ -136,9 +137,11 @@ from langgraph.pregel._checkpoint import (
|
||||
copy_checkpoint,
|
||||
create_checkpoint,
|
||||
create_checkpoint_plan_for_update_state_api,
|
||||
delta_channels_overwritten,
|
||||
delta_channels_with_pending_writes,
|
||||
empty_checkpoint,
|
||||
get_updated_channels_from_tasks,
|
||||
put_writes_accepts_task_path,
|
||||
versions_seen_without_bumps,
|
||||
)
|
||||
from langgraph.pregel._draw import draw_graph
|
||||
@@ -150,6 +153,7 @@ from langgraph.pregel._loop import (
|
||||
from langgraph.pregel._messages import (
|
||||
StreamMessagesHandler,
|
||||
StreamMessagesHandlerV2,
|
||||
ensure_message_ids,
|
||||
)
|
||||
from langgraph.pregel._read import DEFAULT_BOUND, PregelNode
|
||||
from langgraph.pregel._retry import RetryPolicy
|
||||
@@ -1764,6 +1768,7 @@ class Pregel(
|
||||
get_next_version=checkpointer.get_next_version,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = checkpointer.put(
|
||||
checkpoint_config,
|
||||
@@ -1786,6 +1791,22 @@ class Pregel(
|
||||
)
|
||||
|
||||
if input_writes := deque(map_input(self.input_channels, values)):
|
||||
_store_or_fork_delta_writes(
|
||||
checkpointer,
|
||||
config,
|
||||
saved,
|
||||
checkpoint_config,
|
||||
self.channels,
|
||||
[
|
||||
(
|
||||
str(uuid5(UUID(checkpoint["id"]), INPUT)),
|
||||
input_writes,
|
||||
None,
|
||||
)
|
||||
],
|
||||
fork_pending,
|
||||
is_first=is_first,
|
||||
)
|
||||
updated_channels = apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
@@ -1793,6 +1814,9 @@ class Pregel(
|
||||
checkpointer.get_next_version,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
fork_pending |= delta_channels_overwritten(
|
||||
self.channels, input_writes
|
||||
)
|
||||
|
||||
# apply input write to channels
|
||||
next_step = (
|
||||
@@ -1820,6 +1844,7 @@ class Pregel(
|
||||
get_next_version=checkpointer.get_next_version,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = checkpointer.put(
|
||||
checkpoint_config,
|
||||
@@ -1831,13 +1856,6 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
|
||||
# store the writes
|
||||
checkpointer.put_writes(
|
||||
next_config,
|
||||
input_writes,
|
||||
str(uuid5(UUID(checkpoint["id"]), INPUT)),
|
||||
)
|
||||
|
||||
return patch_checkpoint_map(
|
||||
next_config, saved.metadata if saved else None
|
||||
)
|
||||
@@ -1983,13 +2001,13 @@ 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:
|
||||
raise InvalidUpdateError(f"Node {as_node} has no writers")
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
|
||||
task = PregelTaskWrites((INTERRUPT, i), as_node, writes, [INTERRUPT])
|
||||
# get the task ids that were prepared for this node
|
||||
# if a task id was provided in the StateUpdate, we use it
|
||||
# otherwise, we use the next available task id
|
||||
@@ -1997,7 +2015,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)
|
||||
@@ -2031,27 +2049,22 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
updated_channels = get_updated_channels_from_tasks(run_tasks)
|
||||
# The base's other children replay whatever is stored on it, so an
|
||||
# edit of an older checkpoint stores none of its writes there: the
|
||||
# checkpoint written here carries them, its delta channels
|
||||
# snapshotted. Later supersteps address the checkpoint just written.
|
||||
if (
|
||||
is_first
|
||||
and saved is not None
|
||||
and checkpoint_superseded(checkpointer, config, saved)
|
||||
):
|
||||
fork_pending.update(
|
||||
ch
|
||||
for ch in updated_channels
|
||||
if isinstance(self.channels.get(ch), DeltaChannel)
|
||||
)
|
||||
elif saved is not None:
|
||||
for task_id, task in zip(run_task_ids, run_tasks):
|
||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
||||
if channel_writes:
|
||||
checkpointer.put_writes(
|
||||
checkpoint_config, channel_writes, task_id
|
||||
)
|
||||
fork_pending |= delta_channels_overwritten(
|
||||
self.channels, (w for t in run_tasks for w in t.writes)
|
||||
)
|
||||
_store_or_fork_delta_writes(
|
||||
checkpointer,
|
||||
config,
|
||||
saved,
|
||||
checkpoint_config,
|
||||
self.channels,
|
||||
[
|
||||
(task_id, [w for w in task.writes if w[0] != PUSH], task)
|
||||
for task_id, task in zip(run_task_ids, run_tasks)
|
||||
],
|
||||
fork_pending,
|
||||
is_first=is_first,
|
||||
)
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
@@ -2081,6 +2094,7 @@ class Pregel(
|
||||
else None,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = checkpointer.put(
|
||||
checkpoint_config,
|
||||
@@ -2253,6 +2267,7 @@ class Pregel(
|
||||
get_next_version=checkpointer.get_next_version,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = await checkpointer.aput(
|
||||
checkpoint_config,
|
||||
@@ -2275,6 +2290,22 @@ class Pregel(
|
||||
)
|
||||
|
||||
if input_writes := deque(map_input(self.input_channels, values)):
|
||||
await _astore_or_fork_delta_writes(
|
||||
checkpointer,
|
||||
config,
|
||||
saved,
|
||||
checkpoint_config,
|
||||
self.channels,
|
||||
[
|
||||
(
|
||||
str(uuid5(UUID(checkpoint["id"]), INPUT)),
|
||||
input_writes,
|
||||
None,
|
||||
)
|
||||
],
|
||||
fork_pending,
|
||||
is_first=is_first,
|
||||
)
|
||||
updated_channels = apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
@@ -2282,6 +2313,9 @@ class Pregel(
|
||||
checkpointer.get_next_version,
|
||||
self.trigger_to_nodes,
|
||||
)
|
||||
fork_pending |= delta_channels_overwritten(
|
||||
self.channels, input_writes
|
||||
)
|
||||
|
||||
# apply input write to channels
|
||||
next_step = (
|
||||
@@ -2309,6 +2343,7 @@ class Pregel(
|
||||
get_next_version=checkpointer.get_next_version,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = await checkpointer.aput(
|
||||
checkpoint_config,
|
||||
@@ -2320,13 +2355,6 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
|
||||
# store the writes
|
||||
await checkpointer.aput_writes(
|
||||
next_config,
|
||||
input_writes,
|
||||
str(uuid5(UUID(checkpoint["id"]), INPUT)),
|
||||
)
|
||||
|
||||
return patch_checkpoint_map(
|
||||
next_config, saved.metadata if saved else None
|
||||
)
|
||||
@@ -2471,13 +2499,13 @@ 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:
|
||||
raise InvalidUpdateError(f"Node {as_node} has no writers")
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
|
||||
task = PregelTaskWrites((INTERRUPT, i), as_node, writes, [INTERRUPT])
|
||||
# get the task ids that were prepared for this node
|
||||
# if a task id was provided in the StateUpdate, we use it
|
||||
# otherwise, we use the next available task id
|
||||
@@ -2485,7 +2513,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)
|
||||
@@ -2519,27 +2547,22 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
updated_channels = get_updated_channels_from_tasks(run_tasks)
|
||||
# The base's other children replay whatever is stored on it, so an
|
||||
# edit of an older checkpoint stores none of its writes there: the
|
||||
# checkpoint written here carries them, its delta channels
|
||||
# snapshotted. Later supersteps address the checkpoint just written.
|
||||
if (
|
||||
is_first
|
||||
and saved is not None
|
||||
and await acheckpoint_superseded(checkpointer, config, saved)
|
||||
):
|
||||
fork_pending.update(
|
||||
ch
|
||||
for ch in updated_channels
|
||||
if isinstance(self.channels.get(ch), DeltaChannel)
|
||||
)
|
||||
elif saved is not None:
|
||||
for task_id, task in zip(run_task_ids, run_tasks):
|
||||
channel_writes = [w for w in task.writes if w[0] != PUSH]
|
||||
if channel_writes:
|
||||
await checkpointer.aput_writes(
|
||||
checkpoint_config, channel_writes, task_id
|
||||
)
|
||||
fork_pending |= delta_channels_overwritten(
|
||||
self.channels, (w for t in run_tasks for w in t.writes)
|
||||
)
|
||||
await _astore_or_fork_delta_writes(
|
||||
checkpointer,
|
||||
config,
|
||||
saved,
|
||||
checkpoint_config,
|
||||
self.channels,
|
||||
[
|
||||
(task_id, [w for w in task.writes if w[0] != PUSH], task)
|
||||
for task_id, task in zip(run_task_ids, run_tasks)
|
||||
],
|
||||
fork_pending,
|
||||
is_first=is_first,
|
||||
)
|
||||
apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
@@ -2569,6 +2592,7 @@ class Pregel(
|
||||
else None,
|
||||
channels_to_snapshot=channels_to_snapshot,
|
||||
stored_versions=checkpoint_previous_versions,
|
||||
trigger_to_nodes=self.trigger_to_nodes,
|
||||
)
|
||||
next_config = await checkpointer.aput(
|
||||
checkpoint_config,
|
||||
@@ -3706,16 +3730,15 @@ class Pregel(
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
version: Literal["v1", "v2", "v3"] = "v2",
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
control: RunControl | None = None,
|
||||
transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
"""Stream events from this graph.
|
||||
|
||||
For `version="v1"` / `"v2"`, yields `StreamEvent` dicts (see
|
||||
`Runnable.stream_events`). For `version="v3"`, returns a
|
||||
For `version="v1"` / `"v2"`, delegates to
|
||||
`Runnable.(a)stream_events`; synchronous v1/v2 event streaming is
|
||||
not implemented in langchain-core, so use `astream_events` for
|
||||
those versions. For `version="v3"`, returns a
|
||||
`GraphRunStream` whose typed projections the caller drives by
|
||||
iterating — no background thread.
|
||||
|
||||
@@ -3745,20 +3768,27 @@ class Pregel(
|
||||
config: Optional runnable config.
|
||||
version: Streaming-event schema version. `"v3"` selects the
|
||||
content-block-centric streaming protocol.
|
||||
interrupt_before: Nodes to interrupt before, if any. Only
|
||||
used for `version="v3"`.
|
||||
interrupt_after: Nodes to interrupt after, if any. Only
|
||||
used for `version="v3"`.
|
||||
interrupt_before: Nodes to interrupt before, if any.
|
||||
Honored on every version that can run; type-checked
|
||||
only on the `version="v3"` overloads.
|
||||
interrupt_after: Nodes to interrupt after, if any. Honored
|
||||
on every version that can run; type-checked only on
|
||||
the `version="v3"` overloads.
|
||||
control: Optional run control used to request cooperative
|
||||
drain. Only used for `version="v3"`.
|
||||
drain. Honored on every version that can run;
|
||||
type-checked only on the `version="v3"` overloads.
|
||||
transformers: Extra transformer classes or configured
|
||||
factories appended after compile-time
|
||||
`stream_transformers`. Factories are called as
|
||||
`factory(scope)` so they can propagate to subgraph
|
||||
scopes. Only used for `version="v3"`.
|
||||
**kwargs: For `version="v1"`/`"v2"`, forwarded to
|
||||
`Runnable.stream_events`. For `version="v3"`, forwarded
|
||||
to the underlying `stream(...)` call (e.g. `context`,
|
||||
**kwargs: For `version="v1"`/`"v2"` on `astream_events`,
|
||||
forwarded to `Runnable.astream_events`, which passes
|
||||
them through to `astream` — so execution kwargs such as
|
||||
`context`, `durability`, `interrupt_before`,
|
||||
`interrupt_after` and `control` are honored on every
|
||||
version that can run. For `version="v3"`, forwarded to the
|
||||
underlying `stream(...)` call (e.g. `context`,
|
||||
`durability`, `output_keys`, `print_mode`, `debug`).
|
||||
`stream_mode` and `subgraphs` are not accepted under
|
||||
`version="v3"` and raise `TypeError` if supplied; v3
|
||||
@@ -3766,16 +3796,15 @@ class Pregel(
|
||||
|
||||
Returns:
|
||||
For `version="v3"`, a `GraphRunStream` the caller iterates
|
||||
to drive the run. Otherwise an `Iterator[StreamEvent]`.
|
||||
to drive the run. For `version="v1"`/`"v2"`,
|
||||
`astream_events` yields `StreamEvent` dicts; the synchronous
|
||||
v1/v2 path is not implemented in langchain-core.
|
||||
"""
|
||||
if version == "v3":
|
||||
_reject_v3_invariant_kwargs(kwargs)
|
||||
return self._pregel_stream_v3(
|
||||
input,
|
||||
config,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
control=control,
|
||||
transformers=transformers,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -3811,9 +3840,6 @@ class Pregel(
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
version: Literal["v1", "v2", "v3"] = "v2",
|
||||
interrupt_before: All | Sequence[str] | None = None,
|
||||
interrupt_after: All | Sequence[str] | None = None,
|
||||
control: RunControl | None = None,
|
||||
transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[StreamEvent] | Awaitable[Any]:
|
||||
@@ -3836,9 +3862,6 @@ class Pregel(
|
||||
return self._apregel_stream_v3(
|
||||
input,
|
||||
config,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
control=control,
|
||||
transformers=transformers,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -4237,6 +4260,109 @@ class Pregel(
|
||||
await self.cache.aclear(namespaces)
|
||||
|
||||
|
||||
def _task_path_kwarg(put_writes: Callable[..., Any], task: PregelTaskWrites) -> dict:
|
||||
"""Pass the task's path to savers whose `put_writes` takes one.
|
||||
|
||||
Savers that replay a checkpoint's writes in task path order then give back
|
||||
updates applied together in the order they were given.
|
||||
"""
|
||||
if not put_writes_accepts_task_path(put_writes):
|
||||
return {}
|
||||
return {"task_path": task_path_str(task.path)}
|
||||
|
||||
|
||||
_UpdateWrites = Sequence[tuple[str, Sequence[tuple[str, Any]], PregelTaskWrites | None]]
|
||||
|
||||
|
||||
def _delta_writes(
|
||||
channels: Mapping[str, BaseChannel | ManagedValueSpec], writes: _UpdateWrites
|
||||
) -> list[tuple[str, Any]]:
|
||||
"""Give the update's DeltaChannel writes message ids, as the loop's
|
||||
`put_writes` does, so every read of them returns the same ids."""
|
||||
delta = [
|
||||
(ch, value)
|
||||
for _, task_writes, _ in writes
|
||||
for ch, value in task_writes
|
||||
if isinstance(channels.get(ch), DeltaChannel)
|
||||
]
|
||||
for _, value in delta:
|
||||
ensure_message_ids(value)
|
||||
return delta
|
||||
|
||||
|
||||
def _store_or_fork_delta_writes(
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
config: RunnableConfig,
|
||||
saved: CheckpointTuple | None,
|
||||
checkpoint_config: RunnableConfig,
|
||||
channels: Mapping[str, BaseChannel | ManagedValueSpec],
|
||||
writes: _UpdateWrites,
|
||||
fork_pending: set[str],
|
||||
*,
|
||||
is_first: bool,
|
||||
) -> None:
|
||||
"""Save an update's writes where its DeltaChannels will read them.
|
||||
|
||||
A DeltaChannel rebuilds its value from the writes saved on a checkpoint's
|
||||
ancestors, so the writes go on the checkpoint the update builds on. If the
|
||||
thread already moved past that checkpoint, its other children would read
|
||||
them too, so the new checkpoint snapshots those channels instead. `writes`
|
||||
holds `(task_id, writes, task)` per task; `task` is `None` for input.
|
||||
"""
|
||||
delta = _delta_writes(channels, writes)
|
||||
if saved is None:
|
||||
return
|
||||
if is_first and checkpoint_superseded(checkpointer, config, saved):
|
||||
fork_pending.update(ch for ch, _ in delta)
|
||||
return
|
||||
for task_id, task_writes, task in writes:
|
||||
if task_writes:
|
||||
checkpointer.put_writes(
|
||||
checkpoint_config,
|
||||
task_writes,
|
||||
task_id,
|
||||
**(_task_path_kwarg(checkpointer.put_writes, task) if task else {}),
|
||||
)
|
||||
|
||||
|
||||
async def _astore_or_fork_delta_writes(
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
config: RunnableConfig,
|
||||
saved: CheckpointTuple | None,
|
||||
checkpoint_config: RunnableConfig,
|
||||
channels: Mapping[str, BaseChannel | ManagedValueSpec],
|
||||
writes: _UpdateWrites,
|
||||
fork_pending: set[str],
|
||||
*,
|
||||
is_first: bool,
|
||||
) -> None:
|
||||
"""Async `_store_or_fork_delta_writes`."""
|
||||
delta = _delta_writes(channels, writes)
|
||||
if saved is None:
|
||||
return
|
||||
if is_first and await acheckpoint_superseded(checkpointer, config, saved):
|
||||
fork_pending.update(ch for ch, _ in delta)
|
||||
return
|
||||
for task_id, task_writes, task in writes:
|
||||
if task_writes:
|
||||
await checkpointer.aput_writes(
|
||||
checkpoint_config,
|
||||
task_writes,
|
||||
task_id,
|
||||
**(_task_path_kwarg(checkpointer.aput_writes, task) if task else {}),
|
||||
)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
@@ -20,7 +20,7 @@ from warnings import warn
|
||||
from langchain_core.messages import AnyMessage
|
||||
from langchain_core.runnables import Runnable, RunnableConfig
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
|
||||
from pydantic import TypeAdapter
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import (
|
||||
NotRequired,
|
||||
TypeAliasType,
|
||||
@@ -882,6 +882,16 @@ class Command(Generic[N], ToolOutputMixin):
|
||||
PARENT: ClassVar[Literal["__parent__"]] = "__parent__"
|
||||
|
||||
|
||||
def _validate_resume(adapter: TypeAdapter[Any], value: Any) -> Any:
|
||||
from langgraph.errors import _mark_invalid_resume
|
||||
|
||||
try:
|
||||
return adapter.validate_python(value)
|
||||
except ValidationError as exc:
|
||||
_mark_invalid_resume(exc)
|
||||
raise
|
||||
|
||||
|
||||
@overload
|
||||
def interrupt(value: Any, *, response_schema: type[ResponseT]) -> ResponseT: ...
|
||||
|
||||
@@ -989,6 +999,7 @@ def interrupt(
|
||||
Raises:
|
||||
GraphInterrupt: On the first invocation within the node, halts execution and surfaces the provided value to the client.
|
||||
pydantic.ValidationError: When a resume value does not match a Pydantic model, `TypedDict`, or dataclass `response_schema`.
|
||||
Nothing is saved, so the interrupt can be answered again. `is_invalid_resume` identifies it.
|
||||
"""
|
||||
from langgraph._internal._constants import (
|
||||
CONFIG_KEY_CHECKPOINT_NS,
|
||||
@@ -1012,14 +1023,14 @@ def interrupt(
|
||||
if scratchpad.resume:
|
||||
if idx < len(scratchpad.resume):
|
||||
v = scratchpad.resume[idx]
|
||||
validated = adapter.validate_python(v) if adapter else v
|
||||
validated = _validate_resume(adapter, v) if adapter else v
|
||||
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume[: idx + 1])])
|
||||
return validated
|
||||
# find current resume value
|
||||
v = scratchpad.get_null_resume(True)
|
||||
if v is not None:
|
||||
assert len(scratchpad.resume) == idx, (scratchpad.resume, idx)
|
||||
validated = adapter.validate_python(v) if adapter else v
|
||||
validated = _validate_resume(adapter, v) if adapter else v
|
||||
scratchpad.resume.append(v)
|
||||
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume)])
|
||||
return validated
|
||||
|
||||
@@ -25,7 +25,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"langchain-core>=1.4.7,<2",
|
||||
"langgraph-checkpoint>=4.1.0,<5.0.0",
|
||||
"langgraph-checkpoint>=4.3.0,<5.0.0",
|
||||
"langgraph-sdk>=0.4.6,<0.5.0",
|
||||
"langgraph-prebuilt>=1.1.0,<1.2.0",
|
||||
"xxhash>=3.5.0",
|
||||
|
||||
@@ -6,19 +6,23 @@ 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
|
||||
|
||||
from langgraph._internal._constants import NULL_TASK_ID
|
||||
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
|
||||
|
||||
@@ -35,6 +39,7 @@ def test_exit_delta_task_id_is_valid_uuid_and_ordered() -> None:
|
||||
assert id1.split("-")[0] == "00000001"
|
||||
assert id7.split("-")[0] == "00000007"
|
||||
assert id1.endswith("-0270-bf16-1ef8-fb321bef9f3d")
|
||||
assert exit_delta_task_id(0, NULL_TASK_ID) != NULL_TASK_ID
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
uuid.UUID(f"00000001-{tid}")
|
||||
@@ -389,3 +394,234 @@ 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_with_a_command_update_replays_its_write_once(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("done", lambda state: _both("done"))
|
||||
builder.add_node("ask", _ask("ask"))
|
||||
builder.add_edge(START, "done")
|
||||
builder.add_edge(START, "ask")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability=durability)
|
||||
|
||||
graph.invoke(
|
||||
Command(resume="yes", update=_both("cmd")), config, durability=durability
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.values["log"] == state.values["plain"] == ["in", "cmd", "ask", "done"]
|
||||
|
||||
|
||||
def test_command_update_on_an_input_checkpoint_matches_a_plain_channel(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("node", lambda state: _both("node"))
|
||||
builder.add_edge(START, "node")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.update_state(config, _both("in"), as_node="__input__")
|
||||
|
||||
graph.invoke(Command(update=_both("cmd")), config, durability=durability)
|
||||
|
||||
history = list(graph.get_state_history(config))
|
||||
assert [s.values.get("log", []) for s in history] == [
|
||||
s.values.get("plain", []) for s in history
|
||||
]
|
||||
replayed = graph.invoke(None, history[-1].config, durability=durability)
|
||||
assert replayed["log"] == replayed["plain"]
|
||||
|
||||
|
||||
def test_exit_command_update_on_a_new_thread_matches_a_plain_channel(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("node", lambda state: _both("node"))
|
||||
builder.add_edge(START, "node")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
graph.invoke(Command(update=_both("cmd")), config, durability="exit")
|
||||
|
||||
history = list(graph.get_state_history(config))
|
||||
assert [s.values.get("log", []) for s in history] == [
|
||||
s.values.get("plain", []) for s in history
|
||||
]
|
||||
|
||||
|
||||
class _FlagState(_ResumeState, total=False):
|
||||
extra: Annotated[list, DeltaChannel(_append)]
|
||||
flag: bool
|
||||
|
||||
|
||||
def test_addressed_resume_keeps_a_rerun_tasks_write_to_a_new_channel(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
def done(state: _FlagState) -> dict:
|
||||
return {**_both("done"), **({"extra": ["new"]} if state.get("flag") else {})}
|
||||
|
||||
builder = StateGraph(_FlagState)
|
||||
builder.add_node("done", done)
|
||||
builder.add_node("ask", _ask("ask"))
|
||||
builder.add_edge(START, "done")
|
||||
builder.add_edge(START, "ask")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability=durability)
|
||||
head = graph.get_state(config).config
|
||||
|
||||
live = graph.invoke(
|
||||
Command(resume="yes", update={"flag": True}), head, durability=durability
|
||||
)
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert live["extra"] == state.values["extra"] == ["new"]
|
||||
assert state.values["log"] == state.values["plain"]
|
||||
|
||||
|
||||
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"]
|
||||
|
||||
|
||||
class _FailingPutSaver(InMemorySaver):
|
||||
fail = False
|
||||
|
||||
def put(
|
||||
self, config: Any, checkpoint: Any, metadata: Any, new_versions: Any
|
||||
) -> Any:
|
||||
if self.fail:
|
||||
raise RuntimeError("final checkpoint lost")
|
||||
return super().put(config, checkpoint, metadata, new_versions)
|
||||
|
||||
|
||||
def test_exit_resume_retried_after_its_final_checkpoint_fails_reruns_the_resumed_task() -> (
|
||||
None
|
||||
):
|
||||
saver = _FailingPutSaver()
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("done", lambda state: _both("done"))
|
||||
builder.add_node("ask", _ask("ask"))
|
||||
builder.add_edge(START, "done")
|
||||
builder.add_edge(START, "ask")
|
||||
graph = builder.compile(checkpointer=saver)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke(_both("in"), config, durability="exit")
|
||||
saver.fail = True
|
||||
with pytest.raises(RuntimeError, match="final checkpoint lost"):
|
||||
graph.invoke(Command(resume="yes"), config, durability="exit")
|
||||
saver.fail = False
|
||||
|
||||
graph.invoke(Command(resume="yes"), config, durability="exit")
|
||||
|
||||
state = graph.get_state(config)
|
||||
assert state.values["log"] == state.values["plain"]
|
||||
assert sorted(state.values["plain"]) == ["ask", "done", "in"]
|
||||
|
||||
@@ -740,20 +740,6 @@ def _build_deferred_after_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
|
||||
return builder.compile(checkpointer=checkpointer, interrupt_after=["a"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"durability",
|
||||
[
|
||||
"sync",
|
||||
"async",
|
||||
pytest.param(
|
||||
"exit",
|
||||
marks=pytest.mark.xfail(
|
||||
reason="exit durability stores a resumed run's loaded writes twice",
|
||||
strict=True,
|
||||
),
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_resume_on_an_interrupted_head_consumes_its_writes_without_a_snapshot(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
"""`DeltaChannel` replay must apply parallel writes in the order `invoke` did."""
|
||||
|
||||
import asyncio
|
||||
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_sync_get_state_on_async_saver_matches_live_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 asyncio.to_thread(graph.get_state, config)).values["items"]
|
||||
|
||||
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}"
|
||||
@@ -0,0 +1,103 @@
|
||||
"""A run's input to a DeltaChannel input channel reads back on the checkpoints
|
||||
built from it, and on no others."""
|
||||
|
||||
import pytest
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
|
||||
from langgraph.channels.binop import BinaryOperatorAggregate
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.pregel import NodeBuilder, Pregel
|
||||
from langgraph.types import Durability
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
def _sorted_extend(current: list, writes: list) -> list:
|
||||
return sorted([*current, *(item for write in writes for item in write)])
|
||||
|
||||
|
||||
def _delta_input_graph() -> Pregel:
|
||||
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log", "plain")
|
||||
return Pregel(
|
||||
nodes={"n": node},
|
||||
channels={
|
||||
"log": DeltaChannel(_sorted_extend),
|
||||
"plain": BinaryOperatorAggregate(list, lambda a, b: sorted(a + b)),
|
||||
"go": LastValue(int),
|
||||
},
|
||||
input_channels=["log", "plain", "go"],
|
||||
output_channels=["log", "plain"],
|
||||
checkpointer=InMemorySaver(),
|
||||
)
|
||||
|
||||
|
||||
def test_each_run_input_reads_back_on_its_own_checkpoints(
|
||||
durability: Durability,
|
||||
) -> None:
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
graph.invoke({"log": [0], "plain": [0], "go": 1}, config, durability=durability)
|
||||
graph.invoke({"log": [5], "plain": [5], "go": 1}, config, durability=durability)
|
||||
|
||||
for state in graph.get_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
|
||||
async def test_each_run_input_reads_back_on_its_own_checkpoints_async(
|
||||
durability: Durability,
|
||||
) -> None:
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
await graph.ainvoke(
|
||||
{"log": [0], "plain": [0], "go": 1}, config, durability=durability
|
||||
)
|
||||
await graph.ainvoke(
|
||||
{"log": [5], "plain": [5], "go": 1}, config, durability=durability
|
||||
)
|
||||
|
||||
async for state in graph.aget_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
|
||||
OTHER_BRANCH_INPUTS = pytest.mark.parametrize(
|
||||
"other_branch_input",
|
||||
[{"go": 1}, {"log": [5], "plain": [5], "go": 1}],
|
||||
ids=["other-branch-without-delta-input", "other-branch-with-delta-input"],
|
||||
)
|
||||
|
||||
|
||||
@OTHER_BRANCH_INPUTS
|
||||
def test_run_input_from_an_older_checkpoint_stays_out_of_its_other_branch(
|
||||
durability: Durability, other_branch_input: dict
|
||||
) -> None:
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"go": 1}, config, durability=durability)
|
||||
older = graph.get_state(config).config
|
||||
graph.invoke(other_branch_input, config, durability=durability)
|
||||
|
||||
graph.invoke({"log": [7], "plain": [7], "go": 1}, older, durability=durability)
|
||||
|
||||
for state in graph.get_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
|
||||
@OTHER_BRANCH_INPUTS
|
||||
async def test_run_input_from_an_older_checkpoint_stays_out_of_its_other_branch_async(
|
||||
durability: Durability, other_branch_input: dict
|
||||
) -> None:
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"go": 1}, config, durability=durability)
|
||||
older = (await graph.aget_state(config)).config
|
||||
await graph.ainvoke(other_branch_input, config, durability=durability)
|
||||
|
||||
await graph.ainvoke(
|
||||
{"log": [7], "plain": [7], "go": 1}, older, durability=durability
|
||||
)
|
||||
|
||||
async for state in graph.aget_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
@@ -0,0 +1,78 @@
|
||||
import pytest
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.pregel import NodeBuilder, Pregel
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
CONFIG = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
|
||||
def _extend(current: list, writes: list) -> list:
|
||||
return [*current, *(item for write in writes for item in write)]
|
||||
|
||||
|
||||
def _graph(reads: list) -> Pregel:
|
||||
writer = NodeBuilder().subscribe_only("a").do(lambda _: [1]).write_to("d")
|
||||
reader = NodeBuilder().subscribe_only("d").do(lambda d: reads.append(list(d)))
|
||||
return Pregel(
|
||||
nodes={"writer": writer, "reader": reader},
|
||||
channels={"a": LastValue(str), "d": DeltaChannel(_extend)},
|
||||
input_channels="a",
|
||||
output_channels=["d"],
|
||||
checkpointer=InMemorySaver(),
|
||||
)
|
||||
|
||||
|
||||
def test_an_input_update_from_before_the_first_write_starts_only_the_writer() -> None:
|
||||
reads: list = []
|
||||
graph = _graph(reads)
|
||||
graph.invoke("go", CONFIG)
|
||||
first = next(
|
||||
s.config for s in graph.get_state_history(CONFIG) if s.metadata["step"] == -1
|
||||
)
|
||||
|
||||
fork = graph.update_state(first, {"a": "go"}, as_node="__input__")
|
||||
|
||||
assert graph.get_state(fork).next == ("writer",)
|
||||
graph.invoke(None, fork)
|
||||
assert reads == [[1], [1]]
|
||||
|
||||
|
||||
def test_a_replay_from_before_the_first_write_forks_with_only_the_writer_next() -> None:
|
||||
reads: list = []
|
||||
graph = _graph(reads)
|
||||
graph.invoke("go", CONFIG)
|
||||
first = next(
|
||||
s.config for s in graph.get_state_history(CONFIG) if s.metadata["step"] == -1
|
||||
)
|
||||
|
||||
graph.invoke(None, first, durability="sync")
|
||||
|
||||
fork = next(
|
||||
s for s in graph.get_state_history(CONFIG) if s.metadata["source"] == "fork"
|
||||
)
|
||||
assert fork.next == ("writer",)
|
||||
graph.invoke(None, fork.config)
|
||||
assert reads == [[1], [1], [1]]
|
||||
|
||||
|
||||
async def test_an_ainput_update_from_before_the_first_write_starts_only_the_writer() -> (
|
||||
None
|
||||
):
|
||||
reads: list = []
|
||||
graph = _graph(reads)
|
||||
await graph.ainvoke("go", CONFIG)
|
||||
first = [
|
||||
s.config
|
||||
async for s in graph.aget_state_history(CONFIG)
|
||||
if s.metadata["step"] == -1
|
||||
][0]
|
||||
|
||||
fork = await graph.aupdate_state(first, {"a": "go"}, as_node="__input__")
|
||||
|
||||
assert (await graph.aget_state(fork)).next == ("writer",)
|
||||
await graph.ainvoke(None, fork)
|
||||
assert reads == [[1], [1]]
|
||||
@@ -0,0 +1,110 @@
|
||||
"""An `Overwrite` through `update_state` snapshots its DeltaChannel on the
|
||||
checkpoint the update saves, as a node's `Overwrite` does on the loop's."""
|
||||
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langchain_core.messages import HumanMessage
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.graph.message import _messages_delta_reducer
|
||||
from langgraph.pregel import NodeBuilder, Pregel
|
||||
from langgraph.types import Overwrite
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
class _State(TypedDict):
|
||||
messages: Annotated[list, DeltaChannel(_messages_delta_reducer)]
|
||||
|
||||
|
||||
def _messages_graph(saver: InMemorySaver) -> Any:
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("model", lambda state: {})
|
||||
builder.add_edge(START, "model")
|
||||
return builder.compile(checkpointer=saver)
|
||||
|
||||
|
||||
def _extend(current: list, writes: list) -> list:
|
||||
return [*current, *(item for write in writes for item in write)]
|
||||
|
||||
|
||||
def _delta_input_graph(saver: InMemorySaver) -> Pregel:
|
||||
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log")
|
||||
return Pregel(
|
||||
nodes={"n": node},
|
||||
channels={"log": DeltaChannel(_extend), "go": LastValue(int)},
|
||||
input_channels=["log", "go"],
|
||||
output_channels=["log"],
|
||||
checkpointer=saver,
|
||||
)
|
||||
|
||||
|
||||
def test_update_state_with_an_overwrite_snapshots_the_channel() -> None:
|
||||
saver = InMemorySaver()
|
||||
graph = _messages_graph(saver)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="a", id="1")]}, config)
|
||||
|
||||
graph.update_state(
|
||||
config,
|
||||
{"messages": Overwrite([HumanMessage(content="b", id="2")])},
|
||||
as_node="model",
|
||||
)
|
||||
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
|
||||
assert [m.content for m in graph.get_state(config).values["messages"]] == ["b"]
|
||||
|
||||
|
||||
async def test_aupdate_state_with_an_overwrite_snapshots_the_channel() -> None:
|
||||
saver = InMemorySaver()
|
||||
graph = _messages_graph(saver)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage(content="a", id="1")]}, config)
|
||||
|
||||
await graph.aupdate_state(
|
||||
config,
|
||||
{"messages": Overwrite([HumanMessage(content="b", id="2")])},
|
||||
as_node="model",
|
||||
)
|
||||
|
||||
head = await saver.aget_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
|
||||
values = (await graph.aget_state(config)).values
|
||||
assert [m.content for m in values["messages"]] == ["b"]
|
||||
|
||||
|
||||
def test_update_state_as_input_with_an_overwrite_snapshots_the_channel() -> None:
|
||||
saver = InMemorySaver()
|
||||
graph = _delta_input_graph(saver)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"log": [0], "go": 1}, config)
|
||||
|
||||
graph.update_state(config, {"log": Overwrite([1])}, as_node="__input__")
|
||||
|
||||
head = saver.get_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
|
||||
assert graph.get_state(config).values["log"] == [1]
|
||||
|
||||
|
||||
async def test_aupdate_state_as_input_with_an_overwrite_snapshots_the_channel() -> None:
|
||||
saver = InMemorySaver()
|
||||
graph = _delta_input_graph(saver)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"log": [0], "go": 1}, config)
|
||||
|
||||
await graph.aupdate_state(config, {"log": Overwrite([1])}, as_node="__input__")
|
||||
|
||||
head = await saver.aget_tuple(config)
|
||||
assert head is not None
|
||||
assert isinstance(head.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
|
||||
assert (await graph.aget_state(config)).values["log"] == [1]
|
||||
@@ -20,23 +20,28 @@ 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
|
||||
|
||||
from langgraph.channels.binop import BinaryOperatorAggregate
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.graph.message import _messages_delta_reducer
|
||||
from langgraph.types import StateUpdate
|
||||
from langgraph.pregel import NodeBuilder, Pregel
|
||||
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 +68,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)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -304,15 +309,14 @@ 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()
|
||||
graph = _build_graph(saver)
|
||||
config = {"configurable": {"thread_id": "bulk-multi-task"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
base = saver.get_tuple(config)
|
||||
assert base is not None
|
||||
|
||||
graph.bulk_update_state(
|
||||
config,
|
||||
@@ -332,13 +336,155 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
|
||||
],
|
||||
)
|
||||
|
||||
stored = saver.get_tuple(base.config)
|
||||
assert stored is not None
|
||||
assert {task_id for task_id, _, _ in stored.pending_writes or []} == {
|
||||
"task-1",
|
||||
"task-2",
|
||||
}, "explicit task ids must key the stored writes"
|
||||
state = graph.get_state(config)
|
||||
contents = [m.content for m in state.values["messages"]]
|
||||
ids = [m.id for m in state.values["messages"]]
|
||||
assert sorted(contents) == ["first", "second"], (
|
||||
assert sorted(contents) == ["first", "hi", "second"], (
|
||||
f"both updates' writes must persist; got {contents}"
|
||||
)
|
||||
assert sorted(ids) == ["m1", "m2"]
|
||||
assert sorted(ids) == ["hi", "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}"
|
||||
)
|
||||
|
||||
|
||||
class _TaskPathOrderSaver(InMemorySaver):
|
||||
"""Replays each checkpoint's writes by `(task_path, task_id, idx)`."""
|
||||
|
||||
def get_tuple(self, config: Any) -> Any:
|
||||
tup = super().get_tuple(config)
|
||||
if tup is None or not tup.pending_writes:
|
||||
return tup
|
||||
conf = tup.config["configurable"]
|
||||
stored = self.writes[
|
||||
(conf["thread_id"], conf["checkpoint_ns"], conf["checkpoint_id"])
|
||||
]
|
||||
rows = sorted(
|
||||
zip(stored.items(), tup.pending_writes),
|
||||
key=lambda row: (row[0][1][3], *row[0][0]),
|
||||
)
|
||||
return tup._replace(pending_writes=[write for _, write in rows])
|
||||
|
||||
get_delta_channel_history = BaseCheckpointSaver.get_delta_channel_history
|
||||
aget_delta_channel_history = BaseCheckpointSaver.aget_delta_channel_history
|
||||
|
||||
|
||||
GIVEN = ["u1", "u2", "u3", "u4", "u5", "u6"]
|
||||
|
||||
|
||||
def _updates_in_given_order() -> list[list[StateUpdate]]:
|
||||
return [
|
||||
[_update(c, "assistant" if i % 2 else "model") for i, c in enumerate(GIVEN)]
|
||||
]
|
||||
|
||||
|
||||
def test_bulk_update_state_replays_updates_in_the_order_given() -> None:
|
||||
graph = _build_graph(_TaskPathOrderSaver(), two_nodes=True)
|
||||
config = {"configurable": {"thread_id": "bulk-order"}}
|
||||
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
|
||||
graph.bulk_update_state(config, _updates_in_given_order())
|
||||
|
||||
assert _contents(graph.get_state(config)) == ["hi", *GIVEN]
|
||||
|
||||
|
||||
async def test_abulk_update_state_replays_updates_in_the_order_given() -> None:
|
||||
graph = _build_graph(_TaskPathOrderSaver(), two_nodes=True)
|
||||
config = {"configurable": {"thread_id": "bulk-order"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
|
||||
|
||||
await graph.abulk_update_state(config, _updates_in_given_order())
|
||||
|
||||
assert _contents(await graph.aget_state(config)) == ["hi", *GIVEN]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -395,3 +541,189 @@ def test_update_state_that_snapshots_keeps_a_deferred_node_pending() -> None:
|
||||
|
||||
assert [m.content for m in final["messages"]] == ["s", "a", "u", "b"]
|
||||
assert graph.get_state(config).next == ()
|
||||
|
||||
|
||||
def _sorted_extend(current: list, writes: list) -> list:
|
||||
return sorted([*current, *(item for write in writes for item in write)])
|
||||
|
||||
|
||||
def _delta_input_graph(snapshot_frequency: int = 1000) -> Any:
|
||||
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log", "plain")
|
||||
return Pregel(
|
||||
nodes={"n": node},
|
||||
channels={
|
||||
"log": DeltaChannel(_sorted_extend, snapshot_frequency=snapshot_frequency),
|
||||
"plain": BinaryOperatorAggregate(list, lambda a, b: sorted(a + b)),
|
||||
"go": LastValue(int),
|
||||
},
|
||||
input_channels=["log", "plain", "go"],
|
||||
output_channels=["log", "plain"],
|
||||
checkpointer=InMemorySaver(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("snapshot_frequency", [1, 2])
|
||||
def test_update_as_input_reads_back_on_its_checkpoint_and_after_the_next_run(
|
||||
snapshot_frequency: int,
|
||||
) -> None:
|
||||
graph = _delta_input_graph(snapshot_frequency)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"log": [0], "plain": [0], "go": 1}, config)
|
||||
|
||||
graph.update_state(config, {"log": [1], "plain": [1], "go": 1}, as_node="__input__")
|
||||
after_update = graph.get_state(config).values
|
||||
graph.invoke(None, config)
|
||||
after_run = graph.get_state(config).values
|
||||
|
||||
assert after_update["log"] == after_update["plain"]
|
||||
assert after_run["log"] == after_run["plain"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("snapshot_frequency", [1, 2])
|
||||
async def test_aupdate_as_input_reads_back_on_its_checkpoint_and_after_the_next_run(
|
||||
snapshot_frequency: int,
|
||||
) -> None:
|
||||
graph = _delta_input_graph(snapshot_frequency)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"log": [0], "plain": [0], "go": 1}, config)
|
||||
|
||||
await graph.aupdate_state(
|
||||
config, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
|
||||
)
|
||||
after_update = (await graph.aget_state(config)).values
|
||||
await graph.ainvoke(None, config)
|
||||
after_run = (await graph.aget_state(config)).values
|
||||
|
||||
assert after_update["log"] == after_update["plain"]
|
||||
assert after_run["log"] == after_run["plain"]
|
||||
|
||||
|
||||
def test_update_as_input_to_an_older_checkpoint_stays_out_of_its_other_branch() -> None:
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"go": 1}, config)
|
||||
older = graph.get_state(config).config
|
||||
graph.invoke({"go": 1}, config)
|
||||
other_branch = graph.get_state(config)
|
||||
|
||||
edited = graph.update_state(
|
||||
older, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
|
||||
)
|
||||
|
||||
values = graph.get_state(edited).values
|
||||
assert values["log"] == values["plain"]
|
||||
assert graph.get_state(other_branch.config).values == other_branch.values
|
||||
|
||||
|
||||
async def test_aupdate_as_input_to_an_older_checkpoint_stays_out_of_its_other_branch() -> (
|
||||
None
|
||||
):
|
||||
graph = _delta_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"go": 1}, config)
|
||||
older = (await graph.aget_state(config)).config
|
||||
await graph.ainvoke({"go": 1}, config)
|
||||
other_branch = await graph.aget_state(config)
|
||||
|
||||
edited = await graph.aupdate_state(
|
||||
older, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
|
||||
)
|
||||
|
||||
values = (await graph.aget_state(edited)).values
|
||||
assert values["log"] == values["plain"]
|
||||
assert (await graph.aget_state(other_branch.config)).values == other_branch.values
|
||||
|
||||
|
||||
def _message_ids(graph: Any, config: dict) -> list[str | None]:
|
||||
return [m.id for m in graph.get_state(config).values["messages"]]
|
||||
|
||||
|
||||
async def _amessage_ids(graph: Any, config: dict) -> list[str | None]:
|
||||
return [m.id for m in (await graph.aget_state(config)).values["messages"]]
|
||||
|
||||
|
||||
def _messages_input_graph() -> Any:
|
||||
node = (
|
||||
NodeBuilder()
|
||||
.subscribe_only("go")
|
||||
.do(lambda _: [HumanMessage("n", id="n")])
|
||||
.write_to("messages")
|
||||
)
|
||||
return Pregel(
|
||||
nodes={"n": node},
|
||||
channels={
|
||||
"messages": DeltaChannel(_messages_delta_reducer),
|
||||
"go": LastValue(int),
|
||||
},
|
||||
input_channels=["messages", "go"],
|
||||
output_channels=["messages"],
|
||||
checkpointer=InMemorySaver(),
|
||||
)
|
||||
|
||||
|
||||
def test_update_state_gives_a_message_an_id_that_every_read_keeps() -> None:
|
||||
graph = _build_graph(InMemorySaver())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"messages": [HumanMessage("a", id="a")]}, config)
|
||||
|
||||
graph.update_state(config, {"messages": [HumanMessage("b")]})
|
||||
|
||||
first, second = _message_ids(graph, config), _message_ids(graph, config)
|
||||
assert first[-1] is not None
|
||||
assert first == second
|
||||
|
||||
|
||||
async def test_aupdate_state_gives_a_message_an_id_that_every_read_keeps() -> None:
|
||||
graph = _build_graph(InMemorySaver())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage("a", id="a")]}, config)
|
||||
|
||||
await graph.aupdate_state(config, {"messages": [HumanMessage("b")]})
|
||||
|
||||
first, second = (
|
||||
await _amessage_ids(graph, config),
|
||||
await _amessage_ids(graph, config),
|
||||
)
|
||||
assert first[-1] is not None
|
||||
assert first == second
|
||||
|
||||
|
||||
def test_update_state_on_an_older_checkpoint_gives_a_message_an_id() -> None:
|
||||
graph = _build_graph(InMemorySaver())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"messages": [HumanMessage("a", id="a")]}, config)
|
||||
older = graph.get_state(config).config
|
||||
graph.invoke({"messages": [HumanMessage("c", id="c")]}, config)
|
||||
|
||||
branch = graph.update_state(older, {"messages": [HumanMessage("b")]})
|
||||
|
||||
assert _message_ids(graph, branch)[-1] is not None
|
||||
|
||||
|
||||
def test_update_as_input_gives_a_message_an_id_that_every_read_keeps() -> None:
|
||||
graph = _messages_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.invoke({"messages": [HumanMessage("a", id="a")], "go": 1}, config)
|
||||
|
||||
graph.update_state(config, {"messages": [HumanMessage("b")]}, as_node="__input__")
|
||||
|
||||
first, second = _message_ids(graph, config), _message_ids(graph, config)
|
||||
assert first[-1] is not None
|
||||
assert first == second
|
||||
|
||||
|
||||
async def test_aupdate_as_input_gives_a_message_an_id_that_every_read_keeps() -> None:
|
||||
graph = _messages_input_graph()
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
await graph.ainvoke({"messages": [HumanMessage("a", id="a")], "go": 1}, config)
|
||||
|
||||
await graph.aupdate_state(
|
||||
config, {"messages": [HumanMessage("b")]}, as_node="__input__"
|
||||
)
|
||||
|
||||
first, second = (
|
||||
await _amessage_ids(graph, config),
|
||||
await _amessage_ids(graph, config),
|
||||
)
|
||||
assert first[-1] is not None
|
||||
assert first == second
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
"""A checkpoint must never be saved without the `DeltaChannel` writes it reads."""
|
||||
|
||||
import operator
|
||||
import threading
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.types import Durability
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
INPUT = {"log": [], "plain": []}
|
||||
FINAL = {"log": ["a", "b", "c"], "plain": ["a", "b", "c"]}
|
||||
|
||||
# Exit mode saves nothing before the failed write, so its retry starts over.
|
||||
RETRIES = [
|
||||
pytest.param("sync", None, id="sync"),
|
||||
pytest.param("async", None, id="async"),
|
||||
pytest.param("exit", INPUT, id="exit"),
|
||||
]
|
||||
|
||||
|
||||
def _append(current: list, writes: list) -> list:
|
||||
return [*current, *(item for write in writes for item in write)]
|
||||
|
||||
|
||||
class _State(TypedDict):
|
||||
log: Annotated[list, DeltaChannel(_append)]
|
||||
plain: Annotated[list, operator.add]
|
||||
|
||||
|
||||
class _FailsTheWriteOfBOnce(InMemorySaver):
|
||||
failed = False
|
||||
|
||||
def _fail_once(self, writes: Any) -> None:
|
||||
if not self.failed and ("log", ["b"]) in writes:
|
||||
self.failed = True
|
||||
raise ConnectionError("b's write was not saved")
|
||||
|
||||
def put_writes(
|
||||
self, config: Any, writes: Any, task_id: str, task_path: str = ""
|
||||
) -> None:
|
||||
self._fail_once(writes)
|
||||
super().put_writes(config, writes, task_id, task_path)
|
||||
|
||||
async def aput_writes(
|
||||
self, config: Any, writes: Any, task_id: str, task_path: str = ""
|
||||
) -> None:
|
||||
self._fail_once(writes)
|
||||
await super().aput_writes(config, writes, task_id, task_path)
|
||||
|
||||
|
||||
class _FailsTheSaveAfterAOnce(InMemorySaver):
|
||||
failed = False
|
||||
|
||||
def _fail_once(self, checkpoint: Any) -> None:
|
||||
if not self.failed and checkpoint["channel_values"].get("plain") == ["a"]:
|
||||
self.failed = True
|
||||
raise ConnectionError("the checkpoint after a was not saved")
|
||||
|
||||
def put(
|
||||
self, config: Any, checkpoint: Any, metadata: Any, new_versions: Any
|
||||
) -> Any:
|
||||
self._fail_once(checkpoint)
|
||||
return super().put(config, checkpoint, metadata, new_versions)
|
||||
|
||||
async def aput(
|
||||
self, config: Any, checkpoint: Any, metadata: Any, new_versions: Any
|
||||
) -> Any:
|
||||
self._fail_once(checkpoint)
|
||||
return await super().aput(config, checkpoint, metadata, new_versions)
|
||||
|
||||
|
||||
def _a_then_b_then_c(saver: InMemorySaver) -> Any:
|
||||
builder = StateGraph(_State)
|
||||
for name in "abc":
|
||||
builder.add_node(
|
||||
name, lambda state, name=name: {"log": [name], "plain": [name]}
|
||||
)
|
||||
builder.add_edge(START, "a")
|
||||
builder.add_edge("a", "b")
|
||||
builder.add_edge("b", "c")
|
||||
return builder.compile(checkpointer=saver)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("durability", "retry_input"), RETRIES)
|
||||
def test_a_failed_delta_write_is_rerun_not_lost(
|
||||
durability: Durability, retry_input: dict | None
|
||||
) -> None:
|
||||
graph = _a_then_b_then_c(_FailsTheWriteOfBOnce())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
graph.invoke(INPUT, config, durability=durability)
|
||||
for state in graph.get_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
graph.invoke(retry_input, config, durability=durability)
|
||||
assert graph.get_state(config).values == FINAL
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("durability", "retry_input"), RETRIES)
|
||||
async def test_a_failed_delta_write_is_rerun_not_lost_async(
|
||||
durability: Durability, retry_input: dict | None
|
||||
) -> None:
|
||||
graph = _a_then_b_then_c(_FailsTheWriteOfBOnce())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
await graph.ainvoke(INPUT, config, durability=durability)
|
||||
async for state in graph.aget_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
await graph.ainvoke(retry_input, config, durability=durability)
|
||||
assert (await graph.aget_state(config)).values == FINAL
|
||||
|
||||
|
||||
@pytest.mark.parametrize("durability", ["sync", "async"])
|
||||
def test_a_failed_checkpoint_save_is_rerun_not_built_on(
|
||||
durability: Durability,
|
||||
) -> None:
|
||||
graph = _a_then_b_then_c(_FailsTheSaveAfterAOnce())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
graph.invoke(INPUT, config, durability=durability)
|
||||
for state in graph.get_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
graph.invoke(None, config, durability=durability)
|
||||
assert graph.get_state(config).values == FINAL
|
||||
|
||||
|
||||
@pytest.mark.parametrize("durability", ["sync", "async"])
|
||||
async def test_a_failed_checkpoint_save_is_rerun_not_built_on_async(
|
||||
durability: Durability,
|
||||
) -> None:
|
||||
graph = _a_then_b_then_c(_FailsTheSaveAfterAOnce())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
await graph.ainvoke(INPUT, config, durability=durability)
|
||||
async for state in graph.aget_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
await graph.ainvoke(None, config, durability=durability)
|
||||
assert (await graph.aget_state(config)).values == FINAL
|
||||
|
||||
|
||||
def test_a_delta_graph_finishes_on_a_single_background_thread() -> None:
|
||||
graph = _a_then_b_then_c(InMemorySaver())
|
||||
config = {"configurable": {"thread_id": "t"}, "max_concurrency": 1}
|
||||
result: dict = {}
|
||||
run = threading.Thread(
|
||||
target=lambda: result.update(graph.invoke(INPUT, config, durability="async")),
|
||||
daemon=True,
|
||||
)
|
||||
|
||||
run.start()
|
||||
run.join(timeout=10)
|
||||
|
||||
assert not run.is_alive(), "invoke hung"
|
||||
assert result == FINAL
|
||||
@@ -6,6 +6,7 @@ from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.errors import is_invalid_resume
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
from langgraph.types import Command, Durability, Interrupt, interrupt
|
||||
from tests.any_str import AnyStr
|
||||
@@ -202,14 +203,22 @@ def test_interrupt_response_schema_rejects_invalid_resume(
|
||||
def resume(value: dict[str, Any]) -> Command:
|
||||
return Command(resume=value if resume_style == "null" else {pending.id: value})
|
||||
|
||||
with pytest.raises(ValidationError, match="approved"):
|
||||
with pytest.raises(ValidationError, match="approved") as exc_info:
|
||||
graph.invoke(resume({"approved": "nope"}), config)
|
||||
assert is_invalid_resume(exc_info.value)
|
||||
|
||||
assert graph.invoke(resume({"approved": False}), config) == {
|
||||
"answer": Decision(approved=False)
|
||||
}
|
||||
|
||||
|
||||
def test_is_invalid_resume_ignores_other_errors() -> None:
|
||||
with pytest.raises(ValidationError) as exc_info:
|
||||
Decision.model_validate({"approved": "nope"})
|
||||
assert not is_invalid_resume(exc_info.value)
|
||||
assert not is_invalid_resume(ValueError("nope"))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("resume_style", ["null", "id_map"])
|
||||
def test_interrupt_response_schema_invalid_resume_after_earlier_interrupt(
|
||||
sync_checkpointer: BaseCheckpointSaver, resume_style: str
|
||||
|
||||
@@ -248,6 +248,95 @@ def test_drain_from_subgraph_can_resume_parent() -> None:
|
||||
}
|
||||
|
||||
|
||||
class _BreakpointState(TypedDict, total=False):
|
||||
first: str
|
||||
second: str
|
||||
|
||||
|
||||
def _drain_before_breakpoint_builder(
|
||||
second_runs: list[str],
|
||||
) -> StateGraph[_BreakpointState]:
|
||||
def first(state: _BreakpointState, runtime: Runtime) -> _BreakpointState:
|
||||
runtime.control.request_drain("rollout")
|
||||
return {"first": "done"}
|
||||
|
||||
def second(state: _BreakpointState) -> _BreakpointState:
|
||||
second_runs.append("second")
|
||||
return {"second": "done"}
|
||||
|
||||
builder = StateGraph(_BreakpointState)
|
||||
builder.add_node("first", first)
|
||||
builder.add_node("second", second)
|
||||
builder.add_edge(START, "first")
|
||||
builder.add_edge("first", "second")
|
||||
return builder
|
||||
|
||||
|
||||
@pytest.mark.parametrize("at_compile", [True, False])
|
||||
def test_drain_before_interrupt_before_stops_at_breakpoint(at_compile: bool) -> None:
|
||||
# A resume passes the next step's breakpoints, so a drain right before a
|
||||
# breakpoint has to stop there instead, or the resumed run skips it.
|
||||
second_runs: list[str] = []
|
||||
breakpoints = {"interrupt_before": ["second"]}
|
||||
compiled = _drain_before_breakpoint_builder(second_runs).compile(
|
||||
checkpointer=MemorySaver(), **(breakpoints if at_compile else {})
|
||||
)
|
||||
run_kwargs = {} if at_compile else breakpoints
|
||||
config = {"configurable": {"thread_id": "drain-breakpoint"}}
|
||||
|
||||
assert compiled.invoke({}, config, control=RunControl(), **run_kwargs) == {
|
||||
"first": "done"
|
||||
}
|
||||
assert compiled.get_state(config).next == ("second",)
|
||||
assert second_runs == []
|
||||
|
||||
assert compiled.invoke(None, config, **run_kwargs) == {
|
||||
"first": "done",
|
||||
"second": "done",
|
||||
}
|
||||
assert second_runs == ["second"]
|
||||
|
||||
|
||||
def test_drain_before_subgraph_interrupt_before_stops_at_breakpoint() -> None:
|
||||
second_runs: list[str] = []
|
||||
child = _drain_before_breakpoint_builder(second_runs).compile(
|
||||
interrupt_before=["second"]
|
||||
)
|
||||
parent = StateGraph(_BreakpointState)
|
||||
parent.add_node("child", child)
|
||||
parent.add_edge(START, "child")
|
||||
compiled = parent.compile(checkpointer=MemorySaver())
|
||||
config = {"configurable": {"thread_id": "drain-subgraph-breakpoint"}}
|
||||
|
||||
compiled.invoke({}, config, control=RunControl())
|
||||
state = compiled.get_state(config, subgraphs=True)
|
||||
assert state.next == ("child",)
|
||||
assert state.tasks[0].state.next == ("second",)
|
||||
assert second_runs == []
|
||||
|
||||
assert compiled.invoke(None, config) == {"first": "done", "second": "done"}
|
||||
assert second_runs == ["second"]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_drain_before_interrupt_before_stops_at_breakpoint_async() -> None:
|
||||
second_runs: list[str] = []
|
||||
compiled = _drain_before_breakpoint_builder(second_runs).compile(
|
||||
checkpointer=MemorySaver(), interrupt_before=["second"]
|
||||
)
|
||||
config = {"configurable": {"thread_id": "drain-breakpoint-async"}}
|
||||
|
||||
assert await compiled.ainvoke({}, config, control=RunControl()) == {"first": "done"}
|
||||
assert (await compiled.aget_state(config)).next == ("second",)
|
||||
assert second_runs == []
|
||||
|
||||
assert await compiled.ainvoke(None, config) == {
|
||||
"first": "done",
|
||||
"second": "done",
|
||||
}
|
||||
assert second_runs == ["second"]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_drain_requested_in_terminal_step_finishes_normally_async() -> None:
|
||||
class State(TypedDict, total=False):
|
||||
|
||||
@@ -8,20 +8,33 @@ caller kwargs to the inner ``(a)stream`` call but rejects ``stream_mode`` and
|
||||
``subgraphs`` since v3 owns them (``stream_mode`` is built from the
|
||||
transformer mux; ``subgraphs`` is forced True so nested namespaces flow
|
||||
through scoped muxes).
|
||||
|
||||
A second regression is pinned here: #7677 (first released in 1.2.0a3)
|
||||
declared `interrupt_before` / `interrupt_after` / `control` as named
|
||||
parameters on the `Pregel.stream_events` / `astream_events` dispatchers
|
||||
but forwarded them only on the v3 branch, silently dropping them on v1/v2
|
||||
(where they had reached `(a)stream` through `**kwargs` before).
|
||||
`TestAstreamKwargsForwardedOnEveryVersion` and friends pin that they
|
||||
reach `(a)stream` on every version, with the v1/v2 passthrough semantics
|
||||
restored (exactly what the caller passed, including an explicit `None`)
|
||||
and v3's explicit-default semantics preserved.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.constants import END, START
|
||||
from langgraph.errors import GraphDrained
|
||||
from langgraph.graph import StateGraph
|
||||
from langgraph.runtime import Runtime
|
||||
from langgraph.runtime import RunControl, Runtime
|
||||
|
||||
NEEDS_CONTEXTVARS = pytest.mark.skipif(
|
||||
sys.version_info < (3, 11),
|
||||
@@ -106,3 +119,274 @@ class TestKwargForwardingAsync:
|
||||
version="v3",
|
||||
subgraphs=False,
|
||||
)
|
||||
|
||||
|
||||
_KWARG_NAMES = ("control", "interrupt_before", "interrupt_after")
|
||||
|
||||
|
||||
def _build_two_step_graph(
|
||||
first: Callable[[_State], dict[str, Any]] | None = None,
|
||||
) -> Any:
|
||||
"""A `first -> second` graph with a checkpointer, for interrupt/drain tests."""
|
||||
|
||||
def default_first(state: _State) -> dict[str, Any]:
|
||||
return {"message": state["message"] + " first"}
|
||||
|
||||
def second(state: _State) -> dict[str, Any]:
|
||||
return {"message": state["message"] + " second"}
|
||||
|
||||
builder = StateGraph(_State)
|
||||
builder.add_node("first", first or default_first)
|
||||
builder.add_node("second", second)
|
||||
builder.add_edge(START, "first")
|
||||
builder.add_edge("first", "second")
|
||||
builder.add_edge("second", END)
|
||||
return builder.compile(checkpointer=InMemorySaver())
|
||||
|
||||
|
||||
async def _drive_astream_events(
|
||||
graph: Any, config: dict[str, Any], version: str, **kwargs: Any
|
||||
) -> None:
|
||||
"""Consume an astream_events run for `version` to completion."""
|
||||
if version == "v3":
|
||||
run = await graph.astream_events(
|
||||
{"message": "hi"}, config, version="v3", **kwargs
|
||||
)
|
||||
await run.output()
|
||||
else:
|
||||
async for _ in graph.astream_events(
|
||||
{"message": "hi"}, config, version=version, **kwargs
|
||||
):
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.filterwarnings("ignore:astream_events version='v1' is deprecated")
|
||||
@pytest.mark.parametrize("version", ["v1", "v2", "v3"])
|
||||
class TestAstreamKwargsForwardedOnEveryVersion:
|
||||
"""`interrupt_before`/`interrupt_after`/`control` reach `astream` on every
|
||||
version.
|
||||
|
||||
Regression test for #7677 (first released in 1.2.0a3): the dispatchers
|
||||
captured these parameters as named arguments but forwarded them only on
|
||||
the v3 branch, silently dropping them on v1/v2.
|
||||
"""
|
||||
|
||||
async def test_interrupt_before(self, version: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
config = {"configurable": {"thread_id": "ib"}}
|
||||
await _drive_astream_events(graph, config, version, interrupt_before=["second"])
|
||||
state = await graph.aget_state(config)
|
||||
assert state.next == ("second",)
|
||||
assert state.values == {"message": "hi first"}
|
||||
|
||||
async def test_interrupt_after(self, version: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
config = {"configurable": {"thread_id": "ia"}}
|
||||
await _drive_astream_events(graph, config, version, interrupt_after=["first"])
|
||||
state = await graph.aget_state(config)
|
||||
assert state.next == ("second",)
|
||||
assert state.values == {"message": "hi first"}
|
||||
|
||||
async def test_pre_drained_control(self, version: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
config = {"configurable": {"thread_id": "drain"}}
|
||||
control = RunControl()
|
||||
control.request_drain("sigterm")
|
||||
with pytest.raises(GraphDrained, match="sigterm"):
|
||||
await _drive_astream_events(graph, config, version, control=control)
|
||||
|
||||
|
||||
class TestStreamEventsV3SyncInterrupts:
|
||||
"""Sync v3 static interrupts reach `stream()` after the kwargs rewire."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("kwarg", "node"),
|
||||
[("interrupt_before", "second"), ("interrupt_after", "first")],
|
||||
)
|
||||
def test_static_interrupt(self, kwarg: str, node: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
config = {"configurable": {"thread_id": "sync"}}
|
||||
run = graph.stream_events(
|
||||
{"message": "hi"}, config, version="v3", **{kwarg: [node]}
|
||||
)
|
||||
list(run.values)
|
||||
state = graph.get_state(config)
|
||||
assert state.next == ("second",)
|
||||
assert state.values == {"message": "hi first"}
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.filterwarnings("ignore:astream_events version='v1' is deprecated")
|
||||
@pytest.mark.parametrize("version", ["v1", "v2", "v3"])
|
||||
class TestAstreamMidRunDrain:
|
||||
"""A drain requested from inside a node propagates out of `astream_events`.
|
||||
|
||||
This is the graceful-shutdown scenario: `request_drain()` called while
|
||||
the run is in flight (e.g. from a signal handler), with v1/v2 running
|
||||
inside core's event-stream task. The caller's own `RunControl` is used
|
||||
(the drain reason proves identity), `GraphDrained` propagates to the
|
||||
consumer, and the checkpoint keeps the pending step.
|
||||
"""
|
||||
|
||||
async def test_drain_requested_inside_first_node(self, version: str) -> None:
|
||||
control = RunControl()
|
||||
|
||||
def first(state: _State) -> dict[str, Any]:
|
||||
control.request_drain("sigterm-mid")
|
||||
return {"message": state["message"] + " first"}
|
||||
|
||||
graph = _build_two_step_graph(first)
|
||||
config = {"configurable": {"thread_id": "midrun"}}
|
||||
|
||||
with pytest.raises(GraphDrained, match="sigterm-mid"):
|
||||
await _drive_astream_events(graph, config, version, control=control)
|
||||
state = await graph.aget_state(config)
|
||||
assert state.next == ("second",)
|
||||
assert state.values == {"message": "hi first"}
|
||||
|
||||
|
||||
def _record_astream_kwargs(graph: Any) -> list[dict[str, Any]]:
|
||||
"""Patch `graph.astream` to record the kwargs each call receives."""
|
||||
received: list[dict[str, Any]] = []
|
||||
original = graph.astream
|
||||
|
||||
async def recording_astream(
|
||||
input: Any, config: Any = None, **kwargs: Any
|
||||
) -> AsyncIterator[Any]:
|
||||
received.append(kwargs)
|
||||
async for chunk in original(input, config, **kwargs):
|
||||
yield chunk
|
||||
|
||||
graph.astream = recording_astream # type: ignore[method-assign]
|
||||
return received
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.filterwarnings("ignore:astream_events version='v1' is deprecated")
|
||||
@pytest.mark.parametrize("version", ["v1", "v2"])
|
||||
class TestAstreamV1V2KwargsPassthrough:
|
||||
"""v1/v2 forward to `astream` exactly what the caller passed.
|
||||
|
||||
Pre-#7677 semantics: an explicit `None` is forwarded as `None`, and an
|
||||
omitted argument is not forwarded at all (so an override's own default
|
||||
would apply).
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize("name", _KWARG_NAMES)
|
||||
async def test_passed_values_reach_astream(self, version: str, name: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_astream_kwargs(graph)
|
||||
value: Any = RunControl() if name == "control" else ["second"]
|
||||
await _drive_astream_events(
|
||||
graph, {"configurable": {"thread_id": "rec"}}, version, **{name: value}
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert received[0][name] == value
|
||||
if name == "control":
|
||||
assert received[0][name] is value
|
||||
|
||||
@pytest.mark.parametrize("name", _KWARG_NAMES)
|
||||
async def test_explicit_none_is_forwarded(self, version: str, name: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_astream_kwargs(graph)
|
||||
await _drive_astream_events(
|
||||
graph, {"configurable": {"thread_id": "rec-none"}}, version, **{name: None}
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert received[0][name] is None
|
||||
|
||||
@pytest.mark.parametrize("name", _KWARG_NAMES)
|
||||
async def test_omitted_values_are_absent(self, version: str, name: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_astream_kwargs(graph)
|
||||
await _drive_astream_events(
|
||||
graph, {"configurable": {"thread_id": "rec-omit"}}, version
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert name not in received[0]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
class TestAstreamV3KwargsDefaults:
|
||||
"""v3 keeps its since-inception explicit-default semantics (#7519).
|
||||
|
||||
Omitted `interrupt_before`/`interrupt_after`/`control` are supplied to
|
||||
`astream` as `None`; passed values are forwarded as-is.
|
||||
"""
|
||||
|
||||
async def test_omitted_values_arrive_as_none(self) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_astream_kwargs(graph)
|
||||
await _drive_astream_events(
|
||||
graph, {"configurable": {"thread_id": "v3-rec"}}, "v3"
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert received[0]["control"] is None
|
||||
assert received[0]["interrupt_before"] is None
|
||||
assert received[0]["interrupt_after"] is None
|
||||
|
||||
@pytest.mark.parametrize("name", _KWARG_NAMES)
|
||||
async def test_passed_values_reach_astream(self, name: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_astream_kwargs(graph)
|
||||
value: Any = RunControl() if name == "control" else ["second"]
|
||||
await _drive_astream_events(
|
||||
graph, {"configurable": {"thread_id": "v3-rec-2"}}, "v3", **{name: value}
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert received[0][name] == value
|
||||
if name == "control":
|
||||
assert received[0][name] is value
|
||||
|
||||
|
||||
def _record_stream_kwargs(graph: Any) -> list[dict[str, Any]]:
|
||||
"""Patch `graph.stream` to record the kwargs each call receives."""
|
||||
received: list[dict[str, Any]] = []
|
||||
original = graph.stream
|
||||
|
||||
def recording_stream(input: Any, config: Any = None, **kwargs: Any) -> Any:
|
||||
received.append(kwargs)
|
||||
yield from original(input, config, **kwargs)
|
||||
|
||||
graph.stream = recording_stream # type: ignore[method-assign]
|
||||
return received
|
||||
|
||||
|
||||
def _drive_stream_events_v3(graph: Any, config: dict[str, Any], **kwargs: Any) -> None:
|
||||
"""Consume a sync v3 stream_events run to completion."""
|
||||
run = graph.stream_events({"message": "hi"}, config, version="v3", **kwargs)
|
||||
list(run.values)
|
||||
|
||||
|
||||
class TestStreamEventsV3SyncKwargsDefaults:
|
||||
"""Sync v3 keeps its since-inception explicit-default semantics (#7519).
|
||||
|
||||
Mirror of `TestAstreamV3KwargsDefaults`: omitted
|
||||
`interrupt_before`/`interrupt_after`/`control` are supplied to `stream`
|
||||
as `None`; passed values are forwarded as-is. Pins the sync helper
|
||||
against a kwargs-only "simplification" that would change subclass
|
||||
default handling.
|
||||
"""
|
||||
|
||||
def test_omitted_values_arrive_as_none(self) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_stream_kwargs(graph)
|
||||
_drive_stream_events_v3(graph, {"configurable": {"thread_id": "s-rec"}})
|
||||
assert len(received) == 1
|
||||
assert received[0]["control"] is None
|
||||
assert received[0]["interrupt_before"] is None
|
||||
assert received[0]["interrupt_after"] is None
|
||||
|
||||
@pytest.mark.parametrize("name", _KWARG_NAMES)
|
||||
def test_passed_values_reach_stream(self, name: str) -> None:
|
||||
graph = _build_two_step_graph()
|
||||
received = _record_stream_kwargs(graph)
|
||||
value: Any = RunControl() if name == "control" else ["second"]
|
||||
_drive_stream_events_v3(
|
||||
graph, {"configurable": {"thread_id": "s-rec-2"}}, **{name: value}
|
||||
)
|
||||
assert len(received) == 1
|
||||
assert received[0][name] == value
|
||||
if name == "control":
|
||||
assert received[0][name] is value
|
||||
|
||||
Generated
+23
-8
@@ -1217,6 +1217,20 @@ wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/38/64/285f20a31679bf547b75602702f7800e74dbabae36ef324f716c02804753/jupyter-1.1.1-py2.py3-none-any.whl", hash = "sha256:7a59533c22af65439b24bbe60373a4e95af8f16ac65a6c00820ad378e3f7cc83", size = 2657, upload-time = "2024-08-30T07:15:47.045Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jupyter-builder"
|
||||
version = "1.2.3"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "jupyter-core" },
|
||||
{ name = "tomli", marker = "python_full_version < '3.11'" },
|
||||
{ name = "traitlets" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/75/3e/56f593e6a664cd14d441724654ce8243f3cac7d3e45867cf084749c8bc75/jupyter_builder-1.2.3.tar.gz", hash = "sha256:01aba6794eb9b19e0e29ae21137ca60ba4135c70347d4b9f664a586822b8c809", size = 1024218, upload-time = "2026-09-04T18:54:09.411Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/84/aa/be79e87c50698673f196d0633fc1664607da2289f9688d1de7a1c9db9e71/jupyter_builder-1.2.3-py3-none-any.whl", hash = "sha256:c5ea5a7190c2a7b082494abade98eece1b2b5bd5dbd7d610606cbcd10a1d08b3", size = 947541, upload-time = "2026-09-04T18:54:07.644Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jupyter-client"
|
||||
version = "8.8.0"
|
||||
@@ -1342,28 +1356,28 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "jupyterlab"
|
||||
version = "4.5.10"
|
||||
version = "4.6.4"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "async-lru" },
|
||||
{ name = "httpx" },
|
||||
{ name = "ipykernel" },
|
||||
{ name = "jinja2" },
|
||||
{ name = "jupyter-builder" },
|
||||
{ name = "jupyter-core" },
|
||||
{ name = "jupyter-lsp" },
|
||||
{ name = "jupyter-server" },
|
||||
{ name = "jupyterlab-server" },
|
||||
{ name = "notebook-shim" },
|
||||
{ name = "packaging" },
|
||||
{ name = "setuptools" },
|
||||
{ name = "tomli", marker = "python_full_version < '3.11'" },
|
||||
{ name = "tornado" },
|
||||
{ name = "traitlets" },
|
||||
{ name = "typing-extensions", marker = "python_full_version < '3.12'" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/74/24/621aa20ec0d2fe72f52095bda0fc1be7738ac21aabe4129ff623140d5cdf/jupyterlab-4.5.10.tar.gz", hash = "sha256:77e8d80b78be59b2eaba2154562e21caa6e79c2f1281d6f486584f7144ee2f47", size = 23998879, upload-time = "2026-07-21T12:43:27.324Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/33/8d/995cc142f6083346b35e7d3eadc5a6717ee89ce64ea53577010ac493bc3c/jupyterlab-4.6.4.tar.gz", hash = "sha256:404f49b081819378524886c9db66dba57a5565981eff885830df1baba3a17df5", size = 28335647, upload-time = "2026-09-21T15:25:18.779Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/9f/c9/940f95f17ee4e413ad252bf8d4f2ee9a341f18cfeda87775fef3d7847321/jupyterlab-4.5.10-py3-none-any.whl", hash = "sha256:5967ca61e692e67a2f30b5a2b901c941dc6ce56c0b0e357bc6d34fed5ec095f6", size = 12452502, upload-time = "2026-07-21T12:43:23.542Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/3b/e4/072bc0d3c6d45d414062a06af7ee81c02a2c12d769c358ec9919f9995ffd/jupyterlab-4.6.4-py3-none-any.whl", hash = "sha256:15b13f991d3985129c797eb84d9949eeb8b6615e14b444868e642411f2c418b2", size = 17171596, upload-time = "2026-09-21T15:25:14.123Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -1623,7 +1637,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -2146,18 +2160,19 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "notebook"
|
||||
version = "7.5.7"
|
||||
version = "7.6.3"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "jupyter-builder" },
|
||||
{ name = "jupyter-server" },
|
||||
{ name = "jupyterlab" },
|
||||
{ name = "jupyterlab-server" },
|
||||
{ name = "notebook-shim" },
|
||||
{ name = "tornado" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/3e/c4/f71f8716f2903e9e817a47f534b9fd84831e155e2acb32c26691c8e06243/notebook-7.5.7.tar.gz", hash = "sha256:d6d59288a25303b25e1dcb71e9b017ec3a785f7d92f38b9bc288ca1970d5b0a8", size = 14171612, upload-time = "2026-06-04T18:33:45.224Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/42/ac/aedb759dc683dcd129f6c75271bab5eb23020b0f7969163c726343718e1e/notebook-7.6.3.tar.gz", hash = "sha256:e2c08e469c0ae20bb0b3214f0ab77e79653317a2f8e5b34c10361c66874a5b50", size = 5501618, upload-time = "2026-09-21T18:07:24.261Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/e1/4d/b3347f7073a377273531efe4ffc738fc910e93718fd2838c7ebf6736c6af/notebook-7.5.7-py3-none-any.whl", hash = "sha256:1f95f79d117e47d20b5555b5c85a397d2cfecf136978aaab767cf0314b09165b", size = 14583767, upload-time = "2026-06-04T18:33:40.987Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/a1/f7/907b98438cf00bdc4e296570058f29a347214bb42338a48e3b7eae58ae2a/notebook-7.6.3-py3-none-any.whl", hash = "sha256:ad7e0eb765fba836cd4a2ab0c7a3a26cde1d91665fbf6f533b6ae7b2de6d88d2", size = 5548452, upload-time = "2026-09-21T18:07:21.563Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
@@ -99,6 +99,14 @@ if TYPE_CHECKING:
|
||||
from langgraph.runtime import Runtime
|
||||
from pydantic_core import ErrorDetails
|
||||
|
||||
try:
|
||||
from langgraph.errors import is_invalid_resume
|
||||
except ImportError: # `langgraph` before `is_invalid_resume` never marks resume errors
|
||||
|
||||
def is_invalid_resume(error: BaseException) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
# right now we use a dict as the default, can change this to AgentState, but depends
|
||||
# on if this lives in LangChain or LangGraph... ideally would have some typed
|
||||
# messages key
|
||||
@@ -957,6 +965,11 @@ class ToolNode(RunnableCallable):
|
||||
try:
|
||||
response = tool.invoke(call_args, config)
|
||||
except ValidationError as exc:
|
||||
if is_invalid_resume(exc):
|
||||
# An `interrupt()` in this tool, or in a graph it ran, got a resume
|
||||
# value that doesn't match its `response_schema`. That's not a bad
|
||||
# tool argument: fail the run so the interrupt can be answered again.
|
||||
raise
|
||||
# Filter out errors for injected arguments
|
||||
injected = self._injected_args.get(call["name"])
|
||||
filtered_errors = _filter_validation_errors(exc, injected)
|
||||
@@ -982,6 +995,10 @@ class ToolNode(RunnableCallable):
|
||||
except GraphBubbleUp:
|
||||
raise
|
||||
except Exception as e:
|
||||
# The model can't fix a resume value that doesn't match an interrupt's
|
||||
# `response_schema`, so no `handle_tool_errors` setting handles it.
|
||||
if is_invalid_resume(e):
|
||||
raise
|
||||
# Determine which exception types are handled
|
||||
handled_types: tuple[type[Exception], ...]
|
||||
if isinstance(self._handle_tool_errors, type) and issubclass(
|
||||
@@ -1053,9 +1070,13 @@ class ToolNode(RunnableCallable):
|
||||
# Call wrapper with request and execute callable
|
||||
try:
|
||||
return self._wrap_tool_call(tool_request, execute)
|
||||
except GraphBubbleUp:
|
||||
# Interrupts always propagate, as they do without a wrapper.
|
||||
raise
|
||||
except Exception as e:
|
||||
# Wrapper threw an exception
|
||||
if not self._handle_tool_errors:
|
||||
# Wrapper threw an exception. The model can't fix a resume value that
|
||||
# doesn't match an interrupt's `response_schema`, so it's never handled.
|
||||
if not self._handle_tool_errors or is_invalid_resume(e):
|
||||
raise
|
||||
# Convert to error message
|
||||
content = _handle_tool_error(e, flag=self._handle_tool_errors)
|
||||
@@ -1104,6 +1125,11 @@ class ToolNode(RunnableCallable):
|
||||
try:
|
||||
response = await tool.ainvoke(call_args, config)
|
||||
except ValidationError as exc:
|
||||
if is_invalid_resume(exc):
|
||||
# An `interrupt()` in this tool, or in a graph it ran, got a resume
|
||||
# value that doesn't match its `response_schema`. That's not a bad
|
||||
# tool argument: fail the run so the interrupt can be answered again.
|
||||
raise
|
||||
# Filter out errors for injected arguments
|
||||
injected = self._injected_args.get(call["name"])
|
||||
filtered_errors = _filter_validation_errors(exc, injected)
|
||||
@@ -1129,6 +1155,10 @@ class ToolNode(RunnableCallable):
|
||||
except GraphBubbleUp:
|
||||
raise
|
||||
except Exception as e:
|
||||
# The model can't fix a resume value that doesn't match an interrupt's
|
||||
# `response_schema`, so no `handle_tool_errors` setting handles it.
|
||||
if is_invalid_resume(e):
|
||||
raise
|
||||
# Determine which exception types are handled
|
||||
handled_types: tuple[type[Exception], ...]
|
||||
if isinstance(self._handle_tool_errors, type) and issubclass(
|
||||
@@ -1208,9 +1238,13 @@ class ToolNode(RunnableCallable):
|
||||
# None check was performed above already
|
||||
self._wrap_tool_call = cast("ToolCallWrapper", self._wrap_tool_call)
|
||||
return self._wrap_tool_call(tool_request, _sync_execute)
|
||||
except GraphBubbleUp:
|
||||
# Interrupts always propagate, as they do without a wrapper.
|
||||
raise
|
||||
except Exception as e:
|
||||
# Wrapper threw an exception
|
||||
if not self._handle_tool_errors:
|
||||
# Wrapper threw an exception. The model can't fix a resume value that
|
||||
# doesn't match an interrupt's `response_schema`, so it's never handled.
|
||||
if not self._handle_tool_errors or is_invalid_resume(e):
|
||||
raise
|
||||
# Convert to error message
|
||||
content = _handle_tool_error(e, flag=self._handle_tool_errors)
|
||||
|
||||
@@ -25,6 +25,7 @@ from langchain_core.runnables.config import RunnableConfig
|
||||
from langchain_core.tools import BaseTool, InjectedToolArg, ToolException
|
||||
from langchain_core.tools import tool as dec_tool
|
||||
from langchain_core.tools.base import InjectedToolCallId
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.config import get_stream_writer
|
||||
from langgraph.errors import GraphBubbleUp, GraphInterrupt
|
||||
from langgraph.graph import START, MessagesState, StateGraph
|
||||
@@ -32,8 +33,8 @@ from langgraph.graph.message import REMOVE_ALL_MESSAGES, add_messages
|
||||
from langgraph.runtime import ExecutionInfo, ServerInfo
|
||||
from langgraph.store.base import BaseStore
|
||||
from langgraph.store.memory import InMemoryStore
|
||||
from langgraph.types import Command, Send
|
||||
from pydantic import BaseModel
|
||||
from langgraph.types import Command, Send, interrupt
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from pydantic.v1 import BaseModel as BaseModelV1
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
@@ -626,6 +627,149 @@ def test_tool_node_node_interrupt() -> None:
|
||||
assert exc_info.value == "foo"
|
||||
|
||||
|
||||
class _Approval(BaseModel):
|
||||
approved: bool
|
||||
|
||||
|
||||
class _AskState(TypedDict, total=False):
|
||||
answer: str
|
||||
|
||||
|
||||
def _approval_graph():
|
||||
"""A graph that asks a human for approval with a typed interrupt."""
|
||||
|
||||
def ask(state: _AskState) -> _AskState:
|
||||
approval = interrupt("Approve?", response_schema=_Approval)
|
||||
return {"answer": f"approved={approval.approved}"}
|
||||
|
||||
return StateGraph(_AskState).add_node("ask", ask).add_edge(START, "ask").compile()
|
||||
|
||||
|
||||
def _ask_human_call() -> dict[str, list[AnyMessage]]:
|
||||
call = ToolCall(name="ask_human", args={}, id="call_1")
|
||||
return {"messages": [AIMessage("", tool_calls=[call])]}
|
||||
|
||||
|
||||
def _handle_any(e): # no annotation: handles every error
|
||||
return "handled"
|
||||
|
||||
|
||||
# A bad answer to an interrupt must fail the run whatever `handle_tool_errors` is,
|
||||
# including settings that cover `ValidationError` (a `ValueError`), with or without
|
||||
# a wrapper. `create_agent` always runs tools through a wrapper (its middleware).
|
||||
_TOOL_NODES = pytest.mark.parametrize(
|
||||
("wrapped", "handle_tool_errors"),
|
||||
[
|
||||
(wrapped, handler)
|
||||
for wrapped in (False, True)
|
||||
for handler in (None, True, (ValueError,), _handle_any)
|
||||
],
|
||||
ids=[
|
||||
f"{wrapped}-{handler}"
|
||||
for wrapped in ("plain", "wrapped")
|
||||
for handler in ("default", "handle_true", "handle_value_error", "untyped")
|
||||
],
|
||||
)
|
||||
# The interrupt either runs in a graph the tool starts (a subagent) or in the tool.
|
||||
_SHAPE = pytest.mark.parametrize("nested", [True, False], ids=["nested", "direct"])
|
||||
|
||||
|
||||
@_TOOL_NODES
|
||||
@_SHAPE
|
||||
def test_tool_node_reraises_invalid_resume(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
wrapped: bool,
|
||||
nested: bool,
|
||||
handle_tool_errors: Any,
|
||||
) -> None:
|
||||
asker = _approval_graph()
|
||||
|
||||
@dec_tool
|
||||
def ask_human() -> str:
|
||||
"""Ask a human for approval."""
|
||||
if nested:
|
||||
return asker.invoke({})["answer"]
|
||||
approval = interrupt("Approve?", response_schema=_Approval)
|
||||
return f"approved={approval.approved}"
|
||||
|
||||
def pass_through(request, handler):
|
||||
return handler(request)
|
||||
|
||||
errors = (
|
||||
{} if handle_tool_errors is None else {"handle_tool_errors": handle_tool_errors}
|
||||
)
|
||||
graph = (
|
||||
StateGraph(MessagesState)
|
||||
.add_node(
|
||||
"tools",
|
||||
ToolNode(
|
||||
[ask_human], wrap_tool_call=pass_through if wrapped else None, **errors
|
||||
),
|
||||
)
|
||||
.add_edge(START, "tools")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
config: RunnableConfig = {"configurable": {"thread_id": "1"}}
|
||||
[pending] = graph.invoke(_ask_human_call(), config)["__interrupt__"]
|
||||
|
||||
# A bad answer isn't a bad tool argument: the run fails without saving, so
|
||||
# the same interrupt can be answered again.
|
||||
with pytest.raises(ValidationError, match="approved"):
|
||||
graph.invoke(Command(resume={pending.id: {"approved": "maybe"}}), config)
|
||||
assert [i.id for i in graph.get_state(config).interrupts] == [pending.id]
|
||||
|
||||
result = graph.invoke(Command(resume={pending.id: {"approved": True}}), config)
|
||||
assert result["messages"][-1].content == "approved=True"
|
||||
|
||||
|
||||
@_TOOL_NODES
|
||||
@_SHAPE
|
||||
async def test_tool_node_reraises_invalid_resume_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
wrapped: bool,
|
||||
nested: bool,
|
||||
handle_tool_errors: Any,
|
||||
) -> None:
|
||||
asker = _approval_graph()
|
||||
|
||||
@dec_tool
|
||||
async def ask_human() -> str:
|
||||
"""Ask a human for approval."""
|
||||
if nested:
|
||||
return (await asker.ainvoke({}))["answer"]
|
||||
approval = interrupt("Approve?", response_schema=_Approval)
|
||||
return f"approved={approval.approved}"
|
||||
|
||||
async def pass_through(request, handler):
|
||||
return await handler(request)
|
||||
|
||||
errors = (
|
||||
{} if handle_tool_errors is None else {"handle_tool_errors": handle_tool_errors}
|
||||
)
|
||||
graph = (
|
||||
StateGraph(MessagesState)
|
||||
.add_node(
|
||||
"tools",
|
||||
ToolNode(
|
||||
[ask_human], awrap_tool_call=pass_through if wrapped else None, **errors
|
||||
),
|
||||
)
|
||||
.add_edge(START, "tools")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
config: RunnableConfig = {"configurable": {"thread_id": "1"}}
|
||||
[pending] = (await graph.ainvoke(_ask_human_call(), config))["__interrupt__"]
|
||||
|
||||
with pytest.raises(ValidationError, match="approved"):
|
||||
await graph.ainvoke(Command(resume={pending.id: {"approved": "maybe"}}), config)
|
||||
state = await graph.aget_state(config)
|
||||
assert [i.id for i in state.interrupts] == [pending.id]
|
||||
|
||||
resume = Command(resume={pending.id: {"approved": True}})
|
||||
result = await graph.ainvoke(resume, config)
|
||||
assert result["messages"][-1].content == "approved=True"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("input_type", ["dict", "tool_calls"])
|
||||
async def test_tool_node_command(input_type: str) -> None:
|
||||
|
||||
|
||||
Generated
+1
-1
@@ -370,7 +370,7 @@ test = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -22,18 +22,23 @@ RUN pip install --no-cache-dir \
|
||||
"langchain>=1.3.0" \
|
||||
"deepagents>=0.6.2"
|
||||
|
||||
# Swap the published langgraph *core* for this monorepo's local copy, so the
|
||||
# server executes the langgraph under test rather than the latest release.
|
||||
# Swap the published langgraph *core* and *checkpoint* for this monorepo's
|
||||
# local copies, so the server executes the langgraph under test rather than the
|
||||
# latest release.
|
||||
# We keep the rest of the base image (langgraph-api, runtime, Go core server)
|
||||
# on `latest` — that still surfaces upstream regressions — but the core the
|
||||
# server runs is now the PR's, which is what lets this suite catch core
|
||||
# regressions (e.g. ensure_config / runtime changes) before they're published.
|
||||
# `--no-deps` keeps the base image's already-compatible
|
||||
# checkpoint/prebuilt/sdk; we only replace core. The local source comes from
|
||||
# the `langgraph_src` additional build context (see docker-compose.yml).
|
||||
# Checkpoint goes with core because core can need a checkpoint API from the
|
||||
# same release. `--no-deps` keeps the base image's already-compatible
|
||||
# prebuilt/sdk. The local sources come from the `checkpoint_src` and
|
||||
# `langgraph_src` additional build contexts (see docker-compose.yml).
|
||||
COPY --from=checkpoint_src pyproject.toml README.md LICENSE /opt/checkpoint-src/
|
||||
COPY --from=checkpoint_src langgraph /opt/checkpoint-src/langgraph/
|
||||
COPY --from=langgraph_src pyproject.toml README.md LICENSE /opt/langgraph-src/
|
||||
COPY --from=langgraph_src langgraph /opt/langgraph-src/langgraph/
|
||||
RUN pip install --no-cache-dir --force-reinstall --no-deps /opt/langgraph-src
|
||||
RUN pip install --no-cache-dir --force-reinstall --no-deps \
|
||||
/opt/checkpoint-src /opt/langgraph-src
|
||||
|
||||
# Project graphs + registration config.
|
||||
COPY graph/ /app/graph/
|
||||
|
||||
@@ -38,16 +38,17 @@ services:
|
||||
# see ./Dockerfile. The base image bundles langgraph-api +
|
||||
# langgraph_runtime_postgres + langgraph_license + the Go core-server;
|
||||
# on top we add graph deps (deepagents), the graph files, and this
|
||||
# monorepo's local langgraph *core* (so the server runs the code under
|
||||
# test, not the released langgraph).
|
||||
# monorepo's local langgraph *core* and *checkpoint* (so the server runs
|
||||
# the code under test, not the released langgraph).
|
||||
build:
|
||||
context: .
|
||||
dockerfile: Dockerfile
|
||||
additional_contexts:
|
||||
# The monorepo's local langgraph core (libs/langgraph), installed over
|
||||
# the base image so the server runs the langgraph under test. See the
|
||||
# Dockerfile for why.
|
||||
# The monorepo's local langgraph core (libs/langgraph) and checkpoint
|
||||
# (libs/checkpoint), installed over the base image so the server runs
|
||||
# the langgraph under test. See the Dockerfile for why.
|
||||
langgraph_src: ../../langgraph
|
||||
checkpoint_src: ../../checkpoint
|
||||
image: langgraph-v3-integration-api:local
|
||||
depends_on:
|
||||
postgres:
|
||||
|
||||
@@ -1972,7 +1972,7 @@ class AsyncThreadStream:
|
||||
# Mark that we have observed an active run so thread.output
|
||||
# knows a run exists (handles reattach without run.start).
|
||||
self._run_seen = True
|
||||
elif phase in ("completed", "failed"):
|
||||
elif _is_root_terminal_lifecycle(event):
|
||||
# Why: interrupts describe current-run state. Clear on terminal
|
||||
# lifecycle so a subsequent run.respond() can't fire against a
|
||||
# stale prior-run interrupt_id. Acquire `_interrupts_lock` so
|
||||
|
||||
@@ -35,7 +35,11 @@ from langgraph_sdk.stream.decoders import (
|
||||
validate_interleave_channels,
|
||||
)
|
||||
from langgraph_sdk.stream.subscription import compute_union_filter, infer_channel
|
||||
from langgraph_sdk.stream.sync_controller import SyncStreamController, _SyncSubscription
|
||||
from langgraph_sdk.stream.sync_controller import (
|
||||
SyncStreamController,
|
||||
_is_root_terminal_lifecycle,
|
||||
_SyncSubscription,
|
||||
)
|
||||
from langgraph_sdk.stream.transport import (
|
||||
SyncEventStreamHandle,
|
||||
SyncProtocolSseTransport,
|
||||
@@ -1614,7 +1618,7 @@ class SyncThreadStream:
|
||||
phase = data.get("event") if isinstance(data, dict) else None
|
||||
if phase in ("started", "running"):
|
||||
self._run_seen = True
|
||||
elif phase in ("completed", "failed"):
|
||||
elif _is_root_terminal_lifecycle(event):
|
||||
self.interrupted = False
|
||||
self.interrupts = []
|
||||
run_done = self._run_done
|
||||
|
||||
@@ -15,6 +15,7 @@ from langgraph_sdk.stream.transport import EventStreamHandle, ProtocolSseTranspo
|
||||
from streaming._events import (
|
||||
input_requested_event,
|
||||
lifecycle_completed_event,
|
||||
lifecycle_errored_event,
|
||||
lifecycle_event,
|
||||
)
|
||||
from streaming._fake_server import FakeServer, _StreamScript
|
||||
@@ -115,6 +116,25 @@ async def test_terminal_lifecycle_clears_interrupts():
|
||||
assert thread.interrupts == []
|
||||
|
||||
|
||||
async def test_subgraph_completed_event_does_not_end_run():
|
||||
fake = FakeServer()
|
||||
fake.script(
|
||||
[
|
||||
lifecycle_completed_event(seq=0, namespace=["child:1"]),
|
||||
lifecycle_errored_event(seq=1, error="root failed"),
|
||||
]
|
||||
)
|
||||
asgi = httpx.ASGITransport(app=fake.app)
|
||||
async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw:
|
||||
threads = ThreadsClient(HttpClient(raw))
|
||||
async with threads.stream(thread_id="t-1", assistant_id="agent") as thread:
|
||||
run_done = thread._run_done
|
||||
assert run_done is not None
|
||||
terminal = await asyncio.wait_for(run_done, timeout=2.0)
|
||||
assert terminal.status == "errored", "a subgraph's completed event ended the run"
|
||||
assert "root failed" in str(terminal.error)
|
||||
|
||||
|
||||
async def test_lifecycle_error_captured_for_output():
|
||||
"""Lifecycle error terminal state is captured in _run_done with error set."""
|
||||
fake = FakeServer()
|
||||
|
||||
@@ -27,6 +27,7 @@ from streaming._events import (
|
||||
checkpoints_event,
|
||||
custom_event,
|
||||
lifecycle_completed_event,
|
||||
lifecycle_errored_event,
|
||||
lifecycle_event,
|
||||
lifecycle_started_event,
|
||||
message_finish_event,
|
||||
@@ -475,6 +476,23 @@ def test_sync_lifecycle_watcher_reconnects_with_since_after_transport_drop():
|
||||
assert fake.stream_request_bodies[1]["since"] == 1
|
||||
|
||||
|
||||
def test_sync_subgraph_completed_event_does_not_end_run():
|
||||
fake = SyncFakeServer()
|
||||
fake.script(
|
||||
[
|
||||
lifecycle_completed_event(seq=1, namespace=["child:1"]),
|
||||
lifecycle_errored_event(seq=2, error="root failed"),
|
||||
]
|
||||
)
|
||||
with httpx.Client(transport=fake.transport, base_url="http://test") as raw:
|
||||
threads = SyncThreadsClient(SyncHttpClient(raw))
|
||||
with threads.stream(thread_id="existing", assistant_id="agent") as thread:
|
||||
terminal = thread._wait_for_run_done()
|
||||
|
||||
assert terminal.status == "errored", "a subgraph's completed event ended the run"
|
||||
assert "root failed" in str(terminal.error)
|
||||
|
||||
|
||||
def test_sync_threads_stream_accepts_websocket_transport_option():
|
||||
with httpx.Client(base_url="http://test") as raw:
|
||||
threads = SyncThreadsClient(SyncHttpClient(raw))
|
||||
|
||||
Generated
+1
-1
@@ -385,7 +385,7 @@ test = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.2.0"
|
||||
version = "4.3.0"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
Reference in New Issue
Block a user