Compare commits

..
Author SHA1 Message Date
Sydney RunkleandGitHub 530fcabfc3 release: alpha bump (a3) for langgraph, checkpoint, checkpoint-postgres (#7678)
## Summary
- Bumps `langgraph` 1.2.0a1 → 1.2.0a3
- Bumps `langgraph-checkpoint` 4.1.0a1 → 4.1.0a3
- Bumps `langgraph-checkpoint-postgres` 3.1.0a1 → 3.1.0a3
- Bumps min `langgraph-checkpoint` constraint in `langgraph` to
`>=4.1.0a3`
- Refreshes uv locks across the workspace

(Note: `a2` was already cut from another branch.)

## Test plan
- [ ] CI passes
2026-05-01 11:18:52 -04:00
d8b7800183 chore(langgraph): use two phase read to avoid unnecessary data transport (#7660)
## Summary

Replaces the single-roundtrip `UNION ALL` DeltaChannel read with a
two-stage query that avoids fetching unused snapshot blobs, then removes
the old combined path entirely.

### Problem

`_get_channel_writes_history` used a single `UNION ALL` query that
fetched **all** checkpoint metadata, writes, and blobs for a
`(thread_id, channel)` in one shot. With `snapshot_frequency=N`, this
pulled back O(N/freq) full-size snapshot blobs even though only the
nearest one is needed to seed reconstruction. At 500 turns with
`snapshot_frequency=10`, this meant fetching ~100 complete
message-history snapshots per read.

### Solution

Two-stage read:
- **Stage 1** — lightweight scan of `checkpoints` only (no blob bytes):
walks the parent chain from the target checkpoint and stops at the first
ancestor with a snapshot, returning `chain_cids` and `seed_version`
- **Stage 2** — targeted fetch: only the writes for `chain_cids` and the
single seed blob at `seed_version`

The two-stage path is now unconditional — the old combined query and
`LG_DELTA_TWO_STAGE_QUERY` env-var gate have been removed.

### Sentinel cleanup

`DELTA_SENTINEL` is now a pure in-memory signal and is never written to
storage:
- Postgres `put()` already stripped it from `channel_values` before
writing blobs
- Memory saver `put()` now stores `"empty"` instead of serializing the
sentinel
- `EXT_DELTA_SENTINEL` (msgpack ext code 8) removed from
`JsonPlusSerializer`
- `DELTA_SENTINEL` is kept as an in-memory marker:
`DeltaChannel.checkpoint()` returns it so savers know to skip it, and
`_ChannelWritesHistory.seed` uses it to mean "no snapshot found, start
from empty"

## Performance

Benchmarked at `snapshot_frequency=10` on Postgres (`~100 tok/msg`):

| turns | old combined query | two-stage |
|------:|-------------------:|----------:|
| 50    | 6.0ms              | 2.8ms  (2.1x faster) |
| 100   | 10.1ms             | 5.6ms  (1.8x faster) |
| 500   | **216.1ms**        | 15.3ms (**14x faster**) |

The old query's read time grew super-linearly with turn count because
each read fetched O(N/freq) full snapshot blobs. Two-stage keeps read
depth bounded by `snapshot_frequency` regardless of thread length.

## Test plan

- `make test` in `libs/checkpoint`, `libs/checkpoint-postgres`,
`libs/langgraph`
- Removed `test_delta_sentinel_serde_round_trip` (sentinel no longer
serializable)
- Updated `test_memory.py` — delta channel blobs stored as `"empty"`,
not serialized sentinel
- Updated `test_channels.py` — `channel_values` no longer contains
sentinel key for DeltaChannels
- Deleted `test_delta_channel_two_stage_benchmark.py` (one-stage vs
two-stage comparison; path no longer exists)

---------

Co-authored-by: Sydney Runkle <54324534+sydney-runkle@users.noreply.github.com>
Co-authored-by: Claude Sonnet 4.6 (1M context) <noreply@anthropic.com>
2026-05-01 11:06:54 -04:00
Quanzheng LongandGitHub c8c58a0768 fix(langgraph): make NodeTimeoutError retryable by default (#7659)
## Summary

`NodeTimeoutError` previously inherited from `TimeoutError`, which is a
subclass of `OSError`. Since `OSError` is in the default `RetryPolicy`
blocklist, timeout errors from `TimeoutPolicy` were silently **not
retried** unless the user explicitly set `retry_on=NodeTimeoutError`.

This PR changes `NodeTimeoutError` to inherit from `Exception` directly,
so that the default `RetryPolicy` treats it as retryable — matching user
expectations when both `RetryPolicy` and `TimeoutPolicy` are configured
together.

- Change `NodeTimeoutError(TimeoutError)` →
`NodeTimeoutError(Exception)`
- Add test asserting `NodeTimeoutError` is retryable with the default
policy
- Add observer-ordering tests pinning down `finish=error` emission
timing relative to retry backoff, error handler start, and retry
exhaustion

## Breaking change

Code that catches `NodeTimeoutError` via `except TimeoutError` or
`except OSError` will no longer match. Use `except NodeTimeoutError`
instead.

## Test plan

- [x] `test_should_retry_default_retry_on` — asserts `NodeTimeoutError`
is retryable with default `RetryPolicy()`
- [x] Existing timeout+retry tests continue to pass (`test_retry.py`)
2026-04-30 12:17:13 -07:00
Nick HollonandGitHub de9b7c61c3 fix(langgraph): arrival-ordered interleave for StreamChannel projections (#7643) 2026-04-30 10:41:43 -04:00
Quanzheng LongandGitHub 63d861165f feat(langgraph): add node-level error handlers (#7233) 2026-04-29 21:27:06 -07:00
Eugene YurtsevandGitHub 9c1d65695e fix(prebuilt): default ToolRuntime tools to empty list (#7650)
Makes `ToolRuntime.tools` default to an empty list when not provided,
which avoids requiring callers and tests to pass it explicitly. Adds a
focused regression test covering direct `ToolRuntime` construction
without `tools`.

Created with [Deep Agents
CLI](https://docs.langchain.com/oss/python/deepagents/cli/overview)
using gpt-5.4 (provider: openai).
2026-04-30 01:07:00 +00:00
40ab009c62 feat: allow graph to graceful shutdown/drain by request (#7274)
## Summary

Adds cooperative drain support for Pregel runs so a graph can be asked
to stop at the next superstep boundary, persist its checkpoint, and
surface a resumable terminal exception.

- New `RunControl` (in `langgraph.runtime`) — a thread-safe handle whose
`request_drain(reason="shutdown")` sets a single flag.
- New `GraphDrained(GraphBubbleUp)` exception (in `langgraph.errors`)
raised when a run exits early due to drain. Carries the `reason` string.
- New `control: RunControl | None` kwarg on `invoke` / `ainvoke` /
`stream` / `astream` / `stream_v2` / `astream_v2`. Wired through to
`Runtime.control`, so nodes can read `runtime.control.drain_requested` /
`drain_reason` and even call `request_drain()` from inside a node.
- Stream transformers learn `"drained"` as a terminal `SubgraphStatus`.

The intended use is hooking SIGTERM (or any external supervisor signal)
to `control.request_drain("sigterm")` so an in-flight graph run can stop
cleanly and be resumed later from the saved checkpoint.

## Semantics: cooperative, between-superstep

`request_drain()` flips a flag. The Pregel loop checks it at the top of
each `tick()`, **after** the previous superstep's writes have been
applied and checkpointed. It never preempts work that is already
running.

| Scenario | Behavior |
|---|---|
| Node mid-execution (blocking I/O, sleeps, etc.) | Runs to completion.
Drain takes effect on the next superstep. |
| Node with a retry policy currently retrying | Retry loop runs to
exhaustion or success (drain is not checked between retries). Drain
takes effect on the next superstep. |
| Functional API: `@entrypoint` with pending `@task` futures |
Entrypoint and all dispatched tasks complete; drain takes effect after
the entrypoint returns. |
| Graph naturally finishes on the same tick where drain was requested
(no more tasks) | Treated as `done`; returns normally. **No
`GraphDrained` is raised.** The caller can inspect
`control.drain_requested` afterwards to distinguish a
drained-but-completed run from a normal one. |
| More tasks remain | Raises `GraphDrained(reason)`. The checkpoint of
the last completed superstep is saved (also under `durability="exit"`).
Resume with `invoke(None, config)` / `ainvoke(None, config)`. |
| Subgraph requests drain | `GraphDrained` bubbles up through the parent
loop and stops it at its own next superstep boundary; the parent's
checkpoint is saved and resumable. |

Drain does **not** cancel asyncio tasks or kill threads. Pair it with a
graceful timeout + `task.cancel()` (or process exit) if you need a hard
upper bound — see `test_drain_then_cancel_after_graceful_timeout` for
the recommended pattern.

## Usage

```python
from langgraph.runtime import RunControl
from langgraph.errors import GraphDrained

control = RunControl()

# In a signal handler, supervisor, etc.:
# control.request_drain("sigterm")

try:
    result = graph.invoke(input, config, control=control)
    if control.drain_requested:
        # finished naturally on the same tick where drain was requested
        ...
except GraphDrained as e:
    # checkpoint saved; resume later with the same config
    log.info("graph drained: %s", e.reason)
```

## Test plan

- [x] Sync + async drain stops the next superstep
(`test_run_control_request_drain_stops_future_steps[_async]`)
- [x] Drain on the terminal step finishes normally
(`test_drain_requested_in_terminal_step_finishes_normally[_async]`)
- [x] `durability=\"exit\"` persists a resumable checkpoint on drain
(`test_drain_with_exit_durability_persists_resume_checkpoint`)
- [x] Subgraph drain bubbles up and parent resumes correctly
(`test_drain_from_subgraph_can_resume_parent`)
- [x] External thread / task triggering drain mid-run
(`test_external_drain_concurrent_sync` / `_async`)
- [x] Drain + hard cancel after graceful timeout
(`test_drain_then_cancel_after_graceful_timeout`)
- [x] Functional API: in-flight `@task` futures still resolve after
`request_drain()`
(`test_request_drain_allows_inflight_[a]call_scheduling`)
- [x] `control` kwarg wired through `stream_v2`
(`test_stream_v2_accepts_control_for_drain`)
- [x] `Runtime.merge` preserves `control`
(`test_merge_runtime_preserves_run_control`)

---------

Co-authored-by: Quanzheng Long <long@langchain.dev>
Co-authored-by: Will Fu-Hinthorn <will@langchain.dev>
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-29 15:23:31 -07:00
Sydney RunkleandGitHub f95d2309f9 release(checkpoint-postgres): pin to checkpoint 4.1.0a1 (#7648)
## Summary
- Bumps `langgraph-checkpoint-postgres`'s pin on `langgraph-checkpoint`
from `>=4.0.3,<5.0.0` to `>=4.1.0a1,<5.0.0` so the postgres alpha
(3.1.0a1) requires the matching checkpoint alpha released in #7647.
- Mirrors the same pin update applied to `langgraph` (1.2.0a1) on
`wfh/releases/timers`.

## Why
The three alphas (checkpoint 4.1.0a1, checkpoint-postgres 3.1.0a1,
langgraph 1.2.0a1) are meant to be tested as a coherent set. `langgraph`
already pins `>=4.1.0a1`; checkpoint-postgres was missed in that release
and is still letting resolvers fall back to 4.0.x.

## Test plan
- [ ] CI green
- [ ] `uv lock` resolves cleanly across libs (verified locally via `make
lock` — no lock file changes since editable paths already resolved to
4.1.0a1)
2026-04-29 18:00:47 -04:00
William FHGitHubWill Fu-HinthornSydney Runklecopilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
4a5765dd23 release: alpha for timers (#7647)
Co-authored-by: Will Fu-Hinthorn <will@langchain.dev>
Co-authored-by: Sydney Runkle <sydneymarierunkle@gmail.com>
Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
2026-04-29 17:49:01 -04:00
5c18bde0f8 feat(langgraph): DeltaChannel: store sentinel in blobs, reconstruct from checkpoint_writes (#7586)
# DeltaChannel: sentinel-based checkpoint blobs + write-replay
reconstruction

## Summary

`DeltaChannel` is a new fold-reducer channel that stores only a
zero-byte sentinel in checkpoint blobs instead of the full accumulated
value. On restore, the runtime replays ancestor writes through the
reducer to reconstruct state. For long-running threads with large
accumulating state (e.g. message histories), this delivers dramatically
smaller checkpoint blobs with configurable read-depth bounds.

```python
from typing import Annotated
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import _messages_delta_reducer

class State(TypedDict):
    # blob per step: ~60 bytes (sentinel) instead of growing full list
    messages: Annotated[list, DeltaChannel(_messages_delta_reducer)]
    # bound read depth to 10 steps via periodic snapshots
    messages_bounded: Annotated[list, DeltaChannel(_messages_delta_reducer, snapshot_frequency=10)]
```

---

## Storage benchmarks (InMemory, ~400 char/msg)

**Messages blob storage** (`checkpoint_blobs` bytes for the messages
channel):

| turns | add\_messages | delta(inf) | delta(freq=50) | delta(freq=10) |
delta(freq=5) |

|------:|-------------:|-----------:|---------------:|---------------:|--------------:|
| 10 | 91.0 KB | 60 B (1517x) | 60 B (1517x) | 14.4 KB (6x) | 32.6 KB
(3x) |
| 50 | 2.20 MB | 300 B (7347x) | 67.1 KB (33x) | 423 KB (5x) | 864 KB
(3x) |
| 100 | 8.78 MB | 600 B (14636x) | 310 KB (28x) | 1.72 MB (5x) | 3.48 MB
(3x) |
| 250 | 54.80 MB | 1.5 KB (36536x) | 2.09 MB (26x) | 10.87 MB (5x) |
21.84 MB (3x) |
| 500 | 219.19 MB | 3.0 KB (73063x) | 8.56 MB (26x) | 43.67 MB (5x) |
87.50 MB (3x) |

**Total checkpoint storage** (blobs + writes + metadata):

| turns | add\_messages | delta(inf) | delta(freq=50) | delta(freq=10) |
delta(freq=5) |

|------:|-------------:|-----------:|---------------:|---------------:|--------------:|
| 10 | 129.7 KB | 38.7 KB (3.4x) | 38.7 KB (3.4x) | 53.1 KB (2.4x) |
71.2 KB (1.8x) |
| 50 | 2.40 MB | 196 KB (12x) | 263 KB (9x) | 620 KB (3.9x) | 1.06 MB
(2.3x) |
| 100 | 9.18 MB | 394 KB (23x) | 703 KB (13x) | 2.12 MB (4.3x) | 3.87 MB
(2.4x) |
| 250 | 55.79 MB | 987 KB (57x) | 3.07 MB (18x) | 11.86 MB (4.7x) |
22.82 MB (2.4x) |
| 500 | 221.16 MB | 1.98 MB (112x) | 10.53 MB (21x) | 45.64 MB (4.9x) |
89.48 MB (2.5x) |

**Write-phase peak heap**:

| turns | add\_messages | delta(inf) | delta(freq=50) | delta(freq=10) |
delta(freq=5) |

|------:|-------------:|-----------:|---------------:|---------------:|--------------:|
| 10 | 456 KB | 199 KB (2.3x) | 199 KB (2.3x) | 212 KB (2.2x) | 232 KB
(2.0x) |
| 50 | 3.04 MB | 742 KB (4.1x) | 805 KB (3.8x) | 1.21 MB (2.5x) | 1.67
MB (1.8x) |
| 100 | 10.70 MB | 1.41 MB (7.6x) | 1.82 MB (5.9x) | 3.42 MB (3.1x) |
5.25 MB (2.0x) |
| 250 | 60.44 MB | 3.36 MB (18x) | 5.67 MB (11x) | 14.87 MB (4.1x) |
26.31 MB (2.3x) |

**Read-phase avg `get_state` latency** (5 calls, InMemory):

| turns | add\_messages | delta(inf) | delta(freq=50) | delta(freq=10) |
delta(freq=5) |

|------:|-------------:|-----------:|---------------:|---------------:|--------------:|
| 10 | 0.7 ms | 1.1 ms (0.6x) | 1.1 ms (0.6x) | 0.8 ms (0.9x) | 0.6 ms
(1.1x) |
| 50 | 2.7 ms | 5.3 ms (0.5x) | 3.5 ms (0.8x) | 2.7 ms (1.0x) | 2.7 ms
(1.0x) |
| 100 | 5.5 ms | 11.1 ms (0.5x) | 6.0 ms (0.9x) | 5.2 ms (1.1x) | 5.4 ms
(1.0x) |
| 250 | 12.9 ms | 27.2 ms (0.5x) | 13.6 ms (0.9x) | 12.9 ms (1.0x) |
13.0 ms (1.0x) |

**Postgres `get_tuple` read latency** (~100 tok/msg per step):

| steps | full-list | delta(inf) | delta(freq=50) | delta(freq=10) |
delta(freq=5) |

|------:|----------:|-----------:|---------------:|---------------:|--------------:|
| 10 | 0.29 ms | 0.21 ms (1.4x) | 0.19 ms (1.6x) | 0.19 ms (1.5x) | 0.19
ms (1.6x) |
| 50 | 0.19 ms | 0.15 ms (1.3x) | 0.19 ms (1.0x) | 0.22 ms (0.8x) | 0.29
ms (0.7x) |
| 100 | 0.27 ms | 0.17 ms (1.6x) | 0.22 ms (1.2x) | 0.23 ms (1.2x) |
0.21 ms (1.3x) |
| 500 | 0.60 ms | 0.30 ms (2.0x) | 0.66 ms (0.9x) | 0.56 ms (1.1x) |
0.69 ms (0.9x) |

**Takeaway:** `snapshot_frequency=10` matches full-list read latency
while still saving 5x on blob storage and ~4x on total storage.

---

## How it works

### Checkpoint blobs

`checkpoint()` always returns `DELTA_SENTINEL` (a zero-byte msgpack ext
marker) instead of the accumulated value. On restore, the saver's
`_get_channel_writes_history` walks the ancestor chain collecting
`checkpoint_writes` entries and replays them through the reducer:

```python
# blob stored per step: ~1 byte (sentinel)
# vs. full list growing O(N) every step with BinaryOperatorAggregate
```

### Reducer interface

`DeltaChannel` takes a **batch reducer** `(state, list[writes]) ->
state` — all writes for a step arrive in one call, enabling single-pass
implementations:

```python
#  Don't use add_messages directly — it's a binary operator, not a batch reducer
messages: Annotated[list, DeltaChannel(add_messages)]  # wrong

#  Use _messages_delta_reducer — single pass, dedup by ID, RemoveMessage support
messages: Annotated[list, DeltaChannel(_messages_delta_reducer)]

#  Or write your own batch reducer for custom types
def my_dict_reducer(state: dict, writes: list[dict]) -> dict:
    result = dict(state)
    for w in writes:
        result.update(w)
    return result

files: Annotated[dict, DeltaChannel(my_dict_reducer)]
```

### Snapshot frequency

`snapshot_frequency=N` writes a full `_DeltaSnapshot` blob every N
pregel steps, bounding replay depth regardless of thread length.
Snapshots are eager — written even if the channel had no update that
step, so the depth bound always holds:

```python
# Replay walks at most 10 ancestors before hitting a snapshot
messages: Annotated[list, DeltaChannel(_messages_delta_reducer, snapshot_frequency=10)]
```

### Migration from `BinaryOperatorAggregate`

Pre-existing threads written under `BinaryOperatorAggregate` work
transparently after swapping the annotation — the saver detects a
plain-value ancestor blob and uses it as the reconstruction seed:

```python
# Before: BinaryOperatorAggregate stores full list every step
items: Annotated[list, add_messages]

# After: DeltaChannel — existing checkpoints still readable, new steps use sentinel
items: Annotated[list, DeltaChannel(_messages_delta_reducer)]
```

### Async write-ordering safety

In `durability="async"` mode (default), `put_writes` calls are
fire-and-forget. `AsyncPregelLoop` tracks in-flight `aput_writes`
futures for DeltaChannel channels in `_delta_write_futs` and drains them
via `await asyncio.gather()` in `_checkpointer_put_after_previous`
before `aput()` — ensuring `checkpoint_writes` are durable before the
sentinel blob is committed.

---

## What's in scope

- **`libs/langgraph/langgraph/channels/delta.py`** — `DeltaChannel`
implementation
- **`libs/langgraph/langgraph/graph/message.py`** —
`_messages_delta_reducer` (experimental)
- **`libs/checkpoint/`** — `_get_channel_writes_history` ancestor-walk
API on `BaseCheckpointSaver`, `InMemorySaver` optimized override
- **`libs/checkpoint-postgres/`** — `PostgresSaver` /
`AsyncPostgresSaver` single-roundtrip UNION ALL override
- **`libs/langgraph/langgraph/pregel/`** — `channels_from_checkpoint` /
`create_checkpoint` wiring, async write-ordering safety

---

## Follow-ups

- **Batch reconstruction**: each DeltaChannel field issues its own
`_get_channel_writes_history` call; a single walk collecting all
sentinel channels would reduce roundtrips proportionally to the number
of DeltaChannel fields.
- **Sync write ordering**: `BackgroundExecutor.__exit__` guarantees
completion before `invoke()` returns, but within a run there's no
explicit ordering between `put_writes` and `put`. Two-phase commit for
sync would close this gap.
- **`ShallowPostgresSaver` compatibility**: shallow savers keep only the
latest checkpoint and have no parent chain to walk; DeltaChannel is
currently incompatible and should raise or warn at compile time.
- Updating the writes table w/ delta epoch ids for more efficient reads
- follow up w/ LSD checkpointer implementations to support delta
channel! and update prune

---------

Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-authored-by: ccurme <chester.curme@gmail.com>
Co-authored-by: Will Fu-Hinthorn <will@langchain.dev>
2026-04-29 17:26:17 -04:00
47 changed files with 2651 additions and 295 deletions
@@ -19,6 +19,7 @@ from langgraph.checkpoint.base import (
get_serializable_checkpoint_metadata,
)
from langgraph.checkpoint.serde.base import SerializerProtocol
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from psycopg import Capabilities, Connection, Cursor, Pipeline
from psycopg.rows import DictRow, dict_row
from psycopg.types.json import Jsonb
@@ -26,9 +27,11 @@ from psycopg_pool import ConnectionPool
from langgraph.checkpoint.postgres import _internal
from langgraph.checkpoint.postgres.base import (
SELECT_DELTA_COMBINED_SQL,
SELECT_DELTA_STAGE1_SQL,
SELECT_DELTA_STAGE2_SQL,
BasePostgresSaver,
_DeltaCombinedRow,
_DeltaStage1Row,
_DeltaStage2Row,
)
from langgraph.checkpoint.postgres.shallow import ShallowPostgresSaver
@@ -308,7 +311,12 @@ class PostgresSaver(BasePostgresSaver):
# others are stored in blobs table
blob_values = {}
for k, v in checkpoint["channel_values"].items():
if v is None or isinstance(v, (str, int, float, bool)):
if v is DELTA_SENTINEL:
copy["channel_values"].pop(k)
elif isinstance(v, _DeltaSnapshot):
blob_values[k] = copy["channel_values"].pop(k)
copy["channel_values"][k] = True
elif v is None or isinstance(v, (str, int, float, bool)):
pass
else:
blob_values[k] = copy["channel_values"].pop(k)
@@ -441,41 +449,49 @@ class PostgresSaver(BasePostgresSaver):
) -> _ChannelWritesHistory:
"""Fast-path override of `BaseCheckpointSaver._get_channel_writes_history`.
One combined UNION ALL query (`SELECT_DELTA_COMBINED_SQL`) fetches rows
from `checkpoints`, `checkpoint_writes`, and `checkpoint_blobs` in a
single roundtrip; the ancestor walk runs in Python.
Two-stage query: stage 1 scans checkpoint metadata to walk the parent
chain and locate the nearest snapshot; stage 2 fetches only the
chain-limited writes and single seed blob.
"""
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = get_checkpoint_id(config)
if checkpoint_id is None:
# Caller didn't specify a target — resolve to the latest
# checkpoint on the thread. `get_tuple` without `checkpoint_id`
# returns the newest; its config carries the resolved id.
target = self.get_tuple(config)
if target is None:
return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=[])
checkpoint_id = target.config["configurable"]["checkpoint_id"]
with self._cursor() as cur:
cur.execute(
SELECT_DELTA_COMBINED_SQL,
SELECT_DELTA_STAGE1_SQL,
(channel, channel, thread_id, checkpoint_ns),
)
stage1_rows = cur.fetchall()
chain_cids, seed_version = self._walk_stage1(
cast("list[_DeltaStage1Row]", stage1_rows), checkpoint_id
)
seed_versions = [seed_version] if seed_version else []
with self._cursor() as cur:
cur.execute(
SELECT_DELTA_STAGE2_SQL,
(
channel,
thread_id,
checkpoint_ns,
thread_id,
checkpoint_ns,
channel,
chain_cids,
thread_id,
checkpoint_ns,
channel,
seed_versions,
),
)
rows = cur.fetchall()
stage2_rows = cur.fetchall()
return self._build_delta_channel_writes_history(
channel=channel,
target_id=checkpoint_id,
rows=cast("list[_DeltaCombinedRow]", rows),
chain_cids=chain_cids,
seed_version=seed_version,
stage2_rows=cast("list[_DeltaStage2Row]", stage2_rows),
)
def _load_checkpoint_tuple(self, value: DictRow) -> CheckpointTuple:
@@ -19,6 +19,7 @@ from langgraph.checkpoint.base import (
get_serializable_checkpoint_metadata,
)
from langgraph.checkpoint.serde.base import SerializerProtocol
from langgraph.checkpoint.serde.types import _DeltaSnapshot
from psycopg import AsyncConnection, AsyncCursor, AsyncPipeline, Capabilities
from psycopg.rows import DictRow, dict_row
from psycopg.types.json import Jsonb
@@ -26,9 +27,11 @@ from psycopg_pool import AsyncConnectionPool
from langgraph.checkpoint.postgres import _ainternal
from langgraph.checkpoint.postgres.base import (
SELECT_DELTA_COMBINED_SQL,
SELECT_DELTA_STAGE1_SQL,
SELECT_DELTA_STAGE2_SQL,
BasePostgresSaver,
_DeltaCombinedRow,
_DeltaStage1Row,
_DeltaStage2Row,
)
from langgraph.checkpoint.postgres.shallow import AsyncShallowPostgresSaver
@@ -267,7 +270,12 @@ class AsyncPostgresSaver(BasePostgresSaver):
# others are stored in blobs table
blob_values = {}
for k, v in checkpoint["channel_values"].items():
if v is None or isinstance(v, (str, int, float, bool)):
if v is DELTA_SENTINEL:
copy["channel_values"].pop(k)
elif isinstance(v, _DeltaSnapshot):
blob_values[k] = copy["channel_values"].pop(k)
copy["channel_values"][k] = True
elif v is None or isinstance(v, (str, int, float, bool)):
pass
else:
blob_values[k] = copy["channel_values"].pop(k)
@@ -402,10 +410,9 @@ class AsyncPostgresSaver(BasePostgresSaver):
) -> _ChannelWritesHistory:
"""Fast-path override of `BaseCheckpointSaver._aget_channel_writes_history`.
One combined UNION ALL query (`SELECT_DELTA_COMBINED_SQL`) fetches rows
from `checkpoints`, `checkpoint_writes`, and `checkpoint_blobs` in a
single roundtrip; rows are assembled by the shared pure helper on
`BasePostgresSaver`.
Two-stage query: stage 1 scans checkpoint metadata to walk the parent
chain and locate the nearest snapshot; stage 2 fetches only the
chain-limited writes and single seed blob.
"""
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
@@ -415,26 +422,37 @@ class AsyncPostgresSaver(BasePostgresSaver):
if target is None:
return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=[])
checkpoint_id = target.config["configurable"]["checkpoint_id"]
async with self._cursor() as cur:
await cur.execute(
SELECT_DELTA_COMBINED_SQL,
SELECT_DELTA_STAGE1_SQL,
(channel, channel, thread_id, checkpoint_ns),
)
stage1_rows = await cur.fetchall()
chain_cids, seed_version = self._walk_stage1(
cast("list[_DeltaStage1Row]", stage1_rows), checkpoint_id
)
seed_versions = [seed_version] if seed_version else []
async with self._cursor() as cur:
await cur.execute(
SELECT_DELTA_STAGE2_SQL,
(
channel,
thread_id,
checkpoint_ns,
thread_id,
checkpoint_ns,
channel,
chain_cids,
thread_id,
checkpoint_ns,
channel,
seed_versions,
),
)
rows = await cur.fetchall()
stage2_rows = await cur.fetchall()
return self._build_delta_channel_writes_history(
channel=channel,
target_id=checkpoint_id,
rows=cast("list[_DeltaCombinedRow]", rows),
chain_cids=chain_cids,
seed_version=seed_version,
stage2_rows=cast("list[_DeltaStage2Row]", stage2_rows),
)
async def _load_checkpoint_tuple(self, value: DictRow) -> CheckpointTuple:
@@ -156,62 +156,62 @@ INSERT_CHECKPOINT_WRITES_SQL = """
"""
class _DeltaCombinedRow(TypedDict, total=False):
"""One row from `SELECT_DELTA_COMBINED_SQL` (a UNION ALL of three tables).
class _DeltaStage2Row(TypedDict, total=False):
"""One row from `SELECT_DELTA_STAGE2_SQL` (a UNION ALL of writes and blobs)."""
Every row carries `_kind` ("p" / "w" / "b") plus whichever columns are
relevant for that kind; irrelevant columns are NULL and typed as `None`.
"""
_kind: str # always present: "p", "w", or "b"
# checkpoint row ("p")
checkpoint_id: str | None
parent_checkpoint_id: str | None
ver: str | None
# write / blob rows ("w", "b")
_kind: str # "w" or "b"
checkpoint_id: str | None # "w" rows only
type: str | None
blob: bytes | None
# write row only ("w")
task_id: str | None
idx: int | None
# blob row only ("b")
version: str | None
task_id: str | None # "w" rows only
idx: int | None # "w" rows only
version: str | None # "b" rows only
# DeltaChannel reconstruction: one UNION ALL query fetches checkpoints,
# writes, and blobs for `channel` in one roundtrip; the ancestor walk runs
# in Python in `_build_delta_channel_writes_history`.
# Two-stage DeltaChannel reconstruction. Stage 1 scans checkpoint
# metadata (no blob bytes) to walk the parent chain and locate the
# nearest snapshot marker. Stage 2 fetches only the chain-limited
# writes and the single seed snapshot blob.
#
# Parameter order: (channel, thread_id, checkpoint_ns,
# thread_id, checkpoint_ns, channel,
# thread_id, checkpoint_ns, channel)
SELECT_DELTA_COMBINED_SQL = """
SELECT 'p'::text AS _kind,
checkpoint_id,
# Parameter order:
# stage1: (channel, channel, thread_id, checkpoint_ns)
# stage2: (thread_id, checkpoint_ns, channel, chain_cids[],
# thread_id, checkpoint_ns, channel, seed_versions[])
SELECT_DELTA_STAGE1_SQL = """
SELECT checkpoint_id,
parent_checkpoint_id,
checkpoint -> 'channel_versions' ->> %s AS ver,
NULL::text AS type,
NULL::bytea AS blob,
NULL::text AS task_id,
NULL::int AS idx,
NULL::text AS version
(checkpoint -> 'channel_values' -> %s) IS NOT NULL AS has_snapshot
FROM checkpoints
WHERE thread_id = %s AND checkpoint_ns = %s
UNION ALL
SELECT 'w',
checkpoint_id, NULL, NULL,
type, blob, task_id, idx, NULL
"""
SELECT_DELTA_STAGE2_SQL = """
SELECT 'w'::text AS _kind,
checkpoint_id,
type, blob, task_id, idx, NULL::text AS version
FROM checkpoint_writes
WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s
AND checkpoint_id = ANY(%s)
UNION ALL
SELECT 'b',
NULL, NULL, NULL,
SELECT 'b', NULL,
type, blob, NULL, NULL, version
FROM checkpoint_blobs
WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s
AND version = ANY(%s)
"""
class _DeltaStage1Row(TypedDict):
"""One row from `SELECT_DELTA_STAGE1_SQL`."""
checkpoint_id: str
parent_checkpoint_id: str | None
ver: str | None
has_snapshot: bool
class BasePostgresSaver(BaseCheckpointSaver[str]):
SELECT_SQL = SELECT_SQL
SELECT_PENDING_SENDS_SQL = SELECT_PENDING_SENDS_SQL
@@ -254,38 +254,59 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
if t.decode() != "empty"
}
@staticmethod
def _walk_stage1(
stage1_rows: Sequence[_DeltaStage1Row],
target_id: str,
) -> tuple[list[str], str | None]:
"""Walk the parent chain from stage 1 metadata rows.
Returns (chain_cids, seed_version):
chain_cids: ancestor checkpoint IDs from target's parent down to
the seed (or root), in newest-first order.
seed_version: the channel blob version at the nearest ancestor
with has_snapshot=True, or None if pure delta.
"""
parent_of: dict[str, str | None] = {}
ver_of: dict[str, str | None] = {}
snapshot_of: dict[str, bool] = {}
for r in stage1_rows:
cid = r["checkpoint_id"]
parent_of[cid] = r["parent_checkpoint_id"]
ver_of[cid] = r["ver"]
snapshot_of[cid] = r["has_snapshot"]
chain_cids: list[str] = []
seed_version: str | None = None
cur_cid: str | None = parent_of.get(target_id)
while cur_cid is not None:
chain_cids.append(cur_cid)
if snapshot_of.get(cur_cid, False):
seed_version = ver_of.get(cur_cid)
break
cur_cid = parent_of.get(cur_cid)
return chain_cids, seed_version
def _build_delta_channel_writes_history(
self,
*,
channel: str,
target_id: str,
rows: Sequence[_DeltaCombinedRow],
chain_cids: list[str],
seed_version: str | None,
stage2_rows: Sequence[_DeltaStage2Row],
) -> _ChannelWritesHistory:
"""Reconstruct one delta channel's history from the combined UNION ALL rows.
"""Reconstruct delta channel history from two-stage query results.
Pure data transform shared by sync (`PostgresSaver`) and async
(`AsyncPostgresSaver`); both paths run `SELECT_DELTA_COMBINED_SQL`
and feed the tagged rows here.
Walk is newest → oldest from the target's parent. A non-sentinel
blob in `checkpoint_blobs` (a pre-delta snapshot) terminates the
walk and is returned as the seed so replay starts from it.
Writes stored at `target_id` itself are pending writes for the next
step and are excluded — the walk begins at the target's parent.
chain_cids are in newest-first order (target's parent first).
stage2_rows contain only writes for chain_cids and the single
seed blob at seed_version.
"""
parent_of: dict[str, str | None] = {}
ver_of: dict[str, str | None] = {}
writes_by_cid: dict[str, list[tuple[str, bytes, str, int]]] = {}
blob_by_ver: dict[str, tuple[str, bytes]] = {}
seed_blob: tuple[str, bytes] | None = None
for r in rows:
for r in stage2_rows:
kind = r["_kind"]
if kind == "p":
cid = cast(str, r["checkpoint_id"])
parent_of[cid] = r["parent_checkpoint_id"]
ver_of[cid] = r["ver"]
elif kind == "w":
if kind == "w":
cid = cast(str, r["checkpoint_id"])
writes_by_cid.setdefault(cid, []).append(
cast(
@@ -294,42 +315,26 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
)
)
else: # kind == "b"
blob_by_ver[cast(str, r["version"])] = cast(
"tuple[str, bytes]", (r["type"], r["blob"])
)
seed_blob = cast("tuple[str, bytes]", (r["type"], r["blob"]))
# newest write first per ancestor (task_id DESC, idx DESC)
for ws in writes_by_cid.values():
ws.sort(key=lambda w: (w[2], w[3]), reverse=True)
ancestors: list[str] = []
cur_cid: str | None = parent_of.get(target_id)
while cur_cid is not None:
ancestors.append(cur_cid)
cur_cid = parent_of.get(cur_cid)
if not ancestors:
if not chain_cids:
return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=[])
collected: list[PendingWrite] = [] # newest first; reversed at the end
for cid in ancestors:
# Collect writes first — they encode the transition FROM this
# ancestor's state to its child's and must be included even if
# this ancestor is also the seed checkpoint.
collected: list[PendingWrite] = []
for cid in chain_cids:
for type_tag, write_blob, task_id, _idx in writes_by_cid.get(cid, []):
val = self.serde.loads_typed((type_tag, write_blob))
collected.append((task_id, channel, val))
# Then check seed terminator.
ver = ver_of.get(cid)
if ver is not None:
seed_blob = blob_by_ver.get(ver)
if seed_blob is not None and seed_blob[0] != "empty":
blob_value = self.serde.loads_typed(seed_blob)
if blob_value is not DELTA_SENTINEL:
collected.reverse()
return _ChannelWritesHistory(seed=blob_value, writes=collected)
collected.reverse() # oldest → newest
return _ChannelWritesHistory(seed=DELTA_SENTINEL, writes=collected)
seed: Any = DELTA_SENTINEL
if seed_blob is not None and seed_blob[0] != "empty":
seed = self.serde.loads_typed(seed_blob)
collected.reverse()
return _ChannelWritesHistory(seed=seed, writes=collected)
def _dump_blobs(
self,
+2 -2
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-checkpoint-postgres"
version = "3.0.5"
version = "3.1.0a3"
description = "Library with a Postgres implementation of LangGraph checkpoint saver."
authors = []
requires-python = ">=3.10"
@@ -12,7 +12,7 @@ readme = "README.md"
license = "MIT"
license-files = ['LICENSE']
dependencies = [
"langgraph-checkpoint>=4.0.3,<5.0.0",
"langgraph-checkpoint>=4.1.0a3,<5.0.0",
"orjson>=3.11.5",
"psycopg>=3.2.0",
"psycopg-pool>=3.2.0",
+2 -2
View File
@@ -259,7 +259,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.3"
version = "4.1.0a3"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -307,7 +307,7 @@ test = [
[[package]]
name = "langgraph-checkpoint-postgres"
version = "3.0.5"
version = "3.1.0a3"
source = { editable = "." }
dependencies = [
{ name = "langgraph-checkpoint" },
+1 -1
View File
@@ -268,7 +268,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.3"
version = "4.1.0a3"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -452,7 +452,9 @@ class InMemorySaver(
values: dict[str, Any] = c.pop("channel_values") # type: ignore[misc]
for k, v in new_versions.items():
self.blobs[(thread_id, checkpoint_ns, k, v)] = (
self.serde.dumps_typed(values[k]) if k in values else ("empty", b"")
self.serde.dumps_typed(values[k])
if k in values and values[k] is not DELTA_SENTINEL
else ("empty", b"")
)
self.storage[thread_id][checkpoint_ns].update(
{
@@ -34,9 +34,7 @@ from langgraph.checkpoint.serde import _msgpack as _lg_msgpack
from langgraph.checkpoint.serde.base import SerializerProtocol
from langgraph.checkpoint.serde.event_hooks import emit_serde_event
from langgraph.checkpoint.serde.types import (
DELTA_SENTINEL,
SendProtocol,
_DeltaSentinel,
_DeltaSnapshot,
)
from langgraph.store.base import Item
@@ -322,14 +320,11 @@ EXT_PYDANTIC_V1 = 4
EXT_PYDANTIC_V2 = 5
EXT_NUMPY_ARRAY = 6
EXT_DELTA_SNAPSHOT = 7
EXT_DELTA_SENTINEL = 8
def _msgpack_default(obj: Any) -> str | ormsgpack.Ext:
if isinstance(obj, _DeltaSnapshot):
return ormsgpack.Ext(EXT_DELTA_SNAPSHOT, _msgpack_enc(obj.value))
elif isinstance(obj, _DeltaSentinel):
return ormsgpack.Ext(EXT_DELTA_SENTINEL, b"")
elif hasattr(obj, "model_dump") and callable(obj.model_dump): # pydantic v2
return ormsgpack.Ext(
EXT_PYDANTIC_V2,
@@ -656,9 +651,7 @@ def _create_msgpack_ext_hook(
return False
def ext_hook(code: int, data: bytes) -> Any:
if code == EXT_DELTA_SENTINEL:
return DELTA_SENTINEL
elif code == EXT_DELTA_SNAPSHOT:
if code == EXT_DELTA_SNAPSHOT:
return _DeltaSnapshot(
ormsgpack.unpackb(
data, ext_hook=ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
@@ -17,12 +17,10 @@ TASKS = "__pregel_tasks"
class _DeltaSentinel:
"""Singleton marker stored (as zero bytes) in checkpoint_blobs for a
DeltaChannel field. The actual per-step writes live in checkpoint_writes
and are replayed through the reducer at load time.
"""In-memory marker for a DeltaChannel field with no snapshot.
Compare with `is DELTA_SENTINEL` — `loads_typed` always returns the same
module-level instance.
Never serialized to storage — checkpointers strip it before writing.
Compare with `is DELTA_SENTINEL`; always the same module-level instance.
"""
__slots__ = ()
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-checkpoint"
version = "4.0.3"
version = "4.1.0a3"
description = "Library with base interfaces for LangGraph checkpoint savers."
authors = []
requires-python = ">=3.10"
-12
View File
@@ -1048,15 +1048,3 @@ def test_msgpack_nested_pydantic_serializes_as_dict(
# No blocking should occur - inner is serialized as dict, not ext
assert "blocked" not in caplog.text.lower()
assert result == obj
def test_delta_sentinel_serde_round_trip() -> None:
from langgraph.checkpoint.base import DELTA_SENTINEL
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
serde = JsonPlusSerializer()
type_tag, blob = serde.dumps_typed(DELTA_SENTINEL)
assert type_tag == "msgpack"
assert blob # non-empty ext envelope
loaded = serde.loads_typed((type_tag, blob))
assert loaded is DELTA_SENTINEL
+7 -16
View File
@@ -322,26 +322,17 @@ def test_memory_saver_with_allowlist_proxy_isolated() -> None:
class TestInMemorySaverDeltaChannel:
def test_load_blobs_returns_sentinel_for_delta_channel(self) -> None:
"""_load_blobs returns DELTA_SENTINEL for delta channels (reconstruction deferred)."""
def test_load_blobs_omits_delta_channel(self) -> None:
"""_load_blobs omits delta channels (stored as 'empty'); reconstruction deferred."""
saver = InMemorySaver()
serde = JsonPlusSerializer()
thread_id, ns, channel = "t1", "", "messages"
v1 = "00000000000000000000000000000001.0000000000000000"
saver.blobs[(thread_id, ns, channel, v1)] = serde.dumps_typed(DELTA_SENTINEL)
cp1 = empty_checkpoint()
cp1["id"] = "cp1"
cp1["channel_versions"][channel] = v1
saver.storage[thread_id][ns] = {
"cp1": (serde.dumps_typed(cp1), serde.dumps_typed({}), None),
}
saver.blobs[(thread_id, ns, channel, v1)] = ("empty", b"")
result = saver._load_blobs(thread_id, ns, {channel: v1})
assert channel in result
assert result[channel] is DELTA_SENTINEL
assert channel not in result
def test_get_channel_writes_collects_ancestor_writes_only(self) -> None:
"""_get_channel_writes_history collects ancestor writes oldest→newest,
@@ -582,9 +573,9 @@ class TestPreDeltaBlobTerminator:
# Pre-delta: cp1 stored a real blob for the channel.
saver.blobs[(thread_id, ns, channel, v1)] = serde.dumps_typed(["A"])
# Delta-era: cp2 and cp3 store sentinels; real writes in checkpoint_writes.
saver.blobs[(thread_id, ns, channel, v2)] = serde.dumps_typed(DELTA_SENTINEL)
saver.blobs[(thread_id, ns, channel, v3)] = serde.dumps_typed(DELTA_SENTINEL)
# Delta-era: cp2 and cp3 store "empty"; real writes in checkpoint_writes.
saver.blobs[(thread_id, ns, channel, v2)] = ("empty", b"")
saver.blobs[(thread_id, ns, channel, v3)] = ("empty", b"")
cp1 = empty_checkpoint()
cp1["id"] = "cp1"
+1 -1
View File
@@ -286,7 +286,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.3"
version = "4.1.0a3"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
@@ -12,6 +12,9 @@ RESUME = sys.intern("__resume__")
# for values passed to resume a node after an interrupt
ERROR = sys.intern("__error__")
# for errors raised by nodes
ERROR_SOURCE_NODE = sys.intern("__error_source_node__")
# failed source node name for node-level error handlers
# value format in pending writes: `(task_id, ERROR_SOURCE_NODE, node_name: str)`
NO_WRITES = sys.intern("__no_writes__")
# marker to signal node didn't write anything
TASKS = sys.intern("__pregel_tasks")
@@ -71,6 +74,10 @@ CONFIG_KEY_RESUME_MAP = sys.intern("__pregel_resume_map")
CONFIG_KEY_STREAM_MESSAGES_V2 = sys.intern("__pregel_stream_messages_v2")
# when True, attach StreamMessagesHandlerV2 so content-block (v2) events
# flow through stream_mode="messages"; set by StreamingHandler only.
CONFIG_KEY_NODE_ERROR = sys.intern("__pregel_node_error")
# holds a `NodeError` (failed source node + exception) for the current
# node-level error handler invocation, injected when handler signature
# requests `error: NodeError`
# --- Other constants ---
PUSH = sys.intern("__pregel_push")
@@ -98,6 +105,7 @@ RESERVED = {
INTERRUPT,
RESUME,
ERROR,
ERROR_SOURCE_NODE,
NO_WRITES,
# reserved config.configurable keys
CONFIG_KEY_SEND,
@@ -51,9 +51,11 @@ from langgraph._internal._config import (
)
from langgraph._internal._constants import (
CONF,
CONFIG_KEY_NODE_ERROR,
CONFIG_KEY_RUNTIME,
)
from langgraph._internal._typing import MISSING
from langgraph.errors import NodeError
from langgraph.types import StreamWriter
try:
@@ -194,6 +196,15 @@ KWARGS_CONFIG_KEYS: tuple[tuple[str, tuple[Any, ...], str, Any], ...] = (
"N/A",
inspect.Parameter.empty,
),
(
"error",
(NodeError, "NodeError"),
# we never hit this block, we read directly from configurable
"N/A",
# default to None so non-handler nodes that happen to type a parameter
# `error: NodeError` don't blow up; handlers always receive a NodeError.
None,
),
)
"""List of kwargs that can be passed to functions, and their corresponding
config keys, default values and type annotations.
@@ -367,6 +378,8 @@ class RunnableCallable(Runnable):
kw_value: Any = MISSING
if kw == "config":
kw_value = config
elif kw == "error":
kw_value = config.get(CONF, {}).get(CONFIG_KEY_NODE_ERROR, MISSING)
elif runtime:
if kw == "runtime":
kw_value = runtime
@@ -439,6 +452,8 @@ class RunnableCallable(Runnable):
kw_value: Any = MISSING
if kw == "config":
kw_value = config
elif kw == "error":
kw_value = config.get(CONF, {}).get(CONFIG_KEY_NODE_ERROR, MISSING)
elif runtime:
if kw == "runtime":
kw_value = runtime
+43 -8
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
from collections.abc import Sequence
from dataclasses import dataclass
from enum import Enum
from typing import Any, Literal
from warnings import warn
@@ -15,10 +16,12 @@ from langgraph.warnings import LangGraphDeprecatedSinceV10
__all__ = (
"EmptyChannelError",
"ErrorCode",
"GraphDrained",
"GraphRecursionError",
"InvalidUpdateError",
"GraphBubbleUp",
"GraphInterrupt",
"NodeError",
"NodeInterrupt",
"NodeTimeoutError",
"ParentCommand",
@@ -43,6 +46,23 @@ def create_error_message(*, message: str, error_code: ErrorCode) -> str:
)
class GraphBubbleUp(Exception):
pass
class GraphDrained(GraphBubbleUp):
"""Raised when a graph run exits early due to a drain request.
This indicates the graph stopped cooperatively at a superstep boundary
because `RunControl.request_drain()` was called (e.g., in response to
SIGTERM). The checkpoint is saved and the run can be resumed later.
"""
def __init__(self, reason: str = "shutdown") -> None:
self.reason = reason
super().__init__(f"Graph drained: {reason}")
class GraphRecursionError(RecursionError):
"""Raised when the graph has exhausted the maximum number of steps.
@@ -78,10 +98,6 @@ class InvalidUpdateError(Exception):
pass
class GraphBubbleUp(Exception):
pass
class GraphInterrupt(GraphBubbleUp):
"""Raised when a subgraph is interrupted, suppressed by the root graph.
Never raised directly, or surfaced to the user."""
@@ -128,12 +144,31 @@ class TaskNotFound(Exception):
pass
class NodeTimeoutError(TimeoutError):
@dataclass(frozen=True, slots=True)
class NodeError:
"""Failure context passed to a node-level error handler.
Inject by adding a parameter typed `NodeError` to a handler registered via
`StateGraph.add_node(..., error_handler=...)`:
```python
def handler(state: State, error: NodeError) -> Command:
return Command(update={"status": f"recovered from {error.node}: {error.error}"})
```
"""
node: str
"""Name of the node whose execution failed."""
error: BaseException
"""Exception raised by the failed node."""
class NodeTimeoutError(Exception):
"""Raised when a node invocation exceeds one of its configured timeouts.
Subclasses the built-in `TimeoutError`, so existing `except TimeoutError`
handlers keep working. If the node has a `retry_policy` whose `retry_on`
permits `TimeoutError`, the attempt will be retried.
Does **not** inherit from the built-in `TimeoutError` (a subclass of
`OSError`) so that the default `RetryPolicy` treats it as retryable.
Both `idle_timeout` and `run_timeout` reflect the configured policy at the
time of the failure (each is `None` if not configured). `kind` and
+2
View File
@@ -88,6 +88,8 @@ class StateNodeSpec(Generic[NodeInputT, ContextT]):
input_schema: type[NodeInputT]
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None
cache_policy: CachePolicy | None
is_error_handler: bool = False
error_handler_node: str | None = None
ends: tuple[str, ...] | dict[str, str] | None = EMPTY_SEQ
defer: bool = False
timeout: TimeoutPolicy | None = None
+36 -1
View File
@@ -303,6 +303,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema: None = None,
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
error_handler: StateNode[Any, ContextT] | None = None,
destinations: dict[str, str] | tuple[str, ...] | None = None,
timeout: float | timedelta | TimeoutPolicy | None = None,
**kwargs: Unpack[DeprecatedKwargs],
@@ -371,6 +372,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema: type[NodeInputT],
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
error_handler: StateNode[Any, ContextT] | None = None,
destinations: dict[str, str] | tuple[str, ...] | None = None,
timeout: float | timedelta | TimeoutPolicy | None = None,
**kwargs: Unpack[DeprecatedKwargs],
@@ -444,6 +446,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema: None = None,
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
error_handler: StateNode[Any, ContextT] | None = None,
destinations: dict[str, str] | tuple[str, ...] | None = None,
timeout: float | timedelta | TimeoutPolicy | None = None,
**kwargs: Unpack[DeprecatedKwargs],
@@ -512,6 +515,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema: type[NodeInputT],
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
error_handler: StateNode[Any, ContextT] | None = None,
destinations: dict[str, str] | tuple[str, ...] | None = None,
timeout: float | timedelta | TimeoutPolicy | None = None,
**kwargs: Unpack[DeprecatedKwargs],
@@ -587,6 +591,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema: type[NodeInputT] | None = None,
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
error_handler: StateNode[Any, ContextT] | None = None,
destinations: dict[str, str] | tuple[str, ...] | None = None,
timeout: float | timedelta | TimeoutPolicy | None = None,
**kwargs: Unpack[DeprecatedKwargs],
@@ -607,6 +612,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
If a sequence is provided, the first matching policy will be applied.
cache_policy: The cache policy for the node.
error_handler: Optional node-level error handler callable for this node.
destinations: Destinations that indicate where a node can route to.
Useful for edgeless graphs with nodes that return `Command` objects.
@@ -766,6 +772,25 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
if destinations is not None:
ends = destinations
resolved_input_schema: type[Any] = (
input_schema or inferred_input_schema or self.state_schema
)
handler_node_name: str | None = None
if error_handler is not None:
handler_node_name = f"__error_handler__{node}"
if handler_node_name in self.nodes:
raise ValueError(
f"Auto-generated error handler node `{handler_node_name}` already exists."
)
self.nodes[handler_node_name] = StateNodeSpec[Any, ContextT](
coerce_to_runnable(error_handler, name=handler_node_name, trace=False), # type: ignore[arg-type]
metadata=None,
input_schema=resolved_input_schema,
retry_policy=None,
cache_policy=None,
is_error_handler=True,
)
if input_schema is not None:
self.nodes[node] = StateNodeSpec[NodeInputT, ContextT](
coerce_to_runnable(action, name=node, trace=False), # type: ignore[arg-type]
@@ -773,6 +798,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema=input_schema,
retry_policy=retry_policy,
cache_policy=cache_policy,
error_handler_node=handler_node_name,
ends=ends,
defer=defer,
timeout=timeout,
@@ -784,6 +810,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema=inferred_input_schema,
retry_policy=retry_policy,
cache_policy=cache_policy,
error_handler_node=handler_node_name,
ends=ends,
defer=defer,
timeout=timeout,
@@ -795,6 +822,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema=self.state_schema,
retry_policy=retry_policy,
cache_policy=cache_policy,
error_handler_node=handler_node_name,
ends=ends,
defer=defer,
timeout=timeout,
@@ -1052,7 +1080,6 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
for node in interrupt:
if node not in self.nodes:
raise ValueError(f"Interrupt node `{node}` not found")
self.compiled = True
return self
@@ -1166,6 +1193,11 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
key for key, val in self.channels.items() if not is_managed_value(val)
]
)
node_error_handler_map = {
node_name: spec.error_handler_node
for node_name, spec in self.nodes.items()
if spec.error_handler_node is not None
}
compiled = CompiledStateGraph[StateT, ContextT, InputT, OutputT](
builder=self,
@@ -1188,6 +1220,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
debug=debug,
store=store,
cache=cache,
node_error_handler_map=node_error_handler_map,
name=name or "LangGraph",
stream_transformers=transformers,
)
@@ -1362,6 +1395,8 @@ class CompiledStateGraph(
metadata=node.metadata,
retry_policy=node.retry_policy,
cache_policy=node.cache_policy,
is_error_handler=node.is_error_handler,
error_handler_node=node.error_handler_node,
bound=node.runnable, # type: ignore[arg-type]
timeout=node.timeout,
)
+189 -1
View File
@@ -39,6 +39,7 @@ from langgraph._internal._constants import (
CONFIG_KEY_CHECKPOINT_MAP,
CONFIG_KEY_CHECKPOINT_NS,
CONFIG_KEY_CHECKPOINTER,
CONFIG_KEY_NODE_ERROR,
CONFIG_KEY_READ,
CONFIG_KEY_RESUME_MAP,
CONFIG_KEY_RUNTIME,
@@ -47,6 +48,7 @@ from langgraph._internal._constants import (
CONFIG_KEY_TASK_ID,
CONFIG_KEY_THREAD_ID,
ERROR,
ERROR_SOURCE_NODE,
INTERRUPT,
NO_WRITES,
NS_END,
@@ -66,6 +68,7 @@ from langgraph.channels.base import BaseChannel
from langgraph.channels.topic import Topic
from langgraph.channels.untracked_value import UntrackedValue
from langgraph.constants import TAG_HIDDEN
from langgraph.errors import NodeError
from langgraph.managed.base import ManagedValueMapping
from langgraph.pregel._call import get_runnable_for_task, identifier
from langgraph.pregel._io import read_channels
@@ -292,7 +295,15 @@ def apply_writes(
pending_writes_by_channel: dict[str, list[Any]] = defaultdict(list)
for task in tasks:
for chan, val in task.writes:
if chan in (NO_WRITES, PUSH, RESUME, INTERRUPT, RETURN, ERROR):
if chan in (
NO_WRITES,
PUSH,
RESUME,
INTERRUPT,
RETURN,
ERROR,
ERROR_SOURCE_NODE,
):
pass
elif chan in channels:
pending_writes_by_channel[chan].append(val)
@@ -750,6 +761,42 @@ def prepare_single_task(
return PregelTask(task_id, name, task_path[:3])
def _coerce_pending_error(value: Any) -> BaseException:
if isinstance(value, BaseException):
return value
return Exception(str(value))
def _read_errors_from_pending_writes(
pending_writes: list[PendingWrite],
) -> list[BaseException]:
errors: list[BaseException] = []
for _, channel, value in pending_writes:
if channel == ERROR:
errors.append(_coerce_pending_error(value))
return errors
def _read_error_for_task_id_from_pending_writes(
pending_writes: list[PendingWrite], task_id: str
) -> BaseException | None:
for pending_task_id, channel, value in reversed(pending_writes):
if pending_task_id == task_id and channel == ERROR:
return _coerce_pending_error(value)
return None
def _read_error_source_node_from_pending_writes(
pending_writes: list[PendingWrite], task_id: str
) -> str | None:
for pending_task_id, channel, value in reversed(pending_writes):
if pending_task_id == task_id and channel == ERROR_SOURCE_NODE:
if isinstance(value, str):
return value
return str(value)
return None
def prepare_push_task_functional(
task_path: tuple[str, tuple, int, str, Call],
# (PUSH, parent task path, idx of PUSH write, id of parent task, Call)
@@ -1060,6 +1107,147 @@ def prepare_push_task_send(
return PregelTask(task_id, packet.node, translated_task_path)
def prepare_node_error_handler_task(
failed_task: PregelExecutableTask,
*,
handler_node_name: str,
failed_error: BaseException,
checkpoint: Checkpoint,
pending_writes: list[PendingWrite],
processes: Mapping[str, PregelNode],
channels: Mapping[str, BaseChannel],
managed: ManagedValueMapping,
config: RunnableConfig,
step: int,
stop: int,
store: BaseStore | None = None,
checkpointer: BaseCheckpointSaver | None = None,
manager: None | ParentRunManager | AsyncParentRunManager = None,
cache_policy: CachePolicy | None = None,
retry_policy: Sequence[RetryPolicy] = (),
) -> PregelExecutableTask | None:
"""Prepare an immediate node-level error handler task for a failed task."""
if handler_node_name not in processes:
return None
proc = processes[handler_node_name]
proc_node = proc.node
if proc_node is None:
return None
checkpoint_id_bytes = binascii.unhexlify(checkpoint["id"].replace("-", ""))
task_id_func = _xxhash_str if checkpoint["v"] > 1 else _uuid5_str
configurable = config.get(CONF, {})
parent_ns = configurable.get(CONFIG_KEY_CHECKPOINT_NS, "")
checkpoint_ns = (
f"{parent_ns}{NS_SEP}{handler_node_name}" if parent_ns else handler_node_name
)
task_id = task_id_func(
checkpoint_id_bytes,
checkpoint_ns,
str(step),
handler_node_name,
PUSH,
"node_error_handler",
failed_task.id,
)
task_checkpoint_ns = f"{checkpoint_ns}:{task_id}"
translated_task_path = (*failed_task.path[:3], "node_error_handler", False)
metadata = {
"langgraph_step": step,
"langgraph_node": handler_node_name,
"langgraph_triggers": PUSH_TRIGGER,
"langgraph_path": translated_task_path,
"langgraph_checkpoint_ns": task_checkpoint_ns,
}
if proc.metadata:
metadata.update(proc.metadata)
writes: deque[tuple[str, Any]] = deque()
effective_retry_policy = proc.retry_policy or retry_policy
effective_cache_policy = proc.cache_policy or cache_policy
if effective_cache_policy:
args_key = effective_cache_policy.key_func(failed_task.input)
cache_key = CacheKey(
(
CACHE_NS_WRITES,
(identifier(proc) or "__dynamic__"),
handler_node_name,
),
xxh3_128_hexdigest(
args_key.encode() if isinstance(args_key, str) else args_key
),
effective_cache_policy.ttl,
)
else:
cache_key = None
scratchpad = _scratchpad(
config[CONF].get(CONFIG_KEY_SCRATCHPAD),
pending_writes,
task_id,
xxh3_128_hexdigest(task_checkpoint_ns.encode()),
config[CONF].get(CONFIG_KEY_RESUME_MAP),
step,
stop,
)
runtime = cast(Runtime, configurable.get(CONFIG_KEY_RUNTIME, DEFAULT_RUNTIME))
runtime = runtime.override(
store=store, previous=checkpoint["channel_values"].get(PREVIOUS, None)
)
additional_config: RunnableConfig = {
"metadata": metadata,
"tags": proc.tags,
}
return PregelExecutableTask(
handler_node_name,
failed_task.input,
proc_node,
writes,
patch_config(
merge_configs(config, additional_config),
run_name=handler_node_name,
callbacks=manager.get_child(f"graph:step:{step}") if manager else None,
configurable={
CONFIG_KEY_TASK_ID: task_id,
CONFIG_KEY_SEND: writes.extend,
CONFIG_KEY_READ: partial(
local_read,
scratchpad,
channels,
managed,
PregelTaskWrites(
translated_task_path,
handler_node_name,
writes,
PUSH_TRIGGER,
),
),
CONFIG_KEY_CHECKPOINTER: (
checkpointer or configurable.get(CONFIG_KEY_CHECKPOINTER)
),
CONFIG_KEY_CHECKPOINT_MAP: {
**configurable.get(CONFIG_KEY_CHECKPOINT_MAP, {}),
parent_ns: checkpoint["id"],
},
CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_SCRATCHPAD: scratchpad,
CONFIG_KEY_RUNTIME: runtime,
CONFIG_KEY_NODE_ERROR: NodeError(
node=failed_task.name, error=failed_error
),
},
),
PUSH_TRIGGER,
effective_retry_policy,
cache_key,
task_id,
translated_task_path,
writers=proc.flat_writers,
subgraphs=proc.subgraphs,
)
def checkpoint_null_version(
checkpoint: Checkpoint,
) -> V | None:
+108 -3
View File
@@ -45,11 +45,13 @@ from langgraph._internal._constants import (
CONFIG_KEY_REPLAY_STATE,
CONFIG_KEY_RESUME_MAP,
CONFIG_KEY_RESUMING,
CONFIG_KEY_RUNTIME,
CONFIG_KEY_SCRATCHPAD,
CONFIG_KEY_STREAM,
CONFIG_KEY_TASK_ID,
CONFIG_KEY_THREAD_ID,
ERROR,
ERROR_SOURCE_NODE,
INPUT,
INTERRUPT,
NS_END,
@@ -87,6 +89,7 @@ from langgraph.pregel._algo import (
checkpoint_null_version,
increment,
prepare_next_tasks,
prepare_node_error_handler_task,
prepare_single_task,
sanitize_untracked_values_in_send,
should_interrupt,
@@ -119,6 +122,7 @@ from langgraph.pregel.debug import (
map_debug_tasks,
)
from langgraph.pregel.protocol import StreamChunk, StreamProtocol
from langgraph.runtime import RunControl, Runtime
from langgraph.types import (
All,
CachePolicy,
@@ -206,10 +210,12 @@ class PregelLoop:
"input",
"pending",
"done",
"draining",
"interrupt_before",
"interrupt_after",
"out_of_steps",
]
control: RunControl | None
tasks: dict[str, PregelExecutableTask]
output: None | dict[str, Any] | Any = None
updated_channels: set[str] | None = None
@@ -317,6 +323,8 @@ class PregelLoop:
else ()
)
self.prev_checkpoint_config = None
runtime = self.config[CONF].get(CONFIG_KEY_RUNTIME)
self.control = runtime.control if isinstance(runtime, Runtime) else None
def _push_graph_lifecycle_event(
self,
@@ -324,11 +332,16 @@ class PregelLoop:
*,
interrupts: tuple[Interrupt, ...] = (),
) -> None:
# drain status never reaches lifecycle events: tick() returns False
# before pushing, and interrupts are raised through GraphInterrupt
if self.status == "draining":
raise RuntimeError("Draining status cannot emit lifecycle events")
status = self.status
if kind == "resume":
self._graph_lifecycle_events.append(
GraphResumeEvent(
run_id=None,
status=self.status,
status=status,
checkpoint_id=self.checkpoint["id"],
checkpoint_ns=self.checkpoint_ns,
)
@@ -337,7 +350,7 @@ class PregelLoop:
self._graph_lifecycle_events.append(
GraphInterruptEvent(
run_id=None,
status=self.status,
status=status,
checkpoint_id=self.checkpoint["id"],
checkpoint_ns=self.checkpoint_ns,
interrupts=interrupts,
@@ -511,6 +524,16 @@ class PregelLoop:
# return the new task, to be started if not run before
return pushed
def schedule_error_handler(
self, failed_task: PregelExecutableTask, error: BaseException
) -> PregelExecutableTask | None:
raise NotImplementedError
async def aschedule_error_handler(
self, failed_task: PregelExecutableTask, error: BaseException
) -> PregelExecutableTask | None:
raise NotImplementedError
def tick(self) -> bool:
"""Execute a single iteration of the Pregel loop.
@@ -569,6 +592,10 @@ class PregelLoop:
self.status = "done"
return False
if self.control is not None and self.control.drain_requested:
self.status = "draining"
return False
# if there are pending writes from a previous loop, apply them
if not self.is_replaying and self.checkpoint_pending_writes:
self._match_writes(self.tasks)
@@ -635,7 +662,7 @@ class PregelLoop:
def _match_writes(self, tasks: Mapping[str, PregelExecutableTask]) -> None:
for tid, k, v in self.checkpoint_pending_writes:
if k in (ERROR, INTERRUPT, RESUME):
if k in (ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME):
continue
if task := tasks.get(tid):
task.writes.append((k, v))
@@ -1212,6 +1239,45 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
self.output_writes(task.id, task.writes, cached=True)
return pushed
def schedule_error_handler(
self, failed_task: PregelExecutableTask, error: BaseException
) -> PregelExecutableTask | None:
handler_node = self.nodes[failed_task.name].error_handler_node
if not handler_node:
return None
writes = list(failed_task.writes)
writes.append((ERROR_SOURCE_NODE, failed_task.name))
self.put_writes(
failed_task.id,
writes,
)
handler_task = prepare_node_error_handler_task(
failed_task,
handler_node_name=handler_node,
failed_error=error,
checkpoint=self.checkpoint,
pending_writes=self.checkpoint_pending_writes,
processes=self.nodes,
channels=self.channels,
managed=self.managed,
config=failed_task.config,
step=self.step,
stop=self.stop,
store=self.store,
checkpointer=self.checkpointer,
manager=self.manager,
retry_policy=self.retry_policy,
cache_policy=self.cache_policy,
)
if handler_task is None:
return None
self.tasks[handler_task.id] = handler_task
if not self.is_replaying:
self._match_writes({handler_task.id: handler_task})
for task in self.match_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
return handler_task
def put_writes(self, task_id: str, writes: WritesT) -> None:
"""Put writes for a task, to be read by the next tick."""
super().put_writes(task_id, writes)
@@ -1419,6 +1485,45 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
self.output_writes(task.id, task.writes, cached=True)
return pushed
async def aschedule_error_handler(
self, failed_task: PregelExecutableTask, error: BaseException
) -> PregelExecutableTask | None:
handler_node = self.nodes[failed_task.name].error_handler_node
if not handler_node:
return None
writes = list(failed_task.writes)
writes.append((ERROR_SOURCE_NODE, failed_task.name))
self.put_writes(
failed_task.id,
writes,
)
handler_task = prepare_node_error_handler_task(
failed_task,
handler_node_name=handler_node,
failed_error=error,
checkpoint=self.checkpoint,
pending_writes=self.checkpoint_pending_writes,
processes=self.nodes,
channels=self.channels,
managed=self.managed,
config=failed_task.config,
step=self.step,
stop=self.stop,
store=self.store,
checkpointer=self.checkpointer,
manager=self.manager,
retry_policy=self.retry_policy,
cache_policy=self.cache_policy,
)
if handler_task is None:
return None
self.tasks[handler_task.id] = handler_task
if not self.is_replaying:
self._match_writes({handler_task.id: handler_task})
for task in await self.amatch_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
return handler_task
def put_writes(self, task_id: str, writes: WritesT) -> None:
"""Put writes for a task, to be read by the next tick."""
super().put_writes(task_id, writes)
+10
View File
@@ -138,6 +138,12 @@ class PregelNode:
metadata: Mapping[str, Any] | None
"""Metadata to attach to the node for tracing."""
is_error_handler: bool
"""Whether this node is registered as an error handler node."""
error_handler_node: str | None
"""Optional handler node name for failures from this node."""
subgraphs: Sequence[PregelProtocol]
"""Subgraphs used by the node."""
@@ -153,6 +159,8 @@ class PregelNode:
bound: Runnable[Any, Any] | None = None,
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
is_error_handler: bool = False,
error_handler_node: str | None = None,
subgraphs: Sequence[PregelProtocol] | None = None,
timeout: float | timedelta | TimeoutPolicy | None = None,
) -> None:
@@ -169,6 +177,8 @@ class PregelNode:
self.timeout = coerce_timeout_policy(timeout)
self.tags = tags
self.metadata = metadata
self.is_error_handler = is_error_handler
self.error_handler_node = error_handler_node
if subgraphs is not None:
self.subgraphs = subgraphs
elif self.bound is not DEFAULT_BOUND:
+183 -18
View File
@@ -10,8 +10,10 @@ from collections.abc import (
AsyncIterator,
Awaitable,
Callable,
Collection,
Iterable,
Iterator,
Mapping,
Sequence,
)
from functools import partial
@@ -72,6 +74,10 @@ SKIP_RERAISE_SET: weakref.WeakSet[concurrent.futures.Future | asyncio.Future] =
class FuturesDict(Generic[F, E], dict[F, PregelExecutableTask | None]):
event: E
callback: weakref.ref[Callable[[PregelExecutableTask, BaseException | None], None]]
# Stop condition is injected by PregelRunner instead of hard-coded here.
# This lets the runner treat graph-error-handled exceptions as non-fatal
# so `on_done` does not trigger an early stop for those futures.
should_stop: Callable[[set[F]], bool]
counter: int
done: set[F]
lock: threading.Lock
@@ -82,6 +88,7 @@ class FuturesDict(Generic[F, E], dict[F, PregelExecutableTask | None]):
callback: weakref.ref[
Callable[[PregelExecutableTask, BaseException | None], None]
],
should_stop: Callable[[set[F]], bool],
future_type: type[F],
# used for generic typing, newer py supports FutureDict[...](...)
) -> None:
@@ -89,6 +96,7 @@ class FuturesDict(Generic[F, E], dict[F, PregelExecutableTask | None]):
self.lock = threading.Lock()
self.event = event
self.callback = callback
self.should_stop = should_stop
self.counter = 0
self.done: set[F] = set()
@@ -109,6 +117,7 @@ class FuturesDict(Generic[F, E], dict[F, PregelExecutableTask | None]):
task: PregelExecutableTask,
fut: F,
) -> None:
# Called automatically by future.add_done_callback registered in __setitem__.
try:
if cb := self.callback():
cb(task, _exception(fut))
@@ -116,7 +125,9 @@ class FuturesDict(Generic[F, E], dict[F, PregelExecutableTask | None]):
with self.lock:
self.done.add(fut)
self.counter -= 1
if self.counter == 0 or _should_stop_others(self.done):
# Wake waiter when all tracked futures are done, or when runner-level
# stop condition is met (for example, a non-handled fatal exception).
if self.counter == 0 or self.should_stop(self.done):
self.event.set()
@@ -132,11 +143,34 @@ class PregelRunner:
put_writes: weakref.ref[Callable[[str, Sequence[tuple[str, Any]]], None]],
use_astream: bool = False,
node_finished: Callable[[str], None] | None = None,
node_error_handler_map: Mapping[str, str] | None = None,
schedule_error_handler: Callable[
[PregelExecutableTask, BaseException], PregelExecutableTask | None
]
| None = None,
aschedule_error_handler: Callable[
[PregelExecutableTask, BaseException],
Awaitable[PregelExecutableTask | None],
]
| None = None,
) -> None:
self.submit = submit
self.put_writes = put_writes
self.use_astream = use_astream
self.node_finished = node_finished
self.node_error_handler_map = dict(node_error_handler_map or {})
self.error_handler_nodes = set(self.node_error_handler_map.values())
self.schedule_error_handler = schedule_error_handler
self.aschedule_error_handler = aschedule_error_handler
# Exception object ids that are already routed to graph-level error handler.
# These ids are consulted by stop/panic checks to avoid re-raising handled
# exceptions via the normal fatal path in the same run.
self._handled_exception_ids: set[int] = set()
def _should_route_to_error_handler(self, task: PregelExecutableTask) -> bool:
if task.name in self.error_handler_nodes:
return False
return task.name in self.node_error_handler_map
def tick(
self,
@@ -155,6 +189,9 @@ class PregelRunner:
futures = FuturesDict(
callback=weakref.WeakMethod(self.commit),
event=threading.Event(),
should_stop=partial(
_should_stop_others, handled_exception_ids=self._handled_exception_ids
),
future_type=concurrent.futures.Future,
)
# give control back to the caller
@@ -164,6 +201,7 @@ class PregelRunner:
return
elif len(tasks) == 1 and timeout is None and get_waiter is None:
t = tasks[0]
scheduled_error_handler = False
try:
run_with_retry(
t,
@@ -182,12 +220,23 @@ class PregelRunner:
self.commit(t, None)
except Exception as exc:
self.commit(t, exc)
if (
not isinstance(exc, GraphBubbleUp)
and self._should_route_to_error_handler(t)
and self.schedule_error_handler is not None
):
self._handled_exception_ids.add(id(exc))
if handler_task := self.schedule_error_handler(t, exc):
tasks = (handler_task,)
scheduled_error_handler = True
# Continue to the regular scheduling path for handler execution.
if reraise and futures:
# will be re-raised after futures are done
fut: concurrent.futures.Future = concurrent.futures.Future()
fut.set_exception(exc)
futures.done.add(fut)
elif reraise:
if id(exc) not in self._handled_exception_ids:
# will be re-raised after futures are done
fut: concurrent.futures.Future = concurrent.futures.Future()
fut.set_exception(exc)
futures.done.add(fut)
elif reraise and id(exc) not in self._handled_exception_ids:
if tb := exc.__traceback__:
while tb.tb_next is not None and any(
tb.tb_frame.f_code.co_filename.endswith(name)
@@ -196,10 +245,12 @@ class PregelRunner:
tb = tb.tb_next
exc.__traceback__ = tb
raise
if not futures: # maybe `t` scheduled another task
if not futures and not scheduled_error_handler:
# maybe `t` scheduled another task
return
else:
tasks = () # don't reschedule this task
if not scheduled_error_handler:
tasks = () # don't reschedule this task
# add waiter task if requested
if get_waiter is not None:
futures[get_waiter()] = None
@@ -226,6 +277,7 @@ class PregelRunner:
# each task is independent from all other concurrent tasks
# yield updates/debug output as each task finishes
end_time = timeout + time.monotonic() if timeout else None
handled_futures: set[concurrent.futures.Future[Any]] = set()
while len(futures) > (1 if get_waiter is not None else 0):
done, inflight = concurrent.futures.wait(
futures,
@@ -234,17 +286,49 @@ class PregelRunner:
)
if not done:
break # timed out
done_for_stop: set[concurrent.futures.Future[Any]] = set()
for fut in done:
task = futures.pop(fut)
if task is None:
# waiter task finished, schedule another
if inflight and get_waiter is not None:
futures[get_waiter()] = None
elif (
(task_exc := _exception(fut))
and self._should_route_to_error_handler(task)
and not isinstance(task_exc, GraphBubbleUp)
):
self._handled_exception_ids.add(id(task_exc))
SKIP_RERAISE_SET.add(fut)
handled_futures.add(fut)
if self.schedule_error_handler is not None:
if handler_task := self.schedule_error_handler(task, task_exc):
handler_fut = self.submit()( # type: ignore[misc]
run_with_retry,
handler_task,
retry_policy,
configurable={
CONFIG_KEY_CALL: partial(
_call,
weakref.ref(handler_task),
retry_policy=retry_policy,
futures=weakref.ref(futures),
schedule_task=schedule_task,
submit=self.submit,
),
},
__reraise_on_exit__=reraise,
)
futures[handler_fut] = handler_task
else:
done_for_stop.add(fut)
else:
# remove references to loop vars
del fut, task
# maybe stop other tasks
if _should_stop_others(done):
if _should_stop_others(
done_for_stop, handled_exception_ids=self._handled_exception_ids
):
break
# give control back to the caller
yield
@@ -259,6 +343,8 @@ class PregelRunner:
_panic_or_proceed(
futures.done.union(f for f, t in futures.items() if t is not None),
panic=reraise,
handled_exception_ids=self._handled_exception_ids,
handled_futures=handled_futures,
)
except Exception as exc:
if tb := exc.__traceback__:
@@ -292,6 +378,9 @@ class PregelRunner:
futures = FuturesDict(
callback=weakref.WeakMethod(self.commit),
event=asyncio.Event(),
should_stop=partial(
_should_stop_others, handled_exception_ids=self._handled_exception_ids
),
future_type=asyncio.Future,
)
# give control back to the caller
@@ -301,6 +390,7 @@ class PregelRunner:
return
elif len(tasks) == 1 and get_waiter is None and timeout is None:
t = tasks[0]
scheduled_error_handler = False
try:
await arun_with_retry(
t,
@@ -322,12 +412,22 @@ class PregelRunner:
self.commit(t, None)
except Exception as exc:
self.commit(t, exc)
if (
not isinstance(exc, GraphBubbleUp)
and self._should_route_to_error_handler(t)
and self.aschedule_error_handler is not None
):
self._handled_exception_ids.add(id(exc))
if handler_task := await self.aschedule_error_handler(t, exc):
tasks = (handler_task,)
scheduled_error_handler = True
if reraise and futures:
# will be re-raised after futures are done
fut: asyncio.Future = loop.create_future()
fut.set_exception(exc)
futures.done.add(fut)
elif reraise:
if id(exc) not in self._handled_exception_ids:
# will be re-raised after futures are done
fut: asyncio.Future = loop.create_future()
fut.set_exception(exc)
futures.done.add(fut)
elif reraise and id(exc) not in self._handled_exception_ids:
if tb := exc.__traceback__:
while tb.tb_next is not None and any(
tb.tb_frame.f_code.co_filename.endswith(name)
@@ -336,10 +436,12 @@ class PregelRunner:
tb = tb.tb_next
exc.__traceback__ = tb
raise
if not futures: # maybe `t` scheduled another task
if not futures and not scheduled_error_handler:
# maybe `t` scheduled another task
return
else:
tasks = () # don't reschedule this task
if not scheduled_error_handler:
tasks = () # don't reschedule this task
# add waiter task if requested
if get_waiter is not None:
futures[get_waiter()] = None
@@ -374,6 +476,7 @@ class PregelRunner:
# each task is independent from all other concurrent tasks
# yield updates/debug output as each task finishes
end_time = timeout + loop.time() if timeout else None
handled_futures: set[asyncio.Future[Any]] = set()
while len(futures) > (1 if get_waiter is not None else 0):
done, inflight = await asyncio.wait(
futures,
@@ -382,17 +485,59 @@ class PregelRunner:
)
if not done:
break # timed out
done_for_stop: set[asyncio.Future[Any]] = set()
for fut in done:
task = futures.pop(fut)
if task is None:
# waiter task finished, schedule another
if inflight and get_waiter is not None:
futures[get_waiter()] = None
elif (
(task_exc := _exception(fut))
and self._should_route_to_error_handler(task)
and not isinstance(task_exc, GraphBubbleUp)
):
self._handled_exception_ids.add(id(task_exc))
SKIP_RERAISE_SET.add(fut)
handled_futures.add(fut)
if self.aschedule_error_handler is not None:
if handler_task := await self.aschedule_error_handler(
task, task_exc
):
handler_fut = cast(
asyncio.Future,
self.submit()( # type: ignore[misc]
arun_with_retry,
handler_task,
retry_policy,
stream=self.use_astream,
configurable={
CONFIG_KEY_CALL: partial(
_acall,
weakref.ref(handler_task),
retry_policy=retry_policy,
stream=self.use_astream,
futures=weakref.ref(futures),
schedule_task=schedule_task,
submit=self.submit,
loop=loop,
),
},
__name__=handler_task.name,
__cancel_on_exit__=True,
__reraise_on_exit__=reraise,
),
)
futures[handler_fut] = handler_task
else:
done_for_stop.add(fut)
else:
# remove references to loop vars
del fut, task
# maybe stop other tasks
if _should_stop_others(done):
if _should_stop_others(
done_for_stop, handled_exception_ids=self._handled_exception_ids
):
break
# give control back to the caller
yield
@@ -412,6 +557,8 @@ class PregelRunner:
futures.done.union(f for f, t in futures.items() if t is not None),
timeout_exc_cls=asyncio.TimeoutError,
panic=reraise,
handled_exception_ids=self._handled_exception_ids,
handled_futures=handled_futures,
)
except Exception as exc:
if tb := exc.__traceback__:
@@ -447,6 +594,11 @@ class PregelRunner:
else:
# save error to checkpointer
task.writes.append((ERROR, exception))
if self._should_route_to_error_handler(task) and not isinstance(
exception, GraphBubbleUp
):
# Mark early in commit path; loop-side routing may happen later.
self._handled_exception_ids.add(id(exception))
self.put_writes()(task.id, task.writes) # type: ignore[misc]
else:
if self.node_finished and (
@@ -462,6 +614,8 @@ class PregelRunner:
def _should_stop_others(
done: set[F],
*,
handled_exception_ids: set[int] | None = None,
) -> bool:
"""Check if any task failed, if so, cancel all other tasks.
GraphInterrupts are not considered failures."""
@@ -469,7 +623,11 @@ def _should_stop_others(
if fut.cancelled():
continue
elif exc := fut.exception():
if not isinstance(exc, GraphBubbleUp) and fut not in SKIP_RERAISE_SET:
if (
id(exc) not in (handled_exception_ids or set())
and not isinstance(exc, GraphBubbleUp)
and fut not in SKIP_RERAISE_SET
):
return True
return False
@@ -493,6 +651,9 @@ def _panic_or_proceed(
*,
timeout_exc_cls: type[Exception] = TimeoutError,
panic: bool = True,
handled_exception_ids: set[int] | None = None,
handled_futures: Collection[concurrent.futures.Future[Any] | asyncio.Future[Any]]
| None = None,
) -> None:
"""Cancel remaining tasks if any failed, re-raise exception if panic is True."""
done: set[concurrent.futures.Future[Any] | asyncio.Future[Any]] = set()
@@ -509,6 +670,10 @@ def _panic_or_proceed(
# if any task failed
fut = done.pop()
if exc := _exception(fut):
if fut in (handled_futures or set()):
continue
if id(exc) in (handled_exception_ids or set()):
continue
# cancel all pending tasks
while inflight:
inflight.pop().cancel()
+47
View File
@@ -111,6 +111,7 @@ from langgraph.config import get_config
from langgraph.constants import END
from langgraph.errors import (
ErrorCode,
GraphDrained,
GraphRecursionError,
InvalidUpdateError,
create_error_message,
@@ -156,6 +157,7 @@ from langgraph.pregel.protocol import PregelProtocol, StreamChunk, StreamProtoco
from langgraph.runtime import (
DEFAULT_RUNTIME,
BaseUser,
RunControl,
Runtime,
ServerInfo,
)
@@ -728,6 +730,7 @@ class Pregel(
name: str = "LangGraph"
trigger_to_nodes: Mapping[str, Sequence[str]]
node_error_handler_map: Mapping[str, str]
def __init__(
self,
@@ -752,6 +755,7 @@ class Pregel(
context_schema: type[ContextT] | None = None,
config: RunnableConfig | None = None,
trigger_to_nodes: Mapping[str, Sequence[str]] | None = None,
node_error_handler_map: Mapping[str, str] | None = None,
name: str = "LangGraph",
stream_transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
**deprecated_kwargs: Unpack[DeprecatedKwargs],
@@ -799,6 +803,7 @@ class Pregel(
self.context_schema = context_schema
self.config = config
self.trigger_to_nodes = trigger_to_nodes or {}
self.node_error_handler_map = node_error_handler_map or {}
self.name = name
self.stream_transformers: tuple[Callable[[tuple[str, ...]], Any], ...] = tuple(
stream_transformers or ()
@@ -2570,6 +2575,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
subgraphs: bool = False,
debug: bool | None = None,
version: Literal["v2"],
@@ -2589,6 +2595,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
subgraphs: bool = False,
debug: bool | None = None,
version: Literal["v1"] = ...,
@@ -2607,6 +2614,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
subgraphs: bool = False,
debug: bool | None = None,
version: Literal["v1", "v2"] = "v1",
@@ -2651,6 +2659,7 @@ class Pregel(
- `"sync"`: Changes are persisted synchronously before the next step starts.
- `"async"`: Changes are persisted asynchronously while the next step executes.
- `"exit"`: Changes are persisted only when the graph exits.
control: Optional run control used to request cooperative drain.
subgraphs: Whether to stream events from inside subgraphs, defaults to `False`.
If `True`, the events will be emitted as tuples `(namespace, data)`,
@@ -2815,6 +2824,7 @@ class Pregel(
previous=None,
execution_info=None,
server_info=server_info,
control=control or parent_runtime.control or RunControl(),
)
runtime = parent_runtime.merge(runtime)
config[CONF][CONFIG_KEY_RUNTIME] = runtime
@@ -2864,6 +2874,8 @@ class Pregel(
),
put_writes=weakref.WeakMethod(loop.put_writes),
node_finished=config[CONF].get(CONFIG_KEY_NODE_FINISHED),
node_error_handler_map=self.node_error_handler_map,
schedule_error_handler=loop.schedule_error_handler,
)
# enable subgraph streaming
if subgraphs:
@@ -2945,6 +2957,10 @@ class Pregel(
error_code=ErrorCode.GRAPH_RECURSION_LIMIT,
)
raise GraphRecursionError(msg)
elif loop.status == "draining":
if loop.control is None:
raise RuntimeError("Draining status requires run control")
raise GraphDrained(loop.control.drain_reason or "shutdown")
# set final channel values as run output
run_manager.on_chain_end(loop.output)
except BaseException as e:
@@ -2965,6 +2981,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
subgraphs: bool = False,
debug: bool | None = None,
version: Literal["v2"],
@@ -2984,6 +3001,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
subgraphs: bool = False,
debug: bool | None = None,
version: Literal["v1"] = ...,
@@ -3002,6 +3020,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
subgraphs: bool = False,
debug: bool | None = None,
version: Literal["v1", "v2"] = "v1",
@@ -3046,6 +3065,7 @@ class Pregel(
- `"sync"`: Changes are persisted synchronously before the next step starts.
- `"async"`: Changes are persisted asynchronously while the next step executes.
- `"exit"`: Changes are persisted only when the graph exits.
control: Optional run control used to request cooperative drain.
subgraphs: Whether to stream events from inside subgraphs, defaults to `False`.
If `True`, the events will be emitted as tuples `(namespace, data)`,
@@ -3245,6 +3265,7 @@ class Pregel(
previous=None,
execution_info=None,
server_info=server_info,
control=control or parent_runtime.control or RunControl(),
)
runtime = parent_runtime.merge(runtime)
config[CONF][CONFIG_KEY_RUNTIME] = runtime
@@ -3306,6 +3327,8 @@ class Pregel(
put_writes=weakref.WeakMethod(loop.put_writes),
use_astream=do_stream,
node_finished=config[CONF].get(CONFIG_KEY_NODE_FINISHED),
node_error_handler_map=self.node_error_handler_map,
aschedule_error_handler=loop.aschedule_error_handler,
)
# enable subgraph streaming
if subgraphs:
@@ -3413,6 +3436,10 @@ class Pregel(
error_code=ErrorCode.GRAPH_RECURSION_LIMIT,
)
raise GraphRecursionError(msg)
elif loop.status == "draining":
if loop.control is None:
raise RuntimeError("Draining status requires run control")
raise GraphDrained(loop.control.drain_reason or "shutdown")
# set final channel values as run output
await run_manager.on_chain_end(loop.output)
except BaseException as e:
@@ -3427,6 +3454,7 @@ class Pregel(
*,
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,
) -> Any:
"""Start a sync v2 streaming run driven by transformer projections.
@@ -3456,6 +3484,7 @@ class Pregel(
config: Optional runnable config forwarded to the graph.
interrupt_before: Nodes to interrupt before, if any.
interrupt_after: Nodes to interrupt after, if any.
control: Optional run control used to request cooperative drain.
transformers: Extra transformer classes or configured factories
appended after compile-time `stream_transformers`. Factories
are called as `factory(scope)` so they can propagate to
@@ -3490,6 +3519,7 @@ class Pregel(
version="v2",
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
control=control,
)
)
return GraphRunStream(graph_iter, mux)
@@ -3501,6 +3531,7 @@ class Pregel(
*,
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,
) -> Any:
"""Async counterpart to `stream_v2`.
@@ -3523,6 +3554,7 @@ class Pregel(
config: Optional runnable config forwarded to the graph.
interrupt_before: Nodes to interrupt before, if any.
interrupt_after: Nodes to interrupt after, if any.
control: Optional run control used to request cooperative drain.
transformers: Extra transformer classes or configured factories
appended after compile-time `stream_transformers`. Factories
are called as `factory(scope)` so they can propagate to
@@ -3553,6 +3585,7 @@ class Pregel(
version="v2",
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
control=control,
).__aiter__()
return AsyncGraphRunStream(graph_aiter, mux)
@@ -3569,6 +3602,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
version: Literal["v2"],
**kwargs: Any,
) -> GraphOutput[OutputT]: ...
@@ -3586,6 +3620,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
version: Literal["v2"],
**kwargs: Any,
) -> list[StreamPart[StateT, OutputT]]: ...
@@ -3603,6 +3638,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
version: Literal["v1"] = ...,
**kwargs: Any,
) -> dict[str, Any] | Any: ...
@@ -3619,6 +3655,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
version: Literal["v1", "v2"] = "v1",
**kwargs: Any,
) -> dict[str, Any] | Any:
@@ -3643,6 +3680,7 @@ class Pregel(
- `"sync"`: Changes are persisted synchronously before the next step starts.
- `"async"`: Changes are persisted asynchronously while the next step executes.
- `"exit"`: Changes are persisted only when the graph exits.
control: Optional run control used to request cooperative drain.
version: The streaming format version. `"v1"` (default) returns the
traditional format, `"v2"` returns `StreamPart` typed dicts when
`stream_mode` is not `"values"`.
@@ -3670,6 +3708,7 @@ class Pregel(
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
durability=durability,
control=control,
version=version,
**kwargs,
):
@@ -3693,6 +3732,7 @@ class Pregel(
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
durability=durability,
control=control,
**kwargs,
):
if stream_mode == "values":
@@ -3739,6 +3779,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
version: Literal["v2"],
**kwargs: Any,
) -> GraphOutput[OutputT]: ...
@@ -3756,6 +3797,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
version: Literal["v2"],
**kwargs: Any,
) -> list[StreamPart[StateT, OutputT]]: ...
@@ -3773,6 +3815,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
version: Literal["v1"] = ...,
**kwargs: Any,
) -> dict[str, Any] | Any: ...
@@ -3789,6 +3832,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
control: RunControl | None = None,
version: Literal["v1", "v2"] = "v1",
**kwargs: Any,
) -> dict[str, Any] | Any:
@@ -3813,6 +3857,7 @@ class Pregel(
- `"sync"`: Changes are persisted synchronously before the next step starts.
- `"async"`: Changes are persisted asynchronously while the next step executes.
- `"exit"`: Changes are persisted only when the graph exits.
control: Optional run control used to request cooperative drain.
version: The streaming format version. `"v1"` (default) returns the
traditional format, `"v2"` returns `StreamPart` typed dicts when
`stream_mode` is not `"values"`.
@@ -3840,6 +3885,7 @@ class Pregel(
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
durability=durability,
control=control,
version=version,
**kwargs,
):
@@ -3863,6 +3909,7 @@ class Pregel(
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
durability=durability,
control=control,
**kwargs,
):
if stream_mode == "values":
+49 -2
View File
@@ -16,6 +16,7 @@ from langgraph.typing import ContextT
__all__ = (
"BaseUser",
"ExecutionInfo",
"RunControl",
"Runtime",
"ServerInfo",
"get_runtime",
@@ -75,6 +76,34 @@ class ServerInfo:
"""
class RunControl:
"""Run-scoped control surface for cooperative draining.
Intended for a single graph run. Create a fresh `RunControl` per run;
reusing a control after `request_drain()` leaves it drained.
Safe to call from any thread: the drain request is represented by a
single attribute write, so no lock is needed for this signal.
If more mutable state is added here, add synchronization.
"""
__slots__ = ("_drain_reason",)
def __init__(self) -> None:
self._drain_reason: str | None = None
def request_drain(self, reason: str = "shutdown") -> None:
self._drain_reason = reason
@property
def drain_requested(self) -> bool:
return self._drain_reason is not None
@property
def drain_reason(self) -> str | None:
return self._drain_reason
def _no_op_stream_writer(_: Any) -> None: ...
@@ -89,6 +118,7 @@ class _RuntimeOverrides(TypedDict, Generic[ContextT], total=False):
previous: Any
execution_info: ExecutionInfo
server_info: ServerInfo | None
control: RunControl | None
@dataclass(**_DC_KWARGS)
@@ -167,7 +197,7 @@ class Runtime(Generic[ContextT]):
context: ContextT = field(default=None) # type: ignore[assignment]
"""Static context for the graph run, like `user_id`, `db_conn`, etc.
Can also be thought of as 'run dependencies'."""
store: BaseStore | None = field(default=None)
@@ -188,7 +218,7 @@ class Runtime(Generic[ContextT]):
previous: Any = field(default=None)
"""The previous return value for the given thread.
Only available with the functional API when a checkpointer is provided.
"""
@@ -200,6 +230,13 @@ class Runtime(Generic[ContextT]):
server_info: ServerInfo | None = field(default=None)
"""Metadata injected by LangGraph Server. None when running open-source LangGraph without LangSmith deployments."""
control: RunControl | None = field(default=None)
"""Run-scoped control plane for cooperative draining.
Populated automatically during graph runs. None outside an active
graph runtime.
"""
def merge(self, other: Runtime[ContextT]) -> Runtime[ContextT]:
"""Merge two runtimes together.
@@ -217,6 +254,7 @@ class Runtime(Generic[ContextT]):
previous=self.previous if other.previous is None else other.previous,
execution_info=other.execution_info or self.execution_info,
server_info=other.server_info or self.server_info,
control=other.control or self.control,
)
def override(
@@ -235,6 +273,14 @@ class Runtime(Generic[ContextT]):
execution_info=self.execution_info.patch(**overrides),
)
@property
def drain_requested(self) -> bool:
return self.control.drain_requested if self.control is not None else False
@property
def drain_reason(self) -> str | None:
return self.control.drain_reason if self.control is not None else None
DEFAULT_RUNTIME = Runtime(
context=None,
@@ -243,6 +289,7 @@ DEFAULT_RUNTIME = Runtime(
heartbeat=_no_op_heartbeat,
previous=None,
execution_info=None,
control=None,
)
+7
View File
@@ -91,9 +91,11 @@ class StreamMux:
self._assign_seq = _assign_seq
self._events: StreamChannel[ProtocolEvent] = StreamChannel()
self._events._bind(is_async=is_async)
self._events._bind_mux(self)
self._transformers: list[StreamTransformer] = []
self._channels: list[StreamChannel[Any]] = []
self._seq = 0
self._push_seq = 0
self.extensions: dict[str, Any] = {}
self.native_keys: set[str] = set()
@@ -124,6 +126,10 @@ class StreamMux:
"""Return the transformer that contributed `key` to the projection."""
return self._transformer_by_key.get(key)
def _next_push_seq(self) -> int:
self._push_seq += 1
return self._push_seq
# ------------------------------------------------------------------
# Pump wiring + mini-mux nesting
# ------------------------------------------------------------------
@@ -449,6 +455,7 @@ class StreamMux:
for value in projection.values():
if isinstance(value, StreamChannel):
value._bind(is_async=self.is_async)
value._bind_mux(self)
self._channels.append(value)
if value.name is not None:
method = value.name if native else f"custom:{value.name}"
+74 -27
View File
@@ -185,28 +185,25 @@ class GraphRunStream:
return iter(self._mux._events)
def interleave(self, *names: str) -> Iterator[tuple[str, Any]]:
"""Iterate multiple projections round-robin, yielding ``(name, item)``.
"""Iterate multiple projections in arrival order, yielding ``(name, item)``.
Each turn advances one projection's cursor; when a cursor's buffer
is empty, pulling from it drives the pump once, which fans out to
every subscribed projection log. Projections whose items aren't
consumed on this turn sit in their own buffers only until the next
turn reaches them, bounding memory by the skew between projection
rates rather than letting any single log grow to the full run
length.
Projections are exhausted independently; a projection that finishes
early drops out of the rotation while others continue. The overall
iterator ends once all named projections are done.
Items are ordered by a monotonic push stamp assigned when each
transformer pushes into its `StreamChannel`. This gives strict
arrival ordering across projections, unlike round-robin.
Args:
*names: Projection keys to interleave. Must match keys in
``extensions``.
Yields:
``(name, item)`` tuples in round-robin order across the named
``(name, item)`` tuples in arrival order across the named
projections.
Each named channel is locked for the duration of iteration and
released when the generator completes, is closed, or raises.
Channels cannot be subscribed concurrently use `.tee(n)` if
you need fan-out.
Raises:
KeyError: If a name doesn't match a registered projection.
@@ -219,20 +216,70 @@ class GraphRunStream:
print("val:", item)
```
"""
cursors: dict[str, Iterator[Any]] = {
name: iter(self.extensions[name]) for name in names
}
done: set[str] = set()
while len(done) < len(cursors):
for name, cursor in cursors.items():
if name in done:
continue
try:
item = next(cursor)
except StopIteration:
done.add(name)
continue
yield (name, item)
from langgraph.stream.stream_channel import StreamChannel
channels: dict[str, StreamChannel[Any]] = {}
try:
for name in names:
ch = self.extensions[name]
if not isinstance(ch, StreamChannel):
raise TypeError(
f"interleave() requires StreamChannel projections, "
f"got {type(ch).__name__} for {name!r}"
)
if ch._is_async is None:
raise TypeError(
f"StreamChannel {name!r} has not been bound yet. "
"Register the transformer with a StreamMux first."
)
if ch._is_async:
raise TypeError(
f"StreamChannel {name!r} is bound to async mode — "
"sync interleave() cannot consume async channels."
)
if ch._subscribed:
raise RuntimeError(
f"StreamChannel {name!r} already has a subscriber; "
"use .tee(n) for fan-out."
)
ch._subscribed = True
channels[name] = ch
done: set[str] = set()
while len(done) < len(channels):
best: tuple[int, str] | None = None
for name, ch in channels.items():
if name in done:
continue
if ch._closed and not ch._items:
if ch._error is not None:
raise ch._error
done.add(name)
continue
if ch._items:
stamp = ch._items[0][0]
if best is None or stamp < best[0]:
best = (stamp, name)
if best is not None:
_stamp, item = channels[best[1]]._items.popleft()
yield (best[1], item)
else:
pump = self._mux._pump_fn
if pump is None or not pump():
before = len(done)
for name, ch in channels.items():
if name not in done and not ch._items:
if ch._closed:
if ch._error is not None:
raise ch._error
done.add(name)
if len(done) == before:
break
finally:
for ch in channels.values():
ch._subscribed = False
class AsyncGraphRunStream:
@@ -3,7 +3,10 @@ from __future__ import annotations
import asyncio
from collections import deque
from collections.abc import AsyncIterator, Awaitable, Callable, Iterator
from typing import Generic, TypeVar
from typing import TYPE_CHECKING, Generic, TypeVar
if TYPE_CHECKING:
from langgraph.stream._mux import StreamMux
T = TypeVar("T")
@@ -64,7 +67,7 @@ class StreamChannel(Generic[T]):
if maxlen is not None and maxlen <= 0:
raise ValueError("StreamChannel maxlen must be a positive int or None")
self.name = name
self._items: deque[T] = deque()
self._items: deque[tuple[int, T]] = deque()
self._maxlen: int | None = maxlen
self._closed = False
self._error: BaseException | None = None
@@ -77,11 +80,15 @@ class StreamChannel(Generic[T]):
self._arequest_more: Callable[[], Awaitable[bool]] | None = None
self._wire_fn: Callable[[T], None] | None = None
self._mux: StreamMux | None = None
# ------------------------------------------------------------------
# Binding
# ------------------------------------------------------------------
def _bind_mux(self, mux: StreamMux) -> None:
self._mux = mux
def _bind(self, *, is_async: bool) -> None:
"""Bind this channel to sync or async mode.
@@ -117,13 +124,18 @@ class StreamChannel(Generic[T]):
registered, but auto-forwarding always fires so wired events
reach the main event log regardless of subscription state.
Items are stored as `(stamp, item)` tuples where stamp is a
monotonic counter from the owning mux. Stamps are stripped by
the default cursors; raw stamped tuples are visible on `_items`.
Raises:
RuntimeError: If the channel is closed (and subscribed).
"""
if self._subscribed:
if self._closed:
raise RuntimeError("Cannot push to a closed StreamChannel")
self._items.append(item)
stamp = self._mux._next_push_seq() if self._mux is not None else 0
self._items.append((stamp, item))
if self._wire_fn is not None:
self._wire_fn(item)
@@ -170,7 +182,8 @@ class StreamChannel(Generic[T]):
def _sync_cursor(self) -> Iterator[T]:
while True:
if self._items:
yield self._items.popleft()
_stamp, item = self._items.popleft()
yield item
elif self._closed:
if self._error is not None:
raise self._error
@@ -212,7 +225,8 @@ class StreamChannel(Generic[T]):
async def _async_cursor(self) -> AsyncIterator[T]:
while True:
if self._items:
yield self._items.popleft()
_stamp, item = self._items.popleft()
yield item
elif self._closed:
if self._error is not None:
raise self._error
@@ -12,7 +12,7 @@ from langchain_core.messages import AIMessageChunk, BaseMessage
from langchain_protocol.protocol import MessagesData
from typing_extensions import NotRequired, TypedDict
from langgraph.errors import GraphInterrupt
from langgraph.errors import GraphDrained, GraphInterrupt
from langgraph.stream._types import ProtocolEvent, StreamTransformer
from langgraph.stream.run_stream import AsyncSubgraphRunStream, SubgraphRunStream
from langgraph.stream.stream_channel import StreamChannel
@@ -327,7 +327,7 @@ class MessagesTransformer(StreamTransformer):
self._by_run.clear()
SubgraphStatus = Literal["started", "completed", "failed", "interrupted"]
SubgraphStatus = Literal["started", "completed", "failed", "interrupted", "drained"]
def _parse_ns_segment(segment: str) -> tuple[str, str | None]:
@@ -472,10 +472,8 @@ class _TasksLifecycleBase(StreamTransformer):
self._open.clear()
def fail(self, err: BaseException) -> None:
"""Emit `failed` / `interrupted` for any tracked namespace still open."""
is_interrupt = isinstance(err, GraphInterrupt)
status: SubgraphStatus = "interrupted" if is_interrupt else "failed"
error_str = None if is_interrupt else str(err)
"""Emit terminal status for any tracked namespace still open."""
status, error_str = _status_from_exception(err)
for ns in list(self._open):
self._on_terminal(ns, status, error_str)
self._open.clear()
@@ -483,6 +481,8 @@ class _TasksLifecycleBase(StreamTransformer):
def _status_from_exception(err: BaseException) -> tuple[SubgraphStatus, str | None]:
"""Map a run exception to a subgraph terminal status and error string."""
if isinstance(err, GraphDrained):
return "drained", None
if isinstance(err, GraphInterrupt):
return "interrupted", None
return "failed", str(err)
+2 -2
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph"
version = "1.1.10"
version = "1.2.0a3"
description = "Building stateful, multi-actor applications with LLMs"
authors = []
requires-python = ">=3.10"
@@ -25,7 +25,7 @@ classifiers = [
]
dependencies = [
"langchain-core>=1.3.2,<2",
"langgraph-checkpoint>=4.0.3,<5.0.0",
"langgraph-checkpoint>=4.1.0a3,<5.0.0",
"langgraph-sdk>=0.3.0,<0.4.0",
"langgraph-prebuilt>=1.0.12,<1.1.0",
"xxhash>=3.5.0",
+2 -3
View File
@@ -371,8 +371,7 @@ def test_delta_channel_inmemory_saver_assembles_writes() -> None:
saved = saver.get_tuple(config)
assert saved is not None
assert "messages" in saved.checkpoint["channel_values"]
assert saved.checkpoint["channel_values"]["messages"] is DELTA_SENTINEL
assert "messages" not in saved.checkpoint["channel_values"]
state = graph.get_state(config)
assert len(state.values["messages"]) == 4 # 2 human + 2 AI
@@ -562,7 +561,7 @@ def test_delta_channel_dict_reducer_end_to_end_filesystem() -> None:
saved = saver.get_tuple(config)
assert saved is not None
assert saved.checkpoint["channel_values"]["files"] is DELTA_SENTINEL
assert "files" not in saved.checkpoint["channel_values"]
state = graph.get_state(config)
assert state.values["files"] == {
"/doc_1.txt": "content for turn 1",
@@ -0,0 +1,353 @@
"""Tests for arrival-ordered interleave and push stamps."""
from __future__ import annotations
import operator
from typing import Annotated, Any
import pytest
from typing_extensions import TypedDict
from langgraph.constants import END, START
from langgraph.graph import StateGraph
from langgraph.stream import StreamChannel, StreamTransformer
from langgraph.stream._mux import StreamMux
from langgraph.stream._types import ProtocolEvent
from langgraph.stream.run_stream import GraphRunStream
from langgraph.stream.transformers import ValuesTransformer
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
class _TwoChannelTransformer(StreamTransformer):
"""Transformer that exposes two named channels for testing interleave."""
_native = True
def __init__(self, scope: tuple[str, ...] = ()) -> None:
super().__init__(scope)
self._alpha: StreamChannel[str] = StreamChannel("alpha")
self._beta: StreamChannel[str] = StreamChannel("beta")
def init(self) -> dict[str, Any]:
return {"alpha": self._alpha, "beta": self._beta}
def process(self, event: ProtocolEvent) -> bool:
return True
class SimpleState(TypedDict):
value: str
items: Annotated[list[str], operator.add]
def _build_simple_graph():
def node_a(state: SimpleState) -> dict:
return {"value": state["value"] + "A", "items": ["a"]}
def node_b(state: SimpleState) -> dict:
return {"value": state["value"] + "B", "items": ["b"]}
builder = StateGraph(SimpleState)
builder.add_node("node_a", node_a)
builder.add_node("node_b", node_b)
builder.add_edge(START, "node_a")
builder.add_edge("node_a", "node_b")
builder.add_edge("node_b", END)
return builder.compile()
# ---------------------------------------------------------------------------
# Unit tests: push stamps on StreamChannel
# ---------------------------------------------------------------------------
class TestPushStamps:
def test_stamps_are_monotonic_across_channels(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
alpha = mux.extensions["alpha"]
beta = mux.extensions["beta"]
alpha._subscribed = True
beta._subscribed = True
alpha.push("a1")
beta.push("b1")
alpha.push("a2")
beta.push("b2")
all_stamped = list(alpha._items) + list(beta._items)
stamps = [s for s, _ in all_stamped]
assert len(set(stamps)) == 4
items_by_arrival = [item for _, item in sorted(all_stamped)]
assert items_by_arrival == ["a1", "b1", "a2", "b2"]
def test_regular_iter_strips_stamps(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
alpha = mux.extensions["alpha"]
it = iter(alpha)
alpha.push("a1")
alpha.push("a2")
alpha.close()
items = list(it)
assert items == ["a1", "a2"]
assert all(isinstance(item, str) for item in items)
def test_events_channel_gets_real_stamps(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
alpha = mux.extensions["alpha"]
alpha._subscribed = True
alpha.push("a1")
mux._events._subscribed = True
mux._events.push({"method": "test", "data": "x"})
alpha.push("a2")
all_stamps = [s for s, _ in alpha._items] + [s for s, _ in mux._events._items]
assert len(set(all_stamps)) == len(all_stamps), "all stamps should be unique"
assert all(s > 0 for s in all_stamps), "no stamp should be zero"
def test_channel_without_mux_gets_zero_stamp(self) -> None:
ch: StreamChannel[str] = StreamChannel()
ch._bind(is_async=False)
ch._subscribed = True
ch.push("x")
assert list(ch._items) == [(0, "x")]
# ---------------------------------------------------------------------------
# Unit tests: interleave arrival order
# ---------------------------------------------------------------------------
class TestInterleaveArrivalOrder:
def test_arrival_order_not_round_robin(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
alpha = mux.extensions["alpha"]
beta = mux.extensions["beta"]
run = GraphRunStream(None, mux, wire_pump=False)
# interleave() subscribes channels directly and reads _items
# for stamp-ordered iteration. We simulate the pump by wiring
# a custom callback that pushes items in a known order.
push_script = [
("alpha", "a1"),
("alpha", "a2"),
("beta", "b1"),
("alpha", "a3"),
("beta", "b2"),
]
push_iter = iter(push_script)
channels = {"alpha": alpha, "beta": beta}
def fake_pump() -> bool:
try:
name, item = next(push_iter)
channels[name].push(item)
return True
except StopIteration:
mux.close()
return False
mux.bind_pump(fake_pump)
result = list(run.interleave("alpha", "beta"))
names = [name for name, _ in result]
items = [item for _, item in result]
assert items == ["a1", "a2", "b1", "a3", "b2"]
assert names == ["alpha", "alpha", "beta", "alpha", "beta"]
def test_single_projection(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
alpha = mux.extensions["alpha"]
run = GraphRunStream(None, mux, wire_pump=False)
push_script = [("alpha", "a1"), ("alpha", "a2")]
push_iter = iter(push_script)
def fake_pump() -> bool:
try:
_, item = next(push_iter)
alpha.push(item)
return True
except StopIteration:
mux.close()
return False
mux.bind_pump(fake_pump)
result = list(run.interleave("alpha"))
assert result == [("alpha", "a1"), ("alpha", "a2")]
def test_empty_projection(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
alpha = mux.extensions["alpha"]
run = GraphRunStream(None, mux, wire_pump=False)
push_script = [("alpha", "a1"), ("alpha", "a2")]
push_iter = iter(push_script)
channels = {"alpha": alpha}
def fake_pump() -> bool:
try:
name, item = next(push_iter)
channels[name].push(item)
return True
except StopIteration:
mux.close()
return False
mux.bind_pump(fake_pump)
result = list(run.interleave("alpha", "beta"))
assert result == [("alpha", "a1"), ("alpha", "a2")]
def test_unknown_projection_raises(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
run = GraphRunStream(None, mux, wire_pump=False)
mux.close()
with pytest.raises((KeyError, AttributeError)):
list(run.interleave("alpha", "does_not_exist"))
def test_all_empty(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
run = GraphRunStream(None, mux, wire_pump=False)
def fake_pump() -> bool:
mux.close()
return False
mux.bind_pump(fake_pump)
result = list(run.interleave("alpha", "beta"))
assert result == []
def test_error_propagation(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
alpha = mux.extensions["alpha"]
beta = mux.extensions["beta"]
run = GraphRunStream(None, mux, wire_pump=False)
err = RuntimeError("boom")
push_script = [
("alpha", "a1"),
("beta", "b1"),
]
push_iter = iter(push_script)
channels = {"alpha": alpha, "beta": beta}
def fake_pump() -> bool:
try:
name, item = next(push_iter)
channels[name].push(item)
return True
except StopIteration:
alpha.fail(err)
beta.close()
return False
mux.bind_pump(fake_pump)
collected = []
with pytest.raises(RuntimeError, match="boom"):
for pair in run.interleave("alpha", "beta"):
collected.append(pair)
assert ("alpha", "a1") in collected
assert ("beta", "b1") in collected
# ---------------------------------------------------------------------------
# Integration test: interleave with stream_v2
# ---------------------------------------------------------------------------
class TestInterleaveIntegration:
def test_interleave_values_and_messages(self) -> None:
run = _build_simple_graph().stream_v2({"value": "x", "items": []})
tagged = list(run.interleave("values", "messages"))
names = [name for name, _ in tagged]
assert set(names).issubset({"values", "messages"})
assert names.count("values") >= 1
def test_interleave_rejects_already_subscribed(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
alpha = mux.extensions["alpha"]
run = GraphRunStream(None, mux, wire_pump=False)
# Subscribe alpha via iter first
_ = iter(alpha)
mux.close()
with pytest.raises(RuntimeError, match="already has a subscriber"):
list(run.interleave("alpha"))
def test_interleave_releases_projections_on_completion(self) -> None:
run = _build_simple_graph().stream_v2({"value": "x", "items": []})
list(run.interleave("values", "messages"))
# Subscriptions should be released after the generator completes,
# so the channels can be re-iterated (they'll be empty / closed).
assert run.extensions["values"]._subscribed is False
assert run.extensions["messages"]._subscribed is False
def test_interleave_releases_projections_on_early_break(self) -> None:
run = _build_simple_graph().stream_v2({"value": "x", "items": []})
gen = run.interleave("values", "messages")
next(gen)
gen.close()
assert run.extensions["values"]._subscribed is False
assert run.extensions["messages"]._subscribed is False
def test_interleave_releases_projections_on_validation_failure(self) -> None:
mux = StreamMux(
factories=[ValuesTransformer, _TwoChannelTransformer],
is_async=False,
)
alpha = mux.extensions["alpha"]
# Pre-subscribe alpha so that interleave will fail validation when
# it gets to the second name. The first (already-validated) channel
# should still be released.
run = GraphRunStream(None, mux, wire_pump=False)
mux.close()
alpha._subscribed = True
with pytest.raises(RuntimeError, match="already has a subscriber"):
list(run.interleave("values", "alpha"))
assert mux.extensions["values"]._subscribed is False
+23
View File
@@ -122,6 +122,29 @@ def test_graph_validation() -> None:
graph.invoke({"hello": "there"})
def test_request_drain_allows_inflight_call_scheduling(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
from langgraph.runtime import RunControl
@task
def child(x: int) -> int:
return x + 1
control = RunControl()
@entrypoint(checkpointer=sync_checkpointer)
def graph(x: int) -> int:
control.request_drain()
fut = child(x)
return fut.result()
config = {"configurable": {"thread_id": "drain-call-sync"}}
assert graph.invoke(1, config=config, control=control) == 2
assert control.drain_requested
def test_invalid_checkpointer_type() -> None:
class State(TypedDict):
foo: str
+149
View File
@@ -16,6 +16,7 @@ from typing import (
Literal,
Optional,
)
from unittest.mock import patch
from uuid import UUID
import pytest
@@ -48,6 +49,7 @@ from langgraph.channels.topic import Topic
from langgraph.errors import (
GraphRecursionError,
InvalidUpdateError,
NodeError,
ParentCommand,
)
from langgraph.func import entrypoint, task
@@ -215,6 +217,30 @@ async def test_checkpoint_errors() -> None:
pass
@NEEDS_CONTEXTVARS
async def test_request_drain_allows_inflight_acall_scheduling(
async_checkpointer: BaseCheckpointSaver,
) -> None:
from langgraph.runtime import RunControl
@task
async def child(x: int) -> int:
return x + 1
control = RunControl()
@entrypoint(checkpointer=async_checkpointer)
async def graph(x: int) -> int:
control.request_drain()
fut = child(x)
return await fut
config = {"configurable": {"thread_id": "drain-call-async"}}
assert await graph.ainvoke(1, config=config, control=control) == 2
assert control.drain_requested
async def test_py_async_with_cancel_behavior() -> None:
"""This test confirms that in all versions of Python we support, __aexit__
is not cancelled when the coroutine containing the async with block is cancelled."""
@@ -9739,3 +9765,126 @@ async def test_fork_does_not_apply_pending_writes(
# 1 (input) + 20 (forked node_a) + 100 (node_b) = 121
assert result == {"value": 121}
async def test_graph_error_handler_async_runtime_info() -> None:
class State(TypedDict):
foo: str
attempts = 0
captured: dict[str, object] = {}
async def always_failing_node(state: State) -> State:
nonlocal attempts
attempts += 1
raise ValueError("Always fails async")
async def err_handler_node(state: State, error: NodeError) -> State:
captured["from_node_name"] = error.node
captured["from_node_error"] = error.error
return {"foo": "handled_async"}
graph = (
StateGraph(State)
.add_node(
"always_failing",
always_failing_node,
retry_policy=RetryPolicy(
max_attempts=2,
initial_interval=0.01,
jitter=False,
retry_on=ValueError,
),
error_handler=err_handler_node,
)
.add_edge(START, "always_failing")
.compile()
)
with patch("asyncio.sleep"):
result = await graph.ainvoke({"foo": ""})
assert attempts == 2
assert result["foo"] == "handled_async"
assert captured["from_node_name"] == "always_failing"
assert isinstance(captured["from_node_error"], BaseException)
@NEEDS_CONTEXTVARS
async def test_graph_error_handler_does_not_swallow_interrupt_concurrent() -> None:
"""When a graph error handler is configured and a node calls interrupt()
concurrently with other nodes, the interrupt must still be raised not
silently swallowed."""
class State(TypedDict):
foo: str
async def node_a(state: State) -> State:
val = interrupt("need human input")
return {"foo": f"a_{val}"}
async def node_b(state: State) -> State:
return {}
async def err_handler(state: State) -> State:
return {"foo": "handled"}
checkpointer = InMemorySaver()
graph = (
StateGraph(State)
.add_node("node_a", node_a, error_handler=err_handler)
.add_node("node_b", node_b)
.add_edge(START, "node_a")
.add_edge(START, "node_b")
.compile(checkpointer=checkpointer)
)
config = {"configurable": {"thread_id": "test-interrupt-concurrent-async"}}
await graph.ainvoke({"foo": ""}, config)
state = await graph.aget_state(config)
assert len(state.tasks) > 0
interrupts = [t for t in state.tasks if hasattr(t, "interrupts") and t.interrupts]
assert len(interrupts) > 0, (
"GraphInterrupt was swallowed — interrupt() in node_a "
"should have paused execution"
)
async def test_node_error_handler_handles_subgraph_internal_failure_async() -> None:
class SubState(TypedDict):
foo: str
class ParentState(TypedDict):
foo: str
captured: dict[str, object] = {}
async def sub_fail_node(state: SubState) -> SubState:
raise ValueError("async subgraph boom")
async def parent_handler(state: ParentState, error: NodeError) -> ParentState:
captured["from_node_name"] = error.node
captured["from_node_error"] = error.error
return {"foo": "handled_async_subgraph"}
subgraph = (
StateGraph(SubState)
.add_node("sub_fail_node", sub_fail_node)
.add_edge(START, "sub_fail_node")
.compile()
)
parent_graph = (
StateGraph(ParentState)
.add_node("subgraph_node", subgraph, error_handler=parent_handler)
.add_edge(START, "subgraph_node")
.compile()
)
result = await parent_graph.ainvoke({"foo": ""})
assert result["foo"] == "handled_async_subgraph"
assert captured["from_node_name"] == "subgraph_node"
assert isinstance(captured["from_node_error"], BaseException)
@@ -460,8 +460,9 @@ class TestStreamV2Sync:
names = [name for name, _ in tagged]
assert set(names).issubset({"values", "messages"})
assert names.count("values") >= 1
with pytest.raises(RuntimeError, match="already has a subscriber"):
list(run.values)
# interleave releases its subscription on completion.
assert run.extensions["values"]._subscribed is False
assert run.extensions["messages"]._subscribed is False
def test_abort_marks_exhausted_and_closes_mux(self) -> None:
run = _build_simple_graph().stream_v2({"value": "x", "items": []})
+518 -4
View File
@@ -1,4 +1,5 @@
import asyncio
import operator
import sys
import threading
import time
@@ -15,7 +16,7 @@ from langchain_core.language_models.fake_chat_models import GenericFakeChatModel
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, HumanMessage
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult
from langchain_core.runnables import RunnableLambda, RunnableParallel
from langgraph.checkpoint.memory import MemorySaver
from langgraph.checkpoint.memory import InMemorySaver, MemorySaver
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
from typing_extensions import TypedDict
@@ -34,7 +35,7 @@ from langgraph._internal._runnable import RunnableCallable
from langgraph._internal._timeout import coerce_timeout_policy
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.channels.last_value import LastValue
from langgraph.errors import GraphInterrupt, NodeTimeoutError, ParentCommand
from langgraph.errors import GraphInterrupt, NodeError, NodeTimeoutError, ParentCommand
from langgraph.func import entrypoint, task
from langgraph.graph import END, START, StateGraph, add_messages
from langgraph.pregel import NodeBuilder, Pregel
@@ -209,6 +210,15 @@ def test_should_retry_default_retry_on():
req_error_no_resp.response = None
assert _should_retry_on(policy, req_error_no_resp) is True
# NodeTimeoutError should be retryable by default
assert (
_should_retry_on(
policy,
NodeTimeoutError("node", 1.0, kind="run", run_timeout=0.5),
)
is True
)
# Should retry on other exceptions by default
class CustomException(Exception):
pass
@@ -1455,14 +1465,14 @@ async def test_state_graph_add_node_timeout_composes_with_retry():
async def flaky(state: _TimeoutState) -> _TimeoutState:
attempts.append(len(attempts))
if len(attempts) < 2:
await asyncio.sleep(0.5)
await asyncio.sleep(1.0)
return {"x": state["x"] + 1}
builder = StateGraph(_TimeoutState)
builder.add_node(
"flaky",
flaky,
timeout=TimeoutPolicy(idle_timeout=0.1),
timeout=TimeoutPolicy(idle_timeout=0.3),
retry_policy=RetryPolicy(
max_attempts=3,
initial_interval=0.0,
@@ -1755,3 +1765,507 @@ async def test_arun_with_retry_timeout_observer_treats_bubble_up_as_non_error():
assert finish.status == "success"
assert finish.error_type is None
assert finish.error_message is None
# ---------------------------------------------------------------------------
# Watcher invariant: any timeout that retry/error_handler can recover from
# MUST emit `finish=error` BEFORE the in-process recovery work happens. The
# external watchdog (langgraph-api) relies on this so it only kills a worker
# when no `finish` arrives within the deadline. The tests below pin down the
# three recovery paths so a refactor that moves `_finish_timed_attempt` past
# an `await` (or past the final `raise`) trips CI.
# ---------------------------------------------------------------------------
@pytest.mark.anyio
async def test_arun_with_retry_observer_emits_finish_before_retry_backoff():
"""`finish=error` of attempt N must arrive before the retry backoff sleep."""
timeline: list[tuple[float, Any]] = []
class TimingOutOnceProc:
def __init__(self) -> None:
self.calls = 0
async def ainvoke(self, input, config):
self.calls += 1
if self.calls == 1:
await asyncio.sleep(1.0)
return "ok"
backoff = 0.25
policy = RetryPolicy(
max_attempts=2,
initial_interval=backoff,
backoff_factor=1.0,
max_interval=backoff,
jitter=False,
retry_on=NodeTimeoutError,
)
task = _make_task(
TimingOutOnceProc(),
timeout=_idle_timeout(0.05),
retry_policy=(policy,),
name="backoff_watcher",
)
task.config[CONF][CONFIG_KEY_TIMED_ATTEMPT_OBSERVER] = lambda ev: timeline.append(
(time.monotonic(), ev)
)
assert await arun_with_retry(task, retry_policy=None) == "ok"
starts = [(t, ev) for t, ev in timeline if ev.event == "start"]
finishes = [(t, ev) for t, ev in timeline if ev.event == "finish"]
assert [ev.context.attempt for _, ev in starts] == [1, 2]
assert [ev.status for _, ev in finishes] == ["error", "success"]
first_finish_t = finishes[0][0]
second_start_t = starts[1][0]
# The watcher relies on this gap: `finish=error` for attempt 1 must arrive
# before `arun_with_retry` enters `await asyncio.sleep(backoff)`. We give a
# generous slack to keep this stable on slow CI; the structural invariant
# is "finish lands first", not "the gap equals exactly backoff".
assert second_start_t - first_finish_t >= backoff * 0.5, (
f"finish=error appears to be emitted after retry backoff sleep; "
f"gap was {second_start_t - first_finish_t:.3f}s, expected >= {backoff * 0.5:.3f}s"
)
@pytest.mark.anyio
async def test_state_graph_observer_emits_finish_before_error_handler_start():
"""Original task's `finish=error` must arrive before the error_handler task's `start`."""
class State(TypedDict):
foo: str
async def slow_node(state: State) -> State:
await asyncio.sleep(1.0)
return {"foo": "should-not-happen"}
async def handler_node(state: State, error: NodeError) -> State:
return {"foo": "handled"}
events: list = []
graph = (
StateGraph(State)
.add_node(
"slow",
slow_node,
timeout=TimeoutPolicy(idle_timeout=0.05),
error_handler=handler_node,
)
.add_edge(START, "slow")
.compile()
)
result = await graph.ainvoke(
{"foo": ""},
config={
"configurable": {CONFIG_KEY_TIMED_ATTEMPT_OBSERVER: events.append},
},
)
assert result["foo"] == "handled"
# Filter to events from the failing node only — the handler node has no
# timeout configured here, so it doesn't appear in the observer stream.
slow_events = [ev for ev in events if ev.context.task_name == "slow"]
starts = [ev for ev in slow_events if ev.event == "start"]
finishes = [ev for ev in slow_events if ev.event == "finish"]
assert len(starts) == 1
assert len(finishes) == 1
assert finishes[0].status == "error"
assert finishes[0].error_type == "NodeTimeoutError"
# The slow task's finish-error event must precede every event for any
# follow-up task in the same observer stream.
slow_finish_index = events.index(finishes[0])
for ev in events[slow_finish_index + 1 :]:
assert ev.context.task_name == "slow" or ev.event == "start", (
f"unexpected event {ev.event} for {ev.context.task_name} "
f"after slow's finish=error"
)
@pytest.mark.anyio
async def test_arun_with_retry_observer_emits_finish_before_final_raise_on_exhaustion():
"""When retry exhausts and the timeout propagates, the final `finish=error` must
be emitted before `arun_with_retry` re-raises."""
events: list = []
class AlwaysTimingOutProc:
async def ainvoke(self, input, config):
await asyncio.sleep(1.0)
return "never"
policy = RetryPolicy(
max_attempts=2,
initial_interval=0.0,
jitter=False,
retry_on=NodeTimeoutError,
)
task = _make_task(
AlwaysTimingOutProc(),
timeout=_idle_timeout(0.05),
retry_policy=(policy,),
name="never_succeeds",
)
task.config[CONF][CONFIG_KEY_TIMED_ATTEMPT_OBSERVER] = events.append
with pytest.raises(NodeTimeoutError):
await arun_with_retry(task, retry_policy=None)
starts = [ev for ev in events if ev.event == "start"]
finishes = [ev for ev in events if ev.event == "finish"]
assert [ev.context.attempt for ev in starts] == [1, 2]
assert [ev.context.attempt for ev in finishes] == [1, 2]
assert [ev.status for ev in finishes] == ["error", "error"]
assert all(ev.error_type == "NodeTimeoutError" for ev in finishes)
# Both finish events were observed BEFORE arun_with_retry raised, otherwise
# the `with pytest.raises` block would have exited before `events` got
# populated with the second finish.
@pytest.mark.anyio
async def test_sync_sleep_in_async_node_bypasses_timeout_and_emits_finish_success():
"""Sync `time.sleep` inside an async node blocks the event loop so the
in-process watchdog cannot fire. We document the resulting behavior here:
1. `NodeTimeoutError` is NOT raised, even though the sync sleep exceeds
`idle_timeout`.
2. The node's normal return value flows through.
3. `finish=success` is emitted to the observer.
This is the canonical case where the in-process timeout is defeated and
the only safety net is the external watcher (langgraph-api), which
SIGKILLs the worker when no `finish` arrives within its deadline. The
catch is that with a *short* sync sleep the event loop unblocks before
the watcher's deadline expires, so the watcher legitimately does not
kill meaning the configured `idle_timeout` is silently honored at the
process level only when the block is long enough to outlast the
watcher's grace.
This is the documented "Cooperative cancellation" caveat on
`TimeoutPolicy`. The test pins the behavior so any future change that
starts raising `NodeTimeoutError` for sync-blocked async nodes (or stops
emitting `finish=success`) is caught.
"""
events: list = []
class SyncSleepingProc:
async def ainvoke(self, input, config):
time.sleep(0.1)
return "completed_despite_timeout"
task = _make_task(
SyncSleepingProc(),
timeout=_idle_timeout(0.05),
name="sync_sleeper",
)
task.config[CONF][CONFIG_KEY_TIMED_ATTEMPT_OBSERVER] = events.append
result = await arun_with_retry(task, retry_policy=None)
assert result == "completed_despite_timeout"
starts = [ev for ev in events if ev.event == "start"]
finishes = [ev for ev in events if ev.event == "finish"]
assert len(starts) == 1
assert len(finishes) == 1
assert finishes[0].status == "success"
assert finishes[0].error_type is None
def test_graph_error_handler_runs_after_retry_exhaustion():
class State(TypedDict):
foo: str
attempts = 0
captured: dict[str, object] = {}
def always_failing_node(state: State) -> State:
nonlocal attempts
attempts += 1
raise ValueError("Always fails")
def err_handler_node(state: State, error: NodeError) -> Command:
captured["from_node_name"] = error.node
captured["from_node_error"] = error.error
return Command(update={"foo": "handled"}, goto="after_handler")
def after_handler(state: State) -> State:
return {"foo": f"{state['foo']}_after"}
retry_policy = RetryPolicy(
max_attempts=2,
initial_interval=0.01,
jitter=False,
retry_on=ValueError,
)
graph = (
StateGraph(State)
.add_node(
"always_failing",
always_failing_node,
retry_policy=retry_policy,
error_handler=err_handler_node,
)
.add_node("after_handler", after_handler)
.add_edge(START, "always_failing")
.compile()
)
with patch("time.sleep"):
result = graph.invoke({"foo": ""})
assert attempts == 2
assert result["foo"] == "handled_after"
assert captured["from_node_name"] == "always_failing"
assert isinstance(captured["from_node_error"], BaseException)
def test_graph_error_handler_can_route_with_command():
class State(TypedDict):
foo: str
attempts = 0
def always_failing_node(state: State) -> State:
nonlocal attempts
attempts += 1
raise ValueError("Always fails")
def err_handler_node(state: State) -> Command:
return Command(update={"foo": "handled"}, goto="next_node")
def next_node(state: State) -> State:
return {"foo": f"{state['foo']}_next"}
retry_policy = RetryPolicy(
max_attempts=1,
initial_interval=0.01,
jitter=False,
retry_on=ValueError,
)
graph = (
StateGraph(State)
.add_node(
"always_failing",
always_failing_node,
retry_policy=retry_policy,
error_handler=err_handler_node,
)
.add_node("next_node", next_node)
.add_edge(START, "always_failing")
.compile()
)
result = graph.invoke({"foo": ""})
assert attempts == 1
assert result["foo"] == "handled_next"
def test_graph_error_handler_failure_fails_run():
class State(TypedDict):
foo: str
def always_failing_node(state: State) -> State:
raise ValueError("Always fails")
def err_handler_node(state: State) -> State:
raise RuntimeError("handler failed")
graph = (
StateGraph(State)
.add_node("always_failing", always_failing_node, error_handler=err_handler_node)
.add_edge(START, "always_failing")
.compile()
)
with pytest.raises(RuntimeError, match="handler failed"):
graph.invoke({"foo": ""})
def test_graph_error_handler_handles_subgraph_internal_failure():
class SubState(TypedDict):
foo: str
class ParentState(TypedDict):
foo: str
parent_handler_called = False
captured: dict[str, object] = {}
def sub_fail_node(state: SubState) -> SubState:
raise ValueError("subgraph boom")
def parent_handler(state: ParentState, error: NodeError) -> ParentState:
nonlocal parent_handler_called
parent_handler_called = True
captured["from_node_name"] = error.node
captured["from_node_error"] = error.error
return {"foo": "handled_by_parent"}
subgraph = (
StateGraph(SubState)
.add_node("sub_fail_node", sub_fail_node)
.add_edge(START, "sub_fail_node")
.compile()
)
parent_graph = (
StateGraph(ParentState)
.add_node("subgraph_node", subgraph, error_handler=parent_handler)
.add_edge(START, "subgraph_node")
.compile()
)
result = parent_graph.invoke({"foo": ""})
assert result["foo"] == "handled_by_parent"
assert parent_handler_called is True
assert captured["from_node_name"] == "subgraph_node"
assert isinstance(captured["from_node_error"], BaseException)
def test_graph_error_handler_error_context_survives_checkpoint_resume():
class State(TypedDict):
foo: str
captured: dict[str, object] = {}
def always_failing_node(state: State) -> State:
raise RuntimeError("failed before handler")
def err_handler_node(state: State, error: NodeError) -> State:
captured["from_node_name"] = error.node
captured["from_node_error"] = error.error
return {"foo": "handled_after_resume"}
checkpointer = InMemorySaver()
config = {"configurable": {"thread_id": "graph-error-resume"}}
graph = (
StateGraph(State)
.add_node("always_failing", always_failing_node, error_handler=err_handler_node)
.add_edge(START, "always_failing")
.compile(
checkpointer=checkpointer,
interrupt_before=["__error_handler__always_failing"],
)
)
# First run pauses before handler, after failure context is checkpointed.
graph.invoke({"foo": ""}, config)
# Resume should execute handler and recover serialized error context.
result = graph.invoke(None, config)
assert result["foo"] == "handled_after_resume"
assert captured["from_node_name"] == "always_failing"
assert isinstance(captured["from_node_error"], BaseException)
def test_graph_error_handler_does_not_swallow_interrupt_concurrent():
"""When a graph error handler is configured and a node calls interrupt()
concurrently with other nodes, the interrupt must still be raised not
silently swallowed."""
from langgraph.types import interrupt
class State(TypedDict):
foo: str
def node_a(state: State) -> State:
# This node uses interrupt() which raises GraphInterrupt
val = interrupt("need human input")
return {"foo": f"a_{val}"}
def node_b(state: State) -> State:
return {}
def err_handler(state: State) -> State:
return {"foo": "handled"}
checkpointer = InMemorySaver()
graph = (
StateGraph(State)
.add_node("node_a", node_a, error_handler=err_handler)
.add_node("node_b", node_b)
# Fan-out: both node_a and node_b run concurrently
.add_edge(START, "node_a")
.add_edge(START, "node_b")
.compile(checkpointer=checkpointer)
)
config = {"configurable": {"thread_id": "test-interrupt-concurrent"}}
# First invoke should pause at the interrupt, not silently complete
graph.invoke({"foo": ""}, config)
# The graph should have an interrupt pending
state = graph.get_state(config)
assert len(state.tasks) > 0
# There should be a pending interrupt from node_a
interrupts = [t for t in state.tasks if hasattr(t, "interrupts") and t.interrupts]
assert len(interrupts) > 0, (
"GraphInterrupt was swallowed — interrupt() in node_a "
"should have paused execution"
)
def test_node_error_handlers_route_to_matching_handler():
class State(TypedDict):
route: str
foo: Annotated[list[str], operator.add]
def route_node(state: State) -> State:
return {"foo": []}
def choose_node(state: State) -> str:
return state["route"]
def fail_a(state: State) -> State:
raise ValueError("a failed")
def fail_b(state: State) -> State:
raise RuntimeError("b failed")
def handler_a(state: State, error: NodeError) -> State:
assert error.node == "fail_a"
return {"foo": ["handled_a"]}
def handler_b(state: State, error: NodeError) -> State:
assert error.node == "fail_b"
return {"foo": ["handled_b"]}
graph = (
StateGraph(State)
.add_node("route_node", route_node)
.add_node("fail_a", fail_a, error_handler=handler_a)
.add_node("fail_b", fail_b, error_handler=handler_b)
.add_edge(START, "route_node")
.add_conditional_edges("route_node", choose_node, path_map=["fail_a", "fail_b"])
.compile()
)
result_a = graph.invoke({"route": "fail_a", "foo": []})
result_b = graph.invoke({"route": "fail_b", "foo": []})
assert result_a["foo"] == ["handled_a"]
assert result_b["foo"] == ["handled_b"]
def test_node_without_error_handler_still_fails_run():
class State(TypedDict):
foo: str
def fail_without_handler(state: State) -> State:
raise ValueError("no handler")
graph = (
StateGraph(State)
.add_node("fail_without_handler", fail_without_handler)
.add_edge(START, "fail_without_handler")
.compile()
)
with pytest.raises(ValueError, match="no handler"):
graph.invoke({"foo": ""})
+516 -1
View File
@@ -1,3 +1,6 @@
import asyncio
import threading
import time
from dataclasses import dataclass
from typing import Any
@@ -6,8 +9,15 @@ from langgraph.checkpoint.memory import MemorySaver
from pydantic import BaseModel, ValidationError
from typing_extensions import TypedDict
from langgraph.errors import GraphDrained
from langgraph.graph import END, START, StateGraph
from langgraph.runtime import ExecutionInfo, Runtime, ServerInfo, get_runtime
from langgraph.runtime import (
ExecutionInfo,
RunControl,
Runtime,
ServerInfo,
get_runtime,
)
def test_injected_runtime() -> None:
@@ -79,6 +89,183 @@ def test_merge_runtime() -> None:
assert runtime1.merge(runtime3).context.api_key == "abc" # type: ignore
def test_merge_runtime_preserves_run_control() -> None:
control = RunControl()
runtime1 = Runtime(control=control)
runtime2 = Runtime(context=None)
assert runtime1.merge(runtime2).control is control
def test_run_control_request_drain_stops_future_steps() -> None:
class State(TypedDict, total=False):
first: str
second: str
control = RunControl()
def first_node(state: State) -> dict[str, str]:
control.request_drain()
return {"first": "done"}
def second_node(state: State) -> dict[str, str]:
return {"second": "should-not-run"}
graph = StateGraph(State)
graph.add_node("first", first_node)
graph.add_node("second", second_node)
graph.add_edge(START, "first")
graph.add_edge("first", "second")
graph.add_edge("second", END)
with pytest.raises(GraphDrained, match="shutdown"):
graph.compile().invoke({}, control=control)
@pytest.mark.anyio
async def test_run_control_request_drain_stops_future_steps_async() -> None:
class State(TypedDict, total=False):
first: str
second: str
control = RunControl()
async def first_node(state: State) -> dict[str, str]:
control.request_drain()
return {"first": "done"}
async def second_node(state: State) -> dict[str, str]:
return {"second": "should-not-run"}
graph = StateGraph(State)
graph.add_node("first", first_node)
graph.add_node("second", second_node)
graph.add_edge(START, "first")
graph.add_edge("first", "second")
graph.add_edge("second", END)
with pytest.raises(GraphDrained, match="shutdown"):
await graph.compile().ainvoke({}, control=control)
def test_drain_requested_in_terminal_step_finishes_normally() -> None:
class State(TypedDict, total=False):
value: str
control = RunControl()
def node(state: State) -> dict[str, str]:
control.request_drain()
return {"value": "done"}
graph = StateGraph(State)
graph.add_node("node", node)
graph.add_edge(START, "node")
graph.add_edge("node", END)
assert graph.compile().invoke({}, control=control) == {"value": "done"}
assert control.drain_requested
def test_drain_with_exit_durability_persists_resume_checkpoint() -> None:
class State(TypedDict, total=False):
first: str
second: str
control = RunControl()
def first_node(state: State) -> dict[str, str]:
control.request_drain("sigterm")
return {"first": "done"}
def second_node(state: State) -> dict[str, str]:
return {"second": "done"}
graph = StateGraph(State)
graph.add_node("first", first_node)
graph.add_node("second", second_node)
graph.add_edge(START, "first")
graph.add_edge("first", "second")
graph.add_edge("second", END)
compiled = graph.compile(checkpointer=MemorySaver())
config = {"configurable": {"thread_id": "drain-exit"}}
with pytest.raises(GraphDrained, match="sigterm"):
compiled.invoke({}, config, durability="exit", control=control)
assert compiled.invoke(None, config, durability="exit") == {
"first": "done",
"second": "done",
}
def test_drain_from_subgraph_can_resume_parent() -> None:
class State(TypedDict, total=False):
child_first: str
child_second: str
parent_second: str
control = RunControl()
def child_first(state: State) -> dict[str, str]:
control.request_drain("sigterm")
return {"child_first": "done"}
def child_second(state: State) -> dict[str, str]:
return {"child_second": "done"}
child_builder = StateGraph(State)
child_builder.add_node("child_first", child_first)
child_builder.add_node("child_second", child_second)
child_builder.add_edge(START, "child_first")
child_builder.add_edge("child_first", "child_second")
child_builder.add_edge("child_second", END)
child_graph = child_builder.compile(checkpointer=True)
def parent_second(state: State) -> dict[str, str]:
return {"parent_second": "done"}
parent_builder = StateGraph(State)
parent_builder.add_node("child", child_graph)
parent_builder.add_node("parent_second", parent_second)
parent_builder.add_edge(START, "child")
parent_builder.add_edge("child", "parent_second")
parent_builder.add_edge("parent_second", END)
compiled = parent_builder.compile(checkpointer=MemorySaver())
config = {"configurable": {"thread_id": "drain-subgraph"}}
with pytest.raises(GraphDrained, match="sigterm"):
compiled.invoke({}, config, control=control)
assert compiled.invoke(None, config) == {
"child_first": "done",
"child_second": "done",
"parent_second": "done",
}
@pytest.mark.anyio
async def test_drain_requested_in_terminal_step_finishes_normally_async() -> None:
class State(TypedDict, total=False):
value: str
control = RunControl()
async def node(state: State) -> dict[str, str]:
control.request_drain()
return {"value": "done"}
graph = StateGraph(State)
graph.add_node("node", node)
graph.add_edge(START, "node")
graph.add_edge("node", END)
assert await graph.compile().ainvoke({}, control=control) == {"value": "done"}
assert control.drain_requested
def test_runtime_propogated_to_subgraph() -> None:
@dataclass
class Context:
@@ -392,6 +579,334 @@ def test_context_coercion_pydantic_validation_errors() -> None:
)
def test_external_drain_concurrent_sync() -> None:
"""External thread calls request_drain() while graph is mid-execution."""
class State(TypedDict, total=False):
first: str
second: str
started = threading.Event()
def first_node(state: State) -> dict[str, str]:
started.set()
time.sleep(0.05)
return {"first": "done"}
def second_node(state: State) -> dict[str, str]:
return {"second": "should-not-run"}
graph = StateGraph(State)
graph.add_node("first", first_node)
graph.add_node("second", second_node)
graph.add_edge(START, "first")
graph.add_edge("first", "second")
graph.add_edge("second", END)
control = RunControl()
compiled = graph.compile()
exc_holder: list[BaseException | None] = [None]
def run_graph() -> None:
try:
compiled.invoke({}, control=control)
except GraphDrained as e:
exc_holder[0] = e
t = threading.Thread(target=run_graph)
t.start()
started.wait(timeout=5)
control.request_drain("sigterm")
t.join(timeout=10)
exc = exc_holder[0]
assert isinstance(exc, GraphDrained)
assert exc.reason == "sigterm"
@pytest.mark.anyio
async def test_external_drain_concurrent_async() -> None:
"""External task calls request_drain() while graph is mid-execution."""
class State(TypedDict, total=False):
first: str
second: str
started = asyncio.Event()
async def first_node(state: State) -> dict[str, str]:
started.set()
await asyncio.sleep(0.05)
return {"first": "done"}
async def second_node(state: State) -> dict[str, str]:
return {"second": "should-not-run"}
graph = StateGraph(State)
graph.add_node("first", first_node)
graph.add_node("second", second_node)
graph.add_edge(START, "first")
graph.add_edge("first", "second")
graph.add_edge("second", END)
control = RunControl()
compiled = graph.compile()
async def drain_after_start() -> None:
await started.wait()
control.request_drain("sigterm")
drain_task = asyncio.create_task(drain_after_start())
with pytest.raises(GraphDrained, match="sigterm"):
await compiled.ainvoke({}, control=control)
await drain_task
@pytest.mark.anyio
async def test_drain_then_cancel_after_graceful_timeout() -> None:
"""Simulate: drain requested -> node still running -> graceful timeout -> cancel.
This shows what happens when a long-running node doesn't finish within
the graceful period after drain is requested.
"""
class State(TypedDict, total=False):
first: str
second: str
node_started = asyncio.Event()
node_cancelled = asyncio.Event()
node_finished = asyncio.Event()
async def slow_node(state: State) -> dict[str, str]:
node_started.set()
try:
await asyncio.sleep(30) # very long operation
except asyncio.CancelledError:
node_cancelled.set()
raise
node_finished.set()
return {"first": "done"}
async def second_node(state: State) -> dict[str, str]:
return {"second": "should-not-run"}
graph = StateGraph(State)
graph.add_node("first", slow_node)
graph.add_node("second", second_node)
graph.add_edge(START, "first")
graph.add_edge("first", "second")
graph.add_edge("second", END)
control = RunControl()
compiled = graph.compile()
# Phase 1: start graph
graph_task = asyncio.create_task(compiled.ainvoke({}, control=control))
# Phase 2: wait for node to start, then request drain
await node_started.wait()
control.request_drain("sigterm")
# Phase 3: graceful timeout — node is still running, cancel after 1s
graceful_timeout = 1.0
await asyncio.sleep(graceful_timeout)
assert not node_finished.is_set(), "node should still be running"
assert not node_cancelled.is_set(), "node should not be cancelled yet"
# Phase 4: force cancel
graph_task.cancel()
with pytest.raises(asyncio.CancelledError):
await graph_task
# The node received CancelledError at the await point
assert node_cancelled.is_set(), "node should have received CancelledError"
assert not node_finished.is_set(), "node should NOT have finished normally"
@pytest.mark.anyio
async def test_cancel_ainvoke_with_async_node() -> None:
"""Cancel ainvoke running an async node: CancelledError is delivered
at the await point and the node stops immediately."""
class State(TypedDict, total=False):
first: str
second: str
timeline: list[str] = []
node_started = asyncio.Event()
async def slow_async_node(state: State) -> dict[str, str]:
timeline.append(f"async_node:start thread={threading.current_thread().name}")
node_started.set()
try:
await asyncio.sleep(30)
except asyncio.CancelledError:
timeline.append("async_node:cancelled")
raise
timeline.append("async_node:finished")
return {"first": "done"}
async def second_node(state: State) -> dict[str, str]:
timeline.append("second_node:run")
return {"second": "should-not-run"}
graph = StateGraph(State)
graph.add_node("first", slow_async_node)
graph.add_node("second", second_node)
graph.add_edge(START, "first")
graph.add_edge("first", "second")
graph.add_edge("second", END)
compiled = graph.compile()
graph_task = asyncio.create_task(compiled.ainvoke({}))
await node_started.wait()
timeline.append("test:cancel")
graph_task.cancel()
with pytest.raises(asyncio.CancelledError):
await graph_task
timeline.append("test:done")
# async node runs on the event loop thread (MainThread)
assert any("MainThread" in e for e in timeline if "async_node:start" in e)
# CancelledError was delivered at the await point — node stopped
assert "async_node:cancelled" in timeline
# Node did NOT run to completion
assert "async_node:finished" not in timeline
# Second node never ran
assert "second_node:run" not in timeline
@pytest.mark.anyio
async def test_cancel_ainvoke_with_sync_node() -> None:
"""Cancel ainvoke running a sync node.
Sync nodes in ainvoke run on a separate thread (via run_in_executor),
NOT on the event loop thread. Cancelling the asyncio task disconnects
from the thread future, but the thread keeps running as an orphan and
completes on its own.
Key difference from async nodes:
- async node: CancelledError stops the coroutine at an await point
- sync node: cancel only disconnects asyncio; the thread runs to completion
In shutdown case, we will ignore this because the instance will be destroyed soon.
"""
class State(TypedDict, total=False):
first: str
second: str
timeline: list[str] = []
node_started = threading.Event()
node_finished = threading.Event()
def slow_sync_node(state: State) -> dict[str, str]:
timeline.append(f"sync_node:start thread={threading.current_thread().name}")
node_started.set()
time.sleep(1)
timeline.append("sync_node:after_sleep")
node_finished.set()
return {"first": "done"}
def second_node(state: State) -> dict[str, str]:
timeline.append("second_node:run")
return {"second": "should-not-run"}
graph = StateGraph(State)
graph.add_node("first", slow_sync_node)
graph.add_node("second", second_node)
graph.add_edge(START, "first")
graph.add_edge("first", "second")
graph.add_edge("second", END)
control = RunControl()
compiled = graph.compile()
timeline.append(f"test:main thread={threading.current_thread().name}")
graph_task = asyncio.create_task(compiled.ainvoke({}, control=control))
loop = asyncio.get_event_loop()
await loop.run_in_executor(None, node_started.wait, 5)
timeline.append("test:cancel+drain")
graph_task.cancel()
control.request_drain("sigterm")
with pytest.raises(asyncio.CancelledError):
await graph_task
timeline.append("test:exc=CancelledError")
# Sync node runs on a background thread (asyncio_*), NOT MainThread
sync_start = next(e for e in timeline if "sync_node:start" in e)
assert "MainThread" not in sync_start, (
"sync node should run on a background thread, not the event loop thread"
)
# At this point, the asyncio task is done but the thread is orphaned.
# The sync node has NOT finished yet — cancel only disconnected asyncio.
assert not node_finished.is_set(), (
"sync node should still be running in its background thread"
)
# Wait for the orphaned thread to complete on its own.
await loop.run_in_executor(None, node_finished.wait, 5)
assert node_finished.is_set()
# After the orphaned thread finishes, the full timeline looks like:
# test:main thread=MainThread
# sync_node:start thread=asyncio_N <- background thread
# test:cancel+drain <- cancel + drain fired
# test:exc=CancelledError <- asyncio disconnected
# sync_node:after_sleep <- thread ran to completion anyway
assert "sync_node:after_sleep" in timeline
# Second node never ran
assert "second_node:run" not in timeline
# Verify timeline ordering: cancel happened before node finished
cancel_idx = timeline.index("test:cancel+drain")
sleep_idx = timeline.index("sync_node:after_sleep")
assert cancel_idx < sleep_idx, (
"cancel was issued while the sync node was still sleeping"
)
def test_drain_with_control_parameter_sync() -> None:
"""Control parameter is wired through invoke -> stream."""
class State(TypedDict, total=False):
value: str
ran = False
def node(state: State) -> dict[str, str]:
nonlocal ran
ran = True
return {"value": "done"}
graph = StateGraph(State)
graph.add_node("node", node)
graph.add_edge(START, "node")
graph.add_edge("node", END)
# Pre-drained control stops before executing the first pending task.
control = RunControl()
control.request_drain("pre-drained")
with pytest.raises(GraphDrained, match="pre-drained"):
graph.compile().invoke({}, control=control)
assert not ran
# --- ExecutionInfo unit tests ---
@@ -77,8 +77,13 @@ def _arm(mux: StreamMux, transformer: Any) -> None:
transformer._log._subscribed = True
def _unstamped(items):
"""Strip push stamps from a StreamChannel's internal buffer."""
return [item for _stamp, item in items]
def _drain(transformer: Any) -> list[Any]:
return list(transformer._log._items)
return _unstamped(transformer._log._items)
# ---------------------------------------------------------------------------
@@ -140,7 +145,7 @@ def test_custom_does_not_suppress_from_main_log() -> None:
mux.push(_custom_event([], "data"))
methods = [evt["method"] for evt in mux._events._items]
methods = [evt["method"] for evt in _unstamped(mux._events._items)]
assert "custom" in methods
@@ -219,7 +224,7 @@ def test_checkpoints_does_not_suppress_from_main_log() -> None:
mux.push(_checkpoints_event([], {"values": {}}))
methods = [evt["method"] for evt in mux._events._items]
methods = [evt["method"] for evt in _unstamped(mux._events._items)]
assert "checkpoints" in methods
@@ -284,7 +289,7 @@ def test_debug_does_not_suppress_from_main_log() -> None:
mux.push(_debug_event([], {"step": 0}))
methods = [evt["method"] for evt in mux._events._items]
methods = [evt["method"] for evt in _unstamped(mux._events._items)]
assert "debug" in methods
@@ -360,7 +365,7 @@ def test_tasks_does_not_suppress_from_main_log() -> None:
mux.push(_tasks_event([], {"id": "t1"}))
methods = [evt["method"] for evt in mux._events._items]
methods = [evt["method"] for evt in _unstamped(mux._events._items)]
assert "tasks" in methods
@@ -429,7 +434,7 @@ def test_updates_does_not_suppress_from_main_log() -> None:
mux.push(_updates_event([], {"n": {}}))
methods = [evt["method"] for evt in mux._events._items]
methods = [evt["method"] for evt in _unstamped(mux._events._items)]
assert "updates" in methods
@@ -469,7 +474,7 @@ def test_unrelated_events_ignored_by_all() -> None:
)
for t in transformers:
assert list(t._log._items) == []
assert _unstamped(t._log._items) == []
# ---------------------------------------------------------------------------
@@ -674,7 +679,7 @@ def test_tasks_and_lifecycle_coregistration() -> None:
assert _drain(tasks) == [task_data]
methods = [evt["method"] for evt in mux._events._items]
methods = [evt["method"] for evt in _unstamped(mux._events._items)]
assert "tasks" not in methods
@@ -92,11 +92,16 @@ def _arm(mux: StreamMux) -> None:
transformer._channel._subscribed = True
def _unstamped(items):
"""Strip push stamps from a StreamChannel's internal buffer."""
return [item for _stamp, item in items]
def _drain_lifecycle(mux: StreamMux) -> list[LifecyclePayload]:
"""Snapshot the lifecycle channel's buffer."""
transformer = mux.transformer_by_key("lifecycle")
assert isinstance(transformer, LifecycleTransformer)
return list(transformer._channel._items)
return _unstamped(transformer._channel._items)
def _build_lifecycle_mux(*, scope: tuple[str, ...] = ()) -> StreamMux:
@@ -301,7 +306,7 @@ def test_protocol_event_method_is_native() -> None:
mux = _build_lifecycle_mux()
mux.push(_tasks_start(["agent:abc"], task_id="t1", name="tool"))
methods = {evt["method"] for evt in mux._events._items}
methods = {evt["method"] for evt in _unstamped(mux._events._items)}
assert "lifecycle" in methods
assert "custom:lifecycle" not in methods
@@ -312,7 +317,7 @@ def test_tasks_events_suppressed_from_main_log() -> None:
mux.push(_tasks_start(["agent:abc"], task_id="t1", name="tool"))
mux.push(_tasks_result([], task_id="abc", name="agent"))
methods = [evt["method"] for evt in mux._events._items]
methods = [evt["method"] for evt in _unstamped(mux._events._items)]
assert "tasks" not in methods
# Lifecycle events did make it through, though.
assert "lifecycle" in methods
@@ -26,6 +26,11 @@ from langgraph.stream.transformers import MessagesTransformer, ValuesTransformer
TS = int(time.time() * 1000)
def _unstamped(items):
"""Strip push stamps from a StreamChannel's internal buffer."""
return [item for _stamp, item in items]
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@@ -174,7 +179,7 @@ class TestProtocolEventRouting:
)
)
log.close()
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert isinstance(stream, ChatModelStream)
assert stream.message_id == "run-1"
@@ -183,7 +188,7 @@ class TestProtocolEventRouting:
for evt in _lifecycle(text="hello world"):
t.process(_proto_event(evt, run_id="run-1"))
log.close()
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert stream.done
assert stream.output.text == "hello world"
@@ -206,7 +211,7 @@ class TestProtocolEventRouting:
)
)
log.close()
assert list(log._items) == []
assert _unstamped(log._items) == []
def test_concurrent_streams_routed_by_run_id(self) -> None:
t, log = _make_sync_transformer()
@@ -216,7 +221,7 @@ class TestProtocolEventRouting:
t.process(_proto_event(a, run_id="run-a"))
t.process(_proto_event(b, run_id="run-b"))
log.close()
streams = list(log._items)
streams = _unstamped(log._items)
assert len(streams) == 2
by_id = {s.message_id: s for s in streams}
assert by_id["run-a"].output.text == "aaaa"
@@ -227,7 +232,7 @@ class TestProtocolEventRouting:
for evt in _lifecycle(text="abcdef"):
t.process(_proto_event(evt))
log.close()
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert "".join(stream._text_proj._deltas) == "abcdef"
def test_stream_pushed_on_message_start_not_finish(self) -> None:
@@ -250,7 +255,7 @@ class TestProtocolEventRouting:
node="my_llm",
)
)
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert stream.node == "my_llm"
@@ -264,7 +269,7 @@ class TestWholeMessageFallback:
t, log = _make_sync_transformer()
t.process(_whole_msg("the full answer"))
log.close()
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert stream.done
assert stream.output.text == "the full answer"
@@ -272,7 +277,7 @@ class TestWholeMessageFallback:
t, log = _make_sync_transformer()
t.process(_whole_msg("full"))
log.close()
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert [e["event"] for e in stream._events] == [
"message-start",
"content-block-start",
@@ -318,7 +323,7 @@ class TestFiltering:
}
)
log.close()
assert list(log._items) == []
assert _unstamped(log._items) == []
def test_legacy_v1_chunks_ignored(self) -> None:
# v1 AIMessageChunk tuples (from on_llm_new_token) are not streamed
@@ -327,7 +332,7 @@ class TestFiltering:
t.process(_v1_chunk("hello"))
t.process(_v1_chunk(" world", finish=True))
log.close()
assert list(log._items) == []
assert _unstamped(log._items) == []
# ---------------------------------------------------------------------------
@@ -343,7 +348,7 @@ class TestLifecycle:
{"event": "message-start", "message_id": "run-1"}, run_id="run-1"
)
)
streams = list(log._items)
streams = _unstamped(log._items)
err = RuntimeError("graph died")
t.fail(err)
assert t._by_run == {}
@@ -371,14 +376,14 @@ class TestAsyncMode:
t, log = _make_async_transformer()
for evt in _lifecycle(text="async stream"):
t.process(_proto_event(evt))
assert isinstance(list(log._items)[0], AsyncChatModelStream)
assert isinstance(_unstamped(log._items)[0], AsyncChatModelStream)
@pytest.mark.anyio
async def test_text_projection_yields_deltas(self) -> None:
t, log = _make_async_transformer()
for evt in _lifecycle(text="hello world"):
t.process(_proto_event(evt))
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert isinstance(stream, AsyncChatModelStream)
assert "".join([d async for d in stream.text]) == "hello world"
@@ -387,7 +392,7 @@ class TestAsyncMode:
t, log = _make_async_transformer()
for evt in _lifecycle(text="async"):
t.process(_proto_event(evt))
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert (await stream.output).text == "async"
@@ -419,7 +424,7 @@ class TestWireRequestMore:
for evt in _lifecycle():
messages_t.process(_proto_event(evt))
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert stream._request_more is messages_t._pump_fn
@@ -445,14 +450,14 @@ class TestViaMux:
for evt in _lifecycle(text="mux stream"):
mux.push(_proto_event(evt))
mux.close()
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert stream.output.text == "mux stream"
def test_whole_message_via_mux(self) -> None:
t, mux, log = self._make_mux()
mux.push(_whole_msg("result"))
mux.close()
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert stream.output.text == "result"
@pytest.mark.anyio
@@ -466,7 +471,7 @@ class TestViaMux:
for evt in _lifecycle(text="async mux"):
await mux.apush(_proto_event(evt))
(stream,) = list(log._items)
(stream,) = _unstamped(log._items)
assert (await stream.output).text == "async mux"
await mux.aclose()
@@ -159,8 +159,13 @@ def _subgraph_transformer(mux: StreamMux) -> SubgraphTransformer:
return transformer
def _unstamped(items):
"""Strip push stamps from a StreamChannel's internal buffer."""
return [item for _stamp, item in items]
def _drain_subgraphs(mux: StreamMux) -> list[SubgraphRunStream]:
return list(_subgraph_transformer(mux)._log._items)
return _unstamped(_subgraph_transformer(mux)._log._items)
def _child_mux(handle: SubgraphRunStream | AsyncSubgraphRunStream) -> StreamMux:
@@ -169,13 +174,13 @@ def _child_mux(handle: SubgraphRunStream | AsyncSubgraphRunStream) -> StreamMux:
def _event_items(mux: StreamMux) -> list[ProtocolEvent]:
return list(mux._events._items)
return _unstamped(mux._events._items)
def _lifecycle_payloads(mux: StreamMux) -> list[dict[str, Any]]:
lifecycle_t = mux.transformer_by_key("lifecycle")
assert isinstance(lifecycle_t, LifecycleTransformer)
return list(lifecycle_t._channel._items)
return _unstamped(lifecycle_t._channel._items)
# ---------------------------------------------------------------------------
@@ -247,7 +252,7 @@ def test_grandchild_discovered_via_child_mini_mux() -> None:
[child_handle] = _drain_subgraphs(mux)
assert child_handle.path == ("agent:abc",)
# The grandchild appears on the CHILD'S subgraphs projection.
grandchildren = list(child_handle.subgraphs._items)
grandchildren = _unstamped(child_handle.subgraphs._items)
assert len(grandchildren) == 1
assert grandchildren[0].path == ("agent:abc", "tool:def")
+28
View File
@@ -19,9 +19,11 @@ from typing_extensions import TypedDict, assert_type
from langgraph._internal._constants import INTERRUPT
from langgraph.constants import END, START
from langgraph.errors import GraphDrained
from langgraph.func import entrypoint
from langgraph.graph import StateGraph
from langgraph.graph.message import MessagesState
from langgraph.runtime import RunControl
from langgraph.types import (
CheckpointPayload,
CheckpointStreamPart,
@@ -229,6 +231,32 @@ class TestV2Stream:
for c in chunks:
_assert_stream_part_shape(c)
def test_stream_v2_accepts_control_for_drain(self) -> None:
class DrainState(TypedDict, total=False):
value: str
skipped: str
control = RunControl()
def first_node(state: DrainState) -> dict[str, str]:
control.request_drain("sigterm")
return {"value": "done"}
def second_node(state: DrainState) -> dict[str, str]:
return {"skipped": "nope"}
builder = StateGraph(DrainState)
builder.add_node("first", first_node)
builder.add_node("second", second_node)
builder.add_edge(START, "first")
builder.add_edge("first", "second")
builder.add_edge("second", END)
graph = builder.compile()
run = graph.stream_v2({}, control=control)
with pytest.raises(GraphDrained, match="sigterm"):
list(run.values)
def test_subgraphs_ns(self) -> None:
outer = _make_subgraph()
chunks = list(
+3 -3
View File
@@ -1380,7 +1380,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.1.10"
version = "1.2.0a3"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
@@ -1561,7 +1561,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.3"
version = "4.1.0a3"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -1609,7 +1609,7 @@ test = [
[[package]]
name = "langgraph-checkpoint-postgres"
version = "3.0.5"
version = "3.1.0a3"
source = { editable = "../checkpoint-postgres" }
dependencies = [
{ name = "langgraph-checkpoint" },
@@ -44,7 +44,7 @@ import inspect
import json
from collections.abc import Awaitable, Callable
from copy import copy, deepcopy
from dataclasses import dataclass, replace
from dataclasses import dataclass, field, replace
from types import UnionType
from typing import (
TYPE_CHECKING,
@@ -1723,9 +1723,9 @@ class ToolRuntime(_DirectlyInjectedToolArg, Generic[ContextT, StateT]):
context: ContextT
config: RunnableConfig
stream_writer: StreamWriter
tools: list[BaseTool]
tool_call_id: str | None
store: BaseStore | None
tools: list[BaseTool] = field(default_factory=list)
execution_info: ExecutionInfo | None = None
server_info: ServerInfo | None = None
@@ -30,6 +30,11 @@ from langgraph.prebuilt._tool_call_stream import ToolCallStream
TS = int(time.time() * 1000)
def _unstamped(items):
"""Strip push stamps from a StreamChannel's internal buffer."""
return [item for _stamp, item in items]
def _tool_event(
event: str,
tool_call_id: str,
@@ -95,7 +100,7 @@ class TestToolCallTransformerUnit:
input={"text": "hi"},
)
)
handles = list(transformer._log._items)
handles = _unstamped(transformer._log._items)
assert len(handles) == 1
h = handles[0]
assert isinstance(h, ToolCallStream)
@@ -111,7 +116,7 @@ class TestToolCallTransformerUnit:
mux.push(_tool_event("tool-output-delta", "tc1", delta="a"))
mux.push(_tool_event("tool-output-delta", "tc1", delta="b"))
stream = transformer._active["tc1"]
assert list(stream._output_deltas._items) == ["a", "b"]
assert _unstamped(stream._output_deltas._items) == ["a", "b"]
def test_finish_closes_stream(self) -> None:
mux, transformer = _mux()
@@ -142,14 +147,14 @@ class TestToolCallTransformerUnit:
mux.push(_tool_event("tool-output-delta", "a", delta="A1"))
mux.push(_tool_event("tool-output-delta", "b", delta="B1"))
mux.push(_tool_event("tool-output-delta", "a", delta="A2"))
assert list(transformer._active["a"]._output_deltas._items) == ["A1", "A2"]
assert list(transformer._active["b"]._output_deltas._items) == ["B1"]
assert _unstamped(transformer._active["a"]._output_deltas._items) == ["A1", "A2"]
assert _unstamped(transformer._active["b"]._output_deltas._items) == ["B1"]
def test_tools_event_passes_through_main_log(self) -> None:
mux, transformer = _mux()
_subscribe(mux._events)
mux.push(_tool_event("tool-started", "tc1", tool_name="echo"))
kept = [e for e in mux._events._items if e["method"] == "tools"]
kept = [e for e in _unstamped(mux._events._items) if e["method"] == "tools"]
assert len(kept) == 1
+13
View File
@@ -2016,6 +2016,19 @@ async def test_tool_node_inject_runtime_dynamic_tool_via_wrap_tool_call_async()
assert tool_message.tool_call_id == "call_dynamic_2"
def test_tool_runtime_defaults_tools_to_empty_list() -> None:
runtime = ToolRuntime(
state={},
context=None,
config={},
stream_writer=lambda *args, **kwargs: None,
tool_call_id=None,
store=None,
)
assert runtime.tools == []
def test_tool_runtime_forwards_execution_info_server_info_and_tools() -> None:
"""Test that execution_info, server_info, and tools are forwarded from Runtime to ToolRuntime."""
from langgraph.runtime import ExecutionInfo, ServerInfo
+3 -3
View File
@@ -281,7 +281,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.1.10"
version = "1.2.0a3"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },
@@ -365,7 +365,7 @@ test = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.3"
version = "4.1.0a3"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -413,7 +413,7 @@ test = [
[[package]]
name = "langgraph-checkpoint-postgres"
version = "3.0.5"
version = "3.1.0a3"
source = { editable = "../checkpoint-postgres" }
dependencies = [
{ name = "langgraph-checkpoint" },
+2 -2
View File
@@ -298,7 +298,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.1.10"
version = "1.2.0a3"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },
@@ -382,7 +382,7 @@ test = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.3"
version = "4.1.0a3"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },