Compare commits

...
Author SHA1 Message Date
Elior Nataf Lackritz 06b50627a6 fix(langgraph): read a run's DeltaChannel input back only on its own checkpoints
Under sync and async durability, a raw Pregel run saved its input to a
DeltaChannel input channel as a NULL_TASK_ID write on the checkpoint it
started from, which readers apply as that checkpoint's own state. A new
thread lost its first input, a run's last checkpoint read the next run's
input, and a run from an older checkpoint leaked into the branch that grew
from it. The input is now stored there under its own task id, or, on a new
thread or an addressed checkpoint, the input checkpoint snapshots the
channel.
2026-10-09 11:31:58 -04:00
cba111d8d6 fix(langgraph): don't save a checkpoint past a failed DeltaChannel write (#9235)
A DeltaChannel is rebuilt from its writes along the parent chain, so a
checkpoint saved without one of them reads back short for good. If
saving a DeltaChannel write failed (a database blip, say), the next
checkpoints were saved anyway and the channel lost that write, while
plain channels kept it. Three things let that happen:

- the sync loop waited for the delta writes before saving a checkpoint,
but never checked that they succeeded;
- both loops saved a checkpoint even when the previous save had failed,
so it pointed at a parent that was never saved;
- each save took the pending write list only when it started in the
background, and saves can start out of order, so a save could take a
later checkpoint's writes and not wait for its own. Waiting on writes
queued behind it, a save could also hang a sync run with
`durability="async"` (the default) once every worker was such a save: on
a 15-core machine a loop of 100 fast steps hung every time with default
settings, and with `max_concurrency=1` three steps are enough.

Now `_put_checkpoint` hands each save the writes submitted before it,
and a save fails if one of them or the previous save failed. The run
raises the write's error as before, the thread stays at its last
checkpoint that has all its writes, and running it again reruns the
task. For graphs without a DeltaChannel, a failed save now also stops
the later saves of that run, as JS already does, instead of saving
checkpoints whose parent is missing.

Thanks @shivangsharma01 for reproducing this in #8234, and @iroiro147,
whose #8299 had the first two changes before the issue-assignment check
closed it. Thanks also to @HarshShroff for the failed-save test case.

JS twin: langchain-ai/langgraphjs#2969

## Tests

`test_delta_channel_write_failures.py`:

- a saver fails one delta write once; on every durability, with `invoke`
and `ainvoke`, every saved checkpoint matches a plain channel and a
retry gets the write back. The `invoke` cases and `ainvoke` with
`durability="async"` fail on main.
- a saver fails the save of the checkpoint after `a` once; with `invoke`
and `ainvoke`, every saved checkpoint matches a plain channel and a
retry finishes the run. The `durability="async"` cases fail on main,
which saves the later checkpoints on top of the missing one.
- a delta graph run with `max_concurrency=1` finishes. It hangs on main.

Reverting any one of the three changes fails at least one of these.
`make format`, `make lint` and the langgraph suite pass. A 100-step
benchmark (delta and plain channels, every durability, with and without
1 ms of saver latency) makes the same saver calls and runs in the same
time as main, except for the case that hangs on main.

---------

Co-authored-by: iroiro147 <265728356+iroiro147@users.noreply.github.com>
2026-10-09 12:48:24 +00:00
Elior Nataf LackritzandGitHub 93a5a28008 perf(checkpoint-postgres): start the delta walk at the target (#9248)
Reading a DeltaChannel at an older checkpoint (the state at a past
checkpoint, each entry of `get_state_history`, a fork or replay from
one) walked `checkpoints` down from the newest row, so every newer row
was read before the target's chain. A checkpoint's ancestors have
smaller ids, the order `get_tuple` already relies on to find the latest
checkpoint, so the walk now starts at the target: the first page is
`checkpoint_id <= target`, later pages stay `< cursor`.

The cursor is never NULL anymore, so the `(cursor IS NULL OR
checkpoint_id < cursor)` predicate goes too. Once Postgres caches a
generic plan for the prepared statement (`from_conn_string` sets
`prepare_threshold=0`), that predicate is a filter rather than an index
condition, and each later page re-read every newer row.

`get_delta_channel_history` on a 6,000-checkpoint thread with a snapshot
every 200 checkpoints, median of three runs of 60 reads each, same
results before and after:

| target | main | this PR |
|---|---|---|
| newest | 3.1 ms | 2.7 ms |
| 1,000 back | 4.5 ms | 2.9 ms |
| 3,000 back | 8.9 ms | 2.8 ms |
| 5,000 back | 15.7 ms | 2.8 ms |

The target is now always the first row read, so the pagination tests
from #8556 can no longer leave it off the first page; they are renamed
for what they still cover, a walk that continues across pages.

## Test plan
- [x] `test_{sync,async}_walk_reads_nothing_newer_than_the_target`
record the rows the walk reads; both fail on main, which also reads the
three newer checkpoints.
- [x] `checkpoint-postgres` suite on Postgres 16: 281 passed, 3 skipped.
- [x] `langgraph` DeltaChannel tests against this saver (`-k "delta and
not redis"`): 1,154 passed, 1 skipped.
2026-10-09 08:31:15 -04:00
Baihao WangGitHubBaihao Wangopen-swe[bot] <open-swe@users.noreply.github.com>
bfcfea554e fix(checkpoint,checkpoint-postgres,checkpoint-sqlite): validate namespace labels before store dispatch (#9217)
The SQLite and Postgres stores save a namespace as its labels joined
with `.`, so a label containing `.` could resolve to a different
namespace: `("foo.bar",)` and `("foo", "bar")` share the same row. This
rejects those labels at every entry point.

- `BaseStore` and `AsyncBatchedBaseStore`: `get`, `search`, `delete`,
and `list_namespaces` (sync and async) reject labels that are empty, not
strings, or contain `.`. The check runs before an op is queued, so one
caller's bad label can't fail other callers' concurrent requests.
- New `langgraph.store.base.validate_op_namespace(op)` applies the same
label checks to a single op. `SqliteStore`/`AsyncSqliteStore` and
`PostgresStore`/`AsyncPostgresStore` call it from `_group_ops` before
any SQL runs, which covers direct `batch`/`abatch` calls. Other stores
that serialize namespaces as delimited text can call it too.
- Unchanged: hierarchical prefix search, empty-prefix `search(())`, `*`
wildcards in `list_namespaces`, and the existing `put` rules.

## Release note

`get`, `search`, `delete`, and `list_namespaces` now raise
`InvalidNamespaceError` for labels that are empty, not strings, or
contain `.`, where they previously returned `None` or `[]` (or, on the
SQLite and Postgres stores, matched a different namespace). `put`
already rejected these labels, so nothing written through `put` becomes
unreachable.

## Release order

Release `langgraph-checkpoint` first. `langgraph-checkpoint-sqlite` and
`langgraph-checkpoint-postgres` import `validate_op_namespace`, so they
need the `langgraph-checkpoint` release that ships it. `main` already
has `libs/checkpoint` at 4.3.0 and requires
`langgraph-checkpoint>=4.3.0` in both stores (from #8544), so this must
merge before 4.3.0 is released; if 4.3.0 ships first, bump to the next
version and raise both requirements to it. That also guarantees the
pre-queue checks are installed: with an older `langgraph-checkpoint`, an
invalid op would fail inside the shared `abatch` and reject the calls
batched with it.

---------

Co-authored-by: Baihao Wang <byhow@users.noreply.github.com>
Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
2026-10-08 16:10:55 -07:00
Elior Nataf LackritzandGitHub 93f5eaff21 fix(checkpoint): keep serialized data when an object can't be rebuilt (#9251)
Fixes #6970

If a checkpointed object can't be rebuilt on load (its module was
removed or renamed, or its constructor now rejects the stored fields),
`JsonPlusSerializer` returned `None` for it, so the value silently
disappeared from restored state. The constructor and method ext codes
now return the serialized payload instead, the same thing the pydantic
and blocked-type paths already do, and log a one-time warning per type.
The warning names the exception type only, since exception messages can
echo the stored value.

Thanks @yangbaechu for the report and repro. This takes the same
direction as #6972 by @pandego and #7152 by @SinzoL.
2026-10-08 16:01:46 -04:00
Baihao WangandGitHub 5965d72ff7 fix(langgraph,prebuilt): stop ToolNode from swallowing invalid resume values (#9232)
**Problem.** `interrupt(response_schema=...)` checks a human's answer
when the graph resumes. If the answer doesn't match the schema, the
resume fails, nothing is saved, and the human can answer again (#8886).
That works when the interrupt runs in a graph node, but not when it runs
inside a tool: either the tool calls `interrupt()` itself, or it starts
a graph that does, such as a subagent called as a tool. There, the
resume returns normally and the interrupt disappears.

**Why.** `interrupt()` raises a `pydantic.ValidationError` for an
invalid answer. When one comes out of a tool, `ToolNode` assumes the
model passed bad arguments, so it hands the error to the model ("Error
invoking tool … Please fix the error and try again") and the run carries
on. The human never sees the error, the model is blamed for it, and the
interrupt is no longer pending.

**Fix.** `interrupt()` still raises the same `ValidationError`, but
marks it so a new helper, `langgraph.errors.is_invalid_resume(error)`,
can recognize it. `ToolNode` re-raises these errors instead of handing
them to the model, whatever `handle_tool_errors` is set to, so a bad
answer fails the resume the same way inside a tool as outside one:
nothing is saved and the interrupt can be answered again. Every other
`ValidationError` from a tool is handled as before.

**Also fixed, in the same error handling.** With `wrap_tool_call` set
and a `handle_tool_errors` other than the default (for example `True`),
`ToolNode` turned any exception from the wrapper into an error
`ToolMessage`, including `GraphInterrupt`. So an interrupt in that setup
never paused: the model got `Error: GraphInterrupt(...)` instead.
Interrupts now propagate there too, as they already do without a
wrapper. `create_agent` uses the default handler, so it was not
affected.

**Compatibility.** Not a breaking change: the error type is unchanged.
Worth a careful look: a wrapped `ToolNode` with a non-default
`handle_tool_errors` now pauses on interrupts instead of returning them
to the model as errors. With an older `langgraph` that doesn't have
`is_invalid_resume`, `ToolNode` keeps today's behavior.
2026-10-08 12:59:25 -07:00
Elior Nataf LackritzandGitHub 40a2e6d845 fix(langgraph): never store exit-mode delta writes under the null task id (#9229)
In exit mode, a run stores its DeltaChannel writes on the checkpoint it
started from, under synthetic task ids that put the superstep first. For
a `Command`'s writes in a run whose first superstep is 0, that id came
out as `NULL_TASK_ID` itself. Readers take writes under `NULL_TASK_ID`
for the checkpoint's own pending writes, so after
`invoke(Command(update=...), durability="exit")`:

- on a new thread, `get_state` on the empty first checkpoint showed the
update in the DeltaChannel but not in a plain channel;
- on a thread whose first checkpoint came from `update_state(...,
as_node="__input__")`, `get_state` on that checkpoint showed the update
the same way, and a replay or fork from it applied the update to the
DeltaChannel only.

`exit_delta_task_id` now never returns `NULL_TASK_ID`. Threads already
saved this way stay as they are.

## Tests

- `test_command_update_on_an_input_checkpoint_matches_a_plain_channel`
compares every checkpoint in the history, and a replay from the input
checkpoint, with a plain channel, on every checkpointer and durability.
The exit cases fail without the fix.
- `test_exit_command_update_on_a_new_thread_matches_a_plain_channel`
compares every checkpoint in a new thread's history with a plain
channel, on every checkpointer, and fails without the fix. It runs in
exit mode only: in sync and async the two channels already differ on
`main`, because of a separate bug that applies a new thread's `Command`
update twice.

`make format`, `make lint` and the langgraph suite pass.
2026-10-07 19:17:02 -04:00
Elior Nataf LackritzandGitHub a0053bb616 fix(sdk-py): end a thread stream's run only on a root lifecycle event (#9228)
The thread stream's lifecycle watcher ended the run on any `completed`
or `failed` lifecycle event, including the one a subgraph sends when it
finishes. So `thread.output` could read the thread state before the
parent stored the step that ran the subgraph, and a run that failed
after a subgraph completed was reported as completed. The async and sync
watchers now end the run only on a root event, with the
`_is_root_terminal_lifecycle` check the fanout already uses. The JS SDK
checks the root namespace here too.

This is the `sdk-py integration` flake where the final `items` comes
back as `['streamed', 'tool', 'asked']` without `'sub'`: the example
graph's last node runs a subgraph.

## Tests

New async and sync tests send a subgraph `completed` and then a root
`failed`. Without the fix the run ends as completed. `make lint` and
`make test` in `libs/sdk-py` pass.
2026-10-07 19:03:46 -04:00
Igor SoarezandGitHub 87f1c8eb9a fix(langgraph): stop astream_events dropping control/interrupts on v1/v2 (#9219)
### Problem

`Pregel.astream_events` declared `interrupt_before`, `interrupt_after`
and
`control` as keyword-only parameters, but forwarded them only on the
`version="v3"` branch. On async `version="v1"`/`"v2"` the values were
bound to
named parameters and never entered `**kwargs`, so they were silently
dropped
on the way to `astream` — no error, no effect:

- `request_drain()` on a caller-supplied `RunControl` did nothing; the
graph
  created its own control and ran to completion.
- Static interrupts passed as `astream_events(...,
interrupt_before=[...])`
  never fired.

This is a regression introduced by #7677 (first released in `1.2.0a3`).
Before that, `Pregel` did not override `astream_events`, and these
keywords
reached `astream` through `**kwargs` — langchain-core's v1 and v2 event
implementations both forward caller kwargs to `astream`
(`tracers/log_stream.py`, `tracers/event_stream.py`).

Downstream impact: langgraph-api passes run-level `interrupt_before` /
`interrupt_after` into `graph.astream_events(..., version="v2",
**kwargs)` for
runs with `stream_mode="events"` (`models/run.py`, `stream.py`), so
server
runs using that stream mode have silently dropped static interrupts
since
langgraph 1.2.0a3.

Sync `stream_events(version="v1"/"v2")` is unaffected: langchain-core
raises
`NotImplementedError` there, so there was no working sync path to break.

### Fix

Remove the three parameters from the public runtime implementations of
`stream_events` / `astream_events` (and from the explicit passing in
their v3
dispatch calls), so the values travel through `**kwargs` again:

- **v1/v2** recover the pre-#7677 passthrough. `astream` receives
exactly what
the caller passed — a supplied value, an explicit `None`, or nothing for
an
  omitted argument.
- **v3** is unchanged. The values bind on `_pregel_stream_v3` /
`_apregel_stream_v3`, which keep their named parameters and
since-inception
(#7519, `1.2.0a1`) explicit-`None` defaults: an omitted argument arrives
at
  `astream` as `None`, exactly as before.

What `astream` receives after this change:

| caller supplies | async v1/v2 | v3 (sync and async) |
|---|---|---|
| nothing | argument absent | explicit `None` |
| explicit `None` | `None` | `None` |
| a value | the value | the value |

The typed `version="v3"` `@overload` stubs are untouched, so v3 call
sites
keep their static types (invalid values still fail `ty` there, same as
on
`main`). `RemoteGraph` and the v3-only `transformers` parameter are
untouched.

Complexity: net −24 lines in `main.py` with no branches, sentinels or
wrappers added. The public capture-and-drop pattern that caused the bug
is
gone; the two private helpers each have a single call path and forward
unconditionally, so they cannot drop anything.

Docstrings now describe the three parameters as honored on every async
version (type-checked only on the v3 overloads), and no longer promise
synchronous v1/v2 event streaming, which langchain-core does not
implement.

### Compatibility notes

- The three parameters were keyword-only, so no caller can break;
keyword
  callers bind through `**kwargs` identically.
- The names leave the runtime signatures (`inspect.signature`), while
the v3
  overloads keep the typing.
- v3 is unchanged: `_pregel_stream_v3`/`_apregel_stream_v3` still
default all
  three to `None` and pass them explicitly, so an override's non-`None`
default never applied on v3 and still doesn't. The only behavior change
beyond restoring the pre-#7677 passthrough is on v1/v2 with an explicit
  `None`: since 1.2.0a3 the named-parameter capture dropped it, so an
  override's non-`None` default was silently applied instead; now the
explicit `None` reaches `astream`, as it did before 1.2.0 and as it does
on
  v3.

### Testing

New tests in `tests/test_stream_events_v3_kwarg_forwarding.py`:

- `interrupt_before` / `interrupt_after` / pre-drained `control`
parametrized
  over v1/v2/v3 (12 of them fail on unpatched `main`).
- Sync v3 static interrupts.
- Mid-run drain: a node calls `request_drain()`; asserts `GraphDrained`
  propagates, the caller's own `RunControl` was used, and the checkpoint
  keeps the pending step.
- Recording tests pin, for all three names and on both the async and
sync v3
paths, exactly what `(a)stream` receives: supplied values (with
`control`
  identity), explicit `None`, and absence for omitted arguments.

Verified by mutation: replacing the implementation with each rejected
alternative (unpatched parent; injecting `None` for omitted kwargs on
v1/v2;
dropping explicit `None` on v1/v2; kwargs-only private helpers) makes
the
corresponding tests fail.

Suites: 246 passed on the four adjacent streaming test files; 906 passed
/
1 skipped across every test file calling `(a)stream_events` plus
`test_runtime.py`. `ruff check` / `format --check` clean; `ty check
langgraph`
clean.

Not verifiable locally (CI please): Python 3.10 (the new tests carry no
version skip), minimum supported langchain-core, and a live
langgraph-api
server run.

### Release note

> Fixed a regression since 1.2.0a3: `Pregel.astream_events` silently
dropped
> `interrupt_before`, `interrupt_after` and `control` for
> `version="v1"`/`"v2"`. Static interrupts and `RunControl` drains now
work on
> every async version. Server runs using `stream_mode="events"` were
affected
> and get their static interrupts back.
2026-10-07 12:26:50 -04:00
31 changed files with 1566 additions and 105 deletions
@@ -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 = []
+61
View File
@@ -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
+125
View File
@@ -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()
+19
View File
@@ -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:
+50 -31
View File
@@ -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)
+22 -23
View File
@@ -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,
)
+14 -3
View File
@@ -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
+10 -1
View File
@@ -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
+38 -4
View File
@@ -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)
+146 -2
View File
@@ -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:
+1 -1
View File
@@ -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
+6 -2
View File
@@ -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))