mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-02 14:28:46 +02:00
feat(langgraph): add node-level error handlers (#7233)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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": ""})
|
||||
|
||||
|
||||
Reference in New Issue
Block a user