feat(langgraph): add node-level error handlers (#7233)

This commit is contained in:
Quanzheng Long
2026-04-29 21:27:06 -07:00
committed by GitHub
parent 9c1d65695e
commit 63d861165f
12 changed files with 989 additions and 23 deletions
@@ -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
+22
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
@@ -20,6 +21,7 @@ __all__ = (
"InvalidUpdateError",
"GraphBubbleUp",
"GraphInterrupt",
"NodeError",
"NodeInterrupt",
"NodeTimeoutError",
"ParentCommand",
@@ -142,6 +144,26 @@ class TaskNotFound(Exception):
pass
@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(TimeoutError):
"""Raised when a node invocation exceeds one of its configured timeouts.
+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:
+91 -1
View File
@@ -51,6 +51,7 @@ from langgraph._internal._constants import (
CONFIG_KEY_TASK_ID,
CONFIG_KEY_THREAD_ID,
ERROR,
ERROR_SOURCE_NODE,
INPUT,
INTERRUPT,
NS_END,
@@ -88,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,
@@ -522,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.
@@ -650,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))
@@ -1227,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)
@@ -1434,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()
+7
View File
@@ -730,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,
@@ -754,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],
@@ -801,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 ()
@@ -2871,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:
@@ -3322,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:
+125
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
@@ -9763,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)
+301 -2
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
@@ -1755,3 +1756,301 @@ 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
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": ""})