Compare commits

...
Author SHA1 Message Date
ClaudeandIgor Soarez 536059f99b fix(langgraph): stop at an interrupt_before breakpoint when a drain is requested before it
At a step boundary, tick() checked for a drain before the interrupt_before
breakpoints. A drain requested in the step before a breakpoint stopped the
run with GraphDrained, and the resume after it (invoke(None, config)) marks
the next step's breakpoints as passed, so the breakpoint never fired. The
drain check now runs after the breakpoint check, so the run stops at the
breakpoint as it would have without the drain. Subgraph breakpoints share
the fix, since subgraphs share the run's control.
2026-10-11 12:36:24 +03:00
Elior Nataf LackritzandGitHub b4991f1ba3 fix(langgraph): store an update_state input's writes on the checkpoint it updates (#9258)
`update_state(..., as_node="__input__")` stored its writes on the
checkpoint it saved. A DeltaChannel reads the writes stored on a
checkpoint's ancestors, not its own, so that checkpoint read back
without the input unless the channel snapshotted there, and once a later
checkpoint snapshotted, the input was gone for good. Only a raw `Pregel`
graph whose input channel is a DeltaChannel hits it, since a
`StateGraph`'s input goes to its start channel.

The input's writes now go on the checkpoint the update builds on, as a
node update's do. An update on a checkpoint the thread has moved past
stores nothing there and snapshots the DeltaChannels it writes on its
own checkpoint instead, the rule a node update already follows (#9165).

Both kinds of update now go through one helper, which also gives a
message an id before saving it, as the loop's `put_writes` does with
`ensure_message_ids`. Before, a message saved through `update_state`
without an id read back with id `None`, on the node path too, while the
same message from a node got one.

JS twin: langchain-ai/langgraphjs#2973, where review found it.

## Tests

`test_delta_channel_update_state.py`, sync and async: an update as input
reads back on its checkpoint and after the next run, with
`snapshot_frequency` 1 and 2, and an update as input to an older
checkpoint stays out of its other branch. The frequency 2 and
older-checkpoint cases fail on main. Five more check that a message
saved through `update_state` gets an id every read keeps, on the node
path (latest and older checkpoint) and the input path, sync and async;
they fail on main. The langgraph suite passes, and `lint_package` and
`lint_tests` are clean.
2026-10-10 11:57:46 +00:00
Elior Nataf LackritzandGitHub 6aa0afba68 fix(checkpoint): default InMemorySaver delta history to the latest checkpoint (#9257)
`InMemorySaver.get_delta_channel_history` looked its target up by
`config["configurable"].get("checkpoint_id", "")`, so a config without a
`checkpoint_id` matched no checkpoint, and every channel came back with
no seed and no writes, which a `DeltaChannel` reads as empty. The base
implementation, `PostgresSaver` and `SqliteSaver` read from the thread's
latest checkpoint instead, as `get_tuple` does.

A missing id now resolves to the namespace's newest checkpoint id, with
the same `max` that `InMemorySaver.get_tuple` uses. Graphs always pass a
checkpoint id here, so only direct callers of the saver API were
affected, and a call with an id does the same work as before.

Thanks @Hotragn, who pinned this down on #8242, where @longquanzheng's
branch had already fixed it.

JS twin: langchain-ai/langgraphjs#2979

## Tests

- `test_get_channel_writes_without_checkpoint_id_reads_the_latest` fails
on main (`{'messages': {'writes': []}}`).
- The `checkpoint`, `checkpoint-conformance`, `checkpoint-sqlite`,
`checkpoint-postgres` (Postgres 16), `prebuilt` and `langgraph` suites
pass, the last two without their Redis cases. The DeltaChannel
conformance suite also passes against `InMemorySaver`; CI skips that run
in `libs/checkpoint`, where the conformance package isn't installed, so
the new test lives in `test_memory.py`.
2026-10-10 00:59:20 +00:00
Elior Nataf LackritzandGitHub 12aeb0fddb fix(langgraph): read a run's DeltaChannel input back only on its own checkpoints (#9260)
A raw `Pregel` graph whose input channel is a DeltaChannel saved a run's
input to it as a `NULL_TASK_ID` write on the checkpoint the run started
from, under `"sync"` and `"async"` durability. Readers apply a
checkpoint's `NULL_TASK_ID` writes as its own state, so:

- on a new thread that checkpoint is never saved, and the first run's
input was lost for good: after `{"log": [0], "plain": [0], ...}`, `log`
read `[2]` next to `plain` `[0, 2]`;
- the last checkpoint of a run read the next run's input;
- a run with input from an older checkpoint (`checkpoint_id`) leaked its
input into the branch that already grew from it, and read that branch's
input when it had one.

The input's writes now go on the checkpoint the run starts from under a
task id of their own, `uuid5(checkpoint_id, INPUT)`, the one
`update_state` uses for an input update. On a new thread, or a
checkpoint a `checkpoint_id` addressed, which may already have children
that would read them, nothing is stored there and the input checkpoint
snapshots the channel instead. A normal run stores the same write as
before, and `durability="exit"`, which already kept the input right, is
unchanged. Threads saved before this keep the input writes already
stored on their checkpoints.

A `StateGraph` routes its input through its start channel, so only raw
`Pregel` graphs hit this.

JS twin: langchain-ai/langgraphjs#2981

## Tests

`test_delta_channel_run_input.py`, under every durability, with `invoke`
and `ainvoke`: two runs on a new thread, and a run with input from an
older checkpoint whose other branch had DeltaChannel input or none.
Every checkpoint reads the same in the DeltaChannel and a plain channel.
The 12 sync and async cases fail on main, and dropping the task id, the
new-thread snapshot or the addressed-checkpoint snapshot each fails its
own cases. The `langgraph` and `prebuilt` suites pass (without Redis),
and `lint_package` and `lint_tests` are clean.
2026-10-09 15:59:10 -04:00
Elior Nataf LackritzandGitHub d05236f805 fix(langgraph): snapshot a DeltaChannel that an update_state Overwrite resets (#9263)
A node that writes an `Overwrite` to a DeltaChannel makes the loop
snapshot the channel on the checkpoint that superstep saves, so no read
replays across the reset. The same write through `update_state`, as a
node or as input, didn't snapshot: the value read back right, but reads
of the update's checkpoint and the ones after it walked back past the
reset to the last snapshot. `update_state` now snapshots a DeltaChannel
an `Overwrite` resets, as the loop does.

Raised in review of langchain-ai/langgraphjs#2973. JS twin:
langchain-ai/langgraphjs#2982

## Tests

`test_delta_channel_update_overwrite.py`, sync and async: an `Overwrite`
through `update_state` as a node, and as input to a raw `Pregel` graph,
snapshots the channel on the checkpoint the update saves, and the value
reads back as the overwrite. All four fail on main. The `langgraph`
suite passes (without Redis), and `lint_package` and `lint_tests` are
clean.
2026-10-09 15:26:16 -04:00
Elior Nataf LackritzandGitHub 9d92f33cca fix(langgraph): keep a fork's DeltaChannel snapshot from starting nodes that never ran (#9264)
A fork from a checkpoint with pending writes to a DeltaChannel snapshots
the channel on the fork's first checkpoint, so the fork doesn't replay
the other branch's writes. If the channel had never been written there,
the snapshot also gives it its first version, and `_mark_bumps_seen`
only marked the bump seen for nodes that already have a `versions_seen`
entry. A node subscribed to the channel that never ran has none, so the
bump started it, and it read the empty value before the node that writes
the channel had run:

```python
graph.invoke("go", config)  # writer writes [1] to d, reader reads [1]
fork = graph.update_state(first_checkpoint, {"a": "go"}, as_node="__input__")
graph.get_state(fork).next  # ("writer", "reader"): resuming runs reader on [] first
```

A replay from that checkpoint saves the same kind of fork from the loop.
A channel bumped from no version was never written, so
`create_checkpoint` now also marks the bump seen for the nodes the
channel triggers, from the graph's `trigger_to_nodes`, at every call
that can bump. Only raw `Pregel` graphs hit this: `StateGraph` nodes
trigger on their `branch:to:*` channels, not on state keys. An
`as_node="__copy__"` fork still starts such a node, with a plain list
channel as well, so that case isn't DeltaChannel-specific and isn't
changed here.

Raised in review of langchain-ai/langgraphjs#2974. JS twin:
langchain-ai/langgraphjs#2983

## Tests

`test_delta_channel_seal_subscribers.py`: an update as input (sync and
async) and a replay from the checkpoint before a DeltaChannel's first
write leave only the writer next, and the reader then reads the written
value once. All three fail on main. The `langgraph` suite passes
(without Redis), and `lint_package` and `lint_tests` are clean.
2026-10-09 14:36:27 -04:00
Elior Nataf LackritzandGitHub 26356227c4 test(checkpoint): run the delta conformance suite against InMemorySaver in CI (#9261)
The conformance self-test in `libs/checkpoint-conformance` ran every
capability against `InMemorySaver` but asserted only the base ones.
`libs/checkpoint/tests/test_conformance_delta.py`, which ran the delta
conformance tests against it, always skips in CI: `libs/checkpoint`
can't depend on the conformance package, which depends on it. So no CI
job checked `InMemorySaver.get_delta_channel_history` against the delta
conformance tests.

The self-test now asserts every capability `InMemorySaver` implements,
and the skipped copy is gone.

## Tests

- With `InMemorySaver.get_delta_channel_history` broken to return no
history, the self-test now fails on 7 of the 10 delta conformance tests;
on main it still passes.
- `make format`, `make lint` and `make test` pass in
`libs/checkpoint-conformance`.
2026-10-09 13:44:55 -04:00
Elior Nataf LackritzandGitHub a5dbacae0d fix(checkpoint-postgres,checkpoint-sqlite): treat an empty checkpoint_id as the latest in delta history (#9262)
`PostgresSaver` and `SqliteSaver`, sync and async, read a DeltaChannel's
history from the latest checkpoint only when the config's
`checkpoint_id` was missing or `None`. An empty one was looked up as an
id, matched nothing, and came back with no seed and no writes, while
their own `get_tuple` treats an empty id as the latest checkpoint, as
the default implementation does. They now check for a missing or empty
id the way `get_tuple` does. The in-memory saver gets the same in #9257.

Graphs always pass the id of a saved checkpoint, so only direct callers
of the saver API hit this.

## Tests

-
`test_{sync,async}_empty_checkpoint_id_reads_from_the_latest_checkpoint`
in `checkpoint-postgres/tests/test_delta_pagination.py` and
`checkpoint-sqlite/tests/test_delta_parent_walk.py`: all four fail on
main.
- The `checkpoint-postgres` (Postgres 16) and `checkpoint-sqlite` suites
pass, and `lint_package` and `lint_tests` are clean in both.
2026-10-09 13:44:51 -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
gethin-langchainandGitHub 736a18eaab release(cli): 0.4.33 (#9227)
Bumping the CLI to 0.4.33 to cut a release including #9221 (deploy
listeners list) and #9222 (--image-uri).
2026-10-07 09:19:34 -04:00
gethin-langchainandGitHub ec4f8535f8 feat(cli): add --image-uri to deploy an already-pushed image (#9222)
## Summary

`--push-to` always builds (or retags a local `--image`) and pushes, even
when the image already exists at the destination. That breaks CI/CD
pipelines where the build/push step and the deploy step run in separate
jobs: the deploy job has no local copy of the image to retag, and
`--image` requires one (it runs `docker image inspect`, which needs the
image present in the local Docker daemon).

This adds `--image-uri`, which takes a full image reference already in a
registry and deploys it directly via `source_revision_config.image_uri`,
with no local Docker involvement at all: no inspect, no build, no tag,
no push. It gets the same placement behavior as `--push-to`
(`--listener-id`/`--k8s-namespace`, or auto-selecting the sole listener)
and the same update-in-place semantics, since both produce an
`external_docker` deployment and are interchangeable on the same
deployment.

`--image-uri` requires a digest (`repo@sha256:...`) and rejects
empty/blank input. `--push-to` always resolves to a digest after
pushing; `--image-uri` has no push step to resolve one from, so a
mutable tag (or an unvalidated empty string) could let Kubernetes
silently keep running a previously cached image while the CLI reports
success.

Implementation: `CustomerRegistrySource`'s image-acquisition step is
split into two variants, `BuildAndPush` (the existing
`--push-to`/`--image` flow, unchanged behavior) and `PublishedImage`
(the new flag), so create/update/placement logic stays shared. Also
reworded `_ensure_customer_registry_source`'s error, which previously
mentioned only `--push-to` even though it's reachable from `--image-uri`
too.

## Test plan

- [x] `make format && make lint` clean in `libs/cli`
- [x] Unit tests in `test_deploy_helpers.py` for `_select_source`
selecting `PublishedImage` (by digest, trimmed), requiring no local
Docker, carrying placement, and the mutual-exclusivity + validation
errors (`--push-to`, `--image`, `--tag`, `--remote`, empty/blank input,
missing digest)
- [x] CLI-level tests in `test_deploy_command.py`: create and update an
external deployment via `--image-uri` with zero Docker calls, listener
auto-placement, rejecting a non-external existing deployment (now
mentioning both flags), and the conflicting-flag cases end to end
- [x] Full `libs/cli` suite: 524 passed, no regressions
2026-10-07 13:32:46 +01:00
gethin-langchainandGitHub ce14d33687 feat(cli): add 'langgraph deploy listeners list' (#9221)
## Summary

Originally this auto-selected a workspace's sole listener for
self-hosted/hybrid control planes the same way #9056 already does for
Cloud. Per review, that's reverted: in self-hosted environments, leaving
placement blank is the deliberate signal to deploy to the bundled
operator, not an unresolved default. Auto-selecting a sole listener
there would silently take that option away from anyone who has both a
bundled operator and one listener configured, with no way to get it
back.

The actual need behind this was simpler: a way to find a listener's id
without hardcoding it, so it can be passed to `--listener-id`
explicitly. `HostBackendClient.list_listeners()` already existed and is
used internally for placement resolution, but had no CLI surface of its
own.

This adds `langgraph deploy listeners list`, mirroring the existing
`deploy list` / `deploy revisions list` pattern, to look up a
workspace's listeners and their namespaces. No change to default
placement behavior.

## Test plan

- [x] `make format && make lint` clean in `libs/cli`
- [x] `format_listeners_table` unit test in `test_util.py`
- [x] CLI-level tests for `deploy listeners list` (with and without
results) in `test_cli.py`
- [x] Full `libs/cli` suite: 508 passed, no regressions
2026-10-07 13:32:26 +01:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
708aaa4ebc chore(deps): bump notebook from 7.5.7 to 7.6.3 in /libs/langgraph (#9206)
Bumps [notebook](https://github.com/jupyter/notebook) from 7.5.7 to
7.6.3.
<details>
<summary>Release notes</summary>
<p><em>Sourced from <a
href="https://github.com/jupyter/notebook/releases">notebook's
releases</a>.</em></p>
<blockquote>
<h2>v7.6.3</h2>
<h2>7.6.3</h2>
<p>(<a
href="https://github.com/jupyter/notebook/compare/@jupyter-notebook/application-extension@7.6.2...9c8d7c24006554dda8c5fc02c69f91d6449c712a">Full
Changelog</a>)</p>
<h3>Maintenance and upkeep improvements</h3>
<ul>
<li>Update to JupyterLab v4.6.4 <a
href="https://redirect.github.com/jupyter/notebook/pull/8066">#8066</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
</ul>
<h3>Contributors to this release</h3>
<p>The following people contributed discussions, new ideas, code and
documentation contributions, and review.
See <a
href="https://github-activity.readthedocs.io/en/latest/use/#how-does-this-tool-define-contributions-in-the-reports">our
definition of contributors</a>.</p>
<p>(<a
href="https://github.com/jupyter/notebook/graphs/contributors?from=2026-08-11&amp;to=2026-09-21&amp;type=c">GitHub
contributors page for this release</a>)</p>
<p><a href="https://github.com/jtpio"><code>@​jtpio</code></a> (<a
href="https://github.com/search?q=repo%3Ajupyter%2Fnotebook+involves%3Ajtpio+updated%3A2026-08-11..2026-09-21&amp;type=Issues">activity</a>)</p>
<h2>v7.6.2</h2>
<h2>7.6.2</h2>
<p>(<a
href="https://github.com/jupyter/notebook/compare/@jupyter-notebook/application-extension@7.6.1...2f1f1622fc64e09b90ade87ed886fdfe1bc653fb">Full
Changelog</a>)</p>
<h3>Maintenance and upkeep improvements</h3>
<ul>
<li>Add nvd.nist.gov to the check_links ignore list <a
href="https://redirect.github.com/jupyter/notebook/pull/8032">#8032</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
<li>Backport JupyterLab dependency updater improvements <a
href="https://redirect.github.com/jupyter/notebook/pull/8031">#8031</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
<li>Update to JupyterLab v4.6.3 <a
href="https://redirect.github.com/jupyter/notebook/pull/8030">#8030</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
<li>Switch local pre-commit hooks to language: system <a
href="https://redirect.github.com/jupyter/notebook/pull/8002">#8002</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
</ul>
<h3>Documentation improvements</h3>
<ul>
<li>Update 'notebook v7 proposal' link in doc <a
href="https://redirect.github.com/jupyter/notebook/pull/8004">#8004</a>
(<a href="https://github.com/brichet"><code>@​brichet</code></a>)</li>
</ul>
<h3>Contributors to this release</h3>
<p>The following people contributed discussions, new ideas, code and
documentation contributions, and review.
See <a
href="https://github-activity.readthedocs.io/en/latest/use/#how-does-this-tool-define-contributions-in-the-reports">our
definition of contributors</a>.</p>
<p>(<a
href="https://github.com/jupyter/notebook/graphs/contributors?from=2026-07-22&amp;to=2026-08-11&amp;type=c">GitHub
contributors page for this release</a>)</p>
<p><a href="https://github.com/brichet"><code>@​brichet</code></a> (<a
href="https://github.com/search?q=repo%3Ajupyter%2Fnotebook+involves%3Abrichet+updated%3A2026-07-22..2026-08-11&amp;type=Issues">activity</a>)
| <a href="https://github.com/jtpio"><code>@​jtpio</code></a> (<a
href="https://github.com/search?q=repo%3Ajupyter%2Fnotebook+involves%3Ajtpio+updated%3A2026-07-22..2026-08-11&amp;type=Issues">activity</a>)</p>
<h2>v7.6.1</h2>
<!-- raw HTML omitted -->
</blockquote>
<p>... (truncated)</p>
</details>
<details>
<summary>Changelog</summary>
<p><em>Sourced from <a
href="https://github.com/jupyter/notebook/blob/@jupyter-notebook/tree@7.6.3/CHANGELOG.md">notebook's
changelog</a>.</em></p>
<blockquote>
<h2>7.6.3</h2>
<p>(<a
href="https://github.com/jupyter/notebook/compare/@jupyter-notebook/application-extension@7.6.2...9c8d7c24006554dda8c5fc02c69f91d6449c712a">Full
Changelog</a>)</p>
<h3>Maintenance and upkeep improvements</h3>
<ul>
<li>Update to JupyterLab v4.6.4 <a
href="https://redirect.github.com/jupyter/notebook/pull/8066">#8066</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
</ul>
<h3>Contributors to this release</h3>
<p>The following people contributed discussions, new ideas, code and
documentation contributions, and review.
See <a
href="https://github-activity.readthedocs.io/en/latest/use/#how-does-this-tool-define-contributions-in-the-reports">our
definition of contributors</a>.</p>
<p>(<a
href="https://github.com/jupyter/notebook/graphs/contributors?from=2026-08-11&amp;to=2026-09-21&amp;type=c">GitHub
contributors page for this release</a>)</p>
<p><a href="https://github.com/jtpio"><code>@​jtpio</code></a> (<a
href="https://github.com/search?q=repo%3Ajupyter%2Fnotebook+involves%3Ajtpio+updated%3A2026-08-11..2026-09-21&amp;type=Issues">activity</a>)</p>
<!-- raw HTML omitted -->
<h2>7.6.2</h2>
<p>(<a
href="https://github.com/jupyter/notebook/compare/@jupyter-notebook/application-extension@7.6.1...2f1f1622fc64e09b90ade87ed886fdfe1bc653fb">Full
Changelog</a>)</p>
<h3>Maintenance and upkeep improvements</h3>
<ul>
<li>Add nvd.nist.gov to the check_links ignore list <a
href="https://redirect.github.com/jupyter/notebook/pull/8032">#8032</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
<li>Backport JupyterLab dependency updater improvements <a
href="https://redirect.github.com/jupyter/notebook/pull/8031">#8031</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
<li>Update to JupyterLab v4.6.3 <a
href="https://redirect.github.com/jupyter/notebook/pull/8030">#8030</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
<li>Switch local pre-commit hooks to language: system <a
href="https://redirect.github.com/jupyter/notebook/pull/8002">#8002</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
</ul>
<h3>Documentation improvements</h3>
<ul>
<li>Update 'notebook v7 proposal' link in doc <a
href="https://redirect.github.com/jupyter/notebook/pull/8004">#8004</a>
(<a href="https://github.com/brichet"><code>@​brichet</code></a>)</li>
</ul>
<h3>Contributors to this release</h3>
<p>The following people contributed discussions, new ideas, code and
documentation contributions, and review.
See <a
href="https://github-activity.readthedocs.io/en/latest/use/#how-does-this-tool-define-contributions-in-the-reports">our
definition of contributors</a>.</p>
<p>(<a
href="https://github.com/jupyter/notebook/graphs/contributors?from=2026-07-22&amp;to=2026-08-11&amp;type=c">GitHub
contributors page for this release</a>)</p>
<p><a href="https://github.com/brichet"><code>@​brichet</code></a> (<a
href="https://github.com/search?q=repo%3Ajupyter%2Fnotebook+involves%3Abrichet+updated%3A2026-07-22..2026-08-11&amp;type=Issues">activity</a>)
| <a href="https://github.com/jtpio"><code>@​jtpio</code></a> (<a
href="https://github.com/search?q=repo%3Ajupyter%2Fnotebook+involves%3Ajtpio+updated%3A2026-07-22..2026-08-11&amp;type=Issues">activity</a>)</p>
<h2>7.6.1</h2>
<p>(<a
href="https://github.com/jupyter/notebook/compare/@jupyter-notebook/application-extension@7.6.0...72c228d76d1a9e5aa81531f7cf3c6af410c5e53f">Full
Changelog</a>)</p>
<h3>Maintenance and upkeep improvements</h3>
<ul>
<li>Update to JupyterLab v4.6.2 <a
href="https://redirect.github.com/jupyter/notebook/pull/7996">#7996</a>
(<a href="https://github.com/jtpio"><code>@​jtpio</code></a>)</li>
</ul>
<!-- raw HTML omitted -->
</blockquote>
<p>... (truncated)</p>
</details>
<details>
<summary>Commits</summary>
<ul>
<li><a
href="https://github.com/jupyter/notebook/commit/dc710245bf156f11bafc6e7d283695253d6b5cd7"><code>dc71024</code></a>
Publish 7.6.3</li>
<li><a
href="https://github.com/jupyter/notebook/commit/9c8d7c24006554dda8c5fc02c69f91d6449c712a"><code>9c8d7c2</code></a>
Update to JupyterLab v4.6.4 (<a
href="https://redirect.github.com/jupyter/notebook/issues/8066">#8066</a>)</li>
<li><a
href="https://github.com/jupyter/notebook/commit/ffc52152951a52ef4885f12521d7a5f8ebd2f9c1"><code>ffc5215</code></a>
Publish 7.6.2</li>
<li><a
href="https://github.com/jupyter/notebook/commit/2f1f1622fc64e09b90ade87ed886fdfe1bc653fb"><code>2f1f162</code></a>
Update to JupyterLab v4.6.3 (<a
href="https://redirect.github.com/jupyter/notebook/issues/8030">#8030</a>)</li>
<li><a
href="https://github.com/jupyter/notebook/commit/63df75281ec7226826877f2e77dc68670b109c33"><code>63df752</code></a>
Backport JupyterLab dependency updater improvements (<a
href="https://redirect.github.com/jupyter/notebook/issues/8031">#8031</a>)</li>
<li><a
href="https://github.com/jupyter/notebook/commit/66101117b79c92baf364d95647242f490807f18c"><code>6610111</code></a>
Backport PR <a
href="https://redirect.github.com/jupyter/notebook/issues/8032">#8032</a>:
Add nvd.nist.gov to the check_links ignore list (<a
href="https://redirect.github.com/jupyter/notebook/issues/8033">#8033</a>)</li>
<li><a
href="https://github.com/jupyter/notebook/commit/75a02cb28e7ed7a58358b69323be3d9e4800f941"><code>75a02cb</code></a>
Backport PR <a
href="https://redirect.github.com/jupyter/notebook/issues/8004">#8004</a>:
Update 'notebook v7 proposal' link in doc (<a
href="https://redirect.github.com/jupyter/notebook/issues/8005">#8005</a>)</li>
<li><a
href="https://github.com/jupyter/notebook/commit/2b0ebe6584d3f6d59d458a9414f00c985ead4feb"><code>2b0ebe6</code></a>
Backport PR <a
href="https://redirect.github.com/jupyter/notebook/issues/8002">#8002</a>:
Switch local pre-commit hooks to language: system (<a
href="https://redirect.github.com/jupyter/notebook/issues/8003">#8003</a>)</li>
<li><a
href="https://github.com/jupyter/notebook/commit/e058ab3a0d98a4698f201e10d7de723c70150d34"><code>e058ab3</code></a>
Publish 7.6.1</li>
<li><a
href="https://github.com/jupyter/notebook/commit/72c228d76d1a9e5aa81531f7cf3c6af410c5e53f"><code>72c228d</code></a>
Update to JupyterLab v4.6.2 (<a
href="https://redirect.github.com/jupyter/notebook/issues/7996">#7996</a>)</li>
<li>Additional commits viewable in <a
href="https://github.com/jupyter/notebook/compare/@jupyter-notebook/tree@7.5.7...@jupyter-notebook/tree@7.6.3">compare
view</a></li>
</ul>
</details>
<br />


[![Dependabot compatibility
score](https://dependabot-badges.githubapp.com/badges/compatibility_score?dependency-name=notebook&package-manager=uv&previous-version=7.5.7&new-version=7.6.3)](https://docs.github.com/en/github/managing-security-vulnerabilities/about-dependabot-security-updates#about-compatibility-scores)

Dependabot will resolve any conflicts with this PR as long as you don't
alter it yourself. You can also trigger a rebase manually by commenting
`@dependabot rebase`.

[//]: # (dependabot-automerge-start)
[//]: # (dependabot-automerge-end)

---

<details>
<summary>Dependabot commands and options</summary>
<br />

You can trigger Dependabot actions by commenting on this PR:
- `@dependabot rebase` will rebase this PR
- `@dependabot recreate` will recreate this PR, overwriting any edits
that have been made to it
- `@dependabot show <dependency name> ignore conditions` will show all
of the ignore conditions of the specified dependency
- `@dependabot ignore this major version` will close this PR and stop
Dependabot creating any more for this major version (unless you reopen
the PR or upgrade to it yourself)
- `@dependabot ignore this minor version` will close this PR and stop
Dependabot creating any more for this minor version (unless you reopen
the PR or upgrade to it yourself)
- `@dependabot ignore this dependency` will close this PR and stop
Dependabot creating any more for this dependency (unless you reopen the
PR or upgrade to it yourself)
You can disable automated security fix PRs for this repo from the
[Security Alerts
page](https://github.com/langchain-ai/langgraph/network/alerts).

</details>

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-10-07 17:21:37 +09:00
Elior Nataf LackritzandGitHub 39c523eb0a fix(langgraph): give each bulk_update_state update its own task id (#9128)
`bulk_update_state` stores each update's writes on the checkpoint it
updates, under a task id. An update whose node has no pending task to
reuse got `uuid5(checkpoint_id, INTERRUPT)`, so every such update in one
superstep shared an id. Savers keep one write per `(task_id, idx)`, so
all but the first update's writes were dropped.

Plain channels don't notice, since their value is stored in the new
checkpoint. A `DeltaChannel` rebuilds its value from those writes, so it
lost every update after the first, on memory, sqlite and postgres:

```python
graph.bulk_update_state(config, [[
    StateUpdate({"d": ["u1"], "p": ["u1"]}, as_node="a"),
    StateUpdate({"d": ["u2"], "p": ["u2"]}, as_node="m"),
]])
# plain: [..., 'u1', 'u2']
# delta: [..., 'u1']
```

The existing multi-update test only passed because it gave each update
an explicit `task_id`, and its docstring said so.

## Fix

The `i`th update gets `uuid5(checkpoint_id, f"{INTERRUPT}:{i}")`. The
first keeps the old id, so a single `update_state` stores exactly what
it did before. Updates that reuse a pending task's id, or pass their
own, are unchanged.

Each update also gets the task path `(__interrupt__, i)`, stored with
its writes. Live execution applies a superstep's updates in the order
given, and a saver that replays a checkpoint's writes by task path
(#8544) now gives them back in that order instead of by task id. Savers
whose `put_writes` takes no `task_path` get the old call.

## Limits

- On savers that order a checkpoint's writes by task id only (the
released ones, until #8544), several updates in one superstep still
replay in task id order.
- langgraphjs has the same shared id in `bulkUpdateState` and loses the
same update; fixed in langchain-ai/langgraphjs#2923 with the same ids.

## Tests

`test_delta_channel_update_state.py` adds sync, async and
next-to-a-pending-task cases across the checkpointer fixtures, with no
explicit task ids, on a thread that already has history (a fresh thread
snapshots every delta channel and hides the bug). All 19 fail on main.
Two more (sync and async) give six updates in one superstep and check
that `get_state` returns them in that order, on a saver that replays by
task path; they fail without the path.
2026-10-07 01:04:58 +00:00
3a1796ecf0 fix(langgraph): replay a resumed exit-mode run's delta writes in live order (#9114)
Depends on #8544

Resuming an exit-durability run after a parallel interrupt replays the
resumed checkpoint's `DeltaChannel` writes wrong: twice on `main`, out
of order if the duplicate is simply dropped. `p` finished before `q`
interrupted, and `r` runs after `q`:

```
live:            ['in1', 'p', 'r']
main, reload:    ['in1', 'p', 'r', 'p']
```

No fork is involved. Any exit-mode run that starts from a checkpoint
with pending delta writes hits it: resuming an interrupt, or retrying
after a task error.

## Cause

Exit durability stores a run's delta writes on the checkpoint it started
from, under step-prefixed task ids (`exit_delta_task_id`), so they
replay in superstep order. On a resume, that checkpoint already holds
the finished tasks' writes under their real task ids. The accumulator
stored those loaded writes again under a step-prefixed id, so they
replayed twice. Skipping them alone reorders the replay, because the
real ids sort after the step-prefixed ones.

## Fix

At exit:

- writes loaded with the checkpoint are skipped, since they are already
stored on it: `NULL_TASK_ID` writes by task id (this run's own `Command`
writes are recorded separately in `_first`), the rest by `(task_id,
channel)`;
- the checkpoint's own superstep keeps its real task paths, so it
interleaves with the loaded writes as live execution did, but gets a
step-prefixed task id. Under the real id, a resume whose final
checkpoint fails to save would leave the resumed task looking done, and
the retry would skip it and lose its other writes;
- a resume addressed by `checkpoint_id` reruns tasks whose writes are
already stored on the checkpoint; their writes to those channels are
skipped, and a write to a channel the task did not write before is kept;
- later supersteps get task id `ffffffff-<step>-...` and task path
`~~<step>...`, which sort after every real task id and every real task
path, in step order.

So superstep order holds whether a saver orders by task path (#8544) or
only by task id.

## Limits

- Within the resumed superstep, a saver that orders only by task id
replays this run's writes before the ones loaded from the checkpoint.
#8544 fixes that for the OSS savers.
- If the final checkpoint fails to save, the writes of supersteps after
the resumed one stay on the checkpoint and replay once more after the
retry (on `main`, every write does). That comes from exit durability
storing its writes before its checkpoint, not from this change.

## Tests

`test_delta_channel_exit_mode.py`, on the shared checkpointer fixture
(memory, sqlite, postgres) in all three durabilities:

- resume after a parallel interrupt with a later write, plain and with
the head's `checkpoint_id`
- the resumed superstep interleaved by task path: `a_asks` resumes after
`z_done` finished, and live order puts `a` first
- the same runs on a saver that orders writes by task id only
- a resume retried after its final checkpoint failed to save
- a resume with `Command(resume=..., update=...)` that writes the delta
channel replays that write once, in live order
- an addressed resume whose rerun task writes a delta channel it did not
write before keeps that write

The exit cases fail without the change, except the last one, which
guards against skipping by task id alone (an earlier version of this PR
did); sync and async pass either way.

#8548 marked exit durability in
`test_resume_on_an_interrupted_head_consumes_its_writes_without_a_snapshot`
as a strict expected failure for this bug; this PR drops the mark.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

---------

Co-authored-by: ErenAta16 <149434812+ErenAta16@users.noreply.github.com>
Co-authored-by: ragnarok268 <58264829+ragnarok268@users.noreply.github.com>
2026-10-06 22:05:55 +00:00
2b93e8b79f fix: order delta channel replay by task path (#8544)
Fixes langchain-ai/langgraph#8382

`DeltaChannel` rebuilds its value by replaying ancestor writes through
the reducer. Every saver ordered a checkpoint's writes by `(task_id,
idx)`, but live execution applies a super-step's writes in task-path
order (`apply_writes` sorts tasks by `task_path_str(task.path[:3])`).
`task_id` is a hash of the path, so when two or more tasks write one
`DeltaChannel` in the same super-step, replay applies them in an
arbitrary permutation.

Reducers must be batching-invariant, not order-invariant, so the
permutation changes the value: `get_state` disagrees with what `invoke`
returned, and continuing the thread saves the reordered value as the
base for every later write.

## Fix

Replay orders by `(task_path, task_id, idx)`, the order
`SELECT_PENDING_SENDS_SQL` already uses for sends. `apply_writes` and
every saver now sort with one `writes_sort_key` from
`langgraph.checkpoint.base`, in Python, so live and replay can't drift
apart and don't depend on a database collation. `apply_writes` sorts by
the full task path, the string `put_writes` stores, instead of
`path[:3]`. The paths that reach `apply_writes` no longer carry task
ids, so the live order doesn't change.

`BaseCheckpointSaver`'s default `get_delta_channel_history` can't sort
by path, since `PendingWrite` has no task path, so it replays each
checkpoint's writes in `get_tuple` order. `get_tuple` now has to return
`pending_writes` in `writes_sort_key` order, and memory, SQLite and
Postgres do, so the default walk is right for any saver that follows
that. `AsyncPostgresSaver` also gets the sync
`get_delta_channel_history` that `AsyncSqliteSaver` already had; sync
reads used to fall back to the default walk.

Memory and Postgres already stored `task_path`. SQLite accepted it on
`put_writes` and dropped it, so `writes` gains the column and `setup()`
adds it to existing databases.

Writes without a task path (graph input, `update_state` updates, rows
written before the column existed) sort first by `task_id`, which is
also where live execution applies input.

## Compatibility

- Memory and Postgres already had `task_path` on disk, so existing
threads replay in the corrected order after upgrade.
- On SQLite, rows from before the upgrade read back as `''` and keep
their old order. Only new writes get the corrected order, since there is
no stored path to recover.
- SQLite has no `ADD COLUMN IF NOT EXISTS`, so `setup()` runs the
`ALTER` and treats `duplicate column name` as success, which also covers
two processes migrating one file. That error comes while SQLite prepares
the statement, before it asks for the write lock, so `setup()` on an
up-to-date database doesn't wait on other writers. A read-only database
from before the column is read with empty paths instead of failing
setup.
- Downgrade is safe: older code's `INSERT` omits the column and the
column has a default.
- `get_tuple` and `list` return `pending_writes` in `writes_sort_key`
order instead of task id order (SQLite, Postgres) or put order (memory).
Nothing in langgraph depends on the old order.
- Postgres reads `task_path` and `idx` with each pending write, so a
subclass that pairs `SELECT_SQL` with its own `_load_writes`, or the
other way round, has to match the new row shape.
- A third-party saver has to replay in `writes_sort_key` order: in its
own `get_delta_channel_history` if it has one, otherwise by returning
`pending_writes` from `get_tuple` in that order. Otherwise it fails the
new conformance tests; the `get_tuple` one is a base test, so it runs
for every saver.
- langgraph, checkpoint-sqlite and checkpoint-postgres import
`writes_sort_key` from langgraph-checkpoint, so this bumps
langgraph-checkpoint to 4.3.0 and requires it there. 4.3.0 has to be
released first.

## Limits

- Several updates in one `bulk_update_state` super-step share a task id,
so the saver keeps only the first one's writes and the `DeltaChannel`
loses the rest (plain channels are fine). That's on main too; #9128
fixes it and stores a task path per update, so with this PR those
updates also replay in the order given.
- Exit durability stores its writes differently; resume ordering there
is fixed in #9114, stacked on this PR.
- langgraphjs applies `DeltaChannel` writes in task id order, live and
on replay (langchain-ai/langgraphjs#2544), so its reads match too.
Moving JS to task path order is a follow-up.

## Tests

Each fails on main and passes here:

- `test_delta_channel_parallel_order.py`, across memory, SQLite and
Postgres: parallel writers and a `Send` fan-out, comparing `get_state`
with the live result, including sync `get_state` on the async savers.
- Three conformance tests whose task ids sort opposite to their task
paths, so they cannot pass by chance: two on
`get_delta_channel_history`, one on `get_tuple`'s `pending_writes`
order.
- SQLite migration tests from a pre-column database, sync and async,
including a read-only file.

One more that passes on main too: SQLite `setup()` while another
connection holds the write lock, with no busy timeout, sync and async.
It keeps the `ALTER` from waiting on that lock.

Thanks to @ErenAta16 for the report, the reproduction, and for tracing
`task_path` through all three backends, and to @ragnarok268 for the
conformance-test design with task ids ordered opposite to their paths.

---------

Co-authored-by: ErenAta16 <149434812+ErenAta16@users.noreply.github.com>
Co-authored-by: ragnarok268 <58264829+ragnarok268@users.noreply.github.com>
2026-10-06 17:55:04 -04:00
71 changed files with 3864 additions and 409 deletions
@@ -267,6 +267,61 @@ async def test_history_seed_ancestor_own_writes_are_replayed(
)
# Every uuid4 `build_delta_chain` tags its own writes with sorts between these
# two, so task_id order is fixed and always disagrees with task_path order.
TASK_ID_SORTS_FIRST = "00000000-0000-0000-0000-000000000000"
TASK_ID_SORTS_LAST = "ffffffff-ffff-ffff-ffff-ffffffffffff"
async def test_history_orders_parallel_writes_by_task_path(
saver: BaseCheckpointSaver,
) -> None:
"""Writes from parallel tasks replay in task_path order, not task_id order."""
configs = await build_delta_chain(
saver,
thread_id=str(uuid4()),
channel="ch",
snapshots_at_steps=[0],
total_steps=3,
)
step_1, head = configs[1], configs[2]
await saver.aput_writes(
step_1, [("ch", "second")], TASK_ID_SORTS_FIRST, "~pull, 02"
)
await saver.aput_writes(step_1, [("ch", "first")], TASK_ID_SORTS_LAST, "~pull, 01")
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
values = [w[2] for w in result["ch"]["writes"]]
assert values == [1, "first", "second"], (
f"Expected task_path order [1, 'first', 'second'], got {values}. "
"Ordering by (task_id, idx) alone yields [1, 'second', 'first']."
)
async def test_history_orders_pathless_writes_first(
saver: BaseCheckpointSaver,
) -> None:
"""Writes stored without a task_path (graph input) replay before task writes."""
configs = await build_delta_chain(
saver,
thread_id=str(uuid4()),
channel="ch",
snapshots_at_steps=[0],
total_steps=3,
)
step_1, head = configs[1], configs[2]
await saver.aput_writes(
step_1, [("ch", "from_node")], TASK_ID_SORTS_FIRST, "~pull, a"
)
await saver.aput_writes(step_1, [("ch", "from_input")], TASK_ID_SORTS_LAST)
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
values = [w[2] for w in result["ch"]["writes"]]
assert values == [1, "from_input", "from_node"], (
f"Expected pathless writes first, got {values}"
)
ALL_DELTA_CHANNEL_HISTORY_TESTS = [
test_history_returns_writes_oldest_first,
test_history_seed_is_nearest_snapshot,
@@ -276,6 +331,8 @@ ALL_DELTA_CHANNEL_HISTORY_TESTS = [
test_history_walk_to_root_no_seed,
test_history_migration_plain_value_as_seed,
test_history_seed_ancestor_own_writes_are_replayed,
test_history_orders_parallel_writes_by_task_path,
test_history_orders_pathless_writes_first,
]
@@ -176,6 +176,35 @@ async def test_get_tuple_pending_writes(saver: BaseCheckpointSaver) -> None:
)
async def test_get_tuple_pending_writes_in_writes_sort_key_order(
saver: BaseCheckpointSaver,
) -> None:
"""pending_writes come back in writes_sort_key order, not put or task_id order."""
config = generate_config(str(uuid4()))
stored = await saver.aput(config, generate_checkpoint(), generate_metadata(), {})
await saver.aput_writes(
stored, [("ch", "b")], "00000000-0000-0000-0000-000000000000", "~pull, 02"
)
await saver.aput_writes(
stored,
[("ch", "a1"), ("ch", "a2")],
"ffffffff-ffff-ffff-ffff-ffffffffffff",
"~pull, 01",
)
await saver.aput_writes(
stored, [("ch", "input")], "88888888-8888-8888-8888-888888888888"
)
tup = await saver.aget_tuple(stored)
assert tup is not None
values = [w[2] for w in tup.pending_writes or []]
assert values == ["input", "a1", "a2", "b"], (
f"Expected writes_sort_key order ['input', 'a1', 'a2', 'b'], got {values}. "
"Put order gives ['b', 'a1', 'a2', 'input'], task_id order gives "
"['b', 'input', 'a1', 'a2']."
)
async def test_get_tuple_respects_namespace(saver: BaseCheckpointSaver) -> None:
"""checkpoint_ns filtering."""
tid = str(uuid4())
@@ -223,6 +252,7 @@ ALL_GET_TUPLE_TESTS = [
test_get_tuple_metadata,
test_get_tuple_parent_config,
test_get_tuple_pending_writes,
test_get_tuple_pending_writes_in_writes_sort_key_order,
test_get_tuple_respects_namespace,
test_get_tuple_nonexistent_checkpoint_id,
]
@@ -14,8 +14,8 @@ async def memory_checkpointer():
@pytest.mark.asyncio
async def test_validate_memory_base():
"""InMemorySaver passes all base capability tests."""
async def test_validate_memory():
"""InMemorySaver passes the tests of every capability it implements."""
report = await validate(memory_checkpointer)
report.print_report()
assert report.passed_all_base(), f"Base tests failed: {report.to_dict()}"
assert report.passed_all(), f"Capability tests failed: {report.to_dict()}"
+1 -1
View File
@@ -279,7 +279,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.2.0"
version = "4.3.0"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -451,7 +451,8 @@ class PostgresSaver(BasePostgresSaver):
* Stage 1 (paged): dynamic SELECT over `checkpoints` with three
columns per channel: its version, an `EXISTS` probe for a stored
blob at that version, and its inline value. Pages newest-first by
`checkpoint_id` with a cursor; page size is `_DELTA_PAGE_SIZE`.
`checkpoint_id`, starting at the target; page size is
`_DELTA_PAGE_SIZE`.
Stops paging when every channel has found its seed or a page comes
back short.
@@ -468,7 +469,7 @@ class PostgresSaver(BasePostgresSaver):
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = get_checkpoint_id(config)
if checkpoint_id is None:
if not checkpoint_id:
target = self.get_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
@@ -476,7 +477,7 @@ class PostgresSaver(BasePostgresSaver):
# Stage 1: paged K-JSONB-lookup scan, walking the parent chain in
# Python after each page. Stops as soon as every channel has its seed.
stage1_sql = _build_delta_stage1_sql(channels, paged=True)
stage1_sql = _build_delta_stage1_sql(channels, paged=True, include_cursor=True)
parent_of: dict[str, str | None] = {}
ver_by_i_by_cid: list[dict[str, str | None]] = [{} for _ in channels]
hb_by_i_by_cid: list[dict[str, bool]] = [{} for _ in channels]
@@ -486,7 +487,7 @@ class PostgresSaver(BasePostgresSaver):
seed_inline_by_ch: dict[str, Any] = {}
walk_cursor_by_ch: dict[str, str | None] = {}
seeded: set[str] = set()
cursor: str | None = None
cursor: str | None = checkpoint_id
with self._cursor() as cur:
while True:
@@ -495,7 +496,7 @@ class PostgresSaver(BasePostgresSaver):
# ver_i, blob channel, blob version, inline_i
stage1_params.extend([ch, ch, ch, ch])
stage1_params.extend(
[thread_id, checkpoint_ns, cursor, cursor, _DELTA_PAGE_SIZE]
[thread_id, checkpoint_ns, cursor, _DELTA_PAGE_SIZE]
)
cur.execute(stage1_sql, stage1_params)
page = cur.fetchall()
@@ -527,6 +528,7 @@ class PostgresSaver(BasePostgresSaver):
if len(seeded) == len(channels) or len(page) < _DELTA_PAGE_SIZE:
break
cursor = oldest
stage1_sql = _build_delta_stage1_sql(channels, paged=True)
# Stage 2: per-channel UNION ALL — one writes branch per channel
# with non-empty chain, plus one blob branch per seeded channel.
@@ -417,13 +417,13 @@ class AsyncPostgresSaver(BasePostgresSaver):
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = get_checkpoint_id(config)
if checkpoint_id is None:
if not checkpoint_id:
target = await self.aget_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
checkpoint_id = target.config["configurable"]["checkpoint_id"]
stage1_sql = _build_delta_stage1_sql(channels, paged=True)
stage1_sql = _build_delta_stage1_sql(channels, paged=True, include_cursor=True)
parent_of: dict[str, str | None] = {}
ver_by_i_by_cid: list[dict[str, str | None]] = [{} for _ in channels]
hb_by_i_by_cid: list[dict[str, bool]] = [{} for _ in channels]
@@ -433,7 +433,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
seed_inline_by_ch: dict[str, Any] = {}
walk_cursor_by_ch: dict[str, str | None] = {}
seeded: set[str] = set()
cursor: str | None = None
cursor: str | None = checkpoint_id
async with self._cursor() as cur:
while True:
@@ -442,7 +442,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
# ver_i, blob channel, blob version, inline_i
stage1_params.extend([ch, ch, ch, ch])
stage1_params.extend(
[thread_id, checkpoint_ns, cursor, cursor, _DELTA_PAGE_SIZE]
[thread_id, checkpoint_ns, cursor, _DELTA_PAGE_SIZE]
)
await cur.execute(stage1_sql, stage1_params)
page = await cur.fetchall()
@@ -472,6 +472,7 @@ class AsyncPostgresSaver(BasePostgresSaver):
if len(seeded) == len(channels) or len(page) < _DELTA_PAGE_SIZE:
break
cursor = oldest
stage1_sql = _build_delta_stage1_sql(channels, paged=True)
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
channels_with_seed = [ch for ch in channels if seed_ver_by_ch[ch] is not None]
@@ -687,5 +688,23 @@ class AsyncPostgresSaver(BasePostgresSaver):
self.adelete_thread(thread_id), self.loop
).result()
def get_delta_channel_history(
self, *, config: RunnableConfig, channels: Sequence[str]
) -> Mapping[str, DeltaChannelHistory]:
"""Sync bridge to `aget_delta_channel_history`, guarded like `get_tuple`."""
try:
if asyncio.get_running_loop() is self.loop:
raise asyncio.InvalidStateError(
"Synchronous calls to AsyncPostgresSaver are only allowed from a "
"different thread. From the main thread, use the async interface. "
"For example, use `await checkpointer.aget_delta_channel_history(...)`."
)
except RuntimeError:
pass
return asyncio.run_coroutine_threadsafe(
self.aget_delta_channel_history(config=config, channels=channels),
self.loop,
).result()
__all__ = ["AsyncPostgresSaver", "AsyncShallowPostgresSaver", "Conn"]
@@ -14,6 +14,7 @@ from langgraph.checkpoint.base import (
DeltaChannelHistory,
PendingWrite,
get_checkpoint_id,
writes_sort_key,
)
from langgraph.checkpoint.serde.types import TASKS
from psycopg.types.json import Jsonb
@@ -109,7 +110,7 @@ select
) as channel_values,
(
select
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob, convert_to(cw.task_path, 'UTF8'), cw.idx::text::bytea])
from checkpoint_writes cw
where cw.thread_id = checkpoints.thread_id
and cw.checkpoint_ns = checkpoints.checkpoint_ns
@@ -168,6 +169,7 @@ class _DeltaStage2Row(TypedDict, total=False):
type: str | None
blob: bytes | None
task_id: str | None # "w" rows only
task_path: str | None # "w" rows only
idx: int | None # "w" rows only
version: str | None # "b" rows only
@@ -176,10 +178,13 @@ class _DeltaStage2Row(TypedDict, total=False):
# `_build_delta_stage2_sql` document their shapes.
def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
def _build_delta_stage1_sql(
channels: Sequence[str], *, paged: bool, include_cursor: bool = False
) -> str:
"""Build stage 1 SQL with K parallel version lookups + seed probes.
For channels=["messages", "files"] (with `paged=True`) the result is::
For channels=["messages", "files"] (with `paged=True, include_cursor=True`)
the result is::
SELECT checkpoint_id, parent_checkpoint_id,
checkpoint -> 'channel_versions' ->> %s AS ver_0,
@@ -195,7 +200,7 @@ def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
checkpoint -> 'channel_values' -> %s AS inline_1
FROM checkpoints
WHERE thread_id = %s AND checkpoint_ns = %s
AND (%s::text IS NULL OR checkpoint_id < %s)
AND checkpoint_id <= %s
ORDER BY checkpoint_id DESC
LIMIT %s
@@ -236,9 +241,16 @@ def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
and uses safe identifiers).
Caller must extend params with `[ch_0 x4, ch_1 x4, ..., thread_id, ns,
cursor, cursor, page_size]` when `paged=True` — four per channel: the
version lookup, the blob's channel, the version the blob must match, and the
inline lookup.
cursor, page_size]` when `paged=True` — four per channel: the version
lookup, the blob's channel, the version the blob must match, and the inline
lookup.
Pages run newest-first from the target down. The first page passes the
target as the cursor with `include_cursor=True`, so it opens with the
target's own row, whose parent starts the walk; each later page continues
below the oldest row read. A checkpoint's ancestors have smaller ids (uuid6
is time-ordered, which `get_tuple` also relies on to find the latest
checkpoint), so no row newer than the target is part of its chain.
When `paged=False`, the WHERE has no cursor predicate and there's no
LIMIT/ORDER BY — kept as a non-public helper for tests/diagnostics.
@@ -262,7 +274,7 @@ def _build_delta_stage1_sql(channels: Sequence[str], *, paged: bool) -> str:
)
if paged:
sql += (
" AND (%s::text IS NULL OR checkpoint_id < %s)"
f" AND checkpoint_id {'<=' if include_cursor else '<'} %s"
" ORDER BY checkpoint_id DESC LIMIT %s"
)
return sql
@@ -297,7 +309,7 @@ def _build_delta_stage2_sql(
branches.append(
"SELECT 'w'::text AS _kind, "
"checkpoint_id, channel, "
"type, blob, task_id, idx, NULL::text AS version "
"type, blob, task_id, task_path, idx, NULL::text AS version "
"FROM checkpoint_writes "
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
"AND checkpoint_id = ANY(%s)"
@@ -305,7 +317,8 @@ def _build_delta_stage2_sql(
for _ in channels_with_seed:
branches.append(
"SELECT 'b'::text AS _kind, NULL::text AS checkpoint_id, channel, "
"type, blob, NULL::text AS task_id, NULL::int AS idx, version "
"type, blob, NULL::text AS task_id, NULL::text AS task_path, "
"NULL::int AS idx, version "
"FROM checkpoint_blobs "
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
"AND version = %s"
@@ -410,8 +423,8 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
materialized at this point),
(c) the next ancestor cid isn't in `parent_of` yet (waiting for
a later page; the cursor stays put), or
(d) the target's own row isn't in `parent_of` yet (the walk has
not started; no cursor is set, so a later page retries).
(d) the target's own row isn't in `parent_of` (the target doesn't
exist, so the walk never starts).
Mutates `chain_by_ch`, `seed_ver_by_ch`, `seed_inline_by_ch`,
`walk_cursor_by_ch`, and `seeded` in place.
@@ -419,8 +432,9 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
for i, ch in enumerate(channels):
if ch in seeded:
continue
# Pages start at the thread head, so the target may not have
# loaded yet; a `None` cursor would read as "target is a root".
# The first page opens with the target's row, so it's missing only
# when the target doesn't exist; a `None` cursor would read as
# "target is a root".
if ch not in walk_cursor_by_ch:
if target_id not in parent_of:
continue
@@ -473,10 +487,11 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
stored value, or when the seed blob is sentinel "empty" — in both cases
the consumer treats absence as "start empty".
"""
# writes_by_ch_by_cid[channel][cid] = list of (type, blob, task_id, idx)
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
ch: {} for ch in channels
}
# writes_by_ch_by_cid[channel][cid] = list of
# (type, blob, task_id, idx, task_path)
writes_by_ch_by_cid: dict[
str, dict[str, list[tuple[str, bytes, str, int, str]]]
] = {ch: {} for ch in channels}
# seed_blob_by_ver[(channel, version)] = (type, blob)
seed_blob_by_ver: dict[tuple[str, str], tuple[str, bytes]] = {}
@@ -487,8 +502,14 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
cid = cast(str, r["checkpoint_id"])
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
cast(
"tuple[str, bytes, str, int]",
(r["type"], r["blob"], r["task_id"], r["idx"]),
"tuple[str, bytes, str, int, str]",
(
r["type"],
r["blob"],
r["task_id"],
r["idx"],
r["task_path"],
),
)
)
else: # kind == "b"
@@ -497,10 +518,10 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
"tuple[str, bytes]", (r["type"], r["blob"])
)
# Sort writes per (channel, cid) newest-first by (task_id, idx)
# Sort writes per (channel, cid) newest-first
for cid_map in writes_by_ch_by_cid.values():
for ws in cid_map.values():
ws.sort(key=lambda w: (w[2], w[3]), reverse=True)
ws.sort(key=lambda w: writes_sort_key(w[4], w[2], w[3]), reverse=True)
result: dict[str, DeltaChannelHistory] = {}
for ch in channels:
@@ -510,7 +531,9 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
collected: list[PendingWrite] = []
cid_writes = writes_by_ch_by_cid.get(ch, {})
for cid in chain_cids:
for type_tag, write_blob, task_id, _idx in cid_writes.get(cid, []):
for type_tag, write_blob, task_id, _idx, _path in cid_writes.get(
cid, []
):
val = self.serde.loads_typed((type_tag, write_blob))
collected.append((task_id, ch, val))
collected.reverse()
@@ -553,20 +576,15 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
]
def _load_writes(
self, writes: list[tuple[bytes, bytes, bytes, bytes]]
self, writes: list[tuple[bytes, bytes, bytes, bytes, bytes, bytes]] | None
) -> list[tuple[str, str, Any]]:
return (
[
(
tid.decode(),
channel.decode(),
self.serde.loads_typed((t.decode(), v)),
)
for tid, channel, t, v in writes
]
if writes
else []
)
return [
(tid.decode(), channel.decode(), self.serde.loads_typed((t.decode(), v)))
for tid, channel, t, v, _, _ in sorted(
writes or [],
key=lambda w: writes_sort_key(w[4].decode(), w[0].decode(), int(w[5])),
)
]
def _dump_writes(
self,
@@ -97,7 +97,7 @@ select
) as channel_values,
(
select
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob, convert_to(cw.task_path, 'UTF8'), cw.idx::text::bytea])
from checkpoint_writes cw
where cw.thread_id = checkpoints.thread_id
and cw.checkpoint_ns = checkpoints.checkpoint_ns
@@ -36,6 +36,7 @@ from langgraph.store.base import (
ensure_embeddings,
get_text_at_path,
tokenize_path,
validate_op_namespace,
)
from psycopg import Capabilities, Connection, Cursor, Pipeline
from psycopg.rows import DictRow, dict_row
@@ -1386,6 +1387,7 @@ def _group_ops(ops: Iterable[Op]) -> tuple[dict[type, list[tuple[int, Op]]], int
grouped_ops: dict[type, list[tuple[int, Op]]] = defaultdict(list)
tot = 0
for idx, op in enumerate(ops):
validate_op_namespace(op)
grouped_ops[type(op)].append((idx, op))
tot += 1
return grouped_ops, tot
+1 -1
View File
@@ -12,7 +12,7 @@ readme = "README.md"
license = "MIT"
license-files = ['LICENSE']
dependencies = [
"langgraph-checkpoint>=4.1.0,<5.0.0",
"langgraph-checkpoint>=4.3.0,<5.0.0",
"orjson>=3.11.5",
"psycopg>=3.2.0",
"psycopg-pool>=3.2.0",
@@ -13,8 +13,10 @@ import pytest
from langchain_core.embeddings import Embeddings
from langgraph.store.base import (
GetOp,
InvalidNamespaceError,
Item,
ListNamespacesOp,
MatchCondition,
PutOp,
SearchOp,
)
@@ -871,3 +873,65 @@ async def test_omit_expired_search_pagination(store: AsyncPostgresStore) -> None
page2 = await store.asearch(ns, limit=2, offset=2)
assert [i.key for i in page1] == ["a", "b"]
assert [i.key for i in page2] == ["c"]
@pytest.mark.parametrize("namespace", [("foo.bar",), ("foo", ""), ("foo", 1)])
async def test_abatch_rejects_invalid_namespace_labels(
store: AsyncPostgresStore, namespace: tuple
) -> None:
await store.aput(("foo", "bar"), "key", {"original": True})
for op in (
GetOp(namespace, "key"),
GetOp(namespace, "key", refresh_ttl=True),
PutOp(namespace, "key", {"changed": True}),
PutOp(namespace, "key", None),
SearchOp(namespace),
ListNamespacesOp((MatchCondition("prefix", namespace),)),
ListNamespacesOp((MatchCondition("suffix", namespace),)),
):
with pytest.raises(InvalidNamespaceError):
await store.abatch([op])
item = await store.aget(("foo", "bar"), "key")
assert item is not None and item.value == {"original": True}
async def test_invalid_namespace_only_fails_its_own_call(
store: AsyncPostgresStore,
) -> None:
"""Concurrent calls share one `abatch`, which fails every op if it raises.
Labels are checked before an op is queued, so one caller's bad label cannot
fail another caller's request.
"""
await store.aput(("foo", "bar"), "key", {"original": True})
valid, invalid = await asyncio.gather(
store.aget(("foo", "bar"), "key"),
store.aget(("foo.bar",), "key"),
return_exceptions=True,
)
assert isinstance(valid, Item) and valid.value == {"original": True}
assert isinstance(invalid, InvalidNamespaceError)
async def test_sync_methods_reject_invalid_namespace_labels(
store: AsyncPostgresStore,
) -> None:
"""The sync wrappers run off the event loop thread and must validate too."""
await store.aput(("foo", "bar"), "key", {"original": True})
for call in (
lambda: store.get(("foo.bar",), "key"),
lambda: store.search(("foo.bar",)),
lambda: store.delete(("foo.bar",), "key"),
lambda: store.list_namespaces(prefix=("foo.bar",)),
lambda: store.batch([GetOp(("foo.bar",), "key")]),
):
with pytest.raises(InvalidNamespaceError):
await asyncio.to_thread(call)
item = await store.aget(("foo", "bar"), "key")
assert item is not None and item.value == {"original": True}
@@ -1,5 +1,6 @@
from __future__ import annotations
from collections.abc import Mapping, Sequence
from typing import Any
from uuid import uuid4
@@ -14,7 +15,7 @@ from langgraph.checkpoint.serde.types import _DeltaSnapshot
from langgraph.checkpoint.postgres import PostgresSaver
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
from langgraph.checkpoint.postgres.base import _DELTA_PAGE_SIZE
from langgraph.checkpoint.postgres.base import _DELTA_PAGE_SIZE, BasePostgresSaver
from tests.conftest import DEFAULT_URI
CHANNEL = "items"
@@ -23,8 +24,8 @@ SEED_STEP = 1
SEED_VALUE = [10, 20]
TARGET_STEP = 4
# The real page size is the control; the rest leave the target off the first
# page (three checkpoints are newer than it).
# The real page size is the control; the rest split the walk from the target
# to its seed across pages.
PAGE_SIZES = [_DELTA_PAGE_SIZE, 3, 2, 1]
@@ -92,7 +93,7 @@ def _assert_history(entry: DeltaChannelHistory, page_size: int) -> None:
@pytest.mark.parametrize("page_size", PAGE_SIZES)
async def test_async_target_older_than_the_first_page(
async def test_async_walk_continues_across_pages(
page_size: int, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr("langgraph.checkpoint.postgres.aio._DELTA_PAGE_SIZE", page_size)
@@ -106,7 +107,7 @@ async def test_async_target_older_than_the_first_page(
@pytest.mark.parametrize("page_size", PAGE_SIZES)
def test_sync_target_older_than_the_first_page(
def test_sync_walk_continues_across_pages(
page_size: int, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr("langgraph.checkpoint.postgres._DELTA_PAGE_SIZE", page_size)
@@ -119,6 +120,84 @@ def test_sync_target_older_than_the_first_page(
_assert_history(result[CHANNEL], page_size)
def _record_rows_read(monkeypatch: pytest.MonkeyPatch) -> list[str]:
read: list[str] = []
ingest = BasePostgresSaver._ingest_stage1_page
def record(rows: Sequence[Mapping[str, Any]], *args: Any) -> str | None:
read.extend(row["checkpoint_id"] for row in rows)
return ingest(rows, *args)
monkeypatch.setattr(BasePostgresSaver, "_ingest_stage1_page", staticmethod(record))
return read
def _ids_from_target_down(configs: list[dict]) -> list[str]:
return [c["configurable"]["checkpoint_id"] for c in configs[TARGET_STEP::-1]]
async def test_async_walk_reads_nothing_newer_than_the_target(
monkeypatch: pytest.MonkeyPatch,
) -> None:
read = _record_rows_read(monkeypatch)
async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver:
await saver.setup()
configs = await _abuild_chain(saver)
result = await saver.aget_delta_channel_history(
config=configs[TARGET_STEP], channels=[CHANNEL]
)
_assert_history(result[CHANNEL], _DELTA_PAGE_SIZE)
assert read == _ids_from_target_down(configs)
def test_sync_walk_reads_nothing_newer_than_the_target(
monkeypatch: pytest.MonkeyPatch,
) -> None:
read = _record_rows_read(monkeypatch)
with PostgresSaver.from_conn_string(DEFAULT_URI) as saver:
saver.setup()
configs = _build_chain(saver)
result = saver.get_delta_channel_history(
config=configs[TARGET_STEP], channels=[CHANNEL]
)
_assert_history(result[CHANNEL], _DELTA_PAGE_SIZE)
assert read == _ids_from_target_down(configs)
def _empty_checkpoint_id(config: dict) -> dict:
return {"configurable": {**config["configurable"], "checkpoint_id": ""}}
async def test_async_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
async with AsyncPostgresSaver.from_conn_string(DEFAULT_URI) as saver:
await saver.setup()
configs = await _abuild_chain(saver)
latest = await saver.aget_delta_channel_history(
config=configs[-1], channels=[CHANNEL]
)
result = await saver.aget_delta_channel_history(
config=_empty_checkpoint_id(configs[-1]), channels=[CHANNEL]
)
assert latest[CHANNEL]["writes"]
assert result == latest
def test_sync_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
with PostgresSaver.from_conn_string(DEFAULT_URI) as saver:
saver.setup()
configs = _build_chain(saver)
latest = saver.get_delta_channel_history(config=configs[-1], channels=[CHANNEL])
result = saver.get_delta_channel_history(
config=_empty_checkpoint_id(configs[-1]), channels=[CHANNEL]
)
assert latest[CHANNEL]["writes"]
assert result == latest
async def test_root_target_has_no_history_and_still_terminates(
monkeypatch: pytest.MonkeyPatch,
) -> None:
@@ -11,6 +11,7 @@ import pytest
from langchain_core.embeddings import Embeddings
from langgraph.store.base import (
GetOp,
InvalidNamespaceError,
Item,
ListNamespacesOp,
MatchCondition,
@@ -1164,3 +1165,51 @@ def test_namespace_labels_with_trailing_newline(store) -> None:
assert set(store.list_namespaces(prefix=["users", "alice"], limit=100)) == {
("users", "alice"),
}
@pytest.mark.parametrize("namespace", [("foo.bar",), ("foo", ""), ("foo", 1)])
@pytest.mark.parametrize(
"kind", ["get", "put", "delete", "search", "list_prefix", "list_suffix"]
)
def test_batch_rejects_invalid_namespace_labels(
store, namespace: tuple, kind: str
) -> None:
"""Ops passed straight to `batch` must not reach another namespace.
Namespaces are stored dot-joined, so `("foo.bar",)` flattens to the same
text as `("foo", "bar")`. `BaseStore` methods validate labels themselves,
but `batch` takes ops as given.
"""
op = {
"get": GetOp(namespace, "key"),
"put": PutOp(namespace, "key", {"changed": True}),
"delete": PutOp(namespace, "key", None),
"search": SearchOp(namespace),
"list_prefix": ListNamespacesOp((MatchCondition("prefix", namespace),)),
"list_suffix": ListNamespacesOp((MatchCondition("suffix", namespace),)),
}[kind]
store.put(("foo", "bar"), "key", {"original": True})
with pytest.raises(InvalidNamespaceError):
store.batch([PutOp(("valid",), "key", {}), op])
item = store.get(("foo", "bar"), "key")
assert item is not None and item.value == {"original": True}
# The whole batch is rejected before any SQL runs.
assert store.get(("valid",), "key") is None
def test_batch_allows_empty_search_prefix_and_listing_wildcards(
store,
) -> None:
store.put(("foo", "bar"), "key", {"v": 1})
found, listed = store.batch(
[
SearchOp(()),
ListNamespacesOp((MatchCondition("prefix", ("foo", "*")),)),
]
)
assert [item.namespace for item in found] == [("foo", "bar")]
assert listed == [("foo", "bar")]
+1 -1
View File
@@ -278,7 +278,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.2.0"
version = "4.3.0"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -29,7 +29,11 @@ from langgraph.checkpoint.sqlite._delta import (
build_delta_stage2_sql,
step_walk_with_row,
)
from langgraph.checkpoint.sqlite.utils import search_where
from langgraph.checkpoint.sqlite.utils import (
load_pending_writes,
pending_writes_sql,
search_where,
)
_AIO_ERROR_MSG = (
"The SqliteSaver does not support async methods. "
@@ -81,6 +85,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
conn: sqlite3.Connection
is_setup: bool
_has_task_path: bool = True
def __init__(
self,
@@ -154,6 +159,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
checkpoint_ns TEXT NOT NULL DEFAULT '',
checkpoint_id TEXT NOT NULL,
task_id TEXT NOT NULL,
task_path TEXT NOT NULL DEFAULT '',
idx INTEGER NOT NULL,
channel TEXT NOT NULL,
type TEXT,
@@ -162,6 +168,19 @@ class SqliteSaver(BaseCheckpointSaver[str]):
);
"""
)
# sqlite has no ADD COLUMN IF NOT EXISTS; this migrates databases
# created before `task_path` existed and is a no-op on the rest.
try:
self.conn.execute(
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
)
except sqlite3.OperationalError as e:
# A read-only database from before the column can still be read;
# its rows would all read back as '' anyway.
if "readonly database" in str(e):
self._has_task_path = False
elif "duplicate column name" not in str(e):
raise
self.is_setup = True
@@ -260,7 +279,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
}
# find any pending writes
cur.execute(
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
pending_writes_sql(self._has_task_path),
(
str(config["configurable"]["thread_id"]),
checkpoint_ns,
@@ -286,10 +305,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
if parent_checkpoint_id
else None
),
[
(task_id, channel, self.serde.loads_typed((type, value)))
for task_id, channel, type, value in cur
],
load_pending_writes(cur, self.serde),
)
def list(
@@ -351,7 +367,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
metadata,
) in cur:
wcur.execute(
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
pending_writes_sql(self._has_task_path),
(thread_id, checkpoint_ns, checkpoint_id),
)
yield CheckpointTuple(
@@ -378,10 +394,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
if parent_checkpoint_id
else None
),
[
(task_id, channel, self.serde.loads_typed((type, value)))
for task_id, channel, type, value in wcur
],
load_pending_writes(wcur, self.serde),
)
def put(
@@ -460,9 +473,9 @@ class SqliteSaver(BaseCheckpointSaver[str]):
task_path: Path of the task creating the writes.
"""
query = (
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
if all(w[0] in WRITES_IDX_MAP for w in writes)
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
)
with self.cursor() as cur:
cur.executemany(
@@ -473,6 +486,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
str(config["configurable"]["checkpoint_ns"]),
str(config["configurable"]["checkpoint_id"]),
task_id,
task_path,
WRITES_IDX_MAP.get(channel, idx),
channel,
*self.serde.dumps_typed(value),
@@ -525,7 +539,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
thread_id = str(config["configurable"]["thread_id"])
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = get_checkpoint_id(config)
if checkpoint_id is None:
if not checkpoint_id:
target = self.get_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
@@ -559,6 +573,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
stage2_sql = build_delta_stage2_sql(
has_task_path=self._has_task_path,
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
)
if stage2_sql:
@@ -569,7 +584,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
)
cur.execute(stage2_sql, stage2_params)
stage2_rows = cast(
"list[tuple[str, str, str, int, str, bytes]]", cur.fetchall()
"list[tuple[str, str, str, int, str, bytes, str]]", cur.fetchall()
)
else:
stage2_rows = []
@@ -24,7 +24,11 @@ from __future__ import annotations
from collections.abc import Mapping, Sequence
from typing import Any
from langgraph.checkpoint.base import DeltaChannelHistory, PendingWrite
from langgraph.checkpoint.base import (
DeltaChannelHistory,
PendingWrite,
writes_sort_key,
)
# Stage 1 streams target, then its ancestors nearest-first, by following
# `parent_checkpoint_id` rather than id order: ids are only monotonic within
@@ -56,7 +60,9 @@ DELTA_STAGE1_SQL = (
)
def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
def build_delta_stage2_sql(
*, chain_lens: Sequence[int], has_task_path: bool = True
) -> str:
"""Stage-2 per-channel UNION ALL fetching writes from `writes`.
One branch per channel with a non-empty chain. Each branch inlines its
@@ -70,11 +76,12 @@ def build_delta_stage2_sql(*, chain_lens: Sequence[int]) -> str:
of a single `channel = ANY(channels)` filter when channels have
different chain depths — same rationale as postgres.
"""
task_path = "task_path" if has_task_path else "''"
branches: list[str] = []
for n in chain_lens:
cid_placeholders = ",".join("?" * n)
branches.append(
"SELECT checkpoint_id, channel, task_id, idx, type, value "
f"SELECT checkpoint_id, channel, task_id, idx, type, value, {task_path} "
"FROM writes "
"WHERE thread_id = ? AND checkpoint_ns = ? AND channel = ? "
f"AND checkpoint_id IN ({cid_placeholders})"
@@ -141,29 +148,31 @@ def build_delta_channels_writes_history(
chain_by_ch: Mapping[str, list[str]],
seed_val_by_ch: Mapping[str, Any],
seeded: set[str],
stage2_rows: Sequence[tuple[str, str, str, int, str, bytes]],
stage2_rows: Sequence[tuple[str, str, str, int, str, bytes, str]],
serde: Any,
) -> dict[str, DeltaChannelHistory]:
"""Demux stage-2 rows per channel; produce per-channel histories.
Stage-2 rows are `(checkpoint_id, channel, task_id, idx, type, value)`.
Final write order is oldest→newest globally and `(task_id, idx)` within
a checkpoint, matching the contract on `DeltaChannelHistory.writes`.
Stage-2 rows are
`(checkpoint_id, channel, task_id, idx, type, value, task_path)`.
Final write order is oldest→newest globally and `writes_sort_key`
within a checkpoint, matching the contract on
`DeltaChannelHistory.writes`.
`seed` is omitted when the walk reached a true root with no snapshot
found (channel never entered `seeded`); consumers treat absence as
"start empty".
"""
writes_by_ch_by_cid: dict[str, dict[str, list[tuple[str, bytes, str, int]]]] = {
ch: {} for ch in channels
}
for cid, ch, task_id, idx, type_tag, value_blob in stage2_rows:
writes_by_ch_by_cid: dict[
str, dict[str, list[tuple[str, bytes, str, int, str]]]
] = {ch: {} for ch in channels}
for cid, ch, task_id, idx, type_tag, value_blob, task_path in stage2_rows:
writes_by_ch_by_cid.setdefault(ch, {}).setdefault(cid, []).append(
(type_tag, value_blob, task_id, idx)
(type_tag, value_blob, task_id, idx, task_path)
)
for cid_map in writes_by_ch_by_cid.values():
for ws in cid_map.values():
ws.sort(key=lambda w: (w[2], w[3]))
ws.sort(key=lambda w: writes_sort_key(w[4], w[2], w[3]))
result: dict[str, DeltaChannelHistory] = {}
for ch in channels:
@@ -172,7 +181,7 @@ def build_delta_channels_writes_history(
collected: list[PendingWrite] = []
# Chain is newest-first; iterate oldest-first for the public order.
for cid in reversed(chain_cids):
for type_tag, value_blob, task_id, _idx in cid_writes.get(cid, []):
for type_tag, value_blob, task_id, _idx, _path in cid_writes.get(cid, []):
collected.append(
(task_id, ch, serde.loads_typed((type_tag, value_blob)))
)
@@ -30,7 +30,11 @@ from langgraph.checkpoint.sqlite._delta import (
build_delta_stage2_sql,
step_walk_with_row,
)
from langgraph.checkpoint.sqlite.utils import search_where
from langgraph.checkpoint.sqlite.utils import (
load_pending_writes,
pending_writes_sql,
search_where,
)
T = TypeVar("T", bound=Callable)
@@ -114,6 +118,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
lock: asyncio.Lock
is_setup: bool
_has_task_path: bool = True
def __init__(
self,
@@ -331,6 +336,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
checkpoint_ns TEXT NOT NULL DEFAULT '',
checkpoint_id TEXT NOT NULL,
task_id TEXT NOT NULL,
task_path TEXT NOT NULL DEFAULT '',
idx INTEGER NOT NULL,
channel TEXT NOT NULL,
type TEXT,
@@ -341,6 +347,21 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
):
await self.conn.commit()
# sqlite has no ADD COLUMN IF NOT EXISTS; this migrates databases
# created before `task_path` existed and is a no-op on the rest.
try:
await self.conn.execute(
"ALTER TABLE writes ADD COLUMN task_path TEXT NOT NULL DEFAULT ''"
)
await self.conn.commit()
except aiosqlite.OperationalError as e:
# A read-only database from before the column can still be read;
# its rows would all read back as '' anyway.
if "readonly database" in str(e):
self._has_task_path = False
elif "duplicate column name" not in str(e):
raise
self.is_setup = True
async def aget_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
@@ -395,7 +416,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
}
# find any pending writes
await cur.execute(
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
pending_writes_sql(self._has_task_path),
(
str(config["configurable"]["thread_id"]),
checkpoint_ns,
@@ -421,10 +442,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
if parent_checkpoint_id
else None
),
[
(task_id, channel, self.serde.loads_typed((type, value)))
async for task_id, channel, type, value in cur
],
load_pending_writes(await cur.fetchall(), self.serde),
)
async def alist(
@@ -473,7 +491,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
metadata,
) in cur:
await wcur.execute(
"SELECT task_id, channel, type, value FROM writes WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ? ORDER BY task_id, idx",
pending_writes_sql(self._has_task_path),
(thread_id, checkpoint_ns, checkpoint_id),
)
yield CheckpointTuple(
@@ -500,10 +518,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
if parent_checkpoint_id
else None
),
[
(task_id, channel, self.serde.loads_typed((type, value)))
async for task_id, channel, type, value in wcur
],
load_pending_writes(await wcur.fetchall(), self.serde),
)
async def aput(
@@ -576,9 +591,9 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
task_path: Path of the task creating the writes.
"""
query = (
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
"INSERT OR REPLACE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
if all(w[0] in WRITES_IDX_MAP for w in writes)
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?)"
else "INSERT OR IGNORE INTO writes (thread_id, checkpoint_ns, checkpoint_id, task_id, task_path, idx, channel, type, value) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)"
)
await self.setup()
async with self.lock, self.conn.cursor() as cur:
@@ -590,6 +605,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
str(config["configurable"]["checkpoint_ns"]),
str(config["configurable"]["checkpoint_id"]),
task_id,
task_path,
WRITES_IDX_MAP.get(channel, idx),
channel,
*self.serde.dumps_typed(value),
@@ -637,7 +653,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
thread_id = str(config["configurable"]["thread_id"])
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = get_checkpoint_id(config)
if checkpoint_id is None:
if not checkpoint_id:
target = await self.aget_tuple(config)
if target is None:
return {ch: {"writes": []} for ch in channels}
@@ -671,6 +687,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
channels_with_chain = [ch for ch in channels if chain_by_ch[ch]]
stage2_sql = build_delta_stage2_sql(
has_task_path=self._has_task_path,
chain_lens=[len(chain_by_ch[ch]) for ch in channels_with_chain],
)
if stage2_sql:
@@ -681,7 +698,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
)
await cur.execute(stage2_sql, stage2_params)
stage2_rows = cast(
"list[tuple[str, str, str, int, str, bytes]]",
"list[tuple[str, str, str, int, str, bytes, str]]",
await cur.fetchall(),
)
else:
@@ -2,11 +2,12 @@ from __future__ import annotations
import json
import re
from collections.abc import Sequence
from collections.abc import Iterable, Sequence
from typing import Any
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import get_checkpoint_id
from langgraph.checkpoint.base import PendingWrite, get_checkpoint_id, writes_sort_key
from langgraph.checkpoint.serde.base import SerializerProtocol
_FILTER_PATTERN = re.compile(r"^[a-zA-Z0-9_.-]+$")
@@ -114,3 +115,23 @@ def search_where(
param_values.append(get_checkpoint_id(before))
return ("WHERE " + " AND ".join(wheres) if wheres else "", param_values)
def pending_writes_sql(has_task_path: bool) -> str:
task_path = "task_path" if has_task_path else "''"
return (
f"SELECT task_id, channel, type, value, {task_path}, idx FROM writes "
"WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ?"
)
def load_pending_writes(
rows: Iterable[Any], serde: SerializerProtocol
) -> list[PendingWrite]:
"""Deserialize `pending_writes_sql` rows in `writes_sort_key` order."""
return [
(task_id, channel, serde.loads_typed((type_, value)))
for task_id, channel, type_, value, _, _ in sorted(
rows, key=lambda r: writes_sort_key(r[4], r[0], r[5])
)
]
@@ -28,6 +28,7 @@ from langgraph.store.base import (
ensure_embeddings,
get_text_at_path,
tokenize_path,
validate_op_namespace,
)
_AIO_ERROR_MSG = (
@@ -257,6 +258,7 @@ def _group_ops(ops: Iterable[Op]) -> tuple[dict[type, list[tuple[int, Op]]], int
grouped_ops: dict[type, list[tuple[int, Op]]] = defaultdict(list)
tot = 0
for idx, op in enumerate(ops):
validate_op_namespace(op)
grouped_ops[type(op)].append((idx, op))
tot += 1
return grouped_ops, tot
+1 -1
View File
@@ -12,7 +12,7 @@ readme = "README.md"
license = "MIT"
license-files = ['LICENSE']
dependencies = [
"langgraph-checkpoint>=4.1.0,<5.0.0",
"langgraph-checkpoint>=4.3.0,<5.0.0",
"aiosqlite>=0.20",
"sqlite-vec>=0.1.6",
]
@@ -9,8 +9,10 @@ from typing import cast
import pytest
from langgraph.store.base import (
GetOp,
InvalidNamespaceError,
Item,
ListNamespacesOp,
MatchCondition,
PutOp,
SearchOp,
)
@@ -745,3 +747,65 @@ async def test_async_namespace_segment_boundary(store: AsyncSqliteStore) -> None
assert set(await store.alist_namespaces(suffix=["alice"], limit=100)) == {
("uid", "users", "alice"),
}
@pytest.mark.parametrize("namespace", [("foo.bar",), ("foo", ""), ("foo", 1)])
async def test_abatch_rejects_invalid_namespace_labels(
store: AsyncSqliteStore, namespace: tuple
) -> None:
await store.aput(("foo", "bar"), "key", {"original": True})
for op in (
GetOp(namespace, "key"),
GetOp(namespace, "key", refresh_ttl=True),
PutOp(namespace, "key", {"changed": True}),
PutOp(namespace, "key", None),
SearchOp(namespace),
ListNamespacesOp((MatchCondition("prefix", namespace),)),
ListNamespacesOp((MatchCondition("suffix", namespace),)),
):
with pytest.raises(InvalidNamespaceError):
await store.abatch([op])
item = await store.aget(("foo", "bar"), "key")
assert item is not None and item.value == {"original": True}
async def test_invalid_namespace_only_fails_its_own_call(
store: AsyncSqliteStore,
) -> None:
"""Concurrent calls share one `abatch`, which fails every op if it raises.
Labels are checked before an op is queued, so one caller's bad label cannot
fail another caller's request.
"""
await store.aput(("foo", "bar"), "key", {"original": True})
valid, invalid = await asyncio.gather(
store.aget(("foo", "bar"), "key"),
store.aget(("foo.bar",), "key"),
return_exceptions=True,
)
assert isinstance(valid, Item) and valid.value == {"original": True}
assert isinstance(invalid, InvalidNamespaceError)
async def test_sync_methods_reject_invalid_namespace_labels(
store: AsyncSqliteStore,
) -> None:
"""The sync wrappers run off the event loop thread and must validate too."""
await store.aput(("foo", "bar"), "key", {"original": True})
for call in (
lambda: store.get(("foo.bar",), "key"),
lambda: store.search(("foo.bar",)),
lambda: store.delete(("foo.bar",), "key"),
lambda: store.list_namespaces(prefix=("foo.bar",)),
lambda: store.batch([GetOp(("foo.bar",), "key")]),
):
with pytest.raises(InvalidNamespaceError):
await asyncio.to_thread(call)
item = await store.aget(("foo", "bar"), "key")
assert item is not None and item.value == {"original": True}
@@ -65,6 +65,39 @@ async def test_async_walk_reaches_parent_whatever_the_id_order(
assert got[CHANNEL] == EXPECTED
EMPTY_CHECKPOINT_ID: dict[str, Any] = {
"configurable": {**CONFIG["configurable"], "checkpoint_id": ""}
}
def test_sync_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
with SqliteSaver.from_conn_string(":memory:") as saver:
root = saver.put(CONFIG, _checkpoint("a-older", {CHANNEL: "seed"}), {}, {})
saver.put_writes(root, [(CHANNEL, "write-root")], "task")
saver.put(root, _checkpoint("z-newer", {}), {}, {})
got = saver.get_delta_channel_history(
config=EMPTY_CHECKPOINT_ID, channels=[CHANNEL]
)
assert got[CHANNEL] == EXPECTED
async def test_async_empty_checkpoint_id_reads_from_the_latest_checkpoint() -> None:
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
root = await saver.aput(
CONFIG, _checkpoint("a-older", {CHANNEL: "seed"}), {}, {}
)
await saver.aput_writes(root, [(CHANNEL, "write-root")], "task")
await saver.aput(root, _checkpoint("z-newer", {}), {}, {})
got = await saver.aget_delta_channel_history(
config=EMPTY_CHECKPOINT_ID, channels=[CHANNEL]
)
assert got[CHANNEL] == EXPECTED
def test_walk_reaches_root_of_long_chain_with_descending_ids() -> None:
steps = 40
with SqliteSaver.from_conn_string(":memory:") as saver:
@@ -14,6 +14,7 @@ import pytest
from langchain_core.embeddings import Embeddings
from langgraph.store.base import (
GetOp,
InvalidNamespaceError,
Item,
ListNamespacesOp,
MatchCondition,
@@ -1435,3 +1436,51 @@ def test_list_namespaces_metacharacter_labels(store: SqliteStore) -> None:
assert set(store.list_namespaces(prefix=[label, "child"], limit=100)) == {
(label, "child"),
}
@pytest.mark.parametrize("namespace", [("foo.bar",), ("foo", ""), ("foo", 1)])
@pytest.mark.parametrize(
"kind", ["get", "put", "delete", "search", "list_prefix", "list_suffix"]
)
def test_batch_rejects_invalid_namespace_labels(
store: SqliteStore, namespace: tuple, kind: str
) -> None:
"""Ops passed straight to `batch` must not reach another namespace.
Namespaces are stored dot-joined, so `("foo.bar",)` flattens to the same
text as `("foo", "bar")`. `BaseStore` methods validate labels themselves,
but `batch` takes ops as given.
"""
op = {
"get": GetOp(namespace, "key"),
"put": PutOp(namespace, "key", {"changed": True}),
"delete": PutOp(namespace, "key", None),
"search": SearchOp(namespace),
"list_prefix": ListNamespacesOp((MatchCondition("prefix", namespace),)),
"list_suffix": ListNamespacesOp((MatchCondition("suffix", namespace),)),
}[kind]
store.put(("foo", "bar"), "key", {"original": True})
with pytest.raises(InvalidNamespaceError):
store.batch([PutOp(("valid",), "key", {}), op])
item = store.get(("foo", "bar"), "key")
assert item is not None and item.value == {"original": True}
# The whole batch is rejected before any SQL runs.
assert store.get(("valid",), "key") is None
def test_batch_allows_empty_search_prefix_and_listing_wildcards(
store: SqliteStore,
) -> None:
store.put(("foo", "bar"), "key", {"v": 1})
found, listed = store.batch(
[
SearchOp(()),
ListNamespacesOp((MatchCondition("prefix", ("foo", "*")),)),
]
)
assert [item.namespace for item in found] == [("foo", "bar")]
assert listed == [("foo", "bar")]
@@ -0,0 +1,151 @@
import sqlite3
from collections.abc import Iterator
from contextlib import closing
from pathlib import Path
import aiosqlite
import pytest
from langgraph.checkpoint.base import empty_checkpoint
from langgraph.checkpoint.sqlite import SqliteSaver
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
WRITES_BEFORE_TASK_PATH = """
CREATE TABLE writes (
thread_id TEXT NOT NULL,
checkpoint_ns TEXT NOT NULL DEFAULT '',
checkpoint_id TEXT NOT NULL,
task_id TEXT NOT NULL,
idx INTEGER NOT NULL,
channel TEXT NOT NULL,
type TEXT,
value BLOB,
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
);
INSERT INTO writes VALUES ('t', '', 'c', 'old-task', 0, 'ch', 'null', X'');
"""
@pytest.fixture
def legacy_db(tmp_path: Path) -> Path:
db = tmp_path / "legacy.sqlite"
with sqlite3.connect(db) as conn:
conn.executescript(WRITES_BEFORE_TASK_PATH)
return db
def test_setup_migrates_legacy_writes_table_repeatably(legacy_db: Path) -> None:
for _ in range(2):
with SqliteSaver.from_conn_string(str(legacy_db)) as saver:
saver.setup()
rows = saver.conn.execute(
"SELECT task_id, task_path FROM writes"
).fetchall()
assert rows == [("old-task", "")]
@pytest.mark.parametrize("fresh", [True, False], ids=["fresh", "legacy"])
def test_put_writes_persists_task_path(
tmp_path: Path, legacy_db: Path, fresh: bool
) -> None:
db = tmp_path / "fresh.sqlite" if fresh else legacy_db
with SqliteSaver.from_conn_string(str(db)) as saver:
config = saver.put(
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
empty_checkpoint(),
{},
{},
)
saver.put_writes(config, [("ch", "v")], "task-1", "~__pregel_pull, node")
stored = saver.conn.execute(
"SELECT task_path FROM writes WHERE task_id = 'task-1'"
).fetchall()
assert stored == [("~__pregel_pull, node",)]
async def test_async_setup_migrates_legacy_writes_table_repeatably(
legacy_db: Path,
) -> None:
for _ in range(2):
async with AsyncSqliteSaver.from_conn_string(str(legacy_db)) as saver:
await saver.setup()
config = await saver.aput(
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
empty_checkpoint(),
{},
{},
)
await saver.aput_writes(
config, [("ch", "v")], "task-1", "~__pregel_pull, node"
)
async with aiosqlite.connect(legacy_db) as conn:
async with conn.execute(
"SELECT DISTINCT task_id, task_path FROM writes ORDER BY task_id"
) as cur:
assert await cur.fetchall() == [
("old-task", ""),
("task-1", "~__pregel_pull, node"),
]
@pytest.fixture
def busy_db(tmp_path: Path) -> Iterator[Path]:
db = tmp_path / "busy.sqlite"
with SqliteSaver.from_conn_string(str(db)) as saver:
saver.setup()
with closing(sqlite3.connect(db, isolation_level=None)) as writer:
writer.execute("BEGIN IMMEDIATE")
yield db
writer.execute("ROLLBACK")
def test_setup_does_not_wait_on_another_writer(busy_db: Path) -> None:
with closing(sqlite3.connect(busy_db, timeout=0)) as conn:
SqliteSaver(conn).setup()
async def test_async_setup_does_not_wait_on_another_writer(busy_db: Path) -> None:
async with aiosqlite.connect(busy_db, timeout=0) as conn:
await AsyncSqliteSaver(conn).setup()
def _legacy_database_with_history(db: Path) -> dict:
root = empty_checkpoint()
root["channel_values"] = {"ch": "seed"}
root["channel_versions"] = {"ch": 1}
with SqliteSaver.from_conn_string(str(db)) as saver:
root_config = saver.put(
{"configurable": {"thread_id": "t", "checkpoint_ns": ""}},
root,
{},
{"ch": 1},
)
saver.put_writes(root_config, [("ch", "write")], "task", "~__pregel_pull, n")
child = saver.put(root_config, empty_checkpoint(), {}, {})
saver.conn.execute("ALTER TABLE writes DROP COLUMN task_path")
saver.conn.commit()
return child
def test_read_only_legacy_database_still_reads_delta_history(tmp_path: Path) -> None:
db = tmp_path / "legacy.sqlite"
child = _legacy_database_with_history(db)
saver = SqliteSaver(sqlite3.connect(f"file:{db}?mode=ro", uri=True))
got = saver.get_delta_channel_history(config=child, channels=["ch"])
assert got["ch"] == {"seed": "seed", "writes": [("task", "ch", "write")]}
async def test_async_read_only_legacy_database_still_reads_delta_history(
tmp_path: Path,
) -> None:
db = tmp_path / "legacy.sqlite"
child = _legacy_database_with_history(db)
async with aiosqlite.connect(f"file:{db}?mode=ro", uri=True) as conn:
saver = AsyncSqliteSaver(conn)
got = await saver.aget_delta_channel_history(config=child, channels=["ch"])
assert got["ch"] == {"seed": "seed", "writes": [("task", "ch", "write")]}
+1 -1
View File
@@ -285,7 +285,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.2.0"
version = "4.3.0"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -162,6 +162,14 @@ class DeltaChannelHistory(TypedDict):
Always present; possibly empty. Already filtered to one channel.
Writes stored at the target checkpoint itself are pending for the
next super-step and are excluded.
Within a single checkpoint, writes are ordered by `writes_sort_key`,
`(task_path, task_id, idx)`, which is the order live execution applies
a super-step's task writes in. `task_id` is a hash of the path, so
ordering by it permutes parallel tasks writing one channel, and
reducers need not be order-invariant. Writes stored without a
`task_path` (graph input, `update_state` and `Command` updates, rows
predating the column) sort first, by `task_id`.
* `seed` — the stored value at the nearest ancestor whose
`channel_values[ch]` is populated. Omitted if the walk reached the
root without finding any stored value (consumer treats absence as
@@ -174,6 +182,20 @@ class DeltaChannelHistory(TypedDict):
seed: NotRequired[Any]
def writes_sort_key(
task_path: str, task_id: str = "", idx: int = 0
) -> tuple[str, str, int]:
"""Sort key for the writes of one super-step.
Live execution applies a super-step's tasks in this order, so a saver
that replays stored writes, as `get_delta_channel_history` does, must
sort them by it too, or an order-sensitive reducer rebuilds a different
value than the run produced. `task_path` is the string passed to
`put_writes`.
"""
return (task_path, task_id, idx)
class BaseCheckpointSaver(Generic[V]):
"""Base class for creating a graph checkpointer.
@@ -244,7 +266,8 @@ class BaseCheckpointSaver(Generic[V]):
config: Configuration specifying which checkpoint to retrieve.
Returns:
The requested checkpoint tuple, or `None` if not found.
The requested checkpoint tuple, or `None` if not found. Its
`pending_writes` must be in `writes_sort_key` order.
Raises:
NotImplementedError: Implement this method in your custom checkpoint saver.
@@ -434,7 +457,8 @@ class BaseCheckpointSaver(Generic[V]):
config: Configuration specifying which checkpoint to retrieve.
Returns:
The requested checkpoint tuple, or `None` if not found.
The requested checkpoint tuple, or `None` if not found. Its
`pending_writes` must be in `writes_sort_key` order.
Raises:
NotImplementedError: Implement this method in your custom checkpoint saver.
@@ -611,6 +635,10 @@ class BaseCheckpointSaver(Generic[V]):
`PostgresSaver`) override for performance; the return contract is
fixed here.
`PendingWrite` carries no `task_path`, so this default replays each
checkpoint's writes in `get_tuple`'s `pending_writes` order, which
`get_tuple` must return in `writes_sort_key` order.
Args:
config: Configuration identifying the target checkpoint.
channels: Channel names to walk for. Empty → empty mapping.
@@ -25,6 +25,7 @@ from langgraph.checkpoint.base import (
SerializerProtocol,
get_checkpoint_id,
get_checkpoint_metadata,
writes_sort_key,
)
logger = logging.getLogger(__name__)
@@ -139,6 +140,15 @@ class InMemorySaver(
result[k] = self.serde.loads_typed(vv)
return result
def _ordered_writes(
self, thread_id: str, checkpoint_ns: str, checkpoint_id: str
) -> list[tuple[str, str, tuple[str, bytes], str]]:
stored = self.writes.get((thread_id, checkpoint_ns, checkpoint_id), {})
return [
stored[k]
for k in sorted(stored, key=lambda k: writes_sort_key(stored[k][3], *k))
]
def get_delta_channel_history(
self, *, config: RunnableConfig, channels: Sequence[str]
) -> Mapping[str, DeltaChannelHistory]:
@@ -160,8 +170,8 @@ class InMemorySaver(
return {}
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = config["configurable"].get("checkpoint_id", "")
ns_storage = self.storage.get(thread_id, {}).get(checkpoint_ns, {})
checkpoint_id = get_checkpoint_id(config) or max(ns_storage, default="")
chain: list[str] = []
target_entry = ns_storage.get(checkpoint_id)
@@ -198,9 +208,8 @@ class InMemorySaver(
blob_value_by_ch[ch] = self.serde.loads_typed(blob_entry)
terminated_here.add(ch)
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
for (_task_id, _idx), (tid, ch, serialized, _) in sorted(
step_writes.items(), reverse=True
for tid, ch, serialized, _ in reversed(
self._ordered_writes(thread_id, checkpoint_ns, cp_id)
):
if ch not in remaining:
continue
@@ -246,7 +255,7 @@ class InMemorySaver(
if checkpoint_id := get_checkpoint_id(config):
if saved := self.storage[thread_id][checkpoint_ns].get(checkpoint_id):
checkpoint, metadata, parent_checkpoint_id = saved
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
writes = self._ordered_writes(thread_id, checkpoint_ns, checkpoint_id)
checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)
return CheckpointTuple(
config=config,
@@ -276,7 +285,7 @@ class InMemorySaver(
if checkpoints := self.storage[thread_id][checkpoint_ns]:
checkpoint_id = max(checkpoints.keys())
checkpoint, metadata, parent_checkpoint_id = checkpoints[checkpoint_id]
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
writes = self._ordered_writes(thread_id, checkpoint_ns, checkpoint_id)
checkpoint_ = self.serde.loads_typed(checkpoint)
return CheckpointTuple(
config={
@@ -379,9 +388,9 @@ class InMemorySaver(
elif limit is not None:
limit -= 1
writes = self.writes[
(thread_id, checkpoint_ns, checkpoint_id)
].values()
writes = self._ordered_writes(
thread_id, checkpoint_ns, checkpoint_id
)
checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)
@@ -55,6 +55,7 @@ logger = logging.getLogger(__name__)
_MAX_WARNED_TYPES = 1000
_warned_unregistered_types: set[tuple[str, str]] = set()
_warned_blocked_types: set[tuple[str, str]] = set()
_warned_unreconstructable_types: set[tuple[str, str]] = set()
def _is_safe_json_type(id_list: list[str]) -> bool:
@@ -79,6 +80,27 @@ def _warn_once(
logger.warning(msg, *args)
def _reconstruction_fallback(tup: Any, exc: Exception) -> Any:
"""Return the serialized payload of an object that could not be rebuilt.
Returning `None` here would silently erase the value from restored state.
"""
try:
module, name, payload = tup[0], tup[1], tup[2]
except Exception:
return None
_warn_once(
_warned_unreconstructable_types,
(str(module), str(name)),
"Could not reconstruct %s.%s from checkpoint (%s); "
"returning its serialized data instead.",
module,
name,
type(exc).__name__,
)
return payload
class JsonPlusSerializer(SerializerProtocol):
"""Serializer that uses ormsgpack, with optional fallbacks.
@@ -638,6 +660,7 @@ def _create_msgpack_ext_hook(
)
)
elif code == EXT_CONSTRUCTOR_SINGLE_ARG:
tup = None
try:
tup = ormsgpack.unpackb(
data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
@@ -649,9 +672,10 @@ def _create_msgpack_ext_hook(
return tup[2]
# module, name, arg
return getattr(importlib.import_module(tup[0]), tup[1])(tup[2])
except Exception:
return None
except Exception as exc:
return _reconstruction_fallback(tup, exc)
elif code == EXT_CONSTRUCTOR_POS_ARGS:
tup = None
try:
tup = ormsgpack.unpackb(
data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
@@ -662,9 +686,10 @@ def _create_msgpack_ext_hook(
return _send_from_args(tup[2])
# module, name, args
return getattr(importlib.import_module(tup[0]), tup[1])(*tup[2])
except Exception:
return None
except Exception as exc:
return _reconstruction_fallback(tup, exc)
elif code == EXT_CONSTRUCTOR_KW_ARGS:
tup = None
try:
tup = ormsgpack.unpackb(
data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
@@ -673,9 +698,10 @@ def _create_msgpack_ext_hook(
return tup[2]
# module, name, kwargs
return getattr(importlib.import_module(tup[0]), tup[1])(**tup[2])
except Exception:
return None
except Exception as exc:
return _reconstruction_fallback(tup, exc)
elif code == EXT_METHOD_SINGLE_ARG:
tup = None
try:
tup = ormsgpack.unpackb(
data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
@@ -686,8 +712,8 @@ def _create_msgpack_ext_hook(
return getattr(
getattr(importlib.import_module(tup[0]), tup[1]), tup[3]
)(tup[2])
except Exception:
return None
except Exception as exc:
return _reconstruction_fallback(tup, exc)
elif code == EXT_PYDANTIC_V1:
try:
tup = ormsgpack.unpackb(
@@ -771,7 +771,12 @@ class BaseStore(ABC):
Returns:
The retrieved item or `None` if not found.
Raises:
InvalidNamespaceError: If a namespace label is empty, is not a string,
or contains a period (`.`).
"""
_validate_namespace_labels(namespace)
return self.batch(
[GetOp(namespace, str(key), _ensure_refresh(self.ttl_config, refresh_ttl))]
)[0]
@@ -801,6 +806,10 @@ class BaseStore(ABC):
Returns:
List of items matching the search criteria.
Raises:
InvalidNamespaceError: If a `namespace_prefix` label is empty, is not a
string, or contains a period (`.`).
???+ example "Examples"
Basic filtering:
@@ -840,6 +849,7 @@ class BaseStore(ABC):
Natural language search support depends on your store implementation
and requires proper embedding configuration.
"""
_validate_namespace_labels(namespace_prefix)
return self.batch(
[
SearchOp(
@@ -887,6 +897,11 @@ class BaseStore(ABC):
By default, the expiration timer refreshes on both read operations (get/search)
and write operations (put/update), whenever the item is included in the operation.
Raises:
InvalidNamespaceError: If the namespace is empty, its root label is
`"langgraph"`, or a label is empty, is not a string, or contains a
period (`.`).
Note:
Indexing support depends on your store implementation.
If you do not initialize the store with indexing capabilities,
@@ -940,7 +955,12 @@ class BaseStore(ABC):
Args:
namespace: Hierarchical path for the item.
key: Unique identifier within the namespace.
Raises:
InvalidNamespaceError: If a namespace label is empty, is not a string,
or contains a period (`.`).
"""
_validate_namespace_labels(namespace)
self.batch([PutOp(namespace, str(key), None, ttl=None)])
def list_namespaces(
@@ -969,6 +989,10 @@ class BaseStore(ABC):
A list of namespace tuples that match the criteria. Each tuple represents a
full namespace path up to `max_depth`.
Raises:
InvalidNamespaceError: If a `prefix` or `suffix` label is empty, is not a
string, or contains a period (`.`).
???+ example "Examples":
Setting `max_depth=3`. Given the namespaces:
@@ -984,6 +1008,8 @@ class BaseStore(ABC):
# [("a", "b", "c"), ("a", "b", "d"), ("a", "b", "f")]
```
"""
_validate_namespace_labels(prefix or ())
_validate_namespace_labels(suffix or ())
match_conditions = []
if prefix:
match_conditions.append(MatchCondition(match_type="prefix", path=prefix))
@@ -1013,7 +1039,12 @@ class BaseStore(ABC):
Returns:
The retrieved item or `None` if not found.
Raises:
InvalidNamespaceError: If a namespace label is empty, is not a string,
or contains a period (`.`).
"""
_validate_namespace_labels(namespace)
return (
await self.abatch(
[
@@ -1052,6 +1083,10 @@ class BaseStore(ABC):
Returns:
List of items matching the search criteria.
Raises:
InvalidNamespaceError: If a `namespace_prefix` label is empty, is not a
string, or contains a period (`.`).
???+ example "Examples"
Basic filtering:
@@ -1091,6 +1126,7 @@ class BaseStore(ABC):
Natural language search support depends on your store implementation
and requires proper embedding configuration.
"""
_validate_namespace_labels(namespace_prefix)
return (
await self.abatch(
[
@@ -1140,6 +1176,11 @@ class BaseStore(ABC):
By default, the expiration timer refreshes on both read operations (get/search)
and write operations (put/update), whenever the item is included in the operation.
Raises:
InvalidNamespaceError: If the namespace is empty, its root label is
`"langgraph"`, or a label is empty, is not a string, or contains a
period (`.`).
Note:
Indexing support depends on your store implementation.
If you do not initialize the store with indexing capabilities,
@@ -1201,7 +1242,12 @@ class BaseStore(ABC):
Args:
namespace: Hierarchical path for the item.
key: Unique identifier within the namespace.
Raises:
InvalidNamespaceError: If a namespace label is empty, is not a string,
or contains a period (`.`).
"""
_validate_namespace_labels(namespace)
await self.abatch([PutOp(namespace, str(key), None)])
async def alist_namespaces(
@@ -1230,6 +1276,10 @@ class BaseStore(ABC):
A list of namespace tuples that match the criteria. Each tuple represents a
full namespace path up to `max_depth`.
Raises:
InvalidNamespaceError: If a `prefix` or `suffix` label is empty, is not a
string, or contains a period (`.`).
???+ example "Examples"
Setting `max_depth=3` with existing namespaces:
@@ -1245,6 +1295,8 @@ class BaseStore(ABC):
# Returns: [("a", "b", "c"), ("a", "b", "d"), ("a", "b", "f")]
```
"""
_validate_namespace_labels(prefix or ())
_validate_namespace_labels(suffix or ())
match_conditions = []
if prefix:
match_conditions.append(MatchCondition(match_type="prefix", path=prefix))
@@ -1263,6 +1315,14 @@ class BaseStore(ABC):
def _validate_namespace(namespace: tuple[str, ...]) -> None:
if not namespace:
raise InvalidNamespaceError("Namespace cannot be empty.")
_validate_namespace_labels(namespace)
if namespace[0] == "langgraph":
raise InvalidNamespaceError(
f'Root label for namespace cannot be "langgraph". Got: {namespace}'
)
def _validate_namespace_labels(namespace: tuple[str, ...]) -> None:
for label in namespace:
if not isinstance(label, str):
raise InvalidNamespaceError(
@@ -1277,10 +1337,27 @@ def _validate_namespace(namespace: tuple[str, ...]) -> None:
raise InvalidNamespaceError(
f"Namespace labels cannot be empty strings. Got {label} in {namespace}"
)
if namespace[0] == "langgraph":
raise InvalidNamespaceError(
f'Root label for namespace cannot be "langgraph". Got: {namespace}'
)
def validate_op_namespace(op: Op) -> None:
"""Validate the namespace labels an op carries before a store executes it.
`BaseStore` methods check labels before batching, but ops passed directly to
`batch`/`abatch` skip those methods. Stores that serialize namespaces as
delimited text should call this for every op they execute, so a label such
as `"foo.bar"` cannot address the namespace `("foo", "bar")`.
Raises:
InvalidNamespaceError: If a label is empty, is not a string, or contains
a period (`.`).
"""
if isinstance(op, (GetOp, PutOp)):
_validate_namespace_labels(op.namespace)
elif isinstance(op, SearchOp):
_validate_namespace_labels(op.namespace_prefix)
elif isinstance(op, ListNamespacesOp):
for condition in op.match_conditions or ():
_validate_namespace_labels(condition.path)
def _ensure_refresh(
@@ -1319,4 +1396,5 @@ __all__ = [
"ensure_embeddings",
"tokenize_path",
"get_text_at_path",
"validate_op_namespace",
]
@@ -25,6 +25,7 @@ from langgraph.store.base import (
_ensure_refresh,
_ensure_ttl,
_validate_namespace,
_validate_namespace_labels,
)
F = TypeVar("F", bound=Callable)
@@ -86,6 +87,7 @@ class AsyncBatchedBaseStore(BaseStore):
*,
refresh_ttl: bool | None = None,
) -> Item | None:
_validate_namespace_labels(namespace)
self._ensure_task()
fut = self._loop.create_future()
self._aqueue.put_nowait(
@@ -111,6 +113,7 @@ class AsyncBatchedBaseStore(BaseStore):
offset: int = 0,
refresh_ttl: bool | None = None,
) -> list[SearchItem]:
_validate_namespace_labels(namespace_prefix)
self._ensure_task()
fut = self._loop.create_future()
self._aqueue.put_nowait(
@@ -155,6 +158,7 @@ class AsyncBatchedBaseStore(BaseStore):
namespace: tuple[str, ...],
key: str,
) -> None:
_validate_namespace_labels(namespace)
self._ensure_task()
fut = self._loop.create_future()
self._aqueue.put_nowait((fut, PutOp(namespace, key, None)))
@@ -169,6 +173,8 @@ class AsyncBatchedBaseStore(BaseStore):
limit: int = 100,
offset: int = 0,
) -> list[tuple[str, ...]]:
_validate_namespace_labels(prefix or ())
_validate_namespace_labels(suffix or ())
self._ensure_task()
fut = self._loop.create_future()
match_conditions = []
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-checkpoint"
version = "4.2.0"
version = "4.3.0"
description = "Library with base interfaces for LangGraph checkpoint savers."
authors = []
requires-python = ">=3.10"
@@ -1,37 +0,0 @@
"""Run delta-channel conformance capabilities against InMemorySaver."""
from __future__ import annotations
import pytest
conformance = pytest.importorskip(
"langgraph.checkpoint.conformance",
reason="langgraph-checkpoint-conformance not installed",
)
@pytest.mark.asyncio
async def test_delta_channel_conformance():
# Imported inside the test: the module-level importorskip above is what
# makes these safe, so they cannot move to the top of the file.
from langgraph.checkpoint.conformance import validate # noqa: PLC0415
from langgraph.checkpoint.conformance.initializer import ( # noqa: PLC0415
checkpointer_test,
)
from langgraph.checkpoint.memory import InMemorySaver # noqa: PLC0415
@checkpointer_test(name="InMemorySaver")
async def mem_saver():
yield InMemorySaver()
report = await validate(
mem_saver,
capabilities={
"delta_channel_history",
},
)
for cap, result in report.results.items():
if result.passed is False:
details = "\n".join(result.failures or [])
pytest.fail(f"Capability {cap} failed:\n{details}")
+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
+13
View File
@@ -420,6 +420,19 @@ class TestInMemorySaverDeltaChannel:
assert "seed" not in result
assert result["writes"] == []
def test_get_channel_writes_without_checkpoint_id_reads_the_latest(self) -> None:
saver = InMemorySaver()
thread: RunnableConfig = {
"configurable": {"thread_id": "t1", "checkpoint_ns": ""}
}
parent = saver.put(thread, empty_checkpoint(), {}, {})
saver.put_writes(parent, [("messages", "hi")], "task1")
saver.put(parent, empty_checkpoint(), {}, {})
result = saver.get_delta_channel_history(config=thread, channels=["messages"])
assert result == {"messages": {"writes": [("task1", "messages", "hi")]}}
class TestBaseFallbackGetChannelWrites:
"""Exercises the `BaseCheckpointSaver.get_delta_channel_history` default
+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()
+1 -1
View File
@@ -301,7 +301,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.2.0"
version = "4.3.0"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
+1 -1
View File
@@ -1 +1 @@
__version__ = "0.4.32"
__version__ = "0.4.33"
+145 -16
View File
@@ -32,7 +32,7 @@ from langgraph_cli.host_backend import (
HostBackendError,
SourceName,
)
from langgraph_cli.image_reference import ImageReference
from langgraph_cli.image_reference import DIGEST_SEPARATOR, ImageReference
from langgraph_cli.progress import Progress
from langgraph_cli.util import warn_non_wolfi_distro
@@ -624,6 +624,41 @@ def format_deployments_table(deployments: Sequence[dict[str, object]]) -> str:
return "\n".join(lines)
def _extract_listener_namespaces(listener: dict[str, object]) -> str:
compute_config = listener.get("compute_config")
namespaces = (
compute_config.get("k8s_namespaces")
if isinstance(compute_config, dict)
else None
)
if isinstance(namespaces, list) and namespaces:
return ", ".join(str(namespace) for namespace in namespaces)
return "-"
def format_listeners_table(listeners: Sequence[dict[str, object]]) -> str:
headers = ("Listener ID", "Compute ID", "Namespaces")
rows = [
(
str(listener.get("id", "-") or "-"),
str(listener.get("compute_id", "-") or "-"),
_extract_listener_namespaces(listener),
)
for listener in listeners
]
widths = [
max(len(headers[index]), *(len(row[index]) for row in rows))
for index in range(len(headers))
]
def format_row(row: Sequence[str]) -> str:
return " ".join(value.ljust(widths[index]) for index, value in enumerate(row))
lines = [format_row(headers), format_row(tuple("-" * width for width in widths))]
lines.extend(format_row(row) for row in rows)
return "\n".join(lines)
def format_revisions_table(revisions: Sequence[dict[str, object]]) -> str:
headers = ("Revision ID", "Status", "Created At")
latest_deployed_seen = False
@@ -1544,9 +1579,9 @@ def _ensure_customer_registry_source(existing: ExistingDeployment) -> None:
if existing.source != _CUSTOMER_REGISTRY_SOURCE:
raise click.UsageError(
f"Deployment {existing.id} was not created from an external image "
"and cannot be updated with --push-to. Run without --push-to to keep "
"its current build mode, or use a different --name to create a new "
"deployment."
"and cannot be updated with --push-to or --image-uri. Run without "
"either flag to keep its current build mode, or use a different "
"--name to create a new deployment."
)
@@ -1599,9 +1634,22 @@ class RemoteBuildSource:
@dataclass(frozen=True, slots=True)
class CustomerRegistrySource:
class BuildAndPush:
reference: ImageReference
prebuilt_image: str | None
@dataclass(frozen=True, slots=True)
class PublishedImage:
image_uri: str
ImageSource = BuildAndPush | PublishedImage
@dataclass(frozen=True, slots=True)
class CustomerRegistrySource:
image: ImageSource
requested_placement: RequestedPlacement
def run(self, ctx: DeployContext) -> DeployOutcome:
@@ -1686,16 +1734,22 @@ class CustomerRegistrySource:
)
def _publish(self, ctx: DeployContext, step: int) -> tuple[str, int]:
image = str(self.reference)
if isinstance(self.image, PublishedImage):
return self.image.image_uri, step
image = str(self.image.reference)
with Runner() as runner:
if self.prebuilt_image:
_log_deploy_step(step, f"Validating image {self.prebuilt_image}")
if self.image.prebuilt_image:
_log_deploy_step(step, f"Validating image {self.image.prebuilt_image}")
_validate_prebuilt_image(
runner, self.prebuilt_image, verbose=ctx.verbose
runner, self.image.prebuilt_image, verbose=ctx.verbose
)
runner.run(
subp_exec(
"docker", "tag", self.prebuilt_image, image, verbose=ctx.verbose
"docker",
"tag",
self.image.prebuilt_image,
image,
verbose=ctx.verbose,
)
)
else:
@@ -1733,20 +1787,35 @@ def _push_reference(push_to: str, tag: str | None) -> ImageReference:
return reference.with_tag(normalize_image_tag(tag or _DEFAULT_IMAGE_TAG))
def _validate_image_uri(image_uri: str) -> str:
value = image_uri.strip()
if not value:
raise click.UsageError("--image-uri must not be empty.")
if DIGEST_SEPARATOR not in value:
raise click.UsageError(
"--image-uri must pin a digest, e.g. "
f"repository{DIGEST_SEPARATOR}<sha256 hex>. Kubernetes can cache "
"images by tag, so redeploying a mutable tag may silently keep "
"running the previous image."
)
return value
def _select_source(
*,
push_to: str | None,
image: str | None,
image_uri: str | None,
image_name: str | None,
tag: str | None,
remote_build_flag: bool | None,
placement: RequestedPlacement,
selector: DeploymentSelector,
) -> DeploymentSource:
if push_to is None and placement.requested:
if push_to is None and image_uri is None and placement.requested:
raise click.UsageError(
"--listener-id and --k8s-namespace only apply when creating a "
"deployment with --push-to."
"deployment with --push-to or --image-uri."
)
if placement.requested and isinstance(selector, ById):
raise click.UsageError(
@@ -1754,6 +1823,19 @@ def _select_source(
"they cannot be set for an existing --deployment-id. Drop them, or "
"create a new deployment with --name."
)
if image_uri is not None:
if push_to is not None:
raise click.UsageError("--image-uri cannot be combined with --push-to.")
if image is not None:
raise click.UsageError("--image-uri cannot be combined with --image.")
if tag is not None:
raise click.UsageError("--image-uri cannot be combined with --tag.")
if remote_build_flag is not None:
raise click.UsageError("--image-uri cannot be combined with --remote.")
return CustomerRegistrySource(
image=PublishedImage(_validate_image_uri(image_uri)),
requested_placement=placement,
)
if push_to is not None:
if remote_build_flag is True:
raise click.UsageError("--push-to cannot be combined with --remote.")
@@ -1761,8 +1843,7 @@ def _select_source(
if image is None:
_require_local_docker()
return CustomerRegistrySource(
reference=reference,
prebuilt_image=image,
image=BuildAndPush(reference=reference, prebuilt_image=image),
requested_placement=placement,
)
if image and remote_build_flag is True:
@@ -2075,19 +2156,30 @@ def _deploy_base_options(
"Give the tag here or with --tag (default: latest)."
),
),
click.option(
"--image-uri",
help=(
"Deploy an image that's already in a registry you manage, "
"without building, retagging, or pushing anything. For "
"self-hosted and hybrid LangSmith. Give the full reference, "
"e.g. 123456789.dkr.ecr.us-east-1.amazonaws.com/agents/"
"my-agent:v1.2.3 or ...@sha256:<digest>. Cannot be combined "
"with --push-to, --image, --tag, or --remote."
),
),
click.option(
"--listener-id",
help=(
"Listener that will run the deployment, for workspaces that "
"deploy through a listener in your own cluster. Only used when "
"creating a deployment with --push-to."
"creating a deployment with --push-to or --image-uri."
),
),
click.option(
"--k8s-namespace",
help=(
"Kubernetes namespace the listener deploys into. Only used when "
"creating a deployment with --push-to."
"creating a deployment with --push-to or --image-uri."
),
),
click.option(
@@ -2200,6 +2292,7 @@ def _deploy_cmd(
image_name: str | None,
image: str | None,
push_to: str | None,
image_uri: str | None,
listener_id: str | None,
k8s_namespace: str | None,
tag: str | None,
@@ -2273,6 +2366,7 @@ def _deploy_cmd(
source = _select_source(
push_to=push_to,
image=image,
image_uri=image_uri,
image_name=image_name,
tag=tag,
remote_build_flag=remote_build_flag,
@@ -2408,6 +2502,41 @@ def deploy_list(
click.echo(format_deployments_table(deployments))
# ---------------------------------------------------------------------------
# deploy listeners
# ---------------------------------------------------------------------------
@deploy.group(
"listeners",
cls=NestedHelpGroup,
help="[Beta] Inspect listeners available to this workspace.",
)
def deploy_listeners() -> None:
pass
@OPT_HOST_API_KEY
@OPT_HOST_URL
@deploy_listeners.command(
"list",
help=(
"[Beta] List listeners available to this workspace.\n\n"
"Pass a listener's id to `langgraph deploy --push-to ... "
"--listener-id <id>` to deploy through it."
),
)
def deploy_listeners_list(api_key: str | None, host_url: str | None) -> None:
client = _create_host_backend_client(host_url, api_key)
listeners = _call_host_backend_with_optional_tenant(
client, lambda c: c.list_listeners()
)
if not listeners:
click.echo("No listeners found for this workspace.")
return
click.echo(format_listeners_table(listeners))
# ---------------------------------------------------------------------------
# deploy revisions
# ---------------------------------------------------------------------------
+82
View File
@@ -453,6 +453,88 @@ def test_deploy_list_command_no_results(monkeypatch) -> None:
assert result.output.strip() == "No deployments found."
def test_deploy_listeners_list_command(monkeypatch) -> None:
runner = CliRunner()
captured: dict[str, str] = {}
class FakeClient:
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
captured["host_url"] = host_url
captured["api_key"] = api_key
captured["tenant_id"] = tenant_id or ""
def list_listeners(self):
return [
{
"id": "listener-1",
"compute_id": "prod-cluster",
"compute_config": {"k8s_namespaces": ["agents"]},
},
{
"id": "listener-2",
"compute_id": "multi-cluster",
"compute_config": {"k8s_namespaces": ["agents", "agents-staging"]},
},
]
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
result = runner.invoke(
cli,
[
"deploy",
"listeners",
"list",
"--api-key",
"test-key",
"--host-url",
"https://api.example.com",
],
)
assert result.exit_code == 0, result.output
assert captured == {
"host_url": "https://api.example.com",
"api_key": "test-key",
"tenant_id": "",
}
assert "Listener ID" in result.output
assert "Compute ID" in result.output
assert "Namespaces" in result.output
assert "listener-1" in result.output
assert "prod-cluster" in result.output
assert "agents, agents-staging" in result.output
def test_deploy_listeners_list_command_no_results(monkeypatch) -> None:
runner = CliRunner()
class FakeClient:
def __init__(self, host_url: str, api_key: str, tenant_id: str | None = None):
pass
def list_listeners(self):
return []
monkeypatch.setattr(deploy_module, "HostBackendClient", FakeClient)
result = runner.invoke(
cli,
[
"deploy",
"listeners",
"list",
"--api-key",
"test-key",
"--host-url",
"https://api.example.com",
],
)
assert result.exit_code == 0, result.output
assert result.output.strip() == "No listeners found for this workspace."
def test_deploy_revisions_list_command(monkeypatch) -> None:
runner = CliRunner()
captured: dict[str, str] = {}
@@ -695,6 +695,106 @@ def test_push_to_with_deployment_id_fetches_the_deployment_once(
]
def test_image_uri_creates_an_external_deployment_without_any_docker_work(
deploy_project: DeployProject,
) -> None:
result = deploy_project.run("--image-uri", EXTERNAL_DIGEST)
assert result.exit_code == 0, result.output
assert deploy_project.docker.verbs() == []
assert deploy_project.timeline == [LIST_DEPLOYMENTS, CREATE_DEPLOYMENT]
assert deploy_project.control_plane.bodies[CREATE_DEPLOYMENT][
"source_revision_config"
] == {"image_uri": EXTERNAL_DIGEST}
def test_image_uri_updates_an_existing_external_deployment_without_any_docker_work(
deploy_project: DeployProject,
) -> None:
deploy_project.control_plane.existing_deployments = [
{"id": "dep-ext", "name": "my-app", "source": "external_docker"}
]
result = deploy_project.run("--image-uri", EXTERNAL_DIGEST)
assert result.exit_code == 0, result.output
assert deploy_project.docker.verbs() == []
assert deploy_project.timeline == [LIST_DEPLOYMENTS, _patch("dep-ext")]
assert deploy_project.control_plane.bodies[_patch("dep-ext")] == {
"source_revision_config": {"image_uri": EXTERNAL_DIGEST},
"secrets": [],
"tracked_packages": TRACKED_PACKAGES,
}
def test_image_uri_places_a_new_deployment_on_the_only_listener(
deploy_project: DeployProject,
) -> None:
deploy_project.control_plane.listeners = [LISTENER]
result = deploy_project.run(
"--image-uri", EXTERNAL_DIGEST, host_url=CLOUD_CONTROL_PLANE_URL
)
assert result.exit_code == 0, result.output
assert deploy_project.docker.verbs() == []
assert deploy_project.control_plane.bodies[CREATE_DEPLOYMENT]["source_config"] == {
"resource_spec": {},
"listener_id": LISTENER_ID,
"listener_config": {"k8s_namespace": "agents"},
}
def test_image_uri_rejects_a_non_external_deployment_before_any_docker_work(
deploy_project: DeployProject,
) -> None:
deploy_project.control_plane.existing_deployments = [
{"id": "dep-cli", "name": "my-app", "source": "internal_docker"}
]
result = deploy_project.run("--image-uri", EXTERNAL_DIGEST)
assert result.exit_code != 0
assert "cannot be updated with --push-to or --image-uri" in result.output
assert deploy_project.docker.verbs() == []
@pytest.mark.parametrize(
("args", "message"),
[
pytest.param(
("--image-uri", EXTERNAL_DIGEST, "--push-to", PUSH_REPOSITORY),
"--image-uri cannot be combined with --push-to.",
id="with_push_to",
),
pytest.param(
("--image-uri", EXTERNAL_DIGEST, "--image", "local/app:dev"),
"--image-uri cannot be combined with --image.",
id="with_image",
),
pytest.param(
("--image-uri", EXTERNAL_DIGEST, "--tag", "v2"),
"--image-uri cannot be combined with --tag.",
id="with_tag",
),
pytest.param(
("--image-uri", EXTERNAL_DIGEST, "--remote"),
"--image-uri cannot be combined with --remote.",
id="with_remote",
),
],
)
def test_image_uri_conflicting_flags_are_rejected_before_any_docker_work(
deploy_project: DeployProject, args: tuple[str, ...], message: str
) -> None:
result = deploy_project.run(*args)
assert result.exit_code != 0
assert message in result.output
assert deploy_project.docker.verbs() == []
assert deploy_project.timeline == []
def test_invalid_tag_fails_before_any_control_plane_call(
deploy_project: DeployProject,
) -> None:
@@ -13,6 +13,7 @@ import pytest
import langgraph_cli.deploy as deploy_mod
from langgraph_cli.deploy import (
BuildAndPush,
ById,
ByName,
CustomerRegistrySource,
@@ -21,6 +22,7 @@ from langgraph_cli.deploy import (
Listener,
ManagedRegistrySource,
OnListener,
PublishedImage,
RemoteBuildSource,
RequestedPlacement,
Unplaced,
@@ -614,6 +616,7 @@ class TestSelectSource:
OPTIONS = {
"push_to": None,
"image": None,
"image_uri": None,
"image_name": None,
"tag": None,
"remote_build_flag": None,
@@ -629,8 +632,10 @@ class TestSelectSource:
{"push_to": REPOSITORY},
True,
CustomerRegistrySource(
reference=ImageReference(REPOSITORY, "latest"),
prebuilt_image=None,
image=BuildAndPush(
reference=ImageReference(REPOSITORY, "latest"),
prebuilt_image=None,
),
requested_placement=RequestedPlacement(),
),
id="push_to_selects_the_external_source_with_the_default_tag",
@@ -639,8 +644,10 @@ class TestSelectSource:
{"push_to": f"{REPOSITORY}:v2"},
True,
CustomerRegistrySource(
reference=ImageReference(REPOSITORY, "v2"),
prebuilt_image=None,
image=BuildAndPush(
reference=ImageReference(REPOSITORY, "v2"),
prebuilt_image=None,
),
requested_placement=RequestedPlacement(),
),
id="push_to_keeps_a_tag_given_in_the_reference",
@@ -649,8 +656,10 @@ class TestSelectSource:
{"push_to": REPOSITORY, "tag": "v3"},
True,
CustomerRegistrySource(
reference=ImageReference(REPOSITORY, "v3"),
prebuilt_image=None,
image=BuildAndPush(
reference=ImageReference(REPOSITORY, "v3"),
prebuilt_image=None,
),
requested_placement=RequestedPlacement(),
),
id="tag_flag_composes_with_push_to",
@@ -659,8 +668,10 @@ class TestSelectSource:
{"push_to": REPOSITORY, "image": "app:dev"},
False,
CustomerRegistrySource(
reference=ImageReference(REPOSITORY, "latest"),
prebuilt_image="app:dev",
image=BuildAndPush(
reference=ImageReference(REPOSITORY, "latest"),
prebuilt_image="app:dev",
),
requested_placement=RequestedPlacement(),
),
id="prebuilt_image_is_retagged_for_push_to_without_docker_checks",
@@ -672,12 +683,44 @@ class TestSelectSource:
},
True,
CustomerRegistrySource(
reference=ImageReference(REPOSITORY, "latest"),
prebuilt_image=None,
image=BuildAndPush(
reference=ImageReference(REPOSITORY, "latest"),
prebuilt_image=None,
),
requested_placement=RequestedPlacement("listener-1", "agents"),
),
id="push_to_carries_the_requested_placement",
),
pytest.param(
{"image_uri": f"{REPOSITORY}@sha256:abc123"},
True,
CustomerRegistrySource(
image=PublishedImage(f"{REPOSITORY}@sha256:abc123"),
requested_placement=RequestedPlacement(),
),
id="image_uri_selects_the_published_image_source_by_digest",
),
pytest.param(
{"image_uri": f" {REPOSITORY}@sha256:abc123 "},
False,
CustomerRegistrySource(
image=PublishedImage(f"{REPOSITORY}@sha256:abc123"),
requested_placement=RequestedPlacement(),
),
id="image_uri_needs_no_local_docker_and_is_trimmed",
),
pytest.param(
{
"image_uri": f"{REPOSITORY}@sha256:abc123",
"placement": RequestedPlacement("listener-1", "agents"),
},
False,
CustomerRegistrySource(
image=PublishedImage(f"{REPOSITORY}@sha256:abc123"),
requested_placement=RequestedPlacement("listener-1", "agents"),
),
id="image_uri_carries_the_requested_placement",
),
pytest.param(
{"remote_build_flag": True},
True,
@@ -763,6 +806,46 @@ class TestSelectSource:
"only apply when creating a deployment with --push-to",
id="namespace_without_push_to",
),
pytest.param(
{"image_uri": REPOSITORY, "push_to": REPOSITORY},
"--image-uri cannot be combined with --push-to.",
id="image_uri_with_push_to",
),
pytest.param(
{"image_uri": REPOSITORY, "image": "app:dev"},
"--image-uri cannot be combined with --image.",
id="image_uri_with_image",
),
pytest.param(
{"image_uri": REPOSITORY, "tag": "v2"},
"--image-uri cannot be combined with --tag.",
id="image_uri_with_tag",
),
pytest.param(
{"image_uri": REPOSITORY, "remote_build_flag": True},
"--image-uri cannot be combined with --remote.",
id="image_uri_with_remote",
),
pytest.param(
{"image_uri": ""},
"--image-uri must not be empty.",
id="image_uri_empty",
),
pytest.param(
{"image_uri": " "},
"--image-uri must not be empty.",
id="image_uri_blank",
),
pytest.param(
{"image_uri": f"{REPOSITORY}:v1.2.3"},
"--image-uri must pin a digest",
id="image_uri_with_a_mutable_tag",
),
pytest.param(
{"image_uri": REPOSITORY},
"--image-uri must pin a digest",
id="image_uri_without_any_tag_or_digest",
),
],
)
def test_conflicting_flags_are_rejected(self, monkeypatch, flags, message):
@@ -1164,6 +1247,7 @@ def test_a_deployment_id_with_listener_flags_is_refused_without_probing_docker(
_select_source(
push_to="registry.example.com/app",
image=None,
image_uri=None,
image_name=None,
tag=None,
remote_build_flag=None,
+32
View File
@@ -3,6 +3,7 @@ from unittest.mock import patch
from langgraph_cli.deploy import (
_extract_deployment_url,
format_deployments_table,
format_listeners_table,
format_revisions_table,
)
from langgraph_cli.util import clean_empty_lines, warn_non_wolfi_distro
@@ -255,3 +256,34 @@ def test_format_revisions_table():
assert "rev-456" in output
assert "rev-789" in output
assert "REPLACED" in output
def test_format_listeners_table():
output = format_listeners_table(
[
{
"id": "listener-1",
"compute_id": "prod-cluster",
"compute_config": {"k8s_namespaces": ["agents"]},
},
{
"id": "listener-2",
"compute_id": "multi-cluster",
"compute_config": {"k8s_namespaces": ["agents", "agents-staging"]},
},
{
"id": "listener-3",
"compute_id": "broken-cluster",
},
]
)
assert "Listener ID" in output
assert "Compute ID" in output
assert "Namespaces" in output
assert "listener-1" in output
assert "prod-cluster" in output
assert "agents" in output
assert "listener-2" in output
assert "agents, agents-staging" in output
assert "listener-3" in output
assert "broken-cluster" in output
+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
+2 -3
View File
@@ -27,6 +27,7 @@ from langgraph.checkpoint.base import (
Checkpoint,
PendingWrite,
V,
writes_sort_key,
)
from langgraph.store.base import BaseStore
from xxhash import xxh3_128_hexdigest
@@ -251,9 +252,7 @@ def apply_writes(
Set of channels that were updated in this step.
"""
# sort tasks on path, to ensure deterministic order for update application
# any path parts after the 3rd are ignored for sorting
# (we use them for eg. task ids which aren't good for sorting)
tasks = sorted(tasks, key=lambda t: task_path_str(t.path[:3]))
tasks = sorted(tasks, key=lambda t: writes_sort_key(task_path_str(t.path)))
# if no task has triggers this is applying writes from the null task only
# so we don't do anything other than update the channels written to
bump_step = any(t.triggers for t in tasks)
+62 -8
View File
@@ -1,8 +1,9 @@
from __future__ import annotations
import uuid
from collections.abc import Callable, Iterable, Mapping
from collections.abc import Callable, Iterable, Mapping, Sequence
from datetime import datetime, timezone
from inspect import signature
from typing import Any, Literal, cast
from langchain_core.runnables import RunnableConfig
@@ -25,11 +26,13 @@ from langgraph._internal._constants import (
CONFIG_KEY_CHECKPOINT_ID,
NS_END,
NS_SEP,
NULL_TASK_ID,
PUSH,
SNAPSHOT_BUMPS,
)
from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel
from langgraph.channels.binop import _get_overwrite
from langgraph.channels.delta import DeltaChannel
from langgraph.managed.base import ManagedValueMapping, ManagedValueSpec
@@ -49,15 +52,38 @@ def empty_checkpoint() -> Checkpoint:
)
def put_writes_accepts_task_path(put_writes: Callable[..., Any]) -> bool:
"""Whether a saver's `put_writes` or `aput_writes` takes `task_path`.
Savers written before the parameter existed don't, so it is passed only
when this is true.
"""
return signature(put_writes).parameters.get("task_path") is not None
def exit_delta_task_id(step: int, task_id: str) -> str:
"""Synthetic task id for exit-mode DeltaChannel writes.
Embeds the superstep in the first UUID group so `ORDER BY task_id, idx`
preserves chronological order while remaining a valid RFC UUID (required by
Postgres `checkpoint_writes.task_id uuid` columns).
Postgres `checkpoint_writes.task_id uuid` columns). Never `NULL_TASK_ID`:
readers apply writes under it as the anchor checkpoint's own pending writes.
"""
parts = str(uuid.UUID(task_id)).split("-")
return f"{step:08d}-{parts[1]}-{parts[2]}-{parts[3]}-{parts[4]}"
synthetic = f"{step:08d}-{parts[1]}-{parts[2]}-{parts[3]}-{parts[4]}"
if synthetic == NULL_TASK_ID:
return f"{step:08d}-0000-0000-0000-000000000001"
return synthetic
def exit_delta_late_task_id(step: int, task_id: str) -> str:
"""Synthetic task id for exit-mode writes of a superstep after the anchor's own.
Sorts after every real task id, in step order, so replay keeps them after
the anchor's own superstep whether a saver orders by task path or task id.
"""
parts = str(uuid.UUID(task_id)).split("-")
return f"ffffffff-{step >> 16:04x}-{step & 0xFFFF:04x}-{parts[3]}-{parts[4]}"
def delta_channels_to_snapshot(
@@ -127,6 +153,23 @@ def delta_channels_with_pending_writes(
}
def delta_channels_overwritten(
specs: Mapping[str, Any], writes: Iterable[tuple[str, Any]]
) -> set[str]:
"""Return the names of the DeltaChannels that `writes` set with an `Overwrite`.
`update_state` saves a full snapshot of these channels in the checkpoint it
creates, like the loop does when a node returns an `Overwrite`. Otherwise,
reading the channel later starts from an older snapshot and replays the
writes the `Overwrite` threw away.
"""
return {
ch
for ch, value in writes
if isinstance(specs.get(ch), DeltaChannel) and _get_overwrite(value)[0]
}
def checkpoint_superseded(
saver: BaseCheckpointSaver, config: RunnableConfig, saved: CheckpointTuple
) -> bool:
@@ -231,6 +274,7 @@ def create_checkpoint(
get_next_version: GetNextVersion | None = None,
channels_to_snapshot: set[str] | None = None,
stored_versions: ChannelVersions | None = None,
trigger_to_nodes: Mapping[str, Sequence[str]] | None = None,
) -> Checkpoint:
"""Build a new Checkpoint from the previous one and live channel state.
@@ -289,7 +333,9 @@ def create_checkpoint(
id=id or str(uuid6(clock_seq=step)),
channel_values=values,
channel_versions=channel_versions,
versions_seen=_mark_bumps_seen(checkpoint["versions_seen"], bumped),
versions_seen=_mark_bumps_seen(
checkpoint["versions_seen"], bumped, trigger_to_nodes or {}
),
updated_channels=None if updated_channels is None else sorted(updated_channels),
)
@@ -297,19 +343,27 @@ def create_checkpoint(
def _mark_bumps_seen(
versions_seen: dict[str, ChannelVersions],
bumped: Mapping[str, tuple[Any, Any]],
trigger_to_nodes: Mapping[str, Sequence[str]],
) -> dict[str, ChannelVersions]:
"""Advance whoever had seen a bumped channel's old version to the new one.
A bump that only stores a snapshot is not a write. Left unseen, it would
re-fire `interrupt_before` and rerun the channel's subscribers. For each
entry it advances, `SNAPSHOT_BUMPS` keeps the new version and the one the
node really read, so `versions_seen_without_bumps` can put the read back.
re-fire `interrupt_before` and rerun the channel's subscribers. A channel
bumped from no version was never written, so it also goes to the
subscribers that never ran: they have no entry, and would start on the bump.
For each entry it advances, `SNAPSHOT_BUMPS` keeps the new version and the
one the node really read, so `versions_seen_without_bumps` can put the read
back.
"""
if not bumped:
return versions_seen
out = dict(versions_seen)
for k, (old, _) in bumped.items():
if old is None:
for node in trigger_to_nodes.get(k, ()):
out.setdefault(node, {})
marks = dict(versions_seen.get(SNAPSHOT_BUMPS, {}))
for node, seen in versions_seen.items():
for node, seen in out.items():
if node == SNAPSHOT_BUMPS:
continue
for k, (old, new) in bumped.items():
+127 -60
View File
@@ -12,7 +12,6 @@ from contextlib import (
ExitStack,
)
from datetime import datetime, timezone
from inspect import signature
from types import TracebackType
from typing import (
Any,
@@ -20,6 +19,7 @@ from typing import (
TypeVar,
cast,
)
from uuid import UUID, uuid5
from langchain_core.callbacks import AsyncParentRunManager, ParentRunManager
from langchain_core.runnables import RunnableConfig
@@ -106,7 +106,9 @@ from langgraph.pregel._checkpoint import (
delta_channels_to_snapshot,
delta_channels_with_pending_writes,
empty_checkpoint,
exit_delta_late_task_id,
exit_delta_task_id,
put_writes_accepts_task_path,
)
from langgraph.pregel._executor import (
AsyncBackgroundExecutor,
@@ -190,6 +192,7 @@ class PregelLoop:
Callable[
[
concurrent.futures.Future | None,
Sequence[Any],
RunnableConfig,
Checkpoint,
str,
@@ -203,11 +206,13 @@ class PregelLoop:
submit: Submit
channels: Mapping[str, BaseChannel]
# Futures from `checkpointer.put_writes` calls that produced delta-channel
# writes. `_checkpointer_put_after_previous` drains this list (swap to a
# local `futs` then reset to `[]` and wait/gather) before putting the
# next checkpoint, so a checkpoint never becomes durable before the
# writes that produced it. Initialised to `[]` in both sync and async
# `__enter__`; stays `None` only when no checkpointer.
# writes. `_put_checkpoint` hands this list to the save it submits, which
# waits for them first, so a checkpoint never becomes durable before the
# writes that produced it. If a write or the previous save failed, the
# save fails too: a DeltaChannel is rebuilt from its writes along the
# parent chain, so a checkpoint saved past either gap reads back short
# for good. Initialised to `[]` in both sync and async `__enter__`;
# stays `None` only when no checkpointer.
_delta_write_futs: list[Any] | None = None
# Same pattern as `_delta_write_futs` but for error-handler writes.
@@ -221,10 +226,18 @@ class PregelLoop:
# `after_tick`). At exit, `_put_exit_delta_writes` filters out channels
# that will snapshot, then persists the rest under an anchor parent.
# `None` when not in exit mode (so the capture sites are no-ops).
# Each tuple is `(step, task_id, channel, value)` — `step` drives the
# synthetic step-prefixed task_id used to preserve chronological order
# under the saver's `ORDER BY task_id, idx` sorting.
_exit_delta_writes: list[tuple[int, str, str, Any]] | None = None
# Each tuple is `(step, task_id, task_path, channel, value)`; see
# `_put_exit_delta_writes` for how they are ordered.
_exit_delta_writes: list[tuple[int, str, str, str, Any]] | None = None
# The (task_id, channel) pairs whose delta writes are stored on the loaded
# checkpoint, which a resume addressed by `checkpoint_id` can rerun; this
# run's `Command` delta writes, kept apart from the NULL_TASK_ID writes
# loaded with the checkpoint; and the checkpoint's own superstep, the
# first one this run ticks.
_stored_delta_writes: set[tuple[str, str]]
_exit_command_writes: list[tuple[str, Any]]
_exit_first_step: int | None = None
# Delta channels that must snapshot at the next checkpoint, whatever their
# cadence counters say:
@@ -257,6 +270,9 @@ class PregelLoop:
# `_put_exit_delta_writes` uses this to decide between anchoring on
# the existing parent (True) or creating a lazy stub (False).
_has_persisted_parent: bool = False
# True iff `__enter__` loaded the thread's latest checkpoint, not one a
# `checkpoint_id` addressed, so nothing has been built on it yet.
_loaded_latest: bool = False
managed: ManagedValueMapping
checkpoint: Checkpoint
@@ -669,10 +685,6 @@ class PregelLoop:
self.status = "done"
return False
if self.control is not None and self.control.drain_requested:
self.status = "draining"
return False
# if there are pending writes from a previous loop, apply them
if self._reapplies_pending_writes and self.checkpoint_pending_writes:
self._reapply_writes_to_succeeded_nodes(self.tasks)
@@ -685,6 +697,13 @@ class PregelLoop:
self.status = "interrupt_before"
raise GraphInterrupt()
# stop before running the next tasks if a drain was requested. after
# the breakpoint check: a resume passes the next step's breakpoints,
# so draining in front of one would skip it
if self.control is not None and self.control.drain_requested:
self.status = "draining"
return False
# produce debug output
self._emit("tasks", map_debug_tasks, self.tasks.values())
@@ -737,9 +756,25 @@ class PregelLoop:
)
# capture delta-channel writes for exit-mode accumulator before clearing
if self._exit_delta_writes is not None:
# On the first tick the pending writes still hold the ones loaded
# with the checkpoint, which are already stored on it.
first = self._exit_first_step is None
if first:
self._exit_first_step = self.step
self._exit_delta_writes.extend(
(self.step, NULL_TASK_ID, "", ch, v)
for ch, v in self._exit_command_writes
)
for tid, ch, v in self.checkpoint_pending_writes:
if isinstance(self.specs.get(ch), DeltaChannel):
self._exit_delta_writes.append((self.step, tid, ch, v))
if not isinstance(self.specs.get(ch), DeltaChannel):
continue
if first and (
tid == NULL_TASK_ID or (tid, ch) in self._stored_delta_writes
):
continue
task = self.tasks.get(tid)
path = task_path_str(task.path) if task else ""
self._exit_delta_writes.append((self.step, tid, path, ch, v))
# clear pending writes
self.checkpoint_pending_writes.clear()
# only replay (re-execute) done tasks on the first tick
@@ -860,6 +895,12 @@ class PregelLoop:
def _first(
self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None
) -> set[str] | None:
self._stored_delta_writes = {
(tid, ch)
for tid, ch, _ in self.checkpoint_pending_writes
if tid != NULL_TASK_ID and isinstance(self.specs.get(ch), DeltaChannel)
}
self._exit_command_writes = []
# Resuming from a previous checkpoint requires two things:
# 1. A prior checkpoint exists (channel_versions is non-empty)
# 2. The input signals continuation (not a fresh run with new input)
@@ -977,6 +1018,12 @@ class PregelLoop:
carried.extend((tid, c, v) for c, v in ws)
else:
self.put_writes(tid, ws)
if self._exit_delta_writes is not None and tid == NULL_TASK_ID:
self._exit_command_writes.extend(
(c, v)
for c, v in ws
if isinstance(self.specs.get(c), DeltaChannel)
)
self._delta_channels_forced_snapshot.update(
delta_channels_with_pending_writes(self.specs, carried)
)
@@ -1074,17 +1121,29 @@ class PregelLoop:
if self._exit_delta_writes is not None:
for c, v in input_writes:
if isinstance(self.specs.get(c), DeltaChannel):
self._exit_delta_writes.append((self.step, NULL_TASK_ID, c, v))
# Persist delta-channel input writes so sub-freq inputs are
# recoverable via ancestor walk (mirrors the Command input path).
self._exit_delta_writes.append(
(self.step, NULL_TASK_ID, "", c, v)
)
# A DeltaChannel reads its input from the writes stored on the
# checkpoint this run starts from, under a task id of their own:
# readers apply a checkpoint's NULL_TASK_ID writes as its own state.
# A new thread has no checkpoint to store them on, and one a
# `checkpoint_id` addressed may have children that would read them,
# so then the input checkpoint snapshots the channel instead.
if self.durability != "exit":
delta_input = [
(c, v)
for c, v in input_writes
if isinstance(self.specs.get(c), DeltaChannel)
]
if delta_input:
self.put_writes(NULL_TASK_ID, delta_input)
if delta_input and self._has_persisted_parent and self._loaded_latest:
self.put_writes(
str(uuid5(UUID(self.checkpoint["id"]), INPUT)), delta_input
)
else:
self._delta_channels_forced_snapshot.update(
c for c, _ in delta_input
)
# save input checkpoint
self.updated_channels = updated_channels
self._put_checkpoint({"source": "input"})
@@ -1214,6 +1273,7 @@ class PregelLoop:
else None,
channels_to_snapshot=channels_to_snapshot,
stored_versions=self.checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
for k in channels_to_snapshot:
new_counters[k] = (0, 0)
@@ -1260,12 +1320,17 @@ class PregelLoop:
)
self.checkpoint_previous_versions = channel_versions
# Take this checkpoint's writes now: saves run in the background
# and can start out of order, so a save that took them itself
# could get another checkpoint's writes.
delta_write_futs, self._delta_write_futs = self._delta_write_futs, []
# save it, without blocking
# if there's a previous checkpoint save in progress, wait for it
# ensuring checkpointers receive checkpoints in order
self._put_checkpoint_fut = self.submit(
self._checkpointer_put_after_previous,
getattr(self, "_put_checkpoint_fut", None),
delta_write_futs,
self.checkpoint_config,
copy_checkpoint(self.checkpoint),
self.checkpoint_metadata,
@@ -1309,9 +1374,7 @@ class PregelLoop:
)
pending = [
(step, tid, ch, v)
for (step, tid, ch, v) in self._exit_delta_writes
if ch not in channels_to_snapshot
w for w in self._exit_delta_writes if w[3] not in channels_to_snapshot
]
if not pending:
return
@@ -1337,6 +1400,7 @@ class PregelLoop:
self._put_checkpoint_fut = self.submit(
self._checkpointer_put_after_previous,
getattr(self, "_put_checkpoint_fut", None),
(),
stub_put_config,
stub_cp,
{"step": -2},
@@ -1346,11 +1410,21 @@ class PregelLoop:
# sees the stub as its parent.
self.checkpoint_config = anchor_config
# Step-prefixed synthetic task_id preserves chronological superstep
# order under the saver's ORDER BY task_id, idx sorting.
grouped: dict[tuple[int, str], list[tuple[str, Any]]] = {}
for step, tid, ch, v in pending:
grouped.setdefault((step, tid), []).append((ch, v))
# The checkpoint's own superstep keeps its real task paths, so it
# interleaves with the writes a resume loaded from it. Its task ids stay
# synthetic: under the real id, a run whose final checkpoint fails to
# save would leave the resumed task looking done to the next resume.
# Later supersteps sort after every real task path and task id, in step
# order, so this holds whether a saver orders by path or by id.
grouped: dict[tuple[str, str], list[tuple[str, Any]]] = {}
for step, tid, path, ch, v in pending:
if tid == NULL_TASK_ID:
key = (exit_delta_task_id(step, tid), "")
elif step == self._exit_first_step:
key = (exit_delta_task_id(step, tid), path)
else:
key = (exit_delta_late_task_id(step, tid), f"~~{step:010d}{path}")
grouped.setdefault(key, []).append((ch, v))
anchor_write_config = patch_configurable(
anchor_config,
{
@@ -1360,22 +1434,21 @@ class PregelLoop:
CONFIG_KEY_CHECKPOINT_ID: anchor_config[CONF][CONFIG_KEY_CHECKPOINT_ID],
},
)
for (step, tid), entries in grouped.items():
synth_tid = exit_delta_task_id(step, tid)
for (tid, path), entries in grouped.items():
if self.checkpointer_put_writes_accepts_task_path:
fut = self.submit(
self.checkpointer_put_writes,
anchor_write_config,
entries,
synth_tid,
"",
tid,
path,
)
else:
fut = self.submit(
self.checkpointer_put_writes,
anchor_write_config,
entries,
synth_tid,
tid,
)
if self._delta_write_futs is not None:
self._delta_write_futs.append(fut)
@@ -1584,8 +1657,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
self.checkpointer_get_next_version = checkpointer.get_next_version
self.checkpointer_put_writes = checkpointer.put_writes
self.checkpointer_put_writes_accepts_task_path = (
signature(checkpointer.put_writes).parameters.get("task_path")
is not None
put_writes_accepts_task_path(checkpointer.put_writes)
)
else:
self.checkpointer_get_next_version = increment
@@ -1596,21 +1668,19 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
def _checkpointer_put_after_previous(
self,
prev: concurrent.futures.Future | None,
delta_write_futs: Sequence[concurrent.futures.Future],
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> RunnableConfig:
if self._delta_write_futs:
futs, self._delta_write_futs = self._delta_write_futs, []
concurrent.futures.wait(futs)
try:
if prev is not None:
prev.result()
finally:
cast(BaseCheckpointSaver, self.checkpointer).put(
config, checkpoint, metadata, new_versions
)
for fut in delta_write_futs:
fut.result()
if prev is not None:
prev.result()
cast(BaseCheckpointSaver, self.checkpointer).put(
config, checkpoint, metadata, new_versions
)
def match_cached_writes(self) -> Sequence[PregelExecutableTask]:
if self.cache is None:
@@ -1722,6 +1792,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
# Normal case: fetch the most recent checkpoint for this
# graph/thread. Returns None on first invocation.
saved = self.checkpointer.get_tuple(self.checkpoint_config)
self._loaded_latest = True
# Capture before the synthetic-empty fallback below overwrites `saved`.
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
@@ -1840,8 +1911,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
self.checkpointer_get_next_version = checkpointer.get_next_version
self.checkpointer_put_writes = checkpointer.aput_writes
self.checkpointer_put_writes_accepts_task_path = (
signature(checkpointer.aput_writes).parameters.get("task_path")
is not None
put_writes_accepts_task_path(checkpointer.aput_writes)
)
else:
self.checkpointer_get_next_version = increment
@@ -1852,23 +1922,19 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
async def _checkpointer_put_after_previous(
self,
prev: asyncio.Task | None,
delta_write_futs: Sequence[asyncio.Future],
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> RunnableConfig:
# Drain DeltaChannel write futures before committing the checkpoint so
# ancestor walks never see a checkpoint without its backing writes.
if self._delta_write_futs:
futs, self._delta_write_futs = self._delta_write_futs, []
await asyncio.gather(*futs)
try:
if prev is not None:
await prev
finally:
await cast(BaseCheckpointSaver, self.checkpointer).aput(
config, checkpoint, metadata, new_versions
)
if delta_write_futs:
await asyncio.gather(*delta_write_futs)
if prev is not None:
await prev
await cast(BaseCheckpointSaver, self.checkpointer).aput(
config, checkpoint, metadata, new_versions
)
async def amatch_cached_writes(self) -> Sequence[PregelExecutableTask]:
if self.cache is None:
@@ -1983,6 +2049,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
# Normal case: fetch the most recent checkpoint for this
# graph/thread. Returns None on first invocation.
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
self._loaded_latest = True
# Capture before the synthetic-empty fallback below overwrites `saved`.
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
+211 -85
View File
@@ -126,6 +126,7 @@ from langgraph.pregel._algo import (
apply_writes,
local_read,
prepare_next_tasks,
task_path_str,
)
from langgraph.pregel._call import identifier
from langgraph.pregel._checkpoint import (
@@ -136,9 +137,11 @@ from langgraph.pregel._checkpoint import (
copy_checkpoint,
create_checkpoint,
create_checkpoint_plan_for_update_state_api,
delta_channels_overwritten,
delta_channels_with_pending_writes,
empty_checkpoint,
get_updated_channels_from_tasks,
put_writes_accepts_task_path,
versions_seen_without_bumps,
)
from langgraph.pregel._draw import draw_graph
@@ -150,6 +153,7 @@ from langgraph.pregel._loop import (
from langgraph.pregel._messages import (
StreamMessagesHandler,
StreamMessagesHandlerV2,
ensure_message_ids,
)
from langgraph.pregel._read import DEFAULT_BOUND, PregelNode
from langgraph.pregel._retry import RetryPolicy
@@ -1764,6 +1768,7 @@ class Pregel(
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = checkpointer.put(
checkpoint_config,
@@ -1786,6 +1791,22 @@ class Pregel(
)
if input_writes := deque(map_input(self.input_channels, values)):
_store_or_fork_delta_writes(
checkpointer,
config,
saved,
checkpoint_config,
self.channels,
[
(
str(uuid5(UUID(checkpoint["id"]), INPUT)),
input_writes,
None,
)
],
fork_pending,
is_first=is_first,
)
updated_channels = apply_writes(
checkpoint,
channels,
@@ -1793,6 +1814,9 @@ class Pregel(
checkpointer.get_next_version,
self.trigger_to_nodes,
)
fork_pending |= delta_channels_overwritten(
self.channels, input_writes
)
# apply input write to channels
next_step = (
@@ -1820,6 +1844,7 @@ class Pregel(
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = checkpointer.put(
checkpoint_config,
@@ -1831,13 +1856,6 @@ class Pregel(
),
)
# store the writes
checkpointer.put_writes(
next_config,
input_writes,
str(uuid5(UUID(checkpoint["id"]), INPUT)),
)
return patch_checkpoint_map(
next_config, saved.metadata if saved else None
)
@@ -1983,13 +2001,13 @@ class Pregel(
run_tasks: list[PregelTaskWrites] = []
run_task_ids: list[str] = []
for as_node, values, provided_task_id in valid_updates:
for i, (as_node, values, provided_task_id) in enumerate(valid_updates):
# create task to run all writers of the chosen node
writers = self.nodes[as_node].flat_writers
if not writers:
raise InvalidUpdateError(f"Node {as_node} has no writers")
writes: deque[tuple[str, Any]] = deque()
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
task = PregelTaskWrites((INTERRUPT, i), as_node, writes, [INTERRUPT])
# get the task ids that were prepared for this node
# if a task id was provided in the StateUpdate, we use it
# otherwise, we use the next available task id
@@ -1997,7 +2015,7 @@ class Pregel(
task_id = provided_task_id or (
prepared_task_ids.popleft()
if prepared_task_ids
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
else _update_task_id(checkpoint["id"], i)
)
run_tasks.append(task)
run_task_ids.append(task_id)
@@ -2031,27 +2049,22 @@ class Pregel(
),
)
updated_channels = get_updated_channels_from_tasks(run_tasks)
# The base's other children replay whatever is stored on it, so an
# edit of an older checkpoint stores none of its writes there: the
# checkpoint written here carries them, its delta channels
# snapshotted. Later supersteps address the checkpoint just written.
if (
is_first
and saved is not None
and checkpoint_superseded(checkpointer, config, saved)
):
fork_pending.update(
ch
for ch in updated_channels
if isinstance(self.channels.get(ch), DeltaChannel)
)
elif saved is not None:
for task_id, task in zip(run_task_ids, run_tasks):
channel_writes = [w for w in task.writes if w[0] != PUSH]
if channel_writes:
checkpointer.put_writes(
checkpoint_config, channel_writes, task_id
)
fork_pending |= delta_channels_overwritten(
self.channels, (w for t in run_tasks for w in t.writes)
)
_store_or_fork_delta_writes(
checkpointer,
config,
saved,
checkpoint_config,
self.channels,
[
(task_id, [w for w in task.writes if w[0] != PUSH], task)
for task_id, task in zip(run_task_ids, run_tasks)
],
fork_pending,
is_first=is_first,
)
apply_writes(
checkpoint,
channels,
@@ -2081,6 +2094,7 @@ class Pregel(
else None,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = checkpointer.put(
checkpoint_config,
@@ -2253,6 +2267,7 @@ class Pregel(
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = await checkpointer.aput(
checkpoint_config,
@@ -2275,6 +2290,22 @@ class Pregel(
)
if input_writes := deque(map_input(self.input_channels, values)):
await _astore_or_fork_delta_writes(
checkpointer,
config,
saved,
checkpoint_config,
self.channels,
[
(
str(uuid5(UUID(checkpoint["id"]), INPUT)),
input_writes,
None,
)
],
fork_pending,
is_first=is_first,
)
updated_channels = apply_writes(
checkpoint,
channels,
@@ -2282,6 +2313,9 @@ class Pregel(
checkpointer.get_next_version,
self.trigger_to_nodes,
)
fork_pending |= delta_channels_overwritten(
self.channels, input_writes
)
# apply input write to channels
next_step = (
@@ -2309,6 +2343,7 @@ class Pregel(
get_next_version=checkpointer.get_next_version,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = await checkpointer.aput(
checkpoint_config,
@@ -2320,13 +2355,6 @@ class Pregel(
),
)
# store the writes
await checkpointer.aput_writes(
next_config,
input_writes,
str(uuid5(UUID(checkpoint["id"]), INPUT)),
)
return patch_checkpoint_map(
next_config, saved.metadata if saved else None
)
@@ -2471,13 +2499,13 @@ class Pregel(
run_tasks: list[PregelTaskWrites] = []
run_task_ids: list[str] = []
for as_node, values, provided_task_id in valid_updates:
for i, (as_node, values, provided_task_id) in enumerate(valid_updates):
# create task to run all writers of the chosen node
writers = self.nodes[as_node].flat_writers
if not writers:
raise InvalidUpdateError(f"Node {as_node} has no writers")
writes: deque[tuple[str, Any]] = deque()
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
task = PregelTaskWrites((INTERRUPT, i), as_node, writes, [INTERRUPT])
# get the task ids that were prepared for this node
# if a task id was provided in the StateUpdate, we use it
# otherwise, we use the next available task id
@@ -2485,7 +2513,7 @@ class Pregel(
task_id = provided_task_id or (
prepared_task_ids.popleft()
if prepared_task_ids
else str(uuid5(UUID(checkpoint["id"]), INTERRUPT))
else _update_task_id(checkpoint["id"], i)
)
run_tasks.append(task)
run_task_ids.append(task_id)
@@ -2519,27 +2547,22 @@ class Pregel(
),
)
updated_channels = get_updated_channels_from_tasks(run_tasks)
# The base's other children replay whatever is stored on it, so an
# edit of an older checkpoint stores none of its writes there: the
# checkpoint written here carries them, its delta channels
# snapshotted. Later supersteps address the checkpoint just written.
if (
is_first
and saved is not None
and await acheckpoint_superseded(checkpointer, config, saved)
):
fork_pending.update(
ch
for ch in updated_channels
if isinstance(self.channels.get(ch), DeltaChannel)
)
elif saved is not None:
for task_id, task in zip(run_task_ids, run_tasks):
channel_writes = [w for w in task.writes if w[0] != PUSH]
if channel_writes:
await checkpointer.aput_writes(
checkpoint_config, channel_writes, task_id
)
fork_pending |= delta_channels_overwritten(
self.channels, (w for t in run_tasks for w in t.writes)
)
await _astore_or_fork_delta_writes(
checkpointer,
config,
saved,
checkpoint_config,
self.channels,
[
(task_id, [w for w in task.writes if w[0] != PUSH], task)
for task_id, task in zip(run_task_ids, run_tasks)
],
fork_pending,
is_first=is_first,
)
apply_writes(
checkpoint,
channels,
@@ -2569,6 +2592,7 @@ class Pregel(
else None,
channels_to_snapshot=channels_to_snapshot,
stored_versions=checkpoint_previous_versions,
trigger_to_nodes=self.trigger_to_nodes,
)
next_config = await checkpointer.aput(
checkpoint_config,
@@ -3706,16 +3730,15 @@ class Pregel(
config: RunnableConfig | None = None,
*,
version: Literal["v1", "v2", "v3"] = "v2",
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
control: RunControl | None = None,
transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
**kwargs: Any,
) -> Any:
"""Stream events from this graph.
For `version="v1"` / `"v2"`, yields `StreamEvent` dicts (see
`Runnable.stream_events`). For `version="v3"`, returns a
For `version="v1"` / `"v2"`, delegates to
`Runnable.(a)stream_events`; synchronous v1/v2 event streaming is
not implemented in langchain-core, so use `astream_events` for
those versions. For `version="v3"`, returns a
`GraphRunStream` whose typed projections the caller drives by
iterating — no background thread.
@@ -3745,20 +3768,27 @@ class Pregel(
config: Optional runnable config.
version: Streaming-event schema version. `"v3"` selects the
content-block-centric streaming protocol.
interrupt_before: Nodes to interrupt before, if any. Only
used for `version="v3"`.
interrupt_after: Nodes to interrupt after, if any. Only
used for `version="v3"`.
interrupt_before: Nodes to interrupt before, if any.
Honored on every version that can run; type-checked
only on the `version="v3"` overloads.
interrupt_after: Nodes to interrupt after, if any. Honored
on every version that can run; type-checked only on
the `version="v3"` overloads.
control: Optional run control used to request cooperative
drain. Only used for `version="v3"`.
drain. Honored on every version that can run;
type-checked only on the `version="v3"` overloads.
transformers: Extra transformer classes or configured
factories appended after compile-time
`stream_transformers`. Factories are called as
`factory(scope)` so they can propagate to subgraph
scopes. Only used for `version="v3"`.
**kwargs: For `version="v1"`/`"v2"`, forwarded to
`Runnable.stream_events`. For `version="v3"`, forwarded
to the underlying `stream(...)` call (e.g. `context`,
**kwargs: For `version="v1"`/`"v2"` on `astream_events`,
forwarded to `Runnable.astream_events`, which passes
them through to `astream` — so execution kwargs such as
`context`, `durability`, `interrupt_before`,
`interrupt_after` and `control` are honored on every
version that can run. For `version="v3"`, forwarded to the
underlying `stream(...)` call (e.g. `context`,
`durability`, `output_keys`, `print_mode`, `debug`).
`stream_mode` and `subgraphs` are not accepted under
`version="v3"` and raise `TypeError` if supplied; v3
@@ -3766,16 +3796,15 @@ class Pregel(
Returns:
For `version="v3"`, a `GraphRunStream` the caller iterates
to drive the run. Otherwise an `Iterator[StreamEvent]`.
to drive the run. For `version="v1"`/`"v2"`,
`astream_events` yields `StreamEvent` dicts; the synchronous
v1/v2 path is not implemented in langchain-core.
"""
if version == "v3":
_reject_v3_invariant_kwargs(kwargs)
return self._pregel_stream_v3(
input,
config,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
control=control,
transformers=transformers,
**kwargs,
)
@@ -3811,9 +3840,6 @@ class Pregel(
config: RunnableConfig | None = None,
*,
version: Literal["v1", "v2", "v3"] = "v2",
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
control: RunControl | None = None,
transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
**kwargs: Any,
) -> AsyncIterator[StreamEvent] | Awaitable[Any]:
@@ -3836,9 +3862,6 @@ class Pregel(
return self._apregel_stream_v3(
input,
config,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
control=control,
transformers=transformers,
**kwargs,
)
@@ -4237,6 +4260,109 @@ class Pregel(
await self.cache.aclear(namespaces)
def _task_path_kwarg(put_writes: Callable[..., Any], task: PregelTaskWrites) -> dict:
"""Pass the task's path to savers whose `put_writes` takes one.
Savers that replay a checkpoint's writes in task path order then give back
updates applied together in the order they were given.
"""
if not put_writes_accepts_task_path(put_writes):
return {}
return {"task_path": task_path_str(task.path)}
_UpdateWrites = Sequence[tuple[str, Sequence[tuple[str, Any]], PregelTaskWrites | None]]
def _delta_writes(
channels: Mapping[str, BaseChannel | ManagedValueSpec], writes: _UpdateWrites
) -> list[tuple[str, Any]]:
"""Give the update's DeltaChannel writes message ids, as the loop's
`put_writes` does, so every read of them returns the same ids."""
delta = [
(ch, value)
for _, task_writes, _ in writes
for ch, value in task_writes
if isinstance(channels.get(ch), DeltaChannel)
]
for _, value in delta:
ensure_message_ids(value)
return delta
def _store_or_fork_delta_writes(
checkpointer: BaseCheckpointSaver,
config: RunnableConfig,
saved: CheckpointTuple | None,
checkpoint_config: RunnableConfig,
channels: Mapping[str, BaseChannel | ManagedValueSpec],
writes: _UpdateWrites,
fork_pending: set[str],
*,
is_first: bool,
) -> None:
"""Save an update's writes where its DeltaChannels will read them.
A DeltaChannel rebuilds its value from the writes saved on a checkpoint's
ancestors, so the writes go on the checkpoint the update builds on. If the
thread already moved past that checkpoint, its other children would read
them too, so the new checkpoint snapshots those channels instead. `writes`
holds `(task_id, writes, task)` per task; `task` is `None` for input.
"""
delta = _delta_writes(channels, writes)
if saved is None:
return
if is_first and checkpoint_superseded(checkpointer, config, saved):
fork_pending.update(ch for ch, _ in delta)
return
for task_id, task_writes, task in writes:
if task_writes:
checkpointer.put_writes(
checkpoint_config,
task_writes,
task_id,
**(_task_path_kwarg(checkpointer.put_writes, task) if task else {}),
)
async def _astore_or_fork_delta_writes(
checkpointer: BaseCheckpointSaver,
config: RunnableConfig,
saved: CheckpointTuple | None,
checkpoint_config: RunnableConfig,
channels: Mapping[str, BaseChannel | ManagedValueSpec],
writes: _UpdateWrites,
fork_pending: set[str],
*,
is_first: bool,
) -> None:
"""Async `_store_or_fork_delta_writes`."""
delta = _delta_writes(channels, writes)
if saved is None:
return
if is_first and await acheckpoint_superseded(checkpointer, config, saved):
fork_pending.update(ch for ch, _ in delta)
return
for task_id, task_writes, task in writes:
if task_writes:
await checkpointer.aput_writes(
checkpoint_config,
task_writes,
task_id,
**(_task_path_kwarg(checkpointer.aput_writes, task) if task else {}),
)
def _update_task_id(checkpoint_id: str, i: int) -> str:
"""Task id for the `i`th update of a superstep that has no task to reuse.
Savers keep one write per `(task_id, idx)`, so updates sharing an id lose
all but the first one's writes, which a `DeltaChannel` replays from. The
first update keeps the id a lone update has always had.
"""
return str(uuid5(UUID(checkpoint_id), INTERRUPT if i == 0 else f"{INTERRUPT}:{i}"))
def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str]]:
"""Index from a trigger to nodes that depend on it."""
trigger_to_nodes: defaultdict[str, list[str]] = defaultdict(list)
+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
+1 -1
View File
@@ -25,7 +25,7 @@ classifiers = [
]
dependencies = [
"langchain-core>=1.4.7,<2",
"langgraph-checkpoint>=4.1.0,<5.0.0",
"langgraph-checkpoint>=4.3.0,<5.0.0",
"langgraph-sdk>=0.4.6,<0.5.0",
"langgraph-prebuilt>=1.1.0,<1.2.0",
"xxhash>=3.5.0",
@@ -6,19 +6,23 @@ channel), lazy stub creation when no parent exists, and proper read-path
reconstruction via ancestor walks.
"""
import operator
import uuid
from typing import Annotated, Any
import pytest
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from typing_extensions import TypedDict
from langgraph._internal._constants import NULL_TASK_ID
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import _messages_delta_reducer
from langgraph.pregel._checkpoint import exit_delta_task_id
from langgraph.types import Command, Durability, interrupt
pytestmark = pytest.mark.anyio
@@ -35,6 +39,7 @@ def test_exit_delta_task_id_is_valid_uuid_and_ordered() -> None:
assert id1.split("-")[0] == "00000001"
assert id7.split("-")[0] == "00000007"
assert id1.endswith("-0270-bf16-1ef8-fb321bef9f3d")
assert exit_delta_task_id(0, NULL_TASK_ID) != NULL_TASK_ID
with pytest.raises(ValueError):
uuid.UUID(f"00000001-{tid}")
@@ -389,3 +394,234 @@ async def test_exit_snapshot_then_tail_deltas() -> None:
assert "seed-msg" in contents
assert "tail-msg" in contents
assert contents.index("seed-msg") < contents.index("tail-msg")
def _append(current: list, writes: list) -> list:
out = list(current)
for write in writes:
out.extend(write)
return out
class _ResumeState(TypedDict):
log: Annotated[list, DeltaChannel(_append)]
plain: Annotated[list, operator.add]
def _both(marker: str) -> dict:
return {"log": [marker], "plain": [marker]}
def _ask(marker: str) -> Any:
def ask(state: _ResumeState) -> dict:
interrupt("approve?")
return _both(marker)
return ask
@pytest.mark.parametrize("addressed", [False, True])
def test_resume_after_a_parallel_interrupt_replays_in_live_order(
sync_checkpointer: BaseCheckpointSaver, durability: Durability, addressed: bool
) -> None:
builder = StateGraph(_ResumeState)
builder.add_node("done", lambda state: _both("done"))
builder.add_node("ask", _ask("ask"))
builder.add_node("after", lambda state: _both("after"))
builder.add_edge(START, "done")
builder.add_edge(START, "ask")
builder.add_edge("ask", "after")
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "t"}}
graph.invoke(_both("in"), config, durability=durability)
head = graph.get_state(config).config
graph.invoke(
Command(resume="yes"), head if addressed else config, durability=durability
)
state = graph.get_state(config)
assert state.values["log"] == state.values["plain"]
assert sorted(state.values["log"]) == ["after", "ask", "done", "in"]
def test_resume_with_a_command_update_replays_its_write_once(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
builder = StateGraph(_ResumeState)
builder.add_node("done", lambda state: _both("done"))
builder.add_node("ask", _ask("ask"))
builder.add_edge(START, "done")
builder.add_edge(START, "ask")
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "t"}}
graph.invoke(_both("in"), config, durability=durability)
graph.invoke(
Command(resume="yes", update=_both("cmd")), config, durability=durability
)
state = graph.get_state(config)
assert state.values["log"] == state.values["plain"] == ["in", "cmd", "ask", "done"]
def test_command_update_on_an_input_checkpoint_matches_a_plain_channel(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
builder = StateGraph(_ResumeState)
builder.add_node("node", lambda state: _both("node"))
builder.add_edge(START, "node")
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "t"}}
graph.update_state(config, _both("in"), as_node="__input__")
graph.invoke(Command(update=_both("cmd")), config, durability=durability)
history = list(graph.get_state_history(config))
assert [s.values.get("log", []) for s in history] == [
s.values.get("plain", []) for s in history
]
replayed = graph.invoke(None, history[-1].config, durability=durability)
assert replayed["log"] == replayed["plain"]
def test_exit_command_update_on_a_new_thread_matches_a_plain_channel(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
builder = StateGraph(_ResumeState)
builder.add_node("node", lambda state: _both("node"))
builder.add_edge(START, "node")
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "t"}}
graph.invoke(Command(update=_both("cmd")), config, durability="exit")
history = list(graph.get_state_history(config))
assert [s.values.get("log", []) for s in history] == [
s.values.get("plain", []) for s in history
]
class _FlagState(_ResumeState, total=False):
extra: Annotated[list, DeltaChannel(_append)]
flag: bool
def test_addressed_resume_keeps_a_rerun_tasks_write_to_a_new_channel(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
def done(state: _FlagState) -> dict:
return {**_both("done"), **({"extra": ["new"]} if state.get("flag") else {})}
builder = StateGraph(_FlagState)
builder.add_node("done", done)
builder.add_node("ask", _ask("ask"))
builder.add_edge(START, "done")
builder.add_edge(START, "ask")
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "t"}}
graph.invoke(_both("in"), config, durability=durability)
head = graph.get_state(config).config
live = graph.invoke(
Command(resume="yes", update={"flag": True}), head, durability=durability
)
state = graph.get_state(config)
assert live["extra"] == state.values["extra"] == ["new"]
assert state.values["log"] == state.values["plain"]
def test_resume_interleaves_the_resumed_superstep_by_task_path(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
builder = StateGraph(_ResumeState)
builder.add_node("z_done", lambda state: _both("z"))
builder.add_node("a_asks", _ask("a"))
builder.add_edge(START, "z_done")
builder.add_edge(START, "a_asks")
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "t"}}
graph.invoke(_both("in"), config, durability=durability)
graph.invoke(Command(resume="yes"), config, durability=durability)
state = graph.get_state(config)
assert state.values["log"] == state.values["plain"] == ["in", "a", "z"]
class _TaskIdOrderSaver(InMemorySaver):
"""Replays each checkpoint's writes by task id, as savers without task path
ordering do."""
def get_tuple(self, config: Any) -> Any:
tup = super().get_tuple(config)
if tup and tup.pending_writes:
tup = tup._replace(pending_writes=sorted(tup.pending_writes))
return tup
get_delta_channel_history = BaseCheckpointSaver.get_delta_channel_history
def test_exit_run_replays_supersteps_in_order_on_a_task_id_ordered_saver() -> None:
builder = StateGraph(_ResumeState)
builder.add_node("a", lambda state: _both("a"))
builder.add_node("b", lambda state: _both("b"))
builder.add_edge(START, "a")
builder.add_edge("a", "b")
graph = builder.compile(checkpointer=_TaskIdOrderSaver())
config = {"configurable": {"thread_id": "t"}}
graph.invoke(_both("in"), config, durability="exit")
assert graph.get_state(config).values["log"] == ["in", "a", "b"]
def test_exit_resume_replays_supersteps_in_order_on_a_task_id_ordered_saver() -> None:
builder = StateGraph(_ResumeState)
builder.add_node("ask", _ask("ask"))
builder.add_node("after", lambda state: _both("after"))
builder.add_edge(START, "ask")
builder.add_edge("ask", "after")
graph = builder.compile(checkpointer=_TaskIdOrderSaver())
config = {"configurable": {"thread_id": "t"}}
graph.invoke(_both("in"), config, durability="exit")
graph.invoke(Command(resume="yes"), config, durability="exit")
assert graph.get_state(config).values["log"] == ["in", "ask", "after"]
class _FailingPutSaver(InMemorySaver):
fail = False
def put(
self, config: Any, checkpoint: Any, metadata: Any, new_versions: Any
) -> Any:
if self.fail:
raise RuntimeError("final checkpoint lost")
return super().put(config, checkpoint, metadata, new_versions)
def test_exit_resume_retried_after_its_final_checkpoint_fails_reruns_the_resumed_task() -> (
None
):
saver = _FailingPutSaver()
builder = StateGraph(_ResumeState)
builder.add_node("done", lambda state: _both("done"))
builder.add_node("ask", _ask("ask"))
builder.add_edge(START, "done")
builder.add_edge(START, "ask")
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "t"}}
graph.invoke(_both("in"), config, durability="exit")
saver.fail = True
with pytest.raises(RuntimeError, match="final checkpoint lost"):
graph.invoke(Command(resume="yes"), config, durability="exit")
saver.fail = False
graph.invoke(Command(resume="yes"), config, durability="exit")
state = graph.get_state(config)
assert state.values["log"] == state.values["plain"]
assert sorted(state.values["plain"]) == ["ask", "done", "in"]
@@ -740,20 +740,6 @@ def _build_deferred_after_interrupt(checkpointer: BaseCheckpointSaver) -> Any:
return builder.compile(checkpointer=checkpointer, interrupt_after=["a"])
@pytest.mark.parametrize(
"durability",
[
"sync",
"async",
pytest.param(
"exit",
marks=pytest.mark.xfail(
reason="exit durability stores a resumed run's loaded writes twice",
strict=True,
),
),
],
)
def test_resume_on_an_interrupted_head_consumes_its_writes_without_a_snapshot(
sync_checkpointer: BaseCheckpointSaver, durability: Durability
) -> None:
@@ -0,0 +1,126 @@
"""`DeltaChannel` replay must apply parallel writes in the order `invoke` did."""
import asyncio
from typing import Annotated, Any
import pytest
from langgraph.checkpoint.base import BaseCheckpointSaver
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import END, START, StateGraph
from langgraph.types import Send
pytestmark = pytest.mark.anyio
# Sorted, because live execution applies PULL tasks in node-name order.
FAN_OUT_NAMES = ["a", "b", "c", "d", "e", "f", "g", "h"]
SEND_ARGS = [f"send-{i:02d}" for i in range(12)]
def _append_reducer(current: list, updates: list) -> list:
return [*current, *(x for u in updates for x in u)]
def _build_fan_out_graph(checkpointer: BaseCheckpointSaver) -> Any:
class State(TypedDict):
items: Annotated[
list, DeltaChannel(_append_reducer, list, snapshot_frequency=10_000)
]
def make_node(label: str) -> Any:
return lambda state: {"items": [label]}
builder = StateGraph(State)
for name in FAN_OUT_NAMES:
builder.add_node(name, make_node(name))
builder.add_edge(START, name)
builder.add_edge(name, END)
return builder.compile(checkpointer=checkpointer)
def _build_send_fan_out_graph(checkpointer: BaseCheckpointSaver) -> Any:
class State(TypedDict):
items: Annotated[
list, DeltaChannel(_append_reducer, list, snapshot_frequency=10_000)
]
builder = StateGraph(State)
builder.add_node("worker", lambda arg: {"items": [arg]})
builder.add_conditional_edges(
START, lambda state: [Send("worker", n) for n in SEND_ARGS]
)
builder.add_edge("worker", END)
return builder.compile(checkpointer=checkpointer)
async def test_get_state_matches_live_send_order(
async_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_send_fan_out_graph(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
live = (await graph.ainvoke({"items": []}, config))["items"]
replayed = (await graph.aget_state(config)).values["items"]
assert live == SEND_ARGS
assert replayed == live
async def test_get_state_matches_live_invoke_order(
async_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_fan_out_graph(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
live = (await graph.ainvoke({"items": []}, config))["items"]
replayed = (await graph.aget_state(config)).values["items"]
assert live == FAN_OUT_NAMES
assert replayed == live
async def test_sync_get_state_on_async_saver_matches_live_order(
async_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_fan_out_graph(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
live = (await graph.ainvoke({"items": []}, config))["items"]
replayed = (await asyncio.to_thread(graph.get_state, config)).values["items"]
assert replayed == live
async def test_continuing_thread_preserves_committed_prefix(
async_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_fan_out_graph(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
first = (await graph.ainvoke({"items": []}, config))["items"]
second = (await graph.ainvoke({"items": []}, config))["items"]
assert second == first + first
assert (await graph.aget_state(config)).values["items"] == second
async def test_state_history_reports_live_order_at_every_step(
async_checkpointer: BaseCheckpointSaver,
) -> None:
runs = 3
graph = _build_fan_out_graph(async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
for _ in range(runs):
await graph.ainvoke({"items": []}, config)
live = FAN_OUT_NAMES * runs
seen = [
s.values["items"]
async for s in graph.aget_state_history(config)
if "items" in s.values
]
assert max(map(len, seen)) == len(live)
for values in seen:
assert values == live[: len(values)], f"{values} is not a prefix of {live}"
@@ -0,0 +1,103 @@
"""A run's input to a DeltaChannel input channel reads back on the checkpoints
built from it, and on no others."""
import pytest
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.last_value import LastValue
from langgraph.pregel import NodeBuilder, Pregel
from langgraph.types import Durability
pytestmark = pytest.mark.anyio
def _sorted_extend(current: list, writes: list) -> list:
return sorted([*current, *(item for write in writes for item in write)])
def _delta_input_graph() -> Pregel:
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log", "plain")
return Pregel(
nodes={"n": node},
channels={
"log": DeltaChannel(_sorted_extend),
"plain": BinaryOperatorAggregate(list, lambda a, b: sorted(a + b)),
"go": LastValue(int),
},
input_channels=["log", "plain", "go"],
output_channels=["log", "plain"],
checkpointer=InMemorySaver(),
)
def test_each_run_input_reads_back_on_its_own_checkpoints(
durability: Durability,
) -> None:
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"log": [0], "plain": [0], "go": 1}, config, durability=durability)
graph.invoke({"log": [5], "plain": [5], "go": 1}, config, durability=durability)
for state in graph.get_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
async def test_each_run_input_reads_back_on_its_own_checkpoints_async(
durability: Durability,
) -> None:
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke(
{"log": [0], "plain": [0], "go": 1}, config, durability=durability
)
await graph.ainvoke(
{"log": [5], "plain": [5], "go": 1}, config, durability=durability
)
async for state in graph.aget_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
OTHER_BRANCH_INPUTS = pytest.mark.parametrize(
"other_branch_input",
[{"go": 1}, {"log": [5], "plain": [5], "go": 1}],
ids=["other-branch-without-delta-input", "other-branch-with-delta-input"],
)
@OTHER_BRANCH_INPUTS
def test_run_input_from_an_older_checkpoint_stays_out_of_its_other_branch(
durability: Durability, other_branch_input: dict
) -> None:
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"go": 1}, config, durability=durability)
older = graph.get_state(config).config
graph.invoke(other_branch_input, config, durability=durability)
graph.invoke({"log": [7], "plain": [7], "go": 1}, older, durability=durability)
for state in graph.get_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
@OTHER_BRANCH_INPUTS
async def test_run_input_from_an_older_checkpoint_stays_out_of_its_other_branch_async(
durability: Durability, other_branch_input: dict
) -> None:
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"go": 1}, config, durability=durability)
older = (await graph.aget_state(config)).config
await graph.ainvoke(other_branch_input, config, durability=durability)
await graph.ainvoke(
{"log": [7], "plain": [7], "go": 1}, older, durability=durability
)
async for state in graph.aget_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
@@ -0,0 +1,78 @@
import pytest
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.last_value import LastValue
from langgraph.pregel import NodeBuilder, Pregel
pytestmark = pytest.mark.anyio
CONFIG = {"configurable": {"thread_id": "t"}}
def _extend(current: list, writes: list) -> list:
return [*current, *(item for write in writes for item in write)]
def _graph(reads: list) -> Pregel:
writer = NodeBuilder().subscribe_only("a").do(lambda _: [1]).write_to("d")
reader = NodeBuilder().subscribe_only("d").do(lambda d: reads.append(list(d)))
return Pregel(
nodes={"writer": writer, "reader": reader},
channels={"a": LastValue(str), "d": DeltaChannel(_extend)},
input_channels="a",
output_channels=["d"],
checkpointer=InMemorySaver(),
)
def test_an_input_update_from_before_the_first_write_starts_only_the_writer() -> None:
reads: list = []
graph = _graph(reads)
graph.invoke("go", CONFIG)
first = next(
s.config for s in graph.get_state_history(CONFIG) if s.metadata["step"] == -1
)
fork = graph.update_state(first, {"a": "go"}, as_node="__input__")
assert graph.get_state(fork).next == ("writer",)
graph.invoke(None, fork)
assert reads == [[1], [1]]
def test_a_replay_from_before_the_first_write_forks_with_only_the_writer_next() -> None:
reads: list = []
graph = _graph(reads)
graph.invoke("go", CONFIG)
first = next(
s.config for s in graph.get_state_history(CONFIG) if s.metadata["step"] == -1
)
graph.invoke(None, first, durability="sync")
fork = next(
s for s in graph.get_state_history(CONFIG) if s.metadata["source"] == "fork"
)
assert fork.next == ("writer",)
graph.invoke(None, fork.config)
assert reads == [[1], [1], [1]]
async def test_an_ainput_update_from_before_the_first_write_starts_only_the_writer() -> (
None
):
reads: list = []
graph = _graph(reads)
await graph.ainvoke("go", CONFIG)
first = [
s.config
async for s in graph.aget_state_history(CONFIG)
if s.metadata["step"] == -1
][0]
fork = await graph.aupdate_state(first, {"a": "go"}, as_node="__input__")
assert (await graph.aget_state(fork)).next == ("writer",)
await graph.ainvoke(None, fork)
assert reads == [[1], [1]]
@@ -0,0 +1,110 @@
"""An `Overwrite` through `update_state` snapshots its DeltaChannel on the
checkpoint the update saves, as a node's `Overwrite` does on the loop's."""
from typing import Annotated, Any
import pytest
from langchain_core.messages import HumanMessage
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.last_value import LastValue
from langgraph.graph import START, StateGraph
from langgraph.graph.message import _messages_delta_reducer
from langgraph.pregel import NodeBuilder, Pregel
from langgraph.types import Overwrite
pytestmark = pytest.mark.anyio
class _State(TypedDict):
messages: Annotated[list, DeltaChannel(_messages_delta_reducer)]
def _messages_graph(saver: InMemorySaver) -> Any:
builder = StateGraph(_State)
builder.add_node("model", lambda state: {})
builder.add_edge(START, "model")
return builder.compile(checkpointer=saver)
def _extend(current: list, writes: list) -> list:
return [*current, *(item for write in writes for item in write)]
def _delta_input_graph(saver: InMemorySaver) -> Pregel:
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log")
return Pregel(
nodes={"n": node},
channels={"log": DeltaChannel(_extend), "go": LastValue(int)},
input_channels=["log", "go"],
output_channels=["log"],
checkpointer=saver,
)
def test_update_state_with_an_overwrite_snapshots_the_channel() -> None:
saver = InMemorySaver()
graph = _messages_graph(saver)
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"messages": [HumanMessage(content="a", id="1")]}, config)
graph.update_state(
config,
{"messages": Overwrite([HumanMessage(content="b", id="2")])},
as_node="model",
)
head = saver.get_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
assert [m.content for m in graph.get_state(config).values["messages"]] == ["b"]
async def test_aupdate_state_with_an_overwrite_snapshots_the_channel() -> None:
saver = InMemorySaver()
graph = _messages_graph(saver)
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"messages": [HumanMessage(content="a", id="1")]}, config)
await graph.aupdate_state(
config,
{"messages": Overwrite([HumanMessage(content="b", id="2")])},
as_node="model",
)
head = await saver.aget_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
values = (await graph.aget_state(config)).values
assert [m.content for m in values["messages"]] == ["b"]
def test_update_state_as_input_with_an_overwrite_snapshots_the_channel() -> None:
saver = InMemorySaver()
graph = _delta_input_graph(saver)
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"log": [0], "go": 1}, config)
graph.update_state(config, {"log": Overwrite([1])}, as_node="__input__")
head = saver.get_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
assert graph.get_state(config).values["log"] == [1]
async def test_aupdate_state_as_input_with_an_overwrite_snapshots_the_channel() -> None:
saver = InMemorySaver()
graph = _delta_input_graph(saver)
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"log": [0], "go": 1}, config)
await graph.aupdate_state(config, {"log": Overwrite([1])}, as_node="__input__")
head = await saver.aget_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("log"), _DeltaSnapshot)
assert (await graph.aget_state(config)).values["log"] == [1]
@@ -20,23 +20,28 @@ from typing import Annotated, Any
import pytest
from langchain_core.messages import HumanMessage
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from typing_extensions import TypedDict
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.last_value import LastValue
from langgraph.graph import START, StateGraph
from langgraph.graph.message import _messages_delta_reducer
from langgraph.types import StateUpdate
from langgraph.pregel import NodeBuilder, Pregel
from langgraph.types import StateSnapshot, StateUpdate
pytestmark = pytest.mark.anyio
def _build_graph(
checkpointer: InMemorySaver,
checkpointer: BaseCheckpointSaver,
*,
two_nodes: bool = False,
snapshot_frequency: int = 1000,
interrupt_before: list[str] | None = None,
) -> Any:
"""Compile a minimal DeltaChannel-backed `messages` graph.
@@ -63,7 +68,7 @@ def _build_graph(
builder.set_finish_point("assistant")
else:
builder.set_finish_point("model")
return builder.compile(checkpointer=checkpointer)
return builder.compile(checkpointer=checkpointer, interrupt_before=interrupt_before)
# ---------------------------------------------------------------------------
@@ -304,15 +309,14 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
that each call `put_writes`. Guards the regression where moving
`put_writes` outside the per-task loop would persist only the last
task's writes.
Explicit `task_id`s are required to disambiguate writes belonging to
different `StateUpdate`s targeting the same node — otherwise both share
the deterministic interrupt-derived id and collide in the saver.
"""
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": "bulk-multi-task"}}
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
base = saver.get_tuple(config)
assert base is not None
graph.bulk_update_state(
config,
@@ -332,13 +336,155 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
],
)
stored = saver.get_tuple(base.config)
assert stored is not None
assert {task_id for task_id, _, _ in stored.pending_writes or []} == {
"task-1",
"task-2",
}, "explicit task ids must key the stored writes"
state = graph.get_state(config)
contents = [m.content for m in state.values["messages"]]
ids = [m.id for m in state.values["messages"]]
assert sorted(contents) == ["first", "second"], (
assert sorted(contents) == ["first", "hi", "second"], (
f"both updates' writes must persist; got {contents}"
)
assert sorted(ids) == ["m1", "m2"]
assert sorted(ids) == ["hi", "m1", "m2"]
def _update(content: str, as_node: str) -> StateUpdate:
return StateUpdate(
values={"messages": [HumanMessage(content=content, id=content)]},
as_node=as_node,
)
def _contents(state: StateSnapshot) -> list[str]:
return [m.content for m in state.values["messages"]]
def test_bulk_update_state_keeps_every_update_without_task_ids(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_graph(sync_checkpointer, two_nodes=True)
config = {"configurable": {"thread_id": "bulk-no-task-ids"}}
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
graph.bulk_update_state(
config,
[
[
_update("first", "model"),
_update("second", "model"),
_update("third", "assistant"),
]
],
)
contents = _contents(graph.get_state(config))
assert sorted(contents) == ["first", "hi", "second", "third"], (
f"every update's writes must persist; got {contents}"
)
async def test_abulk_update_state_keeps_every_update_without_task_ids(
async_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_graph(async_checkpointer, two_nodes=True)
config = {"configurable": {"thread_id": "bulk-no-task-ids"}}
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
await graph.abulk_update_state(
config,
[
[
_update("first", "model"),
_update("second", "model"),
_update("third", "assistant"),
]
],
)
contents = _contents(await graph.aget_state(config))
assert sorted(contents) == ["first", "hi", "second", "third"], (
f"every update's writes must persist; got {contents}"
)
def test_bulk_update_state_keeps_every_update_next_to_a_pending_task(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
graph = _build_graph(
sync_checkpointer, two_nodes=True, interrupt_before=["assistant"]
)
config = {"configurable": {"thread_id": "bulk-pending-task"}}
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
assert graph.get_state(config).next == ("assistant",)
graph.bulk_update_state(
config,
[
[
_update("first", "assistant"),
_update("second", "model"),
_update("third", "model"),
]
],
)
contents = _contents(graph.get_state(config))
assert sorted(contents) == ["first", "hi", "second", "third"], (
f"every update's writes must persist; got {contents}"
)
class _TaskPathOrderSaver(InMemorySaver):
"""Replays each checkpoint's writes by `(task_path, task_id, idx)`."""
def get_tuple(self, config: Any) -> Any:
tup = super().get_tuple(config)
if tup is None or not tup.pending_writes:
return tup
conf = tup.config["configurable"]
stored = self.writes[
(conf["thread_id"], conf["checkpoint_ns"], conf["checkpoint_id"])
]
rows = sorted(
zip(stored.items(), tup.pending_writes),
key=lambda row: (row[0][1][3], *row[0][0]),
)
return tup._replace(pending_writes=[write for _, write in rows])
get_delta_channel_history = BaseCheckpointSaver.get_delta_channel_history
aget_delta_channel_history = BaseCheckpointSaver.aget_delta_channel_history
GIVEN = ["u1", "u2", "u3", "u4", "u5", "u6"]
def _updates_in_given_order() -> list[list[StateUpdate]]:
return [
[_update(c, "assistant" if i % 2 else "model") for i, c in enumerate(GIVEN)]
]
def test_bulk_update_state_replays_updates_in_the_order_given() -> None:
graph = _build_graph(_TaskPathOrderSaver(), two_nodes=True)
config = {"configurable": {"thread_id": "bulk-order"}}
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
graph.bulk_update_state(config, _updates_in_given_order())
assert _contents(graph.get_state(config)) == ["hi", *GIVEN]
async def test_abulk_update_state_replays_updates_in_the_order_given() -> None:
graph = _build_graph(_TaskPathOrderSaver(), two_nodes=True)
config = {"configurable": {"thread_id": "bulk-order"}}
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
await graph.abulk_update_state(config, _updates_in_given_order())
assert _contents(await graph.aget_state(config)) == ["hi", *GIVEN]
# ---------------------------------------------------------------------------
@@ -395,3 +541,189 @@ def test_update_state_that_snapshots_keeps_a_deferred_node_pending() -> None:
assert [m.content for m in final["messages"]] == ["s", "a", "u", "b"]
assert graph.get_state(config).next == ()
def _sorted_extend(current: list, writes: list) -> list:
return sorted([*current, *(item for write in writes for item in write)])
def _delta_input_graph(snapshot_frequency: int = 1000) -> Any:
node = NodeBuilder().subscribe_only("go").do(lambda _: [2]).write_to("log", "plain")
return Pregel(
nodes={"n": node},
channels={
"log": DeltaChannel(_sorted_extend, snapshot_frequency=snapshot_frequency),
"plain": BinaryOperatorAggregate(list, lambda a, b: sorted(a + b)),
"go": LastValue(int),
},
input_channels=["log", "plain", "go"],
output_channels=["log", "plain"],
checkpointer=InMemorySaver(),
)
@pytest.mark.parametrize("snapshot_frequency", [1, 2])
def test_update_as_input_reads_back_on_its_checkpoint_and_after_the_next_run(
snapshot_frequency: int,
) -> None:
graph = _delta_input_graph(snapshot_frequency)
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"log": [0], "plain": [0], "go": 1}, config)
graph.update_state(config, {"log": [1], "plain": [1], "go": 1}, as_node="__input__")
after_update = graph.get_state(config).values
graph.invoke(None, config)
after_run = graph.get_state(config).values
assert after_update["log"] == after_update["plain"]
assert after_run["log"] == after_run["plain"]
@pytest.mark.parametrize("snapshot_frequency", [1, 2])
async def test_aupdate_as_input_reads_back_on_its_checkpoint_and_after_the_next_run(
snapshot_frequency: int,
) -> None:
graph = _delta_input_graph(snapshot_frequency)
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"log": [0], "plain": [0], "go": 1}, config)
await graph.aupdate_state(
config, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
)
after_update = (await graph.aget_state(config)).values
await graph.ainvoke(None, config)
after_run = (await graph.aget_state(config)).values
assert after_update["log"] == after_update["plain"]
assert after_run["log"] == after_run["plain"]
def test_update_as_input_to_an_older_checkpoint_stays_out_of_its_other_branch() -> None:
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"go": 1}, config)
older = graph.get_state(config).config
graph.invoke({"go": 1}, config)
other_branch = graph.get_state(config)
edited = graph.update_state(
older, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
)
values = graph.get_state(edited).values
assert values["log"] == values["plain"]
assert graph.get_state(other_branch.config).values == other_branch.values
async def test_aupdate_as_input_to_an_older_checkpoint_stays_out_of_its_other_branch() -> (
None
):
graph = _delta_input_graph()
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"go": 1}, config)
older = (await graph.aget_state(config)).config
await graph.ainvoke({"go": 1}, config)
other_branch = await graph.aget_state(config)
edited = await graph.aupdate_state(
older, {"log": [1], "plain": [1], "go": 1}, as_node="__input__"
)
values = (await graph.aget_state(edited)).values
assert values["log"] == values["plain"]
assert (await graph.aget_state(other_branch.config)).values == other_branch.values
def _message_ids(graph: Any, config: dict) -> list[str | None]:
return [m.id for m in graph.get_state(config).values["messages"]]
async def _amessage_ids(graph: Any, config: dict) -> list[str | None]:
return [m.id for m in (await graph.aget_state(config)).values["messages"]]
def _messages_input_graph() -> Any:
node = (
NodeBuilder()
.subscribe_only("go")
.do(lambda _: [HumanMessage("n", id="n")])
.write_to("messages")
)
return Pregel(
nodes={"n": node},
channels={
"messages": DeltaChannel(_messages_delta_reducer),
"go": LastValue(int),
},
input_channels=["messages", "go"],
output_channels=["messages"],
checkpointer=InMemorySaver(),
)
def test_update_state_gives_a_message_an_id_that_every_read_keeps() -> None:
graph = _build_graph(InMemorySaver())
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"messages": [HumanMessage("a", id="a")]}, config)
graph.update_state(config, {"messages": [HumanMessage("b")]})
first, second = _message_ids(graph, config), _message_ids(graph, config)
assert first[-1] is not None
assert first == second
async def test_aupdate_state_gives_a_message_an_id_that_every_read_keeps() -> None:
graph = _build_graph(InMemorySaver())
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"messages": [HumanMessage("a", id="a")]}, config)
await graph.aupdate_state(config, {"messages": [HumanMessage("b")]})
first, second = (
await _amessage_ids(graph, config),
await _amessage_ids(graph, config),
)
assert first[-1] is not None
assert first == second
def test_update_state_on_an_older_checkpoint_gives_a_message_an_id() -> None:
graph = _build_graph(InMemorySaver())
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"messages": [HumanMessage("a", id="a")]}, config)
older = graph.get_state(config).config
graph.invoke({"messages": [HumanMessage("c", id="c")]}, config)
branch = graph.update_state(older, {"messages": [HumanMessage("b")]})
assert _message_ids(graph, branch)[-1] is not None
def test_update_as_input_gives_a_message_an_id_that_every_read_keeps() -> None:
graph = _messages_input_graph()
config = {"configurable": {"thread_id": "t"}}
graph.invoke({"messages": [HumanMessage("a", id="a")], "go": 1}, config)
graph.update_state(config, {"messages": [HumanMessage("b")]}, as_node="__input__")
first, second = _message_ids(graph, config), _message_ids(graph, config)
assert first[-1] is not None
assert first == second
async def test_aupdate_as_input_gives_a_message_an_id_that_every_read_keeps() -> None:
graph = _messages_input_graph()
config = {"configurable": {"thread_id": "t"}}
await graph.ainvoke({"messages": [HumanMessage("a", id="a")], "go": 1}, config)
await graph.aupdate_state(
config, {"messages": [HumanMessage("b")]}, as_node="__input__"
)
first, second = (
await _amessage_ids(graph, config),
await _amessage_ids(graph, config),
)
assert first[-1] is not None
assert first == second
@@ -0,0 +1,168 @@
"""A checkpoint must never be saved without the `DeltaChannel` writes it reads."""
import operator
import threading
from typing import Annotated, Any
import pytest
from langgraph.checkpoint.memory import InMemorySaver
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.types import Durability
pytestmark = pytest.mark.anyio
INPUT = {"log": [], "plain": []}
FINAL = {"log": ["a", "b", "c"], "plain": ["a", "b", "c"]}
# Exit mode saves nothing before the failed write, so its retry starts over.
RETRIES = [
pytest.param("sync", None, id="sync"),
pytest.param("async", None, id="async"),
pytest.param("exit", INPUT, id="exit"),
]
def _append(current: list, writes: list) -> list:
return [*current, *(item for write in writes for item in write)]
class _State(TypedDict):
log: Annotated[list, DeltaChannel(_append)]
plain: Annotated[list, operator.add]
class _FailsTheWriteOfBOnce(InMemorySaver):
failed = False
def _fail_once(self, writes: Any) -> None:
if not self.failed and ("log", ["b"]) in writes:
self.failed = True
raise ConnectionError("b's write was not saved")
def put_writes(
self, config: Any, writes: Any, task_id: str, task_path: str = ""
) -> None:
self._fail_once(writes)
super().put_writes(config, writes, task_id, task_path)
async def aput_writes(
self, config: Any, writes: Any, task_id: str, task_path: str = ""
) -> None:
self._fail_once(writes)
await super().aput_writes(config, writes, task_id, task_path)
class _FailsTheSaveAfterAOnce(InMemorySaver):
failed = False
def _fail_once(self, checkpoint: Any) -> None:
if not self.failed and checkpoint["channel_values"].get("plain") == ["a"]:
self.failed = True
raise ConnectionError("the checkpoint after a was not saved")
def put(
self, config: Any, checkpoint: Any, metadata: Any, new_versions: Any
) -> Any:
self._fail_once(checkpoint)
return super().put(config, checkpoint, metadata, new_versions)
async def aput(
self, config: Any, checkpoint: Any, metadata: Any, new_versions: Any
) -> Any:
self._fail_once(checkpoint)
return await super().aput(config, checkpoint, metadata, new_versions)
def _a_then_b_then_c(saver: InMemorySaver) -> Any:
builder = StateGraph(_State)
for name in "abc":
builder.add_node(
name, lambda state, name=name: {"log": [name], "plain": [name]}
)
builder.add_edge(START, "a")
builder.add_edge("a", "b")
builder.add_edge("b", "c")
return builder.compile(checkpointer=saver)
@pytest.mark.parametrize(("durability", "retry_input"), RETRIES)
def test_a_failed_delta_write_is_rerun_not_lost(
durability: Durability, retry_input: dict | None
) -> None:
graph = _a_then_b_then_c(_FailsTheWriteOfBOnce())
config = {"configurable": {"thread_id": "t"}}
with pytest.raises(ConnectionError):
graph.invoke(INPUT, config, durability=durability)
for state in graph.get_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
graph.invoke(retry_input, config, durability=durability)
assert graph.get_state(config).values == FINAL
@pytest.mark.parametrize(("durability", "retry_input"), RETRIES)
async def test_a_failed_delta_write_is_rerun_not_lost_async(
durability: Durability, retry_input: dict | None
) -> None:
graph = _a_then_b_then_c(_FailsTheWriteOfBOnce())
config = {"configurable": {"thread_id": "t"}}
with pytest.raises(ConnectionError):
await graph.ainvoke(INPUT, config, durability=durability)
async for state in graph.aget_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
await graph.ainvoke(retry_input, config, durability=durability)
assert (await graph.aget_state(config)).values == FINAL
@pytest.mark.parametrize("durability", ["sync", "async"])
def test_a_failed_checkpoint_save_is_rerun_not_built_on(
durability: Durability,
) -> None:
graph = _a_then_b_then_c(_FailsTheSaveAfterAOnce())
config = {"configurable": {"thread_id": "t"}}
with pytest.raises(ConnectionError):
graph.invoke(INPUT, config, durability=durability)
for state in graph.get_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
graph.invoke(None, config, durability=durability)
assert graph.get_state(config).values == FINAL
@pytest.mark.parametrize("durability", ["sync", "async"])
async def test_a_failed_checkpoint_save_is_rerun_not_built_on_async(
durability: Durability,
) -> None:
graph = _a_then_b_then_c(_FailsTheSaveAfterAOnce())
config = {"configurable": {"thread_id": "t"}}
with pytest.raises(ConnectionError):
await graph.ainvoke(INPUT, config, durability=durability)
async for state in graph.aget_state_history(config):
assert state.values.get("log", []) == state.values.get("plain", [])
await graph.ainvoke(None, config, durability=durability)
assert (await graph.aget_state(config)).values == FINAL
def test_a_delta_graph_finishes_on_a_single_background_thread() -> None:
graph = _a_then_b_then_c(InMemorySaver())
config = {"configurable": {"thread_id": "t"}, "max_concurrency": 1}
result: dict = {}
run = threading.Thread(
target=lambda: result.update(graph.invoke(INPUT, config, durability="async")),
daemon=True,
)
run.start()
run.join(timeout=10)
assert not run.is_alive(), "invoke hung"
assert result == FINAL
+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
+89
View File
@@ -248,6 +248,95 @@ def test_drain_from_subgraph_can_resume_parent() -> None:
}
class _BreakpointState(TypedDict, total=False):
first: str
second: str
def _drain_before_breakpoint_builder(
second_runs: list[str],
) -> StateGraph[_BreakpointState]:
def first(state: _BreakpointState, runtime: Runtime) -> _BreakpointState:
runtime.control.request_drain("rollout")
return {"first": "done"}
def second(state: _BreakpointState) -> _BreakpointState:
second_runs.append("second")
return {"second": "done"}
builder = StateGraph(_BreakpointState)
builder.add_node("first", first)
builder.add_node("second", second)
builder.add_edge(START, "first")
builder.add_edge("first", "second")
return builder
@pytest.mark.parametrize("at_compile", [True, False])
def test_drain_before_interrupt_before_stops_at_breakpoint(at_compile: bool) -> None:
# A resume passes the next step's breakpoints, so a drain right before a
# breakpoint has to stop there instead, or the resumed run skips it.
second_runs: list[str] = []
breakpoints = {"interrupt_before": ["second"]}
compiled = _drain_before_breakpoint_builder(second_runs).compile(
checkpointer=MemorySaver(), **(breakpoints if at_compile else {})
)
run_kwargs = {} if at_compile else breakpoints
config = {"configurable": {"thread_id": "drain-breakpoint"}}
assert compiled.invoke({}, config, control=RunControl(), **run_kwargs) == {
"first": "done"
}
assert compiled.get_state(config).next == ("second",)
assert second_runs == []
assert compiled.invoke(None, config, **run_kwargs) == {
"first": "done",
"second": "done",
}
assert second_runs == ["second"]
def test_drain_before_subgraph_interrupt_before_stops_at_breakpoint() -> None:
second_runs: list[str] = []
child = _drain_before_breakpoint_builder(second_runs).compile(
interrupt_before=["second"]
)
parent = StateGraph(_BreakpointState)
parent.add_node("child", child)
parent.add_edge(START, "child")
compiled = parent.compile(checkpointer=MemorySaver())
config = {"configurable": {"thread_id": "drain-subgraph-breakpoint"}}
compiled.invoke({}, config, control=RunControl())
state = compiled.get_state(config, subgraphs=True)
assert state.next == ("child",)
assert state.tasks[0].state.next == ("second",)
assert second_runs == []
assert compiled.invoke(None, config) == {"first": "done", "second": "done"}
assert second_runs == ["second"]
@pytest.mark.anyio
async def test_drain_before_interrupt_before_stops_at_breakpoint_async() -> None:
second_runs: list[str] = []
compiled = _drain_before_breakpoint_builder(second_runs).compile(
checkpointer=MemorySaver(), interrupt_before=["second"]
)
config = {"configurable": {"thread_id": "drain-breakpoint-async"}}
assert await compiled.ainvoke({}, config, control=RunControl()) == {"first": "done"}
assert (await compiled.aget_state(config)).next == ("second",)
assert second_runs == []
assert await compiled.ainvoke(None, config) == {
"first": "done",
"second": "done",
}
assert second_runs == ["second"]
@pytest.mark.anyio
async def test_drain_requested_in_terminal_step_finishes_normally_async() -> None:
class State(TypedDict, total=False):
@@ -8,20 +8,33 @@ caller kwargs to the inner ``(a)stream`` call but rejects ``stream_mode`` and
``subgraphs`` since v3 owns them (``stream_mode`` is built from the
transformer mux; ``subgraphs`` is forced True so nested namespaces flow
through scoped muxes).
A second regression is pinned here: #7677 (first released in 1.2.0a3)
declared `interrupt_before` / `interrupt_after` / `control` as named
parameters on the `Pregel.stream_events` / `astream_events` dispatchers
but forwarded them only on the v3 branch, silently dropping them on v1/v2
(where they had reached `(a)stream` through `**kwargs` before).
`TestAstreamKwargsForwardedOnEveryVersion` and friends pin that they
reach `(a)stream` on every version, with the v1/v2 passthrough semantics
restored (exactly what the caller passed, including an explicit `None`)
and v3's explicit-default semantics preserved.
"""
from __future__ import annotations
import sys
from collections.abc import AsyncIterator, Callable
from dataclasses import dataclass
from typing import Any
import pytest
from langgraph.checkpoint.memory import InMemorySaver
from typing_extensions import TypedDict
from langgraph.constants import END, START
from langgraph.errors import GraphDrained
from langgraph.graph import StateGraph
from langgraph.runtime import Runtime
from langgraph.runtime import RunControl, Runtime
NEEDS_CONTEXTVARS = pytest.mark.skipif(
sys.version_info < (3, 11),
@@ -106,3 +119,274 @@ class TestKwargForwardingAsync:
version="v3",
subgraphs=False,
)
_KWARG_NAMES = ("control", "interrupt_before", "interrupt_after")
def _build_two_step_graph(
first: Callable[[_State], dict[str, Any]] | None = None,
) -> Any:
"""A `first -> second` graph with a checkpointer, for interrupt/drain tests."""
def default_first(state: _State) -> dict[str, Any]:
return {"message": state["message"] + " first"}
def second(state: _State) -> dict[str, Any]:
return {"message": state["message"] + " second"}
builder = StateGraph(_State)
builder.add_node("first", first or default_first)
builder.add_node("second", second)
builder.add_edge(START, "first")
builder.add_edge("first", "second")
builder.add_edge("second", END)
return builder.compile(checkpointer=InMemorySaver())
async def _drive_astream_events(
graph: Any, config: dict[str, Any], version: str, **kwargs: Any
) -> None:
"""Consume an astream_events run for `version` to completion."""
if version == "v3":
run = await graph.astream_events(
{"message": "hi"}, config, version="v3", **kwargs
)
await run.output()
else:
async for _ in graph.astream_events(
{"message": "hi"}, config, version=version, **kwargs
):
pass
@pytest.mark.anyio
@pytest.mark.filterwarnings("ignore:astream_events version='v1' is deprecated")
@pytest.mark.parametrize("version", ["v1", "v2", "v3"])
class TestAstreamKwargsForwardedOnEveryVersion:
"""`interrupt_before`/`interrupt_after`/`control` reach `astream` on every
version.
Regression test for #7677 (first released in 1.2.0a3): the dispatchers
captured these parameters as named arguments but forwarded them only on
the v3 branch, silently dropping them on v1/v2.
"""
async def test_interrupt_before(self, version: str) -> None:
graph = _build_two_step_graph()
config = {"configurable": {"thread_id": "ib"}}
await _drive_astream_events(graph, config, version, interrupt_before=["second"])
state = await graph.aget_state(config)
assert state.next == ("second",)
assert state.values == {"message": "hi first"}
async def test_interrupt_after(self, version: str) -> None:
graph = _build_two_step_graph()
config = {"configurable": {"thread_id": "ia"}}
await _drive_astream_events(graph, config, version, interrupt_after=["first"])
state = await graph.aget_state(config)
assert state.next == ("second",)
assert state.values == {"message": "hi first"}
async def test_pre_drained_control(self, version: str) -> None:
graph = _build_two_step_graph()
config = {"configurable": {"thread_id": "drain"}}
control = RunControl()
control.request_drain("sigterm")
with pytest.raises(GraphDrained, match="sigterm"):
await _drive_astream_events(graph, config, version, control=control)
class TestStreamEventsV3SyncInterrupts:
"""Sync v3 static interrupts reach `stream()` after the kwargs rewire."""
@pytest.mark.parametrize(
("kwarg", "node"),
[("interrupt_before", "second"), ("interrupt_after", "first")],
)
def test_static_interrupt(self, kwarg: str, node: str) -> None:
graph = _build_two_step_graph()
config = {"configurable": {"thread_id": "sync"}}
run = graph.stream_events(
{"message": "hi"}, config, version="v3", **{kwarg: [node]}
)
list(run.values)
state = graph.get_state(config)
assert state.next == ("second",)
assert state.values == {"message": "hi first"}
@pytest.mark.anyio
@pytest.mark.filterwarnings("ignore:astream_events version='v1' is deprecated")
@pytest.mark.parametrize("version", ["v1", "v2", "v3"])
class TestAstreamMidRunDrain:
"""A drain requested from inside a node propagates out of `astream_events`.
This is the graceful-shutdown scenario: `request_drain()` called while
the run is in flight (e.g. from a signal handler), with v1/v2 running
inside core's event-stream task. The caller's own `RunControl` is used
(the drain reason proves identity), `GraphDrained` propagates to the
consumer, and the checkpoint keeps the pending step.
"""
async def test_drain_requested_inside_first_node(self, version: str) -> None:
control = RunControl()
def first(state: _State) -> dict[str, Any]:
control.request_drain("sigterm-mid")
return {"message": state["message"] + " first"}
graph = _build_two_step_graph(first)
config = {"configurable": {"thread_id": "midrun"}}
with pytest.raises(GraphDrained, match="sigterm-mid"):
await _drive_astream_events(graph, config, version, control=control)
state = await graph.aget_state(config)
assert state.next == ("second",)
assert state.values == {"message": "hi first"}
def _record_astream_kwargs(graph: Any) -> list[dict[str, Any]]:
"""Patch `graph.astream` to record the kwargs each call receives."""
received: list[dict[str, Any]] = []
original = graph.astream
async def recording_astream(
input: Any, config: Any = None, **kwargs: Any
) -> AsyncIterator[Any]:
received.append(kwargs)
async for chunk in original(input, config, **kwargs):
yield chunk
graph.astream = recording_astream # type: ignore[method-assign]
return received
@pytest.mark.anyio
@pytest.mark.filterwarnings("ignore:astream_events version='v1' is deprecated")
@pytest.mark.parametrize("version", ["v1", "v2"])
class TestAstreamV1V2KwargsPassthrough:
"""v1/v2 forward to `astream` exactly what the caller passed.
Pre-#7677 semantics: an explicit `None` is forwarded as `None`, and an
omitted argument is not forwarded at all (so an override's own default
would apply).
"""
@pytest.mark.parametrize("name", _KWARG_NAMES)
async def test_passed_values_reach_astream(self, version: str, name: str) -> None:
graph = _build_two_step_graph()
received = _record_astream_kwargs(graph)
value: Any = RunControl() if name == "control" else ["second"]
await _drive_astream_events(
graph, {"configurable": {"thread_id": "rec"}}, version, **{name: value}
)
assert len(received) == 1
assert received[0][name] == value
if name == "control":
assert received[0][name] is value
@pytest.mark.parametrize("name", _KWARG_NAMES)
async def test_explicit_none_is_forwarded(self, version: str, name: str) -> None:
graph = _build_two_step_graph()
received = _record_astream_kwargs(graph)
await _drive_astream_events(
graph, {"configurable": {"thread_id": "rec-none"}}, version, **{name: None}
)
assert len(received) == 1
assert received[0][name] is None
@pytest.mark.parametrize("name", _KWARG_NAMES)
async def test_omitted_values_are_absent(self, version: str, name: str) -> None:
graph = _build_two_step_graph()
received = _record_astream_kwargs(graph)
await _drive_astream_events(
graph, {"configurable": {"thread_id": "rec-omit"}}, version
)
assert len(received) == 1
assert name not in received[0]
@pytest.mark.anyio
class TestAstreamV3KwargsDefaults:
"""v3 keeps its since-inception explicit-default semantics (#7519).
Omitted `interrupt_before`/`interrupt_after`/`control` are supplied to
`astream` as `None`; passed values are forwarded as-is.
"""
async def test_omitted_values_arrive_as_none(self) -> None:
graph = _build_two_step_graph()
received = _record_astream_kwargs(graph)
await _drive_astream_events(
graph, {"configurable": {"thread_id": "v3-rec"}}, "v3"
)
assert len(received) == 1
assert received[0]["control"] is None
assert received[0]["interrupt_before"] is None
assert received[0]["interrupt_after"] is None
@pytest.mark.parametrize("name", _KWARG_NAMES)
async def test_passed_values_reach_astream(self, name: str) -> None:
graph = _build_two_step_graph()
received = _record_astream_kwargs(graph)
value: Any = RunControl() if name == "control" else ["second"]
await _drive_astream_events(
graph, {"configurable": {"thread_id": "v3-rec-2"}}, "v3", **{name: value}
)
assert len(received) == 1
assert received[0][name] == value
if name == "control":
assert received[0][name] is value
def _record_stream_kwargs(graph: Any) -> list[dict[str, Any]]:
"""Patch `graph.stream` to record the kwargs each call receives."""
received: list[dict[str, Any]] = []
original = graph.stream
def recording_stream(input: Any, config: Any = None, **kwargs: Any) -> Any:
received.append(kwargs)
yield from original(input, config, **kwargs)
graph.stream = recording_stream # type: ignore[method-assign]
return received
def _drive_stream_events_v3(graph: Any, config: dict[str, Any], **kwargs: Any) -> None:
"""Consume a sync v3 stream_events run to completion."""
run = graph.stream_events({"message": "hi"}, config, version="v3", **kwargs)
list(run.values)
class TestStreamEventsV3SyncKwargsDefaults:
"""Sync v3 keeps its since-inception explicit-default semantics (#7519).
Mirror of `TestAstreamV3KwargsDefaults`: omitted
`interrupt_before`/`interrupt_after`/`control` are supplied to `stream`
as `None`; passed values are forwarded as-is. Pins the sync helper
against a kwargs-only "simplification" that would change subclass
default handling.
"""
def test_omitted_values_arrive_as_none(self) -> None:
graph = _build_two_step_graph()
received = _record_stream_kwargs(graph)
_drive_stream_events_v3(graph, {"configurable": {"thread_id": "s-rec"}})
assert len(received) == 1
assert received[0]["control"] is None
assert received[0]["interrupt_before"] is None
assert received[0]["interrupt_after"] is None
@pytest.mark.parametrize("name", _KWARG_NAMES)
def test_passed_values_reach_stream(self, name: str) -> None:
graph = _build_two_step_graph()
received = _record_stream_kwargs(graph)
value: Any = RunControl() if name == "control" else ["second"]
_drive_stream_events_v3(
graph, {"configurable": {"thread_id": "s-rec-2"}}, **{name: value}
)
assert len(received) == 1
assert received[0][name] == value
if name == "control":
assert received[0][name] is value
+23 -8
View File
@@ -1217,6 +1217,20 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/38/64/285f20a31679bf547b75602702f7800e74dbabae36ef324f716c02804753/jupyter-1.1.1-py2.py3-none-any.whl", hash = "sha256:7a59533c22af65439b24bbe60373a4e95af8f16ac65a6c00820ad378e3f7cc83", size = 2657, upload-time = "2024-08-30T07:15:47.045Z" },
]
[[package]]
name = "jupyter-builder"
version = "1.2.3"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jupyter-core" },
{ name = "tomli", marker = "python_full_version < '3.11'" },
{ name = "traitlets" },
]
sdist = { url = "https://files.pythonhosted.org/packages/75/3e/56f593e6a664cd14d441724654ce8243f3cac7d3e45867cf084749c8bc75/jupyter_builder-1.2.3.tar.gz", hash = "sha256:01aba6794eb9b19e0e29ae21137ca60ba4135c70347d4b9f664a586822b8c809", size = 1024218, upload-time = "2026-09-04T18:54:09.411Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/84/aa/be79e87c50698673f196d0633fc1664607da2289f9688d1de7a1c9db9e71/jupyter_builder-1.2.3-py3-none-any.whl", hash = "sha256:c5ea5a7190c2a7b082494abade98eece1b2b5bd5dbd7d610606cbcd10a1d08b3", size = 947541, upload-time = "2026-09-04T18:54:07.644Z" },
]
[[package]]
name = "jupyter-client"
version = "8.8.0"
@@ -1342,28 +1356,28 @@ wheels = [
[[package]]
name = "jupyterlab"
version = "4.5.10"
version = "4.6.4"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "async-lru" },
{ name = "httpx" },
{ name = "ipykernel" },
{ name = "jinja2" },
{ name = "jupyter-builder" },
{ name = "jupyter-core" },
{ name = "jupyter-lsp" },
{ name = "jupyter-server" },
{ name = "jupyterlab-server" },
{ name = "notebook-shim" },
{ name = "packaging" },
{ name = "setuptools" },
{ name = "tomli", marker = "python_full_version < '3.11'" },
{ name = "tornado" },
{ name = "traitlets" },
{ name = "typing-extensions", marker = "python_full_version < '3.12'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/74/24/621aa20ec0d2fe72f52095bda0fc1be7738ac21aabe4129ff623140d5cdf/jupyterlab-4.5.10.tar.gz", hash = "sha256:77e8d80b78be59b2eaba2154562e21caa6e79c2f1281d6f486584f7144ee2f47", size = 23998879, upload-time = "2026-07-21T12:43:27.324Z" }
sdist = { url = "https://files.pythonhosted.org/packages/33/8d/995cc142f6083346b35e7d3eadc5a6717ee89ce64ea53577010ac493bc3c/jupyterlab-4.6.4.tar.gz", hash = "sha256:404f49b081819378524886c9db66dba57a5565981eff885830df1baba3a17df5", size = 28335647, upload-time = "2026-09-21T15:25:18.779Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/9f/c9/940f95f17ee4e413ad252bf8d4f2ee9a341f18cfeda87775fef3d7847321/jupyterlab-4.5.10-py3-none-any.whl", hash = "sha256:5967ca61e692e67a2f30b5a2b901c941dc6ce56c0b0e357bc6d34fed5ec095f6", size = 12452502, upload-time = "2026-07-21T12:43:23.542Z" },
{ url = "https://files.pythonhosted.org/packages/3b/e4/072bc0d3c6d45d414062a06af7ee81c02a2c12d769c358ec9919f9995ffd/jupyterlab-4.6.4-py3-none-any.whl", hash = "sha256:15b13f991d3985129c797eb84d9949eeb8b6615e14b444868e642411f2c418b2", size = 17171596, upload-time = "2026-09-21T15:25:14.123Z" },
]
[[package]]
@@ -1623,7 +1637,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.2.0"
version = "4.3.0"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -2146,18 +2160,19 @@ wheels = [
[[package]]
name = "notebook"
version = "7.5.7"
version = "7.6.3"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jupyter-builder" },
{ name = "jupyter-server" },
{ name = "jupyterlab" },
{ name = "jupyterlab-server" },
{ name = "notebook-shim" },
{ name = "tornado" },
]
sdist = { url = "https://files.pythonhosted.org/packages/3e/c4/f71f8716f2903e9e817a47f534b9fd84831e155e2acb32c26691c8e06243/notebook-7.5.7.tar.gz", hash = "sha256:d6d59288a25303b25e1dcb71e9b017ec3a785f7d92f38b9bc288ca1970d5b0a8", size = 14171612, upload-time = "2026-06-04T18:33:45.224Z" }
sdist = { url = "https://files.pythonhosted.org/packages/42/ac/aedb759dc683dcd129f6c75271bab5eb23020b0f7969163c726343718e1e/notebook-7.6.3.tar.gz", hash = "sha256:e2c08e469c0ae20bb0b3214f0ab77e79653317a2f8e5b34c10361c66874a5b50", size = 5501618, upload-time = "2026-09-21T18:07:24.261Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/e1/4d/b3347f7073a377273531efe4ffc738fc910e93718fd2838c7ebf6736c6af/notebook-7.5.7-py3-none-any.whl", hash = "sha256:1f95f79d117e47d20b5555b5c85a397d2cfecf136978aaab767cf0314b09165b", size = 14583767, upload-time = "2026-06-04T18:33:40.987Z" },
{ url = "https://files.pythonhosted.org/packages/a1/f7/907b98438cf00bdc4e296570058f29a347214bb42338a48e3b7eae58ae2a/notebook-7.6.3-py3-none-any.whl", hash = "sha256:ad7e0eb765fba836cd4a2ab0c7a3a26cde1d91665fbf6f533b6ae7b2de6d88d2", size = 5548452, upload-time = "2026-09-21T18:07:21.563Z" },
]
[[package]]
+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
@@ -370,7 +370,7 @@ test = [
[[package]]
name = "langgraph-checkpoint"
version = "4.2.0"
version = "4.3.0"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
+11 -6
View File
@@ -22,18 +22,23 @@ RUN pip install --no-cache-dir \
"langchain>=1.3.0" \
"deepagents>=0.6.2"
# Swap the published langgraph *core* for this monorepo's local copy, so the
# server executes the langgraph under test rather than the latest release.
# Swap the published langgraph *core* and *checkpoint* for this monorepo's
# local copies, so the server executes the langgraph under test rather than the
# latest release.
# We keep the rest of the base image (langgraph-api, runtime, Go core server)
# on `latest` — that still surfaces upstream regressions — but the core the
# server runs is now the PR's, which is what lets this suite catch core
# regressions (e.g. ensure_config / runtime changes) before they're published.
# `--no-deps` keeps the base image's already-compatible
# checkpoint/prebuilt/sdk; we only replace core. The local source comes from
# the `langgraph_src` additional build context (see docker-compose.yml).
# Checkpoint goes with core because core can need a checkpoint API from the
# same release. `--no-deps` keeps the base image's already-compatible
# prebuilt/sdk. The local sources come from the `checkpoint_src` and
# `langgraph_src` additional build contexts (see docker-compose.yml).
COPY --from=checkpoint_src pyproject.toml README.md LICENSE /opt/checkpoint-src/
COPY --from=checkpoint_src langgraph /opt/checkpoint-src/langgraph/
COPY --from=langgraph_src pyproject.toml README.md LICENSE /opt/langgraph-src/
COPY --from=langgraph_src langgraph /opt/langgraph-src/langgraph/
RUN pip install --no-cache-dir --force-reinstall --no-deps /opt/langgraph-src
RUN pip install --no-cache-dir --force-reinstall --no-deps \
/opt/checkpoint-src /opt/langgraph-src
# Project graphs + registration config.
COPY graph/ /app/graph/
+6 -5
View File
@@ -38,16 +38,17 @@ services:
# see ./Dockerfile. The base image bundles langgraph-api +
# langgraph_runtime_postgres + langgraph_license + the Go core-server;
# on top we add graph deps (deepagents), the graph files, and this
# monorepo's local langgraph *core* (so the server runs the code under
# test, not the released langgraph).
# monorepo's local langgraph *core* and *checkpoint* (so the server runs
# the code under test, not the released langgraph).
build:
context: .
dockerfile: Dockerfile
additional_contexts:
# The monorepo's local langgraph core (libs/langgraph), installed over
# the base image so the server runs the langgraph under test. See the
# Dockerfile for why.
# The monorepo's local langgraph core (libs/langgraph) and checkpoint
# (libs/checkpoint), installed over the base image so the server runs
# the langgraph under test. See the Dockerfile for why.
langgraph_src: ../../langgraph
checkpoint_src: ../../checkpoint
image: langgraph-v3-integration-api:local
depends_on:
postgres:
+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))
+1 -1
View File
@@ -385,7 +385,7 @@ test = [
[[package]]
name = "langgraph-checkpoint"
version = "4.2.0"
version = "4.3.0"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },