diff --git a/libs/langgraph/langgraph/_internal/_constants.py b/libs/langgraph/langgraph/_internal/_constants.py index f2c57f3ca..360f7f275 100644 --- a/libs/langgraph/langgraph/_internal/_constants.py +++ b/libs/langgraph/langgraph/_internal/_constants.py @@ -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, diff --git a/libs/langgraph/langgraph/_internal/_runnable.py b/libs/langgraph/langgraph/_internal/_runnable.py index 2c1a55ffa..0c110b96f 100644 --- a/libs/langgraph/langgraph/_internal/_runnable.py +++ b/libs/langgraph/langgraph/_internal/_runnable.py @@ -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 diff --git a/libs/langgraph/langgraph/errors.py b/libs/langgraph/langgraph/errors.py index a99546f10..4e7a0b6a7 100644 --- a/libs/langgraph/langgraph/errors.py +++ b/libs/langgraph/langgraph/errors.py @@ -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. diff --git a/libs/langgraph/langgraph/graph/_node.py b/libs/langgraph/langgraph/graph/_node.py index d464f06f9..d8238c2b3 100644 --- a/libs/langgraph/langgraph/graph/_node.py +++ b/libs/langgraph/langgraph/graph/_node.py @@ -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 diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 63f344c94..2bc1a7239 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -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, ) diff --git a/libs/langgraph/langgraph/pregel/_algo.py b/libs/langgraph/langgraph/pregel/_algo.py index a77e31603..103f6cce0 100644 --- a/libs/langgraph/langgraph/pregel/_algo.py +++ b/libs/langgraph/langgraph/pregel/_algo.py @@ -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: diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 80b9e7b42..aeb8d52da 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -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) diff --git a/libs/langgraph/langgraph/pregel/_read.py b/libs/langgraph/langgraph/pregel/_read.py index d90a69483..7d49fce85 100644 --- a/libs/langgraph/langgraph/pregel/_read.py +++ b/libs/langgraph/langgraph/pregel/_read.py @@ -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: diff --git a/libs/langgraph/langgraph/pregel/_runner.py b/libs/langgraph/langgraph/pregel/_runner.py index 3945bbf01..979935c9a 100644 --- a/libs/langgraph/langgraph/pregel/_runner.py +++ b/libs/langgraph/langgraph/pregel/_runner.py @@ -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() diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index fc9428694..2680fa305 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -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: diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index e61a40ecd..85393c73b 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -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) diff --git a/libs/langgraph/tests/test_retry.py b/libs/langgraph/tests/test_retry.py index af538919b..569811242 100644 --- a/libs/langgraph/tests/test_retry.py +++ b/libs/langgraph/tests/test_retry.py @@ -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": ""}) +