Compare commits

..
Author SHA1 Message Date
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
53 changed files with 1951 additions and 666 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,
]
+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" },
@@ -687,5 +687,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
@@ -297,7 +299,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 +307,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"
@@ -473,10 +476,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 +491,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 +507,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 +520,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 +565,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,7 +36,6 @@ 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
@@ -1387,7 +1386,6 @@ 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,10 +13,8 @@ import pytest
from langchain_core.embeddings import Embeddings
from langgraph.store.base import (
GetOp,
InvalidNamespaceError,
Item,
ListNamespacesOp,
MatchCondition,
PutOp,
SearchOp,
)
@@ -873,65 +871,3 @@ 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}
@@ -11,7 +11,6 @@ import pytest
from langchain_core.embeddings import Embeddings
from langgraph.store.base import (
GetOp,
InvalidNamespaceError,
Item,
ListNamespacesOp,
MatchCondition,
@@ -1165,51 +1164,3 @@ 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),
@@ -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),
@@ -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,7 +28,6 @@ from langgraph.store.base import (
ensure_embeddings,
get_text_at_path,
tokenize_path,
validate_op_namespace,
)
_AIO_ERROR_MSG = (
@@ -258,7 +257,6 @@ 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,10 +9,8 @@ from typing import cast
import pytest
from langgraph.store.base import (
GetOp,
InvalidNamespaceError,
Item,
ListNamespacesOp,
MatchCondition,
PutOp,
SearchOp,
)
@@ -747,65 +745,3 @@ async def test_async_namespace_segment_boundary(store: AsyncSqliteStore) -> None
assert set(await store.alist_namespaces(suffix=["alice"], limit=100)) == {
("uid", "users", "alice"),
}
@pytest.mark.parametrize("namespace", [("foo.bar",), ("foo", ""), ("foo", 1)])
async def test_abatch_rejects_invalid_namespace_labels(
store: AsyncSqliteStore, namespace: tuple
) -> None:
await store.aput(("foo", "bar"), "key", {"original": True})
for op in (
GetOp(namespace, "key"),
GetOp(namespace, "key", refresh_ttl=True),
PutOp(namespace, "key", {"changed": True}),
PutOp(namespace, "key", None),
SearchOp(namespace),
ListNamespacesOp((MatchCondition("prefix", namespace),)),
ListNamespacesOp((MatchCondition("suffix", namespace),)),
):
with pytest.raises(InvalidNamespaceError):
await store.abatch([op])
item = await store.aget(("foo", "bar"), "key")
assert item is not None and item.value == {"original": True}
async def test_invalid_namespace_only_fails_its_own_call(
store: AsyncSqliteStore,
) -> None:
"""Concurrent calls share one `abatch`, which fails every op if it raises.
Labels are checked before an op is queued, so one caller's bad label cannot
fail another caller's request.
"""
await store.aput(("foo", "bar"), "key", {"original": True})
valid, invalid = await asyncio.gather(
store.aget(("foo", "bar"), "key"),
store.aget(("foo.bar",), "key"),
return_exceptions=True,
)
assert isinstance(valid, Item) and valid.value == {"original": True}
assert isinstance(invalid, InvalidNamespaceError)
async def test_sync_methods_reject_invalid_namespace_labels(
store: AsyncSqliteStore,
) -> None:
"""The sync wrappers run off the event loop thread and must validate too."""
await store.aput(("foo", "bar"), "key", {"original": True})
for call in (
lambda: store.get(("foo.bar",), "key"),
lambda: store.search(("foo.bar",)),
lambda: store.delete(("foo.bar",), "key"),
lambda: store.list_namespaces(prefix=("foo.bar",)),
lambda: store.batch([GetOp(("foo.bar",), "key")]),
):
with pytest.raises(InvalidNamespaceError):
await asyncio.to_thread(call)
item = await store.aget(("foo", "bar"), "key")
assert item is not None and item.value == {"original": True}
@@ -14,7 +14,6 @@ import pytest
from langchain_core.embeddings import Embeddings
from langgraph.store.base import (
GetOp,
InvalidNamespaceError,
Item,
ListNamespacesOp,
MatchCondition,
@@ -1436,51 +1435,3 @@ 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]:
@@ -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)
@@ -771,12 +771,7 @@ 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]
@@ -806,10 +801,6 @@ 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:
@@ -849,7 +840,6 @@ 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(
@@ -897,11 +887,6 @@ 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,
@@ -955,12 +940,7 @@ 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(
@@ -989,10 +969,6 @@ 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:
@@ -1008,8 +984,6 @@ 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))
@@ -1039,12 +1013,7 @@ 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(
[
@@ -1083,10 +1052,6 @@ 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:
@@ -1126,7 +1091,6 @@ 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(
[
@@ -1176,11 +1140,6 @@ 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,
@@ -1242,12 +1201,7 @@ 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(
@@ -1276,10 +1230,6 @@ 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:
@@ -1295,8 +1245,6 @@ 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))
@@ -1315,14 +1263,6 @@ 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(
@@ -1337,27 +1277,10 @@ def _validate_namespace_labels(namespace: tuple[str, ...]) -> None:
raise InvalidNamespaceError(
f"Namespace labels cannot be empty strings. Got {label} in {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)
if namespace[0] == "langgraph":
raise InvalidNamespaceError(
f'Root label for namespace cannot be "langgraph". Got: {namespace}'
)
def _ensure_refresh(
@@ -25,7 +25,6 @@ from langgraph.store.base import (
_ensure_refresh,
_ensure_ttl,
_validate_namespace,
_validate_namespace_labels,
)
F = TypeVar("F", bound=Callable)
@@ -87,7 +86,6 @@ 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(
@@ -113,7 +111,6 @@ 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(
@@ -158,7 +155,6 @@ 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)))
@@ -173,8 +169,6 @@ 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"
-125
View File
@@ -13,14 +13,10 @@ 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
@@ -532,127 +528,6 @@ 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
+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)
+27 -2
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
import uuid
from collections.abc import Callable, Iterable, Mapping
from datetime import datetime, timezone
from inspect import signature
from typing import Any, Literal, cast
from langchain_core.runnables import RunnableConfig
@@ -25,6 +26,7 @@ from langgraph._internal._constants import (
CONFIG_KEY_CHECKPOINT_ID,
NS_END,
NS_SEP,
NULL_TASK_ID,
PUSH,
SNAPSHOT_BUMPS,
)
@@ -49,15 +51,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(
+69 -25
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,
@@ -106,7 +105,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,
@@ -221,10 +222,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:
@@ -737,9 +746,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 +885,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 +1008,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,7 +1111,9 @@ class PregelLoop:
if self._exit_delta_writes is not None:
for c, v in input_writes:
if isinstance(self.specs.get(c), DeltaChannel):
self._exit_delta_writes.append((self.step, NULL_TASK_ID, c, v))
self._exit_delta_writes.append(
(self.step, NULL_TASK_ID, "", c, v)
)
# Persist delta-channel input writes so sub-freq inputs are
# recoverable via ancestor walk (mirrors the Command input path).
if self.durability != "exit":
@@ -1309,9 +1348,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
@@ -1346,11 +1383,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 +1407,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 +1630,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
@@ -1840,8 +1885,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
+59 -31
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 (
@@ -139,6 +140,7 @@ from langgraph.pregel._checkpoint import (
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
@@ -1983,13 +1985,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 +1999,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)
@@ -2050,7 +2052,10 @@ class Pregel(
channel_writes = [w for w in task.writes if w[0] != PUSH]
if channel_writes:
checkpointer.put_writes(
checkpoint_config, channel_writes, task_id
checkpoint_config,
channel_writes,
task_id,
**_task_path_kwarg(checkpointer.put_writes, task),
)
apply_writes(
checkpoint,
@@ -2471,13 +2476,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 +2490,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)
@@ -2538,7 +2543,10 @@ class Pregel(
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
checkpoint_config,
channel_writes,
task_id,
**_task_path_kwarg(checkpointer.aput_writes, task),
)
apply_writes(
checkpoint,
@@ -3706,16 +3714,15 @@ class Pregel(
config: RunnableConfig | None = None,
*,
version: Literal["v1", "v2", "v3"] = "v2",
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
control: RunControl | None = None,
transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
**kwargs: Any,
) -> Any:
"""Stream events from this graph.
For `version="v1"` / `"v2"`, yields `StreamEvent` dicts (see
`Runnable.stream_events`). For `version="v3"`, returns a
For `version="v1"` / `"v2"`, delegates to
`Runnable.(a)stream_events`; synchronous v1/v2 event streaming is
not implemented in langchain-core, so use `astream_events` for
those versions. For `version="v3"`, returns a
`GraphRunStream` whose typed projections the caller drives by
iterating — no background thread.
@@ -3745,20 +3752,27 @@ class Pregel(
config: Optional runnable config.
version: Streaming-event schema version. `"v3"` selects the
content-block-centric streaming protocol.
interrupt_before: Nodes to interrupt before, if any. Only
used for `version="v3"`.
interrupt_after: Nodes to interrupt after, if any. Only
used for `version="v3"`.
interrupt_before: Nodes to interrupt before, if any.
Honored on every version that can run; type-checked
only on the `version="v3"` overloads.
interrupt_after: Nodes to interrupt after, if any. Honored
on every version that can run; type-checked only on
the `version="v3"` overloads.
control: Optional run control used to request cooperative
drain. Only used for `version="v3"`.
drain. Honored on every version that can run;
type-checked only on the `version="v3"` overloads.
transformers: Extra transformer classes or configured
factories appended after compile-time
`stream_transformers`. Factories are called as
`factory(scope)` so they can propagate to subgraph
scopes. Only used for `version="v3"`.
**kwargs: For `version="v1"`/`"v2"`, forwarded to
`Runnable.stream_events`. For `version="v3"`, forwarded
to the underlying `stream(...)` call (e.g. `context`,
**kwargs: For `version="v1"`/`"v2"` on `astream_events`,
forwarded to `Runnable.astream_events`, which passes
them through to `astream` — so execution kwargs such as
`context`, `durability`, `interrupt_before`,
`interrupt_after` and `control` are honored on every
version that can run. For `version="v3"`, forwarded to the
underlying `stream(...)` call (e.g. `context`,
`durability`, `output_keys`, `print_mode`, `debug`).
`stream_mode` and `subgraphs` are not accepted under
`version="v3"` and raise `TypeError` if supplied; v3
@@ -3766,16 +3780,15 @@ class Pregel(
Returns:
For `version="v3"`, a `GraphRunStream` the caller iterates
to drive the run. Otherwise an `Iterator[StreamEvent]`.
to drive the run. For `version="v1"`/`"v2"`,
`astream_events` yields `StreamEvent` dicts; the synchronous
v1/v2 path is not implemented in langchain-core.
"""
if version == "v3":
_reject_v3_invariant_kwargs(kwargs)
return self._pregel_stream_v3(
input,
config,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
control=control,
transformers=transformers,
**kwargs,
)
@@ -3811,9 +3824,6 @@ class Pregel(
config: RunnableConfig | None = None,
*,
version: Literal["v1", "v2", "v3"] = "v2",
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
control: RunControl | None = None,
transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
**kwargs: Any,
) -> AsyncIterator[StreamEvent] | Awaitable[Any]:
@@ -3836,9 +3846,6 @@ class Pregel(
return self._apregel_stream_v3(
input,
config,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
control=control,
transformers=transformers,
**kwargs,
)
@@ -4237,6 +4244,27 @@ 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)}
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)
+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}"
@@ -20,6 +20,7 @@ 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
@@ -27,16 +28,17 @@ from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import _messages_delta_reducer
from langgraph.types import StateUpdate
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 +65,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 +306,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 +333,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]
# ---------------------------------------------------------------------------
@@ -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]]
+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" },