mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-09 18:05:10 +02:00
Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
06b50627a6 | ||
|
|
cba111d8d6 | ||
|
|
93a5a28008 | ||
|
|
bfcfea554e | ||
|
|
93f5eaff21 | ||
|
|
5965d72ff7 | ||
|
|
40a2e6d845 | ||
|
|
a0053bb616 | ||
|
|
87f1c8eb9a |
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -423,7 +423,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
|
||||
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]
|
||||
|
||||
@@ -178,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,
|
||||
@@ -197,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
|
||||
|
||||
@@ -238,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.
|
||||
@@ -264,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
|
||||
@@ -413,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.
|
||||
@@ -422,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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,50 @@ 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)
|
||||
|
||||
|
||||
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")]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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")]
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -26,6 +26,7 @@ from langgraph._internal._constants import (
|
||||
CONFIG_KEY_CHECKPOINT_ID,
|
||||
NS_END,
|
||||
NS_SEP,
|
||||
NULL_TASK_ID,
|
||||
PUSH,
|
||||
SNAPSHOT_BUMPS,
|
||||
)
|
||||
@@ -64,10 +65,14 @@ def exit_delta_task_id(step: int, task_id: str) -> str:
|
||||
|
||||
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:
|
||||
|
||||
@@ -19,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
|
||||
@@ -191,6 +192,7 @@ class PregelLoop:
|
||||
Callable[
|
||||
[
|
||||
concurrent.futures.Future | None,
|
||||
Sequence[Any],
|
||||
RunnableConfig,
|
||||
Checkpoint,
|
||||
str,
|
||||
@@ -204,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.
|
||||
@@ -266,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
|
||||
@@ -1114,16 +1121,26 @@ class PregelLoop:
|
||||
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).
|
||||
# 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"})
|
||||
@@ -1299,12 +1316,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,
|
||||
@@ -1374,6 +1396,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},
|
||||
@@ -1641,21 +1664,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:
|
||||
@@ -1767,6 +1788,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)
|
||||
@@ -1896,23 +1918,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:
|
||||
@@ -2027,6 +2045,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)
|
||||
|
||||
@@ -3714,16 +3714,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.
|
||||
|
||||
@@ -3753,20 +3752,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
|
||||
@@ -3774,16 +3780,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,
|
||||
)
|
||||
@@ -3819,9 +3824,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]:
|
||||
@@ -3844,9 +3846,6 @@ class Pregel(
|
||||
return self._apregel_stream_v3(
|
||||
input,
|
||||
config,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
control=control,
|
||||
transformers=transformers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -17,6 +17,7 @@ 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
|
||||
@@ -38,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}")
|
||||
@@ -463,6 +465,43 @@ def test_resume_with_a_command_update_replays_its_write_once(
|
||||
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
|
||||
|
||||
@@ -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,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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user