Compare commits

...
Author SHA1 Message Date
lc-arjun 03ef1a26e0 feat: support list of hooks for pre_model_hook and post_model_hook
Allow `pre_model_hook` and `post_model_hook` in `create_react_agent`
to accept a list of `RunnableLike` callables in addition to a single
hook. When a list is provided the hooks are composed in order: each
hook receives the graph state merged with all prior hooks' updates,
and the final merged update is returned to the graph.

This mirrors the middleware-stack pattern familiar from HTTP frameworks
(Express, FastAPI, Starlette) and allows hook logic to be written as
small, reusable units that can be composed without manually threading
state between them.
2026-05-08 16:02:24 -04:00
ccurmeandGitHub dc0d992b90 fix(checkpoint): specify allowed_objects in Reviver (#7743) 2026-05-08 15:08:29 -04:00
open-swe[bot]GitHubopen-swe[bot] <open-swe@users.noreply.github.com>Parker J. Rule
ed168deb97 feat(sdk-py): support metadata filter for crons search/count (#7737)
## Description
Mirrors a server-side change by accepting an optional \`metadata\`
filter on \`crons.search\` and \`crons.count\` in both the async and
sync Python SDK clients. Matches the existing pattern used for
assistants/threads search.

## Release Note
Python SDK: \`crons.search\` and \`crons.count\` now accept an optional
\`metadata\` filter that is forwarded to the server.

## Test Plan
- [ ] New unit tests in \`tests/test_crons_client.py\` verify metadata
is forwarded for both async and sync clients and omitted when not
provided.

---------

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
Co-authored-by: Parker J. Rule <pjrule@users.noreply.github.com>
2026-05-08 14:21:45 -04:00
Quanzheng LongandGitHub fb6e5c2bce chore: remove keepset helper (#7745)
Comment follow up
2026-05-08 18:00:31 +00:00
6dade64aa7 chore: remove unnecessary missing-typed-dict-key suppression (#7744)
Address a comment as follow up

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-05-08 10:39:54 -07:00
398d6cc59d chore(langgraph): add guide/conformance for delta channel checkpointer (#7736)
## Summary

Add a user-facing design doc and `get_delta_channel_keepset` helper for
third-party `BaseCheckpointSaver` authors who need to support graphs
using `DeltaChannel`.

**Deliverables:**

1. ~~**`docs/delta-channel-checkpointer-guide.md`** — comprehensive
guide covering~~:
moved to docs repo

2. **`BaseCheckpointSaver.get_delta_channel_keepset` /
`aget_delta_channel_keepset`** — returns the minimum set of ancestor
`checkpoint_id`s that must survive deletion for a given head's
`DeltaChannel` reconstruction to remain intact. Enables safe `prune`
implementations without silently corrupting delta history.

3. **Docstring warnings** on `prune`, `aprune`, `delete_for_runs`,
`adelete_for_runs`, `copy_thread`, `acopy_thread` explaining the
DeltaChannel pitfall (silent data loss if ancestor writes/snapshots are
deleted).

4. **Three new conformance capabilities** in
`libs/checkpoint-conformance`:
- `delta_channel_history` — validates the `aget_delta_channel_history`
walk contract
   - `delta_channel_keepset` — validates the keep-set contract
- `delta_channel_reconstruction` — end-to-end round-trip (aput +
aput_writes + history + reconstruct)

## Test plan

- [x] `make format lint` passes in `libs/checkpoint`,
`libs/checkpoint-conformance`
- [x] All three new conformance capabilities pass against
`InMemorySaver`
- [x] Run conformance against SQLite saver 
- [x] Run conformance against Postgres saver

---------

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-05-07 17:34:46 -07:00
d736564eb1 test(langgraph): de-flake heartbeat progress test (#7735)
## Summary

De-flake
`test_arun_with_retry_timeout_observer_emits_progress_on_heartbeat` —
the test was hitting a CI-runner-load-sensitive race where the
idle-timeout watchdog could fire before the task body's first await ran.

## Root cause

`_TimedAttemptScope.__init__` sets `_last_progress = time.monotonic()`
immediately, but the watchdog itself doesn't start polling until *after*
`wrap_config` and task scheduling. Under heavy CI load that gap can grow
large enough that:

```
T₀  scope.__init__()  →  _last_progress = T₀
… some scheduling slack …
Tₙ  watchdog runs, computes  remaining = T₀ + 0.2 − Tₙ ≤ 0  →  TimeoutError fires
```

The error reports `elapsed: 0.000s` because `elapsed` is measured from
the post-scheduling `start` (≈Tₙ), not from `_last_progress` (T₀). The
previous test set `idle_timeout=0.2s`, which left almost no headroom for
that scheduling slack.

## Fix (test-side only — no production change)

- **Heartbeat at task-body entry**: `runtime.heartbeat()` is now called
before the first `await asyncio.sleep(...)`, which resets
`_last_progress` to "now" the moment the task body actually starts
running. This eliminates the scope-init-to-first-await gap as a flake
source.
- **Idle timeout 0.2s → 1.0s**: gives ~5× headroom over the ~400ms task
duration, so scheduling pressure stays comfortably within budget.

## Why test-side instead of fixing the production race

The proper production fix would be to set `_last_progress` at
watchdog-entry time rather than at scope-init time. That's a behaviour
change in the retry/timeout machinery and out of scope for a flaky-test
fix. The two test-side defenses make this particular test stable without
touching production semantics; the underlying race in
`_TimedAttemptScope` is worth a separate follow-up.

## Test plan

- 10/10 repeated local runs pass:
  ```
uv run pytest
tests/test_retry.py::test_arun_with_retry_timeout_observer_emits_progress_on_heartbeat
--count=10
  ```
- All assertions still meaningful: still verifies start/finish events,
at least one progress event, rate-limited progress count (≤ total
events), and per-event metadata (task_name, attempt, idle_timeout_secs,
progress_at).

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-05-07 10:31:18 -07:00
69f2d3a430 chore(langgraph): re-implement exit mode for delta channel (#7730)
## Summary

Replaces `durability="exit"`'s blanket force-snapshot of every
`DeltaChannel` with proper write persistence that honors per-channel
`snapshot_frequency`, plus closes two latent bugs the force-snapshot was
masking.

Before: every exit-mode run wrote a full `_DeltaSnapshot` blob for every
delta channel, even when the channel had zero updates this run and was
nowhere near its `snapshot_frequency`. After: the same count-based
decision used by `durability="sync"`/`"async"` applies — channels at or
above `snapshot_frequency` snapshot; channels below it persist their
accumulated writes via a lazy "stub" anchor; untouched channels write
nothing.

## What changed

**Core redesign** (`pregel/_loop.py`, `pregel/_checkpoint.py`)

- Drop `force_delta_snapshot` from `create_checkpoint` and
`_should_snapshot_delta`.
- Add `decide_delta_snapshots(channels, counts)` pure helper used by
both `create_checkpoint` and the new exit-mode peek-ahead path.
- Add `_exit_delta_writes` accumulator: every delta-channel write
produced during a `durability="exit"` run (input writes from `_first` +
per-superstep writes captured before `pending_writes.clear()` in
`after_tick`) is collected into this list.
- Add `_put_exit_delta_writes` (sync + async): runs from
`_suppress_interrupt` BEFORE `_put_checkpoint(exiting=True)`. Filters
out channels that will snapshot, then persists remaining writes to
`checkpoint_writes` under an anchor parent. The anchor is the existing
saved parent on resumed runs, or a lazily-created empty stub on first
runs.
- Visibility ordering: stub put goes onto `_put_checkpoint_fut` (becomes
the next put's `prev`); exit-write futures go onto `_delta_write_futs`.
The existing `_checkpointer_put_after_previous` already drains both
before calling `saver.put`, so `final_checkpoint` is structurally
guaranteed to land last — readers never see a partial view.

**Latent bugs fixed (previously masked by force-snapshot)**

- **Sync drain race**: `SyncPregelLoop` now initializes
`_delta_write_futs = []` in `__enter__` and drains it in sync
`_checkpointer_put_after_previous` before `put`, mirroring the async
version. Without this, a multi-worker `BackgroundExecutor` could publish
a checkpoint before the writes that produced it.
- **Count double-bump in exit mode**: in `_put_checkpoint`,
`delta_updates_since_snapshot` was being incremented twice for the last
superstep — once by the intermediate `after_tick` call, once by
`_suppress_interrupt`. Force-snapshot used to reset all counts to 0 so
this never persisted; without it, snapshots would fire one superstep
early after every exit-mode run. Fixed by gating the count-bump behind
`not exiting`.

**Pre-existing input-durability gap**

- In the plain (non-Command) input path of `_first`, delta-channel input
writes are now persisted via `put_writes` (mirroring the Command path),
so sub-frequency inputs survive a `get_state` on resumed runs in
`sync`/`async` durability. Note: first-run `sync`/`async` still has the
same gap (writes orphan on the synthetic-empty parent id). That's
flagged as a follow-up — out of scope for this PR.

## Test plan

- Existing `tests/test_pregel.py` and `tests/test_pregel_async.py` pass
unchanged.
- Existing `tests/test_channels.py` (29 tests) and
`tests/test_delta_channel_migration.py` pass unchanged.
- New `tests/test_exit_delta_persistence.py` (11 tests) covers:
- **Write-path**: zero-write exit (no stub), all-snapshot first run (no
stub), sub-freq first run (single shared stub), sub-freq resumed run
(anchor on saved parent), sync-vs-exit count parity, mixed
snapshot/non-snapshot channels, snapshot fires at frequency.
- **Read-path**: K-run replay chain reads correctly across
stub→saved-parent transition; metadata `delta_updates_since_snapshot`
round-trips correctly; mixed sync/exit durability alternation produces
correct final state; snapshot+tail-deltas combination reads correctly.
- `make format && make lint && make test` in `libs/langgraph/`.

---------

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: Sydney Runkle <sydneymarierunkle@gmail.com>
2026-05-07 09:47:01 -07:00
20 changed files with 2140 additions and 622 deletions
+1
View File
@@ -63,6 +63,7 @@ The suite tests **base** capabilities (required) and **extended** capabilities (
| `delete_for_runs` | no | `adelete_for_runs` |
| `copy_thread` | no | `acopy_thread` |
| `prune` | no | `aprune` |
| `delta_channel_history` | no | `aget_delta_channel_history` |
Extended capabilities are detected by checking whether the method is overridden from `BaseCheckpointSaver`. If not overridden, those tests are skipped.
@@ -23,6 +23,7 @@ class Capability(str, Enum):
DELETE_FOR_RUNS = "delete_for_runs"
COPY_THREAD = "copy_thread"
PRUNE = "prune"
DELTA_CHANNEL_HISTORY = "delta_channel_history"
# Capabilities that every checkpointer must support.
@@ -42,6 +43,7 @@ EXTENDED_CAPABILITIES = frozenset(
Capability.DELETE_FOR_RUNS,
Capability.COPY_THREAD,
Capability.PRUNE,
Capability.DELTA_CHANNEL_HISTORY,
}
)
@@ -57,6 +59,7 @@ _CAPABILITY_METHOD_MAP: dict[Capability, str] = {
Capability.DELETE_FOR_RUNS: "adelete_for_runs",
Capability.COPY_THREAD: "acopy_thread",
Capability.PRUNE: "aprune",
Capability.DELTA_CHANNEL_HISTORY: "aget_delta_channel_history",
}
@@ -9,6 +9,9 @@ from langgraph.checkpoint.conformance.spec.test_delete_for_runs import (
from langgraph.checkpoint.conformance.spec.test_delete_thread import (
run_delete_thread_tests,
)
from langgraph.checkpoint.conformance.spec.test_delta_channel_history import (
run_delta_channel_history_tests,
)
from langgraph.checkpoint.conformance.spec.test_get_tuple import run_get_tuple_tests
from langgraph.checkpoint.conformance.spec.test_list import run_list_tests
from langgraph.checkpoint.conformance.spec.test_prune import run_prune_tests
@@ -24,4 +27,5 @@ __all__ = [
"run_delete_for_runs_tests",
"run_copy_thread_tests",
"run_prune_tests",
"run_delta_channel_history_tests",
]
@@ -0,0 +1,99 @@
"""Shared fixtures for delta-channel conformance tests.
Builds a parent chain with `_DeltaSnapshot` blobs at known positions via
direct `aput` / `aput_writes` calls. No langgraph or Pregel dependency.
"""
from __future__ import annotations
from collections.abc import Sequence
from typing import Any
from uuid import uuid4
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import BaseCheckpointSaver, Checkpoint
from langgraph.checkpoint.base.id import uuid6
from langgraph.checkpoint.conformance.test_utils import generate_metadata
async def build_delta_chain(
saver: BaseCheckpointSaver,
*,
thread_id: str | None = None,
checkpoint_ns: str = "",
channel: str = "messages",
snapshots_at_steps: Sequence[int] = (0,),
total_steps: int = 6,
write_value_fn: Any | None = None,
) -> list[RunnableConfig]:
"""Build a parent chain with `_DeltaSnapshot` at known positions.
Args:
saver: Checkpointer instance.
thread_id: Defaults to a random UUID.
checkpoint_ns: Namespace (default root).
channel: Channel name used for snapshots and writes.
snapshots_at_steps: Steps at which a `_DeltaSnapshot` blob is stored
in `channel_values[channel]`. Step 0 is the oldest checkpoint.
total_steps: Number of checkpoints in the chain.
write_value_fn: Callable(step) -> write value. Defaults to step index.
Returns:
List of stored configs (oldest first), one per step.
"""
if write_value_fn is None:
def write_value_fn(step: int) -> Any:
return step
from langgraph.checkpoint.serde.types import _DeltaSnapshot
thread_id = thread_id or str(uuid4())
snapshot_set = set(snapshots_at_steps)
stored: list[RunnableConfig] = []
parent_cfg: RunnableConfig | None = None
for step in range(total_steps):
config: RunnableConfig = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": checkpoint_ns,
}
}
if parent_cfg:
config["configurable"]["checkpoint_id"] = parent_cfg["configurable"][
"checkpoint_id"
]
channel_values: dict[str, Any] = {}
channel_versions: dict[str, int] = {}
if step in snapshot_set:
channel_values[channel] = _DeltaSnapshot(
write_value_fn(step),
)
channel_versions[channel] = step + 1
cp = Checkpoint(
v=1,
id=str(uuid6(clock_seq=-1)),
ts="",
channel_values=channel_values,
channel_versions=channel_versions,
versions_seen={},
updated_channels=None,
)
new_versions = dict(channel_versions)
parent_cfg = await saver.aput(
config, cp, generate_metadata(step=step), new_versions
)
stored.append(parent_cfg)
# Write a pending write for non-snapshot steps so the walk has
# something to collect.
if step not in snapshot_set:
await saver.aput_writes(
parent_cfg, [(channel, write_value_fn(step))], str(uuid4())
)
return stored
@@ -0,0 +1,247 @@
"""DELTA_CHANNEL_HISTORY capability tests — aget_delta_channel_history contract."""
from __future__ import annotations
import traceback
from collections.abc import Callable
from uuid import uuid4
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.checkpoint.conformance.spec._delta_fixtures import build_delta_chain
async def test_history_returns_writes_oldest_first(
saver: BaseCheckpointSaver,
) -> None:
"""Writes are returned oldest-to-newest."""
tid = str(uuid4())
# 5 steps: snapshot at 0, writes at 1,2,3,4.
# Head is step 4. Walk starts at step 3 (parent of head).
# Collects writes from steps 1,2,3 (between snapshot at 0 and head's parent).
configs = await build_delta_chain(
saver, thread_id=tid, channel="ch", snapshots_at_steps=[0], total_steps=5
)
head = configs[-1]
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
writes = result["ch"]["writes"]
values = [w[2] for w in writes]
assert values == [1, 2, 3], f"Expected [1,2,3], got {values}"
async def test_history_seed_is_nearest_snapshot(
saver: BaseCheckpointSaver,
) -> None:
"""Seed is the value from the nearest ancestor with channel_values populated."""
tid = str(uuid4())
# 6 steps: snapshots at 0 and 3, writes at 1,2,4,5.
# Head is step 5. Walk from step 4 backward stops at step 3 (snapshot).
# Collects writes from step 4 only (between step 3 and head's parent step 4).
configs = await build_delta_chain(
saver,
thread_id=tid,
channel="ch",
snapshots_at_steps=[0, 3],
total_steps=6,
)
head = configs[-1]
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
assert "seed" in result["ch"], "Expected seed from snapshot at step 3"
seed = result["ch"]["seed"]
from langgraph.checkpoint.serde.types import _DeltaSnapshot
actual_value = seed.value if isinstance(seed, _DeltaSnapshot) else seed
assert actual_value == 3, f"Expected seed value 3 (step 3), got {actual_value}"
writes = result["ch"]["writes"]
values = [w[2] for w in writes]
assert values == [4], f"Expected [4], got {values}"
async def test_history_excludes_target_pending_writes(
saver: BaseCheckpointSaver,
) -> None:
"""Target's own pending_writes are NOT included in the history."""
tid = str(uuid4())
configs = await build_delta_chain(
saver, thread_id=tid, channel="ch", snapshots_at_steps=[0], total_steps=3
)
head = configs[-1]
# Add writes directly to the head checkpoint
await saver.aput_writes(head, [("ch", "extra")], str(uuid4()))
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
writes = result["ch"]["writes"]
values = [w[2] for w in writes]
assert "extra" not in values, f"Target's writes should be excluded, got {values}"
async def test_history_multi_channel(
saver: BaseCheckpointSaver,
) -> None:
"""Multiple channels have independent walk termination."""
tid = str(uuid4())
configs: list = []
parent_cfg = None
from langgraph.checkpoint.base import Checkpoint
from langgraph.checkpoint.base.id import uuid6
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from langgraph.checkpoint.conformance.test_utils import generate_metadata
for step in range(5):
config = {"configurable": {"thread_id": tid, "checkpoint_ns": ""}}
if parent_cfg:
config["configurable"]["checkpoint_id"] = parent_cfg["configurable"][
"checkpoint_id"
]
cv: dict = {}
cvs: dict = {}
if step == 1:
cv["a"] = _DeltaSnapshot("snap_a")
cvs["a"] = step + 1
if step == 3:
cv["b"] = _DeltaSnapshot("snap_b")
cvs["b"] = step + 1
cp = Checkpoint(
v=1,
id=str(uuid6(clock_seq=-1)),
ts="",
channel_values=cv,
channel_versions=cvs,
versions_seen={},
updated_channels=None,
)
parent_cfg = await saver.aput(config, cp, generate_metadata(step=step), cvs)
configs.append(parent_cfg)
await saver.aput_writes(parent_cfg, [("a", step), ("b", step)], str(uuid4()))
head = configs[-1]
result = await saver.aget_delta_channel_history(config=head, channels=["a", "b"])
a_writes = [w[2] for w in result["a"]["writes"]]
b_writes = [w[2] for w in result["b"]["writes"]]
assert a_writes == [1, 2, 3], f"Expected a writes [1,2,3], got {a_writes}"
assert b_writes == [3], f"Expected b writes [3], got {b_writes}"
async def test_history_empty_channels_returns_empty(
saver: BaseCheckpointSaver,
) -> None:
"""Empty channels list returns empty mapping."""
tid = str(uuid4())
configs = await build_delta_chain(
saver, thread_id=tid, channel="ch", snapshots_at_steps=[0], total_steps=3
)
result = await saver.aget_delta_channel_history(config=configs[-1], channels=[])
assert result == {}
async def test_history_walk_to_root_no_seed(
saver: BaseCheckpointSaver,
) -> None:
"""Walk reaches root without finding seed — no 'seed' key in result."""
tid = str(uuid4())
configs = await build_delta_chain(
saver,
thread_id=tid,
channel="ch",
snapshots_at_steps=[],
total_steps=4,
)
head = configs[-1]
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
assert "seed" not in result["ch"], f"Expected no seed, got {result['ch']}"
async def test_history_migration_plain_value_as_seed(
saver: BaseCheckpointSaver,
) -> None:
"""Pre-delta plain value in channel_values acts as seed (migration case).
When a thread was originally using a regular channel (BinaryOperatorAggregate)
and later switches to DeltaChannel, the old checkpoint has a plain value in
channel_values[ch] (not a _DeltaSnapshot). The walk should treat it as the
seed and terminate there.
"""
from langgraph.checkpoint.base import Checkpoint
from langgraph.checkpoint.base.id import uuid6
from langgraph.checkpoint.conformance.test_utils import generate_metadata
tid = str(uuid4())
configs: list = []
parent_cfg = None
for step in range(4):
config = {"configurable": {"thread_id": tid, "checkpoint_ns": ""}}
if parent_cfg:
config["configurable"]["checkpoint_id"] = parent_cfg["configurable"][
"checkpoint_id"
]
cv: dict = {}
cvs: dict = {}
# Step 1: plain value (migration case — old checkpoint before delta)
if step == 1:
cv["ch"] = [10, 20, 30]
cvs["ch"] = step + 1
cp = Checkpoint(
v=1,
id=str(uuid6(clock_seq=-1)),
ts="",
channel_values=cv,
channel_versions=cvs,
versions_seen={},
updated_channels=None,
)
parent_cfg = await saver.aput(config, cp, generate_metadata(step=step), cvs)
configs.append(parent_cfg)
if step != 1:
await saver.aput_writes(parent_cfg, [("ch", step)], str(uuid4()))
head = configs[-1]
result = await saver.aget_delta_channel_history(config=head, channels=["ch"])
# Seed should be the plain value from step 1
assert "seed" in result["ch"], "Expected seed from migration plain value at step 1"
seed = result["ch"]["seed"]
assert seed == [10, 20, 30], f"Expected plain value [10,20,30], got {seed}"
# Writes should be from step 2 only (between seed at step 1 and head's parent step 2)
writes = result["ch"]["writes"]
values = [w[2] for w in writes]
assert values == [2], f"Expected [2], got {values}"
ALL_DELTA_CHANNEL_HISTORY_TESTS = [
test_history_returns_writes_oldest_first,
test_history_seed_is_nearest_snapshot,
test_history_excludes_target_pending_writes,
test_history_multi_channel,
test_history_empty_channels_returns_empty,
test_history_walk_to_root_no_seed,
test_history_migration_plain_value_as_seed,
]
async def run_delta_channel_history_tests(
saver: BaseCheckpointSaver,
on_test_result: Callable[[str, str, bool, str | None], None] | None = None,
) -> tuple[int, int, list[str]]:
"""Run all delta_channel_history tests. Returns (passed, failed, failure_names)."""
passed = 0
failed = 0
failures: list[str] = []
for test_fn in ALL_DELTA_CHANNEL_HISTORY_TESTS:
try:
await test_fn(saver)
passed += 1
if on_test_result:
on_test_result("delta_channel_history", test_fn.__name__, True, None)
except Exception:
failed += 1
msg = f"{test_fn.__name__}: {traceback.format_exc()}"
failures.append(msg)
if on_test_result:
on_test_result(
"delta_channel_history",
test_fn.__name__,
False,
traceback.format_exc(),
)
return passed, failed, failures
@@ -19,6 +19,9 @@ from langgraph.checkpoint.conformance.spec.test_delete_for_runs import (
from langgraph.checkpoint.conformance.spec.test_delete_thread import (
run_delete_thread_tests,
)
from langgraph.checkpoint.conformance.spec.test_delta_channel_history import (
run_delta_channel_history_tests,
)
from langgraph.checkpoint.conformance.spec.test_get_tuple import run_get_tuple_tests
from langgraph.checkpoint.conformance.spec.test_list import run_list_tests
from langgraph.checkpoint.conformance.spec.test_prune import run_prune_tests
@@ -35,6 +38,7 @@ _RUNNERS = {
Capability.DELETE_FOR_RUNS: run_delete_for_runs_tests,
Capability.COPY_THREAD: run_copy_thread_tests,
Capability.PRUNE: run_prune_tests,
Capability.DELTA_CHANNEL_HISTORY: run_delta_channel_history_tests,
}
@@ -43,7 +43,11 @@ asyncio_mode = "auto"
# The extended methods (acopy_thread, adelete_for_runs, aprune) are checked
# at runtime via capability detection and may not exist on the installed
# base class. Dict literal inference is also overly strict for RunnableConfig.
# Delta-channel tests import from `langgraph` (not a declared dep of this
# package — at test time it is installed alongside); private `_DeltaSnapshot`
# imports are intentional (beta surface).
unresolved-attribute = "ignore"
unresolved-import = "ignore"
invalid-argument-type = "ignore"
invalid-return-type = "ignore"
@@ -58,6 +62,9 @@ lint.select = [
lint.ignore = ["E501", "B008"]
target-version = "py310"
[tool.uv.sources]
langgraph-checkpoint = {path = "../checkpoint", editable = true}
[[tool.uv.index]]
name = "testpypi"
url = "https://test.pypi.org/simple/"
+605 -504
View File
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,35 @@
"""Run delta-channel conformance capabilities against AsyncSqliteSaver."""
from __future__ import annotations
import pytest
pytest.importorskip(
"langgraph.checkpoint.conformance",
reason="langgraph-checkpoint-conformance not installed",
)
pytest.importorskip("aiosqlite", reason="aiosqlite not installed")
@pytest.mark.asyncio
async def test_delta_channel_conformance():
from langgraph.checkpoint.conformance import validate
from langgraph.checkpoint.conformance.initializer import checkpointer_test
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
@checkpointer_test(name="AsyncSqliteSaver")
async def sqlite_saver():
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
yield saver
report = await validate(
sqlite_saver,
capabilities={
"delta_channel_history",
},
)
for cap, result in report.results.items():
if result.passed is False:
details = "\n".join(result.failures or [])
pytest.fail(f"Capability {cap} failed:\n{details}")
@@ -327,6 +327,14 @@ class BaseCheckpointSaver(Generic[V]):
Args:
run_ids: The run IDs whose checkpoints should be deleted.
!!! warning "DeltaChannel"
Deleting a run that produced ancestor `checkpoint_writes` — or
the only `_DeltaSnapshot` blob — for a still-live thread will
break reconstruction of any `DeltaChannel` whose history
depended on those rows. See the `DeltaChannel` note on `prune`
for safe-recovery strategies.
"""
raise NotImplementedError
@@ -340,6 +348,17 @@ class BaseCheckpointSaver(Generic[V]):
Args:
source_thread_id: The thread ID to copy from.
target_thread_id: The thread ID to copy to.
!!! warning "DeltaChannel"
Implementations must copy the **complete** parent chain (all
ancestor checkpoints and their `checkpoint_writes`) — copying
only the head checkpoint will leave the target thread with
`DeltaChannel` state that cannot be reconstructed (no path back
to a `_DeltaSnapshot` ancestor). Equivalently, the copy must
include enough ancestors that every `DeltaChannel`-backed key
has either a `_DeltaSnapshot` in `channel_values` somewhere in
the chain, or a complete write history back to the chain root.
"""
raise NotImplementedError
@@ -355,6 +374,34 @@ class BaseCheckpointSaver(Generic[V]):
thread_ids: The thread IDs to prune.
strategy: The pruning strategy. `"keep_latest"` retains only the most
recent checkpoint per namespace. `"delete"` removes all checkpoints.
!!! warning "DeltaChannel"
Custom implementations must be `DeltaChannel`-aware. `DeltaChannel`
stores only a sentinel in `channel_values` for non-snapshot steps;
reconstruction walks the parent chain via
`get_delta_channel_history`, accumulating rows from
`checkpoint_writes` until it reaches an ancestor whose
`channel_values` contains a `_DeltaSnapshot` blob (written every
`snapshot_frequency` updates).
A naive `"keep_latest"` that drops intermediate checkpoints and
their writes can sever that chain: the surviving "latest"
checkpoint is rarely a snapshot point itself, so its delta
channels would silently reconstruct as empty (no error raised —
`get_delta_channel_history` simply returns no `seed`). Safe
options when the graph uses `DeltaChannel`:
* Walk back from each kept checkpoint and preserve every
ancestor (plus its `checkpoint_writes`) up to the nearest one
whose `channel_values` already contains a `_DeltaSnapshot` for
every `DeltaChannel`-backed key.
* Force a fresh snapshot on the kept checkpoint before deleting
ancestors — rewrite `channel_values[k] = _DeltaSnapshot(value)`
for each delta channel `k` (resolving `value` via the existing
ancestor walk first), then prune.
* Skip pruning threads whose graph uses `DeltaChannel` until one
of the above is implemented.
"""
raise NotImplementedError
@@ -471,6 +518,13 @@ class BaseCheckpointSaver(Generic[V]):
Args:
run_ids: The run IDs whose checkpoints should be deleted.
!!! warning "DeltaChannel"
See `delete_for_runs` — deleting rows a still-live thread's
`DeltaChannel` reconstruction depends on (writes between the
head and its nearest `_DeltaSnapshot` ancestor) will silently
corrupt that channel's state.
"""
raise NotImplementedError
@@ -484,6 +538,13 @@ class BaseCheckpointSaver(Generic[V]):
Args:
source_thread_id: The thread ID to copy from.
target_thread_id: The thread ID to copy to.
!!! warning "DeltaChannel"
See `copy_thread` — the copy must carry the complete parent
chain (or at least back to a `_DeltaSnapshot` ancestor for every
`DeltaChannel`) so the target thread can reconstruct delta
state.
"""
raise NotImplementedError
@@ -499,6 +560,13 @@ class BaseCheckpointSaver(Generic[V]):
thread_ids: The thread IDs to prune.
strategy: The pruning strategy. `"keep_latest"` retains only the most
recent checkpoint per namespace. `"delete"` removes all checkpoints.
!!! warning "DeltaChannel"
See `prune` for the full `DeltaChannel` caveat. In short:
`"keep_latest"` must not drop ancestor checkpoints / writes that
sit between the kept checkpoint and the nearest `_DeltaSnapshot`
ancestor, or delta channels will silently reconstruct as empty.
"""
raise NotImplementedError
@@ -44,7 +44,7 @@ if TYPE_CHECKING:
AllowedMsgpackModules,
)
LC_REVIVER = Reviver()
LC_REVIVER = Reviver(allowed_objects="core")
EMPTY_BYTES = b""
logger = logging.getLogger(__name__)
@@ -0,0 +1,33 @@
"""Run delta-channel conformance capabilities against InMemorySaver."""
from __future__ import annotations
import pytest
conformance = pytest.importorskip(
"langgraph.checkpoint.conformance",
reason="langgraph-checkpoint-conformance not installed",
)
@pytest.mark.asyncio
async def test_delta_channel_conformance():
from langgraph.checkpoint.conformance import validate
from langgraph.checkpoint.conformance.initializer import checkpointer_test
from langgraph.checkpoint.memory import InMemorySaver
@checkpointer_test(name="InMemorySaver")
async def mem_saver():
yield InMemorySaver()
report = await validate(
mem_saver,
capabilities={
"delta_channel_history",
},
)
for cap, result in report.results.items():
if result.passed is False:
details = "\n".join(result.failures or [])
pytest.fail(f"Capability {cap} failed:\n{details}")
+36 -58
View File
@@ -34,28 +34,23 @@ def empty_checkpoint() -> Checkpoint:
)
def _should_snapshot_delta(
name: str,
ch: DeltaChannel,
updates_since_snapshot: Mapping[str, int],
*,
force: bool,
) -> bool:
"""Decide whether `ch` should write a `_DeltaSnapshot` this step.
def delta_channels_to_snapshot(
channels: Mapping[str, BaseChannel],
counts: Mapping[str, int],
) -> set[str]:
"""Return the set of DeltaChannel names that should snapshot now.
Triggers:
* `force` — always snapshot (used by `durability="exit"`).
* Update-count: this channel has accumulated at least
`snapshot_frequency` updates since its last snapshot. The count
is supplied by the caller via `updates_since_snapshot[name]` and
is reset to `0` whenever a snapshot fires.
Version-format-independent: works for `int`, `float`, and `str`
versioning schemes alike.
A channel snapshots when its accumulated update count (since the last
snapshot) reaches or exceeds `snapshot_frequency`. This is a pure
predicate — no mutation.
"""
if force:
return True
return updates_since_snapshot.get(name, 0) >= ch.snapshot_frequency
return {
name
for name, ch in channels.items()
if isinstance(ch, DeltaChannel)
and ch.is_available()
and counts.get(name, 0) >= ch.snapshot_frequency
}
def create_checkpoint(
@@ -66,34 +61,19 @@ def create_checkpoint(
id: str | None = None,
updated_channels: set[str] | None = None,
get_next_version: GetNextVersion | None = None,
force_delta_snapshot: bool = False,
updates_since_snapshot: Mapping[str, int] | None = None,
new_updates_since_snapshot: dict[str, int] | None = None,
channels_to_snapshot: set[str] | None = None,
) -> Checkpoint:
"""Create a checkpoint for the given channels.
"""Build a new Checkpoint from the previous one and live channel state.
For each `DeltaChannel`, a `_DeltaSnapshot(value)` blob is written into
`channel_values[k]` when this channel has accumulated at least
`snapshot_frequency` updates since its last snapshot (counter supplied
via `updates_since_snapshot`). Otherwise the channel is omitted from
`channel_values`; its `channel_versions` entry still bumps so that the
saver tracks the channel and the ancestor walk can replay writes.
Snapshots are eager: even if the channel had no write this step, a
version bump is forced (via `get_next_version`) so `put()` includes
the channel in `new_versions` and stores the blob.
`force_delta_snapshot` ignores the cadence and always snapshots —
used by `durability="exit"` where intermediate writes are not stored
as ancestor `checkpoint_writes`.
If `new_updates_since_snapshot` is provided, the function resets the
counter to `0` for any channel that snapshotted this step. Counters
for channels that did not snapshot are left untouched (the caller is
responsible for incrementing them based on `updated_channels`).
For each name in `channels_to_snapshot`, a `_DeltaSnapshot(value)` blob
is written into `channel_values[k]`. Other delta channels are omitted
from `channel_values` — the ancestor walk reconstructs their state
from `checkpoint_writes`. Callers compute the set via
`delta_channels_to_snapshot(channels, counts)`; defaults to empty
(no snapshots) when not provided.
"""
ts = datetime.now(timezone.utc).isoformat()
counts = updates_since_snapshot or {}
channels_to_snapshot = channels_to_snapshot or set()
if channels is None:
values = checkpoint["channel_values"]
channel_versions = checkpoint["channel_versions"]
@@ -104,25 +84,23 @@ def create_checkpoint(
if k not in channel_versions:
continue
ch = channels[k]
if (
isinstance(ch, DeltaChannel)
and ch.is_available()
and _should_snapshot_delta(
k,
ch,
counts,
force=force_delta_snapshot,
)
):
# Eager snapshot: bump version if not already written this step
# so put() includes this channel in new_versions and stores blob.
if k in channels_to_snapshot:
# In exit mode, the snapshot decision is deferred to exit
# time (intermediate steps have do_checkpoint=False). The
# channel's count may have reached snapshot_frequency over
# several supersteps, but the LAST superstep may not have
# written to this channel. In that case apply_writes()
# (in _algo.py) didn't bump this channel's version, so
# saver.put() wouldn't include it in new_versions and
# the snapshot blob would be silently dropped. The manual
# bump below closes the gap. In sync/async durability this
# branch is effectively dead code (the step that pushes
# the count to freq always writes the channel).
if get_next_version is not None and (
updated_channels is None or k not in updated_channels
):
channel_versions[k] = get_next_version(channel_versions[k], None)
values[k] = _DeltaSnapshot(ch.get())
if new_updates_since_snapshot is not None:
new_updates_since_snapshot[k] = 0
else:
v = ch.checkpoint()
if v is not MISSING:
+216 -21
View File
@@ -100,6 +100,7 @@ from langgraph.pregel._checkpoint import (
channels_from_checkpoint,
copy_checkpoint,
create_checkpoint,
delta_channels_to_snapshot,
empty_checkpoint,
)
from langgraph.pregel._executor import (
@@ -194,8 +195,40 @@ class PregelLoop:
_migrate_checkpoint: Callable[[Checkpoint], None] | None
submit: Submit
channels: Mapping[str, BaseChannel]
# Only set on AsyncPregelLoop; sync loops keep this as None.
# Futures from `checkpointer.put_writes` calls that produced delta-channel
# writes. `_checkpointer_put_after_previous` drains this list (swap to a
# local `futs` then reset to `[]` and wait/gather) before putting the
# next checkpoint, so a checkpoint never becomes durable before the
# writes that produced it. Initialised to `[]` in both sync and async
# `__enter__`; stays `None` only when no checkpointer.
_delta_write_futs: list[Any] | None = None
# Exit-mode accumulator: every delta-channel write produced during this
# run (input writes from `_first` + per-superstep writes captured in
# `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
# The checkpoint_config that points at the parent loaded at `__enter__`
# (or the synthetic-empty checkpoint, on first run). We capture it
# eagerly because every `_put_checkpoint` advances `self.checkpoint_config`
# to the newly-saved checkpoint's id — by exit time the original parent
# config would otherwise be lost. `_put_exit_delta_writes` uses this:
# on resumed runs as the anchor for exit delta writes; on first runs
# to derive the lazy stub's config (its `checkpoint_id` is the
# synthetic-empty id we want the stub persisted under).
_initial_checkpoint_config: RunnableConfig
# True iff the saver actually returned a tuple at `__enter__`. False
# on the first-ever run for a thread (no parent persisted yet).
# `_put_exit_delta_writes` uses this to decide between anchoring on
# the existing parent (True) or creating a lazy stub (False).
_has_persisted_parent: bool = False
managed: ManagedValueMapping
checkpoint: Checkpoint
checkpoint_id_saved: str
@@ -637,6 +670,11 @@ class PregelLoop:
self._emit(
"values", map_output_values, self.output_keys, writes, self.channels
)
# capture delta-channel writes for exit-mode accumulator before clearing
if self._exit_delta_writes is not None:
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))
# clear pending writes
self.checkpoint_pending_writes.clear()
# only replay (re-execute) done tasks on the first tick
@@ -854,6 +892,27 @@ class PregelLoop:
self.checkpointer_get_next_version,
self.trigger_to_nodes,
)
# Input writes go through `apply_writes` directly (above) — they
# never enter `checkpoint_pending_writes`, so the after_tick
# capture site does not see them. In exit mode, capture them
# here so `_exit_delta_writes` includes the input's delta writes
# alongside per-superstep writes; otherwise the input would be
# lost on read (it's not in final_checkpoint.channel_values for
# sub-freq channels, and walks ignore target.pending_writes).
if self._exit_delta_writes is not None:
for c, v in input_writes:
if isinstance(self.specs.get(c), DeltaChannel):
self._exit_delta_writes.append((self.step, NULL_TASK_ID, c, v))
# Persist delta-channel input writes so sub-freq inputs are
# recoverable via ancestor walk (mirrors the Command input path).
if self.durability != "exit":
delta_input = [
(c, v)
for c, v in input_writes
if isinstance(self.specs.get(c), DeltaChannel)
]
if delta_input:
self.put_writes(NULL_TASK_ID, delta_input)
# save input checkpoint
self.updated_channels = updated_channels
self._put_checkpoint({"source": "input"})
@@ -905,36 +964,60 @@ class PregelLoop:
return updated_channels
def _put_checkpoint(self, metadata: CheckpointMetadata) -> None:
# assign step and parents
# `is` (object identity) — not `==`. Three of four call sites pass a
# fresh dict ({"source":"input"|"loop"|"fork"}); only
# `_suppress_interrupt`(will rename to _on_loop_exit soon)
# at exit reuses the existing `self.checkpoint_metadata` instance. So
# `metadata is self.checkpoint_metadata` is True only on the exit call,
# which is what we use to gate exit-only behaviour (skip count-bump,
# don't replace metadata). Could be replaced by an explicit
# `exiting: bool = False` parameter; left as-is to match the existing
# idiom in this file.
# TODO: replace with an explicit `exiting: bool = False` parameter.
exiting = metadata is self.checkpoint_metadata
if exiting and self.checkpoint["id"] == self.checkpoint_id_saved:
# checkpoint already saved
return
# Carry per-delta-channel update bookkeeping forward across
# supersteps. Capture from the OLD metadata before potentially
# replacing it with a fresh dict that wouldn't contain it. Then
# increment for any delta channel updated this step (so the count
# reflects "supersteps that wrote to this channel since last
# snapshot"). create_checkpoint will reset entries to 0 for any
# channel that fires a snapshot this step.
prev_counts = dict(
self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
)
new_counts = dict(prev_counts)
if self.updated_channels:
for ch_name in self.updated_channels:
ch_obj = self.channels.get(ch_name)
if isinstance(ch_obj, DeltaChannel):
new_counts[ch_name] = new_counts.get(ch_name, 0) + 1
# Per-delta-channel update bookkeeping.
#
# `_put_checkpoint` is called once per superstep with a fresh
# metadata dict (source="input"|"loop"|"fork") — those are the
# intermediate calls that bump the count by +1 for each delta
# channel touched that step. In exit mode,
# `_suppress_interrupt`(will rename to _on_loop_exit soon)
# additionally calls `_put_checkpoint(self.checkpoint_metadata)` AT
# EXIT to commit the final checkpoint — this runs *after* the last
# intermediate call already counted the last superstep. So the
# exit call must NOT bump again or it would double-count the last
# superstep. (Sync/async durability does not call `_put_checkpoint`
# at exit, so the issue only surfaces in exit mode. force_delta_snapshot
# used to mask this latent bug by resetting every count to 0.)
if not exiting:
prev_counts = dict(
self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
)
new_counts = dict(prev_counts)
if self.updated_channels:
for ch_name in self.updated_channels:
if isinstance(self.channels.get(ch_name), DeltaChannel):
new_counts[ch_name] = new_counts.get(ch_name, 0) + 1
metadata["step"] = self.step
metadata["parents"] = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_MAP, {})
self.checkpoint_metadata = metadata
else:
new_counts = dict(
self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
)
# do checkpoint?
do_checkpoint = self._checkpointer_put_after_previous is not None and (
exiting or self.durability != "exit"
)
# create new checkpoint
channels_to_snapshot = (
delta_channels_to_snapshot(self.channels, new_counts)
if do_checkpoint
else set()
)
self.checkpoint = create_checkpoint(
self.checkpoint,
self.channels if do_checkpoint else None,
@@ -944,10 +1027,10 @@ class PregelLoop:
get_next_version=self.checkpointer_get_next_version
if do_checkpoint
else None,
force_delta_snapshot=exiting and self.durability == "exit",
updates_since_snapshot=new_counts,
new_updates_since_snapshot=new_counts,
channels_to_snapshot=channels_to_snapshot,
)
for k in channels_to_snapshot:
new_counts[k] = 0
if new_counts:
self.checkpoint_metadata["delta_updates_since_snapshot"] = new_counts
elif "delta_updates_since_snapshot" in self.checkpoint_metadata:
@@ -1010,6 +1093,97 @@ class PregelLoop:
# increment step
self.step += 1
def _put_exit_delta_writes(self) -> None:
"""Stage stub + accumulated delta writes so final_checkpoint's put
waits on them (visibility invariant: both must be durable before
final_checkpoint becomes visible to readers).
Stub is created lazily — only when no persisted parent exists AND at
least one delta channel has writes that won't be snapshotted.
"""
if (
not self._exit_delta_writes
or self.checkpointer is None
or self._checkpointer_put_after_previous is None
or self.checkpointer_put_writes is None
):
return
counts = self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {}
channels_to_snapshot = delta_channels_to_snapshot(self.channels, counts)
pending = [
(step, tid, ch, v)
for (step, tid, ch, v) in self._exit_delta_writes
if ch not in channels_to_snapshot
]
if not pending:
return
if self._has_persisted_parent:
# _initial_checkpoint_config's checkpoint_id is the saved parent's
# id (saver returned a real tuple at __enter__).
anchor_config = self._initial_checkpoint_config
else:
stub_cp = empty_checkpoint()
stub_cp["id"] = self.checkpoint_id_saved
stub_cp["ts"] = datetime.now(timezone.utc).isoformat()
# Stub has no parent (checkpoint_id=None in config).
stub_put_config = patch_configurable(
self._initial_checkpoint_config,
{CONFIG_KEY_CHECKPOINT_ID: None},
)
# Anchor config for put_writes: checkpoint_id = stub's id.
anchor_config = patch_configurable(
self._initial_checkpoint_config,
{CONFIG_KEY_CHECKPOINT_ID: stub_cp["id"]},
)
self._put_checkpoint_fut = self.submit(
self._checkpointer_put_after_previous,
getattr(self, "_put_checkpoint_fut", None),
stub_put_config,
stub_cp,
{"step": -2},
{},
)
# Set checkpoint_config so final_checkpoint's _put_checkpoint
# 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))
anchor_write_config = patch_configurable(
anchor_config,
{
CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get(
CONFIG_KEY_CHECKPOINT_NS, ""
),
CONFIG_KEY_CHECKPOINT_ID: anchor_config[CONF][CONFIG_KEY_CHECKPOINT_ID],
},
)
for (step, tid), entries in grouped.items():
synth_tid = f"{step:08d}-{tid}"
if self.checkpointer_put_writes_accepts_task_path:
fut = self.submit(
self.checkpointer_put_writes,
anchor_write_config,
entries,
synth_tid,
"",
)
else:
fut = self.submit(
self.checkpointer_put_writes,
anchor_write_config,
entries,
synth_tid,
)
if self._delta_write_futs is not None:
self._delta_write_futs.append(fut)
def _suppress_interrupt(
self,
exc_type: type[BaseException] | None,
@@ -1025,6 +1199,7 @@ class PregelLoop:
# or a nested graph with checkpointer=True
or all(NS_END not in part for part in self.checkpoint_ns)
):
self._put_exit_delta_writes()
self._put_checkpoint(self.checkpoint_metadata)
self._put_pending_writes()
# suppress interrupt
@@ -1230,6 +1405,9 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> RunnableConfig:
if self._delta_write_futs:
futs, self._delta_write_futs = self._delta_write_futs, []
concurrent.futures.wait(futs)
try:
if prev is not None:
prev.result()
@@ -1347,6 +1525,10 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
# graph/thread. Returns None on first invocation.
saved = self.checkpointer.get_tuple(self.checkpoint_config)
# Capture before the synthetic-empty fallback below overwrites `saved`.
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
# to lazy-create a stub instead of anchoring delta writes on a parent.
self._has_persisted_parent = saved is not None
if saved is None:
saved = CheckpointTuple(
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
@@ -1362,6 +1544,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
**saved.config.get(CONF, {}),
},
}
self._initial_checkpoint_config = self.checkpoint_config
self.prev_checkpoint_config = saved.parent_config
self.checkpoint_id_saved = saved.checkpoint["id"]
self.checkpoint = saved.checkpoint
@@ -1371,6 +1554,10 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
if saved.pending_writes is not None
else []
)
self._delta_write_futs = []
self._exit_delta_writes = (
[] if self.durability == "exit" and self.checkpointer is not None else None
)
self.submit = self.stack.enter_context(BackgroundExecutor(self.config))
self.channels, self.managed = channels_from_checkpoint(
self.specs,
@@ -1596,6 +1783,10 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
# graph/thread. Returns None on first invocation.
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
# Capture before the synthetic-empty fallback below overwrites `saved`.
# `_put_exit_delta_writes` uses this on first run (no persisted parent)
# to lazy-create a stub instead of anchoring delta writes on a parent.
self._has_persisted_parent = saved is not None
if saved is None:
saved = CheckpointTuple(
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
@@ -1611,6 +1802,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
**saved.config.get(CONF, {}),
},
}
self._initial_checkpoint_config = self.checkpoint_config
self.prev_checkpoint_config = saved.parent_config
self.checkpoint_id_saved = saved.checkpoint["id"]
self.checkpoint = saved.checkpoint
@@ -1621,6 +1813,9 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
else []
)
self._delta_write_futs = []
self._exit_delta_writes = (
[] if self.durability == "exit" and self.checkpointer is not None else None
)
self.submit = await self.stack.enter_async_context(
AsyncBackgroundExecutor(self.config)
)
@@ -0,0 +1,365 @@
"""Tests for exit-mode delta channel persistence redesign.
Validates that `durability="exit"` correctly persists delta-channel writes
using count-based snapshot decisions (rather than force-snapshotting every
channel), lazy stub creation when no parent exists, and proper read-path
reconstruction via ancestor walks.
"""
from typing import Annotated, Any
import pytest
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import _messages_delta_reducer
pytestmark = pytest.mark.anyio
def _build_graph(
checkpointer: InMemorySaver,
*,
freq: int = 1000,
) -> Any:
channel = DeltaChannel(_messages_delta_reducer, snapshot_frequency=freq)
# Functional TypedDict form: class form can't reference `channel` (a
# local variable) inside Annotated due to forward-ref evaluation rules.
State = TypedDict("State", {"messages": Annotated[list, channel]}) # type: ignore[call-overload] # noqa: UP013
def respond(state: dict) -> dict:
i = len(state["messages"])
return {"messages": [AIMessage(content=f"reply-{i}", id=f"ai{i}")]}
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_edge(START, "respond")
return builder.compile(checkpointer=checkpointer)
# ---------------------------------------------------------------------------
# 8a. Write-path / structural tests
# ---------------------------------------------------------------------------
async def test_exit_first_run_no_delta_writes() -> None:
"""Graph with delta channel invoked with input that doesn't touch it.
Only one checkpoint row, no stub."""
State = TypedDict( # noqa: UP013
"State",
{
"messages": Annotated[list, DeltaChannel(_messages_delta_reducer)],
"value": str,
},
) # type: ignore[call-overload]
def noop(state: dict) -> dict:
return {"value": "done"}
saver = InMemorySaver()
builder = StateGraph(State)
builder.add_node("noop", noop)
builder.add_edge(START, "noop")
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "no-delta-writes"}}
graph.invoke({"value": "start"}, config, durability="exit")
checkpoints = list(saver.list(config))
assert len(checkpoints) == 1
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
assert len(stubs) == 0
async def test_exit_first_run_all_snapshot() -> None:
"""snapshot_frequency=1 forces every channel to snapshot.
No stub needed; final_checkpoint has _DeltaSnapshot."""
saver = InMemorySaver()
graph = _build_graph(saver, freq=1)
config = {"configurable": {"thread_id": "all-snapshot"}}
result = graph.invoke(
{"messages": [HumanMessage(content="hi", id="h1")]},
config,
durability="exit",
)
assert len(result["messages"]) == 2
checkpoints = list(saver.list(config))
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
assert len(stubs) == 0
head = saver.get_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
state = graph.get_state(config)
assert [m.content for m in state.values["messages"]] == ["hi", "reply-1"]
async def test_exit_first_run_sub_freq_with_writes() -> None:
"""First run with default snapshot_frequency (1000), writes below threshold.
A stub is created; writes are anchored under it; get_state reconstructs."""
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": "sub-freq-first"}}
result = graph.invoke(
{"messages": [HumanMessage(content="hello", id="h1")]},
config,
durability="exit",
)
assert [m.content for m in result["messages"]] == ["hello", "reply-1"]
checkpoints = list(saver.list(config))
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
assert len(stubs) == 1, f"Expected 1 stub, got {len(stubs)}"
head = saver.get_tuple(config)
assert head is not None
assert "messages" not in head.checkpoint["channel_values"]
assert "messages" in head.checkpoint["channel_versions"]
state = graph.get_state(config)
assert [m.content for m in state.values["messages"]] == ["hello", "reply-1"]
async def test_exit_resumed_run_sub_freq() -> None:
"""Two consecutive exit runs. Second run anchors on the first's
final_checkpoint (no new stub). Ordering preserved."""
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": "resumed-sub-freq"}}
graph.invoke(
{"messages": [HumanMessage(content="msg1", id="h1")]},
config,
durability="exit",
)
graph.invoke(
{"messages": [HumanMessage(content="msg2", id="h2")]},
config,
durability="exit",
)
checkpoints = list(saver.list(config))
stubs = [t for t in checkpoints if t.metadata.get("step") == -2]
assert len(stubs) == 1
state = graph.get_state(config)
contents = [m.content for m in state.values["messages"]]
assert len(contents) == 4
assert contents[0] == "msg1"
assert contents[2] == "msg2"
assert contents[0:4:2] == ["msg1", "msg2"]
async def test_exit_count_parity_sync_vs_exit() -> None:
"""Sync and exit durability produce the same delta_updates_since_snapshot
after an equivalent run."""
for durability in ("sync", "exit"):
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": f"parity-{durability}"}}
graph.invoke(
{"messages": [HumanMessage(content="hi", id="h1")]},
config,
durability=durability,
)
head = saver.get_tuple(config)
assert head is not None
counts = head.metadata.get("delta_updates_since_snapshot", {})
assert counts.get("messages") == 2, (
f"durability={durability}: expected count=2, got {counts}"
)
async def test_exit_snapshot_fires_at_frequency() -> None:
"""With snapshot_frequency=3, after 3 exit runs (each incrementing count
by 2: input + superstep), the 2nd run hits count=4>=3, triggering snapshot.
After that run, count resets to 0 and channel_values has _DeltaSnapshot."""
saver = InMemorySaver()
graph = _build_graph(saver, freq=3)
config = {"configurable": {"thread_id": "snapshot-at-freq"}}
graph.invoke(
{"messages": [HumanMessage(content="m1", id="h1")]},
config,
durability="exit",
)
head = saver.get_tuple(config)
assert head is not None
count1 = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0)
assert count1 == 2
graph.invoke(
{"messages": [HumanMessage(content="m2", id="h2")]},
config,
durability="exit",
)
head = saver.get_tuple(config)
assert head is not None
count2 = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0)
assert count2 == 0, f"Expected reset to 0 after snapshot, got {count2}"
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
async def test_exit_mixed_snapshot_and_non_snapshot() -> None:
"""One delta channel at freq=1 (always snapshot) and one at freq=1000
(never snapshot within this test). Verify correct behavior for both."""
fast_ch = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1)
slow_ch = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1000)
State = TypedDict( # noqa: UP013
"State",
{"fast": Annotated[list, fast_ch], "slow": Annotated[list, slow_ch]},
) # type: ignore[call-overload]
def respond(state: dict) -> dict:
return {
"fast": [AIMessage(content="fast-reply", id="f1")],
"slow": [AIMessage(content="slow-reply", id="s1")],
}
saver = InMemorySaver()
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_edge(START, "respond")
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "mixed-freq"}}
graph.invoke(
{
"fast": [HumanMessage(content="fast-in", id="fi")],
"slow": [HumanMessage(content="slow-in", id="si")],
},
config,
durability="exit",
)
head = saver.get_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("fast"), _DeltaSnapshot)
assert "slow" not in head.checkpoint["channel_values"]
state = graph.get_state(config)
assert [m.content for m in state.values["fast"]] == ["fast-in", "fast-reply"]
assert [m.content for m in state.values["slow"]] == ["slow-in", "slow-reply"]
# ---------------------------------------------------------------------------
# 8b. Read-path tests
# ---------------------------------------------------------------------------
async def test_exit_multi_run_replay_chain() -> None:
"""K=4 consecutive exit runs, each adding a message. After each run,
get_state returns all messages in chronological order."""
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": "replay-chain"}}
for i in range(4):
graph.invoke(
{"messages": [HumanMessage(content=f"user-{i}", id=f"h{i}")]},
config,
durability="exit",
)
state = graph.get_state(config)
contents = [m.content for m in state.values["messages"]]
user_msgs = [c for c in contents if c.startswith("user-")]
assert user_msgs == [f"user-{j}" for j in range(i + 1)], (
f"After run {i}: user messages out of order: {user_msgs}"
)
assert len(contents) == (i + 1) * 2
async def test_exit_metadata_round_trip() -> None:
"""K=5 consecutive exit runs with snapshot_frequency=5. Verify metadata
delta_updates_since_snapshot increments correctly across runs."""
freq = 5
saver = InMemorySaver()
graph = _build_graph(saver, freq=freq)
config = {"configurable": {"thread_id": "metadata-rt"}}
for i in range(1, 6):
graph.invoke(
{"messages": [HumanMessage(content=f"m{i}", id=f"h{i}")]},
config,
durability="exit",
)
head = saver.get_tuple(config)
assert head is not None
count = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0)
cumulative = i * 2
if cumulative >= freq:
assert count == 0 or count == cumulative % freq or count < freq, (
f"After run {i}: count={count} should have reset or be partial"
)
else:
assert count == cumulative, (
f"After run {i}: expected {cumulative}, got {count}"
)
async def test_exit_mixed_durability_round_trip() -> None:
"""Alternate sync and exit durability; verify counts stay monotonic
and state accumulates correctly."""
saver = InMemorySaver()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": "mixed-durability"}}
for i, dur in enumerate(["sync", "exit", "sync", "exit"]):
graph.invoke(
{"messages": [HumanMessage(content=f"msg-{i}", id=f"h{i}")]},
config,
durability=dur,
)
state = graph.get_state(config)
contents = [m.content for m in state.values["messages"]]
user_msgs = [c for c in contents if c.startswith("msg-")]
assert user_msgs == [f"msg-{j}" for j in range(i + 1)], (
f"After run {i} (durability={dur}): {user_msgs}"
)
assert len(contents) == (i + 1) * 2
async def test_exit_snapshot_then_tail_deltas() -> None:
"""Run 1 forces snapshot (freq=1). Run 2 at freq=1000 adds more writes
that don't snapshot. Reading after run 2 must combine the snapshot seed
with the tail deltas."""
saver = InMemorySaver()
graph1 = _build_graph(saver, freq=1)
config = {"configurable": {"thread_id": "snapshot-then-tail"}}
graph1.invoke(
{"messages": [HumanMessage(content="seed-msg", id="h1")]},
config,
durability="exit",
)
head = saver.get_tuple(config)
assert head is not None
assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot)
graph2 = _build_graph(saver, freq=1000)
graph2.invoke(
{"messages": [HumanMessage(content="tail-msg", id="h2")]},
config,
durability="exit",
)
state = graph2.get_state(config)
contents = [m.content for m in state.values["messages"]]
assert "seed-msg" in contents
assert "tail-msg" in contents
assert contents.index("seed-msg") < contents.index("tail-msg")
+17 -4
View File
@@ -1674,15 +1674,28 @@ async def test_arun_with_retry_timeout_observer_tracks_attempts():
async def test_arun_with_retry_timeout_observer_emits_progress_on_heartbeat():
events: list = []
# `_TimedAttemptScope.__init__` sets `_last_progress` to `time.monotonic()`,
# but the watchdog itself doesn't start running until after `wrap_config`
# and task scheduling — under CI load that gap can be large enough to eat
# the entire idle window before the task body's first await even runs. We
# defend against that by:
# 1. Using a generous idle_timeout so scheduling slack stays well within it.
# 2. Calling `runtime.heartbeat()` BEFORE the first sleep, which resets
# `_last_progress` to "now" the moment the task body actually starts.
idle_timeout_s = 1.0
class HeartbeatProc:
async def ainvoke(self, input, config):
runtime = config[CONF][CONFIG_KEY_RUNTIME]
runtime.heartbeat() # reset the idle clock at task-body entry
for _ in range(8):
await asyncio.sleep(0.05)
runtime.heartbeat()
return "ok"
task = _make_task(HeartbeatProc(), timeout=_idle_timeout(0.2), name="heartbeat")
task = _make_task(
HeartbeatProc(), timeout=_idle_timeout(idle_timeout_s), name="heartbeat"
)
task.config[CONF][CONFIG_KEY_TIMED_ATTEMPT_OBSERVER] = events.append
assert await arun_with_retry(task, retry_policy=None) == "ok"
@@ -1691,13 +1704,13 @@ async def test_arun_with_retry_timeout_observer_emits_progress_on_heartbeat():
assert by_event[-1] == "finish"
progress = [ev for ev in events if ev.event == "progress"]
assert progress, "expected at least one progress event from heartbeat"
# Rate limit is `idle_timeout / 4` = 0.05s; with 8 heartbeats spaced ~0.05s
# we should see at most ~one progress event per heartbeat (well below 8).
# Rate limit is `idle_timeout / 4` = 0.25s; with the task running for
# ~400ms we expect 1–2 progress events (well below the 9 heartbeats).
assert len(progress) <= len(by_event)
for ev in progress:
assert ev.context.task_name == "heartbeat"
assert ev.context.attempt == 1
assert ev.context.idle_timeout_secs == 0.2
assert ev.context.idle_timeout_secs == idle_timeout_s
assert isinstance(ev.progress_at, datetime)
@@ -6,6 +6,7 @@ from typing import (
Any,
Literal,
TypeVar,
Union,
cast,
get_type_hints,
)
@@ -125,6 +126,9 @@ Prompt = (
| Runnable[StateSchema, LanguageModelInput]
)
# A single hook or a list of hooks to be composed in order.
HookLike = Union[RunnableLike, Sequence[RunnableLike]]
def _get_state_value(state: StateSchema, key: str, default: Any = None) -> Any:
return (
@@ -134,6 +138,131 @@ def _get_state_value(state: StateSchema, key: str, default: Any = None) -> Any:
)
def _set_state_value(state: StateSchema, key: str, value: Any) -> None:
"""Set a value in the state, supporting both dict and Pydantic model states."""
if isinstance(state, dict):
state[key] = value
else:
setattr(state, key, value)
def _merge_state_update(state: StateSchema, update: dict) -> StateSchema:
"""Return a shallow copy of *state* with *update* applied.
This is used when chaining multiple hooks: each hook receives the state
as it would look after all previous hooks have run, so that hooks later
in the chain can observe updates made by earlier ones.
Note: only simple key-level merging is performed here (no reducer logic).
The full reducer logic is applied by the graph engine when the final
combined update dict is written back to the state.
"""
if isinstance(state, dict):
return {**state, **update} # type: ignore[return-value]
else:
# Pydantic / dataclass – make a shallow copy and patch fields
try:
merged = state.model_copy() # pydantic v2
except AttributeError:
merged = state.copy() # pydantic v1 / dataclass fallback
for k, v in update.items():
setattr(merged, k, v)
return merged # type: ignore[return-value]
def _coerce_to_runnable(hook: RunnableLike) -> RunnableCallable:
"""Wrap a plain callable into a RunnableCallable if necessary."""
if isinstance(hook, RunnableCallable):
return hook
if isinstance(hook, Runnable):
# Already a Runnable – wrap so we get a uniform interface
sync_fn = hook.invoke
async_fn = hook.ainvoke
return RunnableCallable(sync_fn, async_fn)
if inspect.iscoroutinefunction(hook):
return RunnableCallable(None, hook)
if callable(hook):
return RunnableCallable(hook)
raise TypeError(f"Expected a callable or Runnable, got {type(hook)!r}")
def _chain_hooks(hooks: Sequence[RunnableLike]) -> RunnableCallable:
"""Compose multiple hook callables into a single hook.
Each hook is called in order. After each hook the returned update dict is
merged into a running copy of the graph state so that subsequent hooks can
observe the changes made by earlier ones. The accumulated update dict
(union of all individual update dicts, with later hooks winning on key
conflicts) is returned as the final state update.
Args:
hooks: A sequence of :data:`RunnableLike` objects. Each must accept
the graph state as its first positional argument and return a
``dict`` of state updates.
Returns:
A :class:`~langgraph._internal._runnable.RunnableCallable` that behaves
like a single hook but applies all of *hooks* in sequence.
"""
if not hooks:
raise ValueError("_chain_hooks requires at least one hook")
if len(hooks) == 1:
return _coerce_to_runnable(hooks[0])
runnables = [_coerce_to_runnable(h) for h in hooks]
def _sync_chained(state: Any, **kwargs: Any) -> dict:
accumulated: dict = {}
current_state = state
for runnable in runnables:
# Pass extra kwargs (e.g. config, store) through if the hook
# accepts them; RunnableCallable handles introspection.
update = runnable.invoke(current_state, **kwargs)
if update:
accumulated.update(update)
current_state = _merge_state_update(current_state, update)
return accumulated
async def _async_chained(state: Any, **kwargs: Any) -> dict:
accumulated: dict = {}
current_state = state
for runnable in runnables:
update = await runnable.ainvoke(current_state, **kwargs)
if update:
accumulated.update(update)
current_state = _merge_state_update(current_state, update)
return accumulated
return RunnableCallable(_sync_chained, _async_chained, name="chained_hooks")
def _resolve_hook(hook: HookLike | None) -> RunnableLike | None:
"""Normalise *hook* to a single ``RunnableLike`` (or ``None``).
* If *hook* is ``None`` → return ``None``.
* If *hook* is already a ``RunnableLike`` → return it unchanged.
* If *hook* is a non-empty :class:`~collections.abc.Sequence` of
``RunnableLike`` → chain them with :func:`_chain_hooks`.
"""
if hook is None:
return None
# A Sequence[RunnableLike] but NOT a single Runnable/callable
if (
isinstance(hook, Sequence)
and not isinstance(hook, str)
and not isinstance(hook, Runnable)
and not callable(hook)
):
hooks_list: list[RunnableLike] = list(hook)
if not hooks_list:
return None
if len(hooks_list) == 1:
return hooks_list[0]
return _chain_hooks(hooks_list)
# Single hook – return as-is
return hook # type: ignore[return-value]
def _get_prompt_runnable(prompt: Prompt | None) -> Runnable:
prompt_runnable: Runnable
if prompt is None:
@@ -293,8 +422,8 @@ def create_react_agent(
response_format: StructuredResponseSchema
| tuple[str, StructuredResponseSchema]
| None = None,
pre_model_hook: RunnableLike | None = None,
post_model_hook: RunnableLike | None = None,
pre_model_hook: HookLike | None = None,
post_model_hook: HookLike | None = None,
state_schema: StateSchemaType | None = None,
context_schema: type[Any] | None = None,
checkpointer: Checkpointer | None = None,
@@ -393,10 +522,21 @@ def create_react_agent(
The graph will make a separate call to the LLM to generate the structured response after the agent loop is finished.
This is not the only strategy to get structured responses, see more options in [this guide](https://langchain-ai.github.io/langgraph/how-tos/react-agent-structured-output/).
pre_model_hook: An optional node to add before the `agent` node (i.e., the node that calls the LLM).
Useful for managing long message histories (e.g., message trimming, summarization, etc.).
Pre-model hook must be a callable or a runnable that takes in current graph state and returns a state update in the form of
```python
pre_model_hook: An optional node (or list of nodes) to add before the
``agent`` node (i.e., the node that calls the LLM). Useful for
managing long message histories (e.g., message trimming,
summarization, etc.) or for composing multiple pre-processing
steps.
A single hook **or a list of hooks** may be provided. When a list
is given the hooks are executed in order: each hook receives the
graph state as updated by all preceding hooks, and the union of
all their return dicts is applied to the graph state before the
agent node runs.
Each hook must be a callable or a runnable that takes the current
graph state and returns a state update::
# At least one of `messages` or `llm_input_messages` MUST be provided
{
# If provided, will UPDATE the `messages` in the state
@@ -407,27 +547,63 @@ def create_react_agent(
# Any other state keys that need to be propagated
...
}
```
!!! Important
At least one of `messages` or `llm_input_messages` MUST be provided and will be used as an input to the `agent` node.
The rest of the keys will be added to the graph state.
At least one of `messages` or `llm_input_messages` MUST be
provided (by at least one hook in the chain) and will be used
as an input to the ``agent`` node. The rest of the keys will
be added to the graph state.
!!! Warning
If you are returning `messages` in the pre-model hook, you should OVERWRITE the `messages` key by doing the following:
If you are returning `messages` in the pre-model hook, you
should OVERWRITE the `messages` key::
{
"messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES), *new_messages]
...
}
!!! Example "Composing multiple pre-model hooks"
```python
{
"messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES), *new_messages]
...
}
from langchain_core.messages import RemoveMessage
from langgraph.graph.message import REMOVE_ALL_MESSAGES
def trim_messages(state):
# Keep only the last 10 messages
return {
"messages": [
RemoveMessage(id=REMOVE_ALL_MESSAGES),
*state["messages"][-10:],
]
}
def inject_system_prompt(state):
return {
"llm_input_messages": [
SystemMessage("You are a helpful assistant."),
*state["messages"],
]
}
agent = create_react_agent(
model,
tools,
pre_model_hook=[trim_messages, inject_system_prompt],
)
```
post_model_hook: An optional node to add after the `agent` node (i.e., the node that calls the LLM).
Useful for implementing human-in-the-loop, guardrails, validation, or other post-processing.
Post-model hook must be a callable or a runnable that takes in current graph state and returns a state update.
post_model_hook: An optional node (or list of nodes) to add after the
``agent`` node (i.e., the node that calls the LLM). Useful for
implementing human-in-the-loop, guardrails, validation, or other
post-processing steps.
Accepts the same single-hook-or-list-of-hooks form as
``pre_model_hook``.
!!! Note
Only available with `version="v2"`.
Only available with ``version="v2"``.
state_schema: An optional state schema that defines graph state.
Must have `messages` and `remaining_steps` keys.
Defaults to `AgentState` that defines those two keys.
@@ -551,6 +727,10 @@ def create_react_agent(
else AgentState
)
# Normalise hook arguments: a list of hooks is composed into a single hook.
resolved_pre_model_hook: RunnableLike | None = _resolve_hook(pre_model_hook)
resolved_post_model_hook: RunnableLike | None = _resolve_hook(post_model_hook)
llm_builtin_tools: list[dict] = []
if isinstance(tools, ToolNode):
tool_classes = list(tools.tools_by_name.values())
@@ -634,7 +814,7 @@ def create_react_agent(
return False
def _get_model_input_state(state: StateSchema) -> StateSchema:
if pre_model_hook is not None:
if resolved_pre_model_hook is not None:
messages = (
_get_state_value(state, "llm_input_messages")
) or _get_state_value(state, "messages")
@@ -721,7 +901,7 @@ def create_react_agent(
return {"messages": [response]}
input_schema: StateSchemaType
if pre_model_hook is not None:
if resolved_pre_model_hook is not None:
# Dynamically create a schema that inherits from state_schema and adds 'llm_input_messages'
if isinstance(state_schema, type) and issubclass(state_schema, BaseModel):
# For Pydantic schemas
@@ -792,8 +972,8 @@ def create_react_agent(
RunnableCallable(call_model, acall_model),
input_schema=input_schema,
)
if pre_model_hook is not None:
workflow.add_node("pre_model_hook", pre_model_hook) # type: ignore[arg-type]
if resolved_pre_model_hook is not None:
workflow.add_node("pre_model_hook", resolved_pre_model_hook) # type: ignore[arg-type]
workflow.add_edge("pre_model_hook", "agent")
entrypoint = "pre_model_hook"
else:
@@ -801,8 +981,8 @@ def create_react_agent(
workflow.set_entry_point(entrypoint)
if post_model_hook is not None:
workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type]
if resolved_post_model_hook is not None:
workflow.add_node("post_model_hook", resolved_post_model_hook) # type: ignore[arg-type]
workflow.add_edge("agent", "post_model_hook")
if response_format is not None:
@@ -813,7 +993,7 @@ def create_react_agent(
agenerate_structured_response,
),
)
if post_model_hook is not None:
if resolved_post_model_hook is not None:
workflow.add_edge("post_model_hook", "generate_structured_response")
else:
workflow.add_edge("agent", "generate_structured_response")
@@ -833,7 +1013,7 @@ def create_react_agent(
last_message = messages[-1]
# If there is no function call, then we finish
if not isinstance(last_message, AIMessage) or not last_message.tool_calls:
if post_model_hook is not None:
if resolved_post_model_hook is not None:
return "post_model_hook"
elif response_format is not None:
return "generate_structured_response"
@@ -844,7 +1024,7 @@ def create_react_agent(
if version == "v1":
return "tools"
elif version == "v2":
if post_model_hook is not None:
if resolved_post_model_hook is not None:
return "post_model_hook"
return [
Send(
@@ -873,8 +1053,8 @@ def create_react_agent(
# Optionally add a pre-model hook node that will be called
# every time before the "agent" (LLM-calling node)
if pre_model_hook is not None:
workflow.add_node("pre_model_hook", pre_model_hook) # type: ignore[arg-type]
if resolved_pre_model_hook is not None:
workflow.add_node("pre_model_hook", resolved_pre_model_hook) # type: ignore[arg-type]
workflow.add_edge("pre_model_hook", "agent")
entrypoint = "pre_model_hook"
else:
@@ -888,8 +1068,8 @@ def create_react_agent(
post_model_hook_paths = [entrypoint, "tools"]
# Add a post model hook node if post_model_hook is provided
if post_model_hook is not None:
workflow.add_node("post_model_hook", post_model_hook) # type: ignore[arg-type]
if resolved_post_model_hook is not None:
workflow.add_node("post_model_hook", resolved_post_model_hook) # type: ignore[arg-type]
agent_paths.append("post_model_hook")
workflow.add_edge("agent", "post_model_hook")
else:
@@ -904,17 +1084,17 @@ def create_react_agent(
agenerate_structured_response,
),
)
if post_model_hook is not None:
if resolved_post_model_hook is not None:
post_model_hook_paths.append("generate_structured_response")
else:
agent_paths.append("generate_structured_response")
else:
if post_model_hook is not None:
if resolved_post_model_hook is not None:
post_model_hook_paths.append(END)
else:
agent_paths.append(END)
if post_model_hook is not None:
if resolved_post_model_hook is not None:
def post_model_hook_router(state: StateSchema) -> str | list[Send]:
"""Route to the next node after post_model_hook.
@@ -1012,4 +1192,5 @@ __all__ = [
"AgentStatePydantic",
"AgentStateWithStructuredResponse",
"AgentStateWithStructuredResponsePydantic",
"HookLike",
]
+11
View File
@@ -18,6 +18,7 @@ from langgraph_sdk.schema import (
CronSortBy,
Durability,
Input,
Json,
OnCompletionBehavior,
QueryParamTypes,
Run,
@@ -413,6 +414,7 @@ class CronClient:
assistant_id: str | None = None,
thread_id: str | None = None,
enabled: bool | None = None,
metadata: Json = None,
limit: int = 10,
offset: int = 0,
sort_by: CronSortBy | None = None,
@@ -427,6 +429,8 @@ class CronClient:
assistant_id: The assistant ID or graph name to search for.
thread_id: the thread ID to search for.
enabled: The enabled status to search for.
metadata: Metadata to filter by. Exact match filter for each KV pair.
!!! version-added "Added in Agent Server version 0.9.0"
limit: The maximum number of results to return.
offset: The number of results to skip.
headers: Optional custom headers to include with the request.
@@ -481,6 +485,8 @@ class CronClient:
"limit": limit,
"offset": offset,
}
if metadata:
payload["metadata"] = metadata
if sort_by:
payload["sort_by"] = sort_by
if sort_order:
@@ -497,6 +503,7 @@ class CronClient:
*,
assistant_id: str | None = None,
thread_id: str | None = None,
metadata: Json = None,
headers: Mapping[str, str] | None = None,
params: QueryParamTypes | None = None,
) -> int:
@@ -505,6 +512,8 @@ class CronClient:
Args:
assistant_id: Assistant ID to filter by.
thread_id: Thread ID to filter by.
metadata: Metadata to filter by. Exact match filter for each KV pair.
!!! version-added "Added in Agent Server version 0.9.0"
headers: Optional custom headers to include with the request.
params: Optional query parameters to include with the request.
@@ -516,6 +525,8 @@ class CronClient:
payload["assistant_id"] = assistant_id
if thread_id:
payload["thread_id"] = thread_id
if metadata:
payload["metadata"] = metadata
return await self.http.post(
"/runs/crons/count", json=payload, headers=headers, params=params
)
+11
View File
@@ -18,6 +18,7 @@ from langgraph_sdk.schema import (
CronSortBy,
Durability,
Input,
Json,
OnCompletionBehavior,
QueryParamTypes,
Run,
@@ -402,6 +403,7 @@ class SyncCronClient:
assistant_id: str | None = None,
thread_id: str | None = None,
enabled: bool | None = None,
metadata: Json = None,
limit: int = 10,
offset: int = 0,
sort_by: CronSortBy | None = None,
@@ -416,6 +418,8 @@ class SyncCronClient:
assistant_id: The assistant ID or graph name to search for.
thread_id: the thread ID to search for.
enabled: Whether the cron job is enabled.
metadata: Metadata to filter by. Exact match filter for each KV pair.
!!! version-added "Added in Agent Server version 0.9.0"
limit: The maximum number of results to return.
offset: The number of results to skip.
headers: Optional custom headers to include with the request.
@@ -468,6 +472,8 @@ class SyncCronClient:
"limit": limit,
"offset": offset,
}
if metadata:
payload["metadata"] = metadata
if sort_by:
payload["sort_by"] = sort_by
if sort_order:
@@ -484,6 +490,7 @@ class SyncCronClient:
*,
assistant_id: str | None = None,
thread_id: str | None = None,
metadata: Json = None,
headers: Mapping[str, str] | None = None,
params: QueryParamTypes | None = None,
) -> int:
@@ -492,6 +499,8 @@ class SyncCronClient:
Args:
assistant_id: Assistant ID to filter by.
thread_id: Thread ID to filter by.
metadata: Metadata to filter by. Exact match filter for each KV pair.
!!! version-added "Added in Agent Server version 0.9.0"
headers: Optional custom headers to include with the request.
params: Optional query parameters to include with the request.
@@ -503,6 +512,8 @@ class SyncCronClient:
payload["assistant_id"] = assistant_id
if thread_id:
payload["thread_id"] = thread_id
if metadata:
payload["metadata"] = metadata
return self.http.post(
"/runs/crons/count", json=payload, headers=headers, params=params
)
+162
View File
@@ -485,3 +485,165 @@ def test_sync_update_with_enabled_parameter(enabled_value):
)
assert result == cron
@pytest.mark.asyncio
async def test_async_search_with_metadata():
"""Test that CronClient.search forwards metadata in the request body."""
cron = _cron_response()
async def handler(request: httpx.Request) -> httpx.Response:
assert request.method == "POST"
assert request.url.path == "/runs/crons/search"
body = json.loads(request.content)
assert body["metadata"] == {"owner": "alice"}
assert body["limit"] == 10
assert body["offset"] == 0
return httpx.Response(200, json=[cron])
transport = httpx.MockTransport(handler)
async with httpx.AsyncClient(
transport=transport, base_url="https://example.com"
) as client:
http_client = HttpClient(client)
cron_client = CronClient(http_client)
result = await cron_client.search(metadata={"owner": "alice"})
assert result == [cron]
@pytest.mark.asyncio
async def test_async_search_omits_empty_metadata():
"""Test that CronClient.search does not send metadata when not provided."""
cron = _cron_response()
async def handler(request: httpx.Request) -> httpx.Response:
body = json.loads(request.content)
assert "metadata" not in body
return httpx.Response(200, json=[cron])
transport = httpx.MockTransport(handler)
async with httpx.AsyncClient(
transport=transport, base_url="https://example.com"
) as client:
http_client = HttpClient(client)
cron_client = CronClient(http_client)
await cron_client.search()
@pytest.mark.asyncio
async def test_async_count_with_metadata():
"""Test that CronClient.count forwards metadata in the request body."""
async def handler(request: httpx.Request) -> httpx.Response:
assert request.method == "POST"
assert request.url.path == "/runs/crons/count"
body = json.loads(request.content)
assert body["metadata"] == {"team": "infra"}
return httpx.Response(200, json=2)
transport = httpx.MockTransport(handler)
async with httpx.AsyncClient(
transport=transport, base_url="https://example.com"
) as client:
http_client = HttpClient(client)
cron_client = CronClient(http_client)
result = await cron_client.count(metadata={"team": "infra"})
assert result == 2
@pytest.mark.asyncio
async def test_async_count_omits_empty_metadata():
"""Test that CronClient.count does not send metadata when not provided."""
async def handler(request: httpx.Request) -> httpx.Response:
body = json.loads(request.content)
assert "metadata" not in body
return httpx.Response(200, json=0)
transport = httpx.MockTransport(handler)
async with httpx.AsyncClient(
transport=transport, base_url="https://example.com"
) as client:
http_client = HttpClient(client)
cron_client = CronClient(http_client)
await cron_client.count()
def test_sync_search_with_metadata():
"""Test that SyncCronClient.search forwards metadata in the request body."""
cron = _cron_response()
def handler(request: httpx.Request) -> httpx.Response:
assert request.method == "POST"
assert request.url.path == "/runs/crons/search"
body = json.loads(request.content)
assert body["metadata"] == {"owner": "alice"}
return httpx.Response(200, json=[cron])
transport = httpx.MockTransport(handler)
with httpx.Client(transport=transport, base_url="https://example.com") as client:
http_client = SyncHttpClient(client)
cron_client = SyncCronClient(http_client)
result = cron_client.search(metadata={"owner": "alice"})
assert result == [cron]
def test_sync_search_omits_empty_metadata():
"""Test that SyncCronClient.search does not send metadata when not provided."""
cron = _cron_response()
def handler(request: httpx.Request) -> httpx.Response:
body = json.loads(request.content)
assert "metadata" not in body
return httpx.Response(200, json=[cron])
transport = httpx.MockTransport(handler)
with httpx.Client(transport=transport, base_url="https://example.com") as client:
http_client = SyncHttpClient(client)
cron_client = SyncCronClient(http_client)
cron_client.search()
def test_sync_count_with_metadata():
"""Test that SyncCronClient.count forwards metadata in the request body."""
def handler(request: httpx.Request) -> httpx.Response:
assert request.method == "POST"
assert request.url.path == "/runs/crons/count"
body = json.loads(request.content)
assert body["metadata"] == {"team": "infra"}
return httpx.Response(200, json=2)
transport = httpx.MockTransport(handler)
with httpx.Client(transport=transport, base_url="https://example.com") as client:
http_client = SyncHttpClient(client)
cron_client = SyncCronClient(http_client)
result = cron_client.count(metadata={"team": "infra"})
assert result == 2
def test_sync_count_omits_empty_metadata():
"""Test that SyncCronClient.count does not send metadata when not provided."""
def handler(request: httpx.Request) -> httpx.Response:
body = json.loads(request.content)
assert "metadata" not in body
return httpx.Response(200, json=0)
transport = httpx.MockTransport(handler)
with httpx.Client(transport=transport, base_url="https://example.com") as client:
http_client = SyncHttpClient(client)
cron_client = SyncCronClient(http_client)
cron_client.count()