diff --git a/libs/langgraph/langgraph/_internal/_constants.py b/libs/langgraph/langgraph/_internal/_constants.py index 68cb48fe8..c51c5c77b 100644 --- a/libs/langgraph/langgraph/_internal/_constants.py +++ b/libs/langgraph/langgraph/_internal/_constants.py @@ -56,6 +56,8 @@ CONFIG_KEY_CHECKPOINT_NS = sys.intern("checkpoint_ns") # holds the current checkpoint_ns, "" for root graph CONFIG_KEY_NODE_FINISHED = sys.intern("__pregel_node_finished") # holds a callback to be called when a node is finished +CONFIG_KEY_TIMED_ATTEMPT_OBSERVER = sys.intern("__pregel_timed_attempt_observer") +# holds a callback to be called when a timed node attempt starts or finishes CONFIG_KEY_SCRATCHPAD = sys.intern("__pregel_scratchpad") # holds a mutable dict for temporary storage scoped to the current task CONFIG_KEY_RUNNER_SUBMIT = sys.intern("__pregel_runner_submit") @@ -106,6 +108,7 @@ RESERVED = { CONFIG_KEY_CHECKPOINT_MAP, CONFIG_KEY_CHECKPOINT_ID, CONFIG_KEY_CHECKPOINT_NS, + CONFIG_KEY_TIMED_ATTEMPT_OBSERVER, CONFIG_KEY_RESUME_MAP, # other constants PUSH, diff --git a/libs/langgraph/langgraph/_internal/_runnable.py b/libs/langgraph/langgraph/_internal/_runnable.py index 63e03f544..eebb02183 100644 --- a/libs/langgraph/langgraph/_internal/_runnable.py +++ b/libs/langgraph/langgraph/_internal/_runnable.py @@ -706,7 +706,12 @@ class RunnableSeq(Runnable): step.ainvoke(input, config, **kwargs), context=context ) else: - input = await step.ainvoke(input, config, **kwargs) + with set_config_context(config) as context: + input = await context.run( + lambda: asyncio.create_task( + step.ainvoke(input, config, **kwargs) + ) + ) else: input = await step.ainvoke(input, config) # finish the root run diff --git a/libs/langgraph/langgraph/_internal/_timeout.py b/libs/langgraph/langgraph/_internal/_timeout.py new file mode 100644 index 000000000..3eac4d589 --- /dev/null +++ b/libs/langgraph/langgraph/_internal/_timeout.py @@ -0,0 +1,26 @@ +from __future__ import annotations + +from datetime import timedelta +from typing import Literal + +_SYNC_TIMEOUT_PREFIX = ( + "Node timeouts are only supported for async nodes because sync Python " + "execution cannot be safely cancelled in-process." +) + + +def coerce_timeout(value: float | timedelta | None) -> float | None: + """Normalize a timeout to positive seconds, or None if unset.""" + if value is None: + return None + seconds = value.total_seconds() if isinstance(value, timedelta) else float(value) + if seconds <= 0: + raise ValueError("timeout must be greater than 0") + return seconds + + +def sync_timeout_unsupported( + name: str, *, kind: Literal["Node", "Task"] = "Node" +) -> ValueError: + """Build the canonical error for using `timeout` with a sync target.""" + return ValueError(f"{_SYNC_TIMEOUT_PREFIX} {kind} {name!r} is sync.") diff --git a/libs/langgraph/langgraph/errors.py b/libs/langgraph/langgraph/errors.py index fb648879e..53a3660b7 100644 --- a/libs/langgraph/langgraph/errors.py +++ b/libs/langgraph/langgraph/errors.py @@ -20,6 +20,7 @@ __all__ = ( "GraphBubbleUp", "GraphInterrupt", "NodeInterrupt", + "NodeTimeoutError", "ParentCommand", "EmptyInputError", "TaskNotFound", @@ -125,3 +126,25 @@ class TaskNotFound(Exception): """Raised when the executor is unable to find a task (for distributed mode).""" pass + + +class NodeTimeoutError(TimeoutError): + """Raised when a node invocation exceeds its configured `timeout`. + + Subclasses the built-in `TimeoutError`, so existing `except TimeoutError` + handlers keep working. If the node has a `retry_policy` whose `retry_on` + permits `TimeoutError`, the attempt will be retried. + """ + + node: str + timeout: float + elapsed: float + + def __init__(self, node: str, timeout: float, elapsed: float) -> None: + super().__init__( + f"Node '{node}' exceeded its timeout of {timeout:.3f}s " + f"(elapsed: {elapsed:.3f}s)." + ) + self.node = node + self.timeout = timeout + self.elapsed = elapsed diff --git a/libs/langgraph/langgraph/func/__init__.py b/libs/langgraph/langgraph/func/__init__.py index c7443e0a2..e0b0a955a 100644 --- a/libs/langgraph/langgraph/func/__init__.py +++ b/libs/langgraph/langgraph/func/__init__.py @@ -5,6 +5,7 @@ import inspect import warnings from collections.abc import Awaitable, Callable, Sequence from dataclasses import dataclass +from datetime import timedelta from typing import ( Any, Generic, @@ -22,6 +23,8 @@ from typing_extensions import Unpack from langgraph._internal import _serde from langgraph._internal._constants import CACHE_NS_WRITES, PREVIOUS +from langgraph._internal._runnable import is_async_callable +from langgraph._internal._timeout import coerce_timeout, sync_timeout_unsupported from langgraph._internal._typing import MISSING, DeprecatedKwargs from langgraph.channels.ephemeral_value import EphemeralValue from langgraph.channels.last_value import LastValue @@ -51,6 +54,7 @@ class _TaskFunction(Generic[P, T]): *, retry_policy: Sequence[RetryPolicy], cache_policy: CachePolicy[Callable[P, str | bytes]] | None = None, + timeout: float | None = None, name: str | None = None, ) -> None: if name is not None: @@ -67,6 +71,7 @@ class _TaskFunction(Generic[P, T]): self.func = func self.retry_policy = retry_policy self.cache_policy = cache_policy + self.timeout = timeout functools.update_wrapper(self, func) def __call__(self, *args: P.args, **kwargs: P.kwargs) -> SyncAsyncFuture[T]: @@ -74,6 +79,7 @@ class _TaskFunction(Generic[P, T]): self.func, retry_policy=self.retry_policy, cache_policy=self.cache_policy, + timeout=self.timeout, *args, **kwargs, ) @@ -98,6 +104,7 @@ def task( name: str | None = None, retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy[Callable[P, str | bytes]] | None = None, + timeout: float | timedelta | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> Callable[ [Callable[P, Awaitable[T]] | Callable[P, T]], @@ -119,6 +126,7 @@ def task( name: str | None = None, retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy[Callable[P, str | bytes]] | None = None, + timeout: float | timedelta | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> ( Callable[[Callable[P, Awaitable[T]] | Callable[P, T]], _TaskFunction[P, T]] @@ -142,6 +150,9 @@ def task( name: An optional name for the task. If not provided, the function name will be used. retry_policy: An optional retry policy (or list of policies) to use for the task in case of a failure. cache_policy: An optional cache policy to use for the task. This allows caching of the task results. + timeout: Maximum wall-clock duration for a single task attempt, in seconds + (or as a `timedelta`). If exceeded, `NodeTimeoutError` is raised. + Supported only for async tasks. Returns: A callable function when used as a decorator. @@ -196,6 +207,7 @@ def task( ) if retry_policy is None: retry_policy = retry # type: ignore[assignment] + timeout_s = coerce_timeout(timeout) retry_policies: Sequence[RetryPolicy] = ( () @@ -208,8 +220,15 @@ def task( def decorator( func: Callable[P, Awaitable[T]] | Callable[P, T], ) -> Callable[P, SyncAsyncFuture[T]]: + if timeout_s is not None and not is_async_callable(func): + name_ = name or getattr(func, "__name__", func.__class__.__name__) + raise sync_timeout_unsupported(str(name_), kind="Task") return _TaskFunction( - func, retry_policy=retry_policies, cache_policy=cache_policy, name=name + func, + retry_policy=retry_policies, + cache_policy=cache_policy, + timeout=timeout_s, + name=name, ) if __func_or_none__ is not None: @@ -400,6 +419,7 @@ class entrypoint(Generic[ContextT]): context_schema: type[ContextT] | None = None, cache_policy: CachePolicy | None = None, retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, + timeout: float | timedelta | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> None: """Initialize the entrypoint decorator.""" @@ -426,6 +446,7 @@ class entrypoint(Generic[ContextT]): self.cache = cache self.cache_policy = cache_policy self.retry_policy = retry_policy + self.timeout = coerce_timeout(timeout) self.context_schema = context_schema @dataclass(**_DC_KWARGS) @@ -535,6 +556,7 @@ class entrypoint(Generic[ContextT]): bound=bound, triggers=[START], channels=START, + timeout=self.timeout, writers=[ ChannelWrite( [ diff --git a/libs/langgraph/langgraph/graph/_node.py b/libs/langgraph/langgraph/graph/_node.py index cadf097d9..b5c4d4a3e 100644 --- a/libs/langgraph/langgraph/graph/_node.py +++ b/libs/langgraph/langgraph/graph/_node.py @@ -90,3 +90,4 @@ class StateNodeSpec(Generic[NodeInputT, ContextT]): cache_policy: CachePolicy | None ends: tuple[str, ...] | dict[str, str] | None = EMPTY_SEQ defer: bool = False + timeout: float | None = None diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index b1c24de2b..8d96508ff 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -7,6 +7,7 @@ import warnings from collections import defaultdict from collections.abc import Awaitable, Callable, Hashable, Sequence from dataclasses import is_dataclass +from datetime import timedelta from functools import partial from inspect import isclass, isfunction, ismethod, signature from types import FunctionType @@ -45,6 +46,7 @@ from langgraph._internal._fields import ( ) from langgraph._internal._pydantic import create_model from langgraph._internal._runnable import coerce_to_runnable +from langgraph._internal._timeout import coerce_timeout from langgraph._internal._typing import EMPTY_SEQ, MISSING, DeprecatedKwargs from langgraph.channels.base import BaseChannel from langgraph.channels.binop import BinaryOperatorAggregate @@ -300,6 +302,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy | None = None, destinations: dict[str, str] | tuple[str, ...] | None = None, + timeout: float | timedelta | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> Self: """Add a new node to the `StateGraph`, input schema is inferred as the state schema. @@ -367,6 +370,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy | None = None, destinations: dict[str, str] | tuple[str, ...] | None = None, + timeout: float | timedelta | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> Self: """Add a new node to the `StateGraph` where input schema is specified. @@ -439,6 +443,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy | None = None, destinations: dict[str, str] | tuple[str, ...] | None = None, + timeout: float | timedelta | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> Self: """Add a new node to the `StateGraph`, input schema is inferred as the state schema. @@ -506,6 +511,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy | None = None, destinations: dict[str, str] | tuple[str, ...] | None = None, + timeout: float | timedelta | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> Self: """Add a new node to the `StateGraph`, input schema is specified. @@ -580,6 +586,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy | None = None, destinations: dict[str, str] | tuple[str, ...] | None = None, + timeout: float | timedelta | None = None, **kwargs: Unpack[DeprecatedKwargs], ) -> Self: """Add a new node to the `StateGraph`. @@ -609,6 +616,12 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): !!! warning This is only used for graph rendering and doesn't have any effect on the graph execution. + timeout: Maximum wall-clock duration for a single invocation of this + node, in seconds (or as a `timedelta`). When exceeded, a + [`NodeTimeoutError`][langgraph.errors.NodeTimeoutError] is raised + and the retry policy (if any) decides whether to retry. Timeouts + are supported only for async nodes; sync nodes cannot be safely + cancelled in-process. Example: ```python @@ -662,6 +675,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): ) if input_schema is None: input_schema = cast(type[NodeInputT] | None, input_) + timeout = coerce_timeout(timeout) if not isinstance(node, str): action = node @@ -757,6 +771,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): cache_policy=cache_policy, ends=ends, defer=defer, + timeout=timeout, ) elif inferred_input_schema is not None: self.nodes[node] = StateNodeSpec( @@ -767,6 +782,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): cache_policy=cache_policy, ends=ends, defer=defer, + timeout=timeout, ) else: self.nodes[node] = StateNodeSpec[StateT, ContextT]( @@ -777,6 +793,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): cache_policy=cache_policy, ends=ends, defer=defer, + timeout=timeout, ) input_schema = input_schema or inferred_input_schema @@ -1332,6 +1349,7 @@ class CompiledStateGraph( retry_policy=node.retry_policy, cache_policy=node.cache_policy, bound=node.runnable, # type: ignore[arg-type] + timeout=node.timeout, ) else: raise RuntimeError diff --git a/libs/langgraph/langgraph/pregel/_algo.py b/libs/langgraph/langgraph/pregel/_algo.py index d7157e239..d8f83d519 100644 --- a/libs/langgraph/langgraph/pregel/_algo.py +++ b/libs/langgraph/langgraph/pregel/_algo.py @@ -7,6 +7,7 @@ import threading from collections import defaultdict, deque from collections.abc import Callable, Iterable, Mapping, Sequence from copy import copy +from datetime import timedelta from functools import partial from hashlib import sha1 from typing import ( @@ -61,6 +62,7 @@ from langgraph._internal._constants import ( TASKS, ) from langgraph._internal._scratchpad import PregelScratchpad +from langgraph._internal._timeout import coerce_timeout from langgraph._internal._typing import EMPTY_SEQ, MISSING from langgraph.channels.base import BaseChannel from langgraph.channels.topic import Topic @@ -114,13 +116,21 @@ class PregelTaskWrites(NamedTuple): class Call: - __slots__ = ("func", "input", "retry_policy", "cache_policy", "callbacks") + __slots__ = ( + "func", + "input", + "retry_policy", + "cache_policy", + "callbacks", + "timeout", + ) func: Callable input: tuple[tuple[Any, ...], dict[str, Any]] retry_policy: Sequence[RetryPolicy] | None cache_policy: CachePolicy | None callbacks: Callbacks + timeout: float | None def __init__( self, @@ -130,12 +140,14 @@ class Call: retry_policy: Sequence[RetryPolicy] | None, cache_policy: CachePolicy | None, callbacks: Callbacks, + timeout: float | timedelta | None = None, ) -> None: self.func = func self.input = input self.retry_policy = retry_policy self.cache_policy = cache_policy self.callbacks = callbacks + self.timeout = coerce_timeout(timeout) def should_interrupt( @@ -733,6 +745,7 @@ def prepare_single_task( task_path[:3], writers=proc.flat_writers, subgraphs=proc.subgraphs, + timeout=proc.timeout, ) else: return PregelTask(task_id, name, task_path[:3]) @@ -870,6 +883,7 @@ def prepare_push_task_functional( cache_key, task_id, in_progress_task_path, + timeout=call.timeout, ) else: return PregelTask(task_id, name, in_progress_task_path) @@ -1041,6 +1055,7 @@ def prepare_push_task_send( translated_task_path, writers=proc.flat_writers, subgraphs=proc.subgraphs, + timeout=proc.timeout, ) else: return PregelTask(task_id, packet.node, translated_task_path) diff --git a/libs/langgraph/langgraph/pregel/_call.py b/libs/langgraph/langgraph/pregel/_call.py index 0cd007042..0b2b5af4c 100644 --- a/libs/langgraph/langgraph/pregel/_call.py +++ b/libs/langgraph/langgraph/pregel/_call.py @@ -8,6 +8,7 @@ import inspect import sys import types from collections.abc import Awaitable, Callable, Generator, Sequence +from datetime import timedelta from typing import Any, Generic, TypeVar, cast from langchain_core.runnables import Runnable @@ -20,6 +21,7 @@ from langgraph._internal._runnable import ( is_async_callable, run_in_executor, ) +from langgraph._internal._timeout import coerce_timeout, sync_timeout_unsupported from langgraph.config import get_config from langgraph.pregel._write import ChannelWrite, ChannelWriteEntry from langgraph.types import CachePolicy, RetryPolicy @@ -255,8 +257,13 @@ def call( *args: Any, retry_policy: Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy | None = None, + timeout: float | timedelta | None = None, **kwargs: Any, ) -> SyncAsyncFuture[T]: + timeout_s = coerce_timeout(timeout) + if timeout_s is not None and not is_async_callable(func): + name = getattr(func, "__name__", func.__class__.__name__) + raise sync_timeout_unsupported(name, kind="Task") config = get_config() impl = config[CONF][CONFIG_KEY_CALL] fut = impl( @@ -265,5 +272,6 @@ def call( retry_policy=retry_policy, cache_policy=cache_policy, callbacks=config["callbacks"], + timeout=timeout_s, ) return fut diff --git a/libs/langgraph/langgraph/pregel/_read.py b/libs/langgraph/langgraph/pregel/_read.py index 8d4c21135..2c53f4485 100644 --- a/libs/langgraph/langgraph/pregel/_read.py +++ b/libs/langgraph/langgraph/pregel/_read.py @@ -1,6 +1,7 @@ from __future__ import annotations from collections.abc import AsyncIterator, Callable, Iterator, Mapping, Sequence +from datetime import timedelta from functools import cached_property from typing import ( Any, @@ -11,6 +12,7 @@ from langchain_core.runnables import Runnable, RunnableConfig from langgraph._internal._config import merge_configs from langgraph._internal._constants import CONF, CONFIG_KEY_READ from langgraph._internal._runnable import RunnableCallable, RunnableSeq +from langgraph._internal._timeout import coerce_timeout from langgraph.pregel._utils import find_subgraph_pregel from langgraph.pregel._write import ChannelWrite from langgraph.pregel.protocol import PregelProtocol @@ -123,6 +125,11 @@ class PregelNode: cache_policy: CachePolicy | None """The cache policy to use when invoking the node.""" + timeout: float | None + """Maximum time in seconds allowed for a single invocation of this node. + If exceeded, `NodeTimeoutError` is raised and the retry policy (if any) + decides whether to retry. Supported only for async nodes.""" + tags: Sequence[str] | None """Tags to attach to the node for tracing.""" @@ -145,6 +152,7 @@ class PregelNode: retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy | None = None, subgraphs: Sequence[PregelProtocol] | None = None, + timeout: float | timedelta | None = None, ) -> None: self.channels = channels self.triggers = list(triggers) @@ -156,6 +164,7 @@ class PregelNode: self.retry_policy = (retry_policy,) else: self.retry_policy = retry_policy + self.timeout = coerce_timeout(timeout) self.tags = tags self.metadata = metadata if subgraphs is not None: diff --git a/libs/langgraph/langgraph/pregel/_retry.py b/libs/langgraph/langgraph/pregel/_retry.py index 538d6b915..9772ad62f 100644 --- a/libs/langgraph/langgraph/pregel/_retry.py +++ b/libs/langgraph/langgraph/pregel/_retry.py @@ -4,12 +4,16 @@ import asyncio import logging import random import sys +import threading import time -from collections.abc import Awaitable, Callable, Sequence +from collections.abc import Awaitable, Callable, Coroutine, Sequence +from contextlib import suppress from dataclasses import replace -from typing import Any +from datetime import datetime, timedelta, timezone +from typing import Any, Literal from langchain_core.runnables import RunnableConfig +from typing_extensions import NotRequired, TypedDict from langgraph._internal._config import patch_configurable, recast_checkpoint_ns from langgraph._internal._constants import ( @@ -18,11 +22,14 @@ from langgraph._internal._constants import ( CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_RESUMING, CONFIG_KEY_RUNTIME, + CONFIG_KEY_SEND, CONFIG_KEY_TASK_ID, CONFIG_KEY_THREAD_ID, + CONFIG_KEY_TIMED_ATTEMPT_OBSERVER, NS_SEP, ) -from langgraph.errors import GraphBubbleUp, ParentCommand +from langgraph._internal._timeout import sync_timeout_unsupported +from langgraph.errors import GraphBubbleUp, NodeTimeoutError, ParentCommand from langgraph.runtime import ExecutionInfo, Runtime from langgraph.types import Command, PregelExecutableTask, RetryPolicy @@ -30,6 +37,182 @@ logger = logging.getLogger(__name__) SUPPORTS_EXC_NOTES = sys.version_info >= (3, 11) +class _TimedAttemptPayload(TypedDict): + execution_id: str + task_id: str + task_name: str + attempt: int + run_id: str | None + thread_id: str | None + checkpoint_ns: str | None + started_at: datetime + deadline_at: datetime + timeout_secs: float + event: Literal["start", "finish"] + finished_at: NotRequired[datetime] + status: NotRequired[Literal["success", "error"]] + error_type: NotRequired[str | None] + error_message: NotRequired[str | None] + + +class _TimedAttemptScope: + """Guarded-config window for timed attempts. + + `close()` and the guarded send are serialized so writes from a cancelled + background task cannot slip past the timeout boundary. + """ + + __slots__ = ("_active", "_lock") + + def __init__(self) -> None: + self._active = True + self._lock = threading.Lock() + + def wrap_config(self, config: RunnableConfig) -> RunnableConfig: + configurable = config.get(CONF, {}) + if (send := configurable.get(CONFIG_KEY_SEND)) is not None: + return patch_configurable(config, {CONFIG_KEY_SEND: self._guard_send(send)}) + return config + + def close(self) -> None: + with self._lock: + self._active = False + + def _guard_send( + self, send: Callable[[Sequence[tuple[str, Any]]], None] + ) -> Callable[[Sequence[tuple[str, Any]]], None]: + def guarded_send(writes: Sequence[tuple[str, Any]]) -> None: + with self._lock: + if self._active: + send(writes) + + return guarded_send + + +def _drain_cancelled(task: asyncio.Task[Any]) -> None: + # Mark the abandoned task's exception as retrieved so asyncio doesn't log it. + with suppress(asyncio.CancelledError): + task.exception() + + +def _create_task_with_config_context( + run: Callable[[], Coroutine[Any, Any, Any]], config: RunnableConfig +) -> asyncio.Task[Any]: + from langgraph._internal._runnable import set_config_context + + with set_config_context(config) as context: + return context.run(lambda: asyncio.create_task(run())) + + +def _start_timed_attempt( + task: PregelExecutableTask, config: RunnableConfig, timeout_s: float +) -> _TimedAttemptPayload | None: + configurable = config.get(CONF, {}) + callback = configurable.get(CONFIG_KEY_TIMED_ATTEMPT_OBSERVER) + if callback is None: + return None + runtime = configurable.get(CONFIG_KEY_RUNTIME) + execution_info = runtime.execution_info if isinstance(runtime, Runtime) else None + attempt = execution_info.node_attempt if execution_info is not None else 1 + run_id = execution_info.run_id if execution_info is not None else None + thread_id = ( + execution_info.thread_id + if execution_info is not None + else configurable.get(CONFIG_KEY_THREAD_ID) + ) + checkpoint_ns = ( + execution_info.checkpoint_ns + if execution_info is not None + else configurable.get(CONFIG_KEY_CHECKPOINT_NS) + ) + started_at = datetime.now(timezone.utc) + payload: _TimedAttemptPayload = { + "execution_id": f"run:{run_id or '-'}|task:{task.id}|attempt:{attempt}", + "task_id": task.id, + "task_name": task.name, + "attempt": attempt, + "run_id": run_id, + "thread_id": thread_id, + "checkpoint_ns": checkpoint_ns, + "started_at": started_at, + "deadline_at": started_at + timedelta(seconds=timeout_s), + "timeout_secs": timeout_s, + "event": "start", + } + _dispatch_observer(callback, payload) + return payload + + +def _finish_timed_attempt( + config: RunnableConfig, + payload: _TimedAttemptPayload | None, + error: BaseException | None = None, +) -> None: + if payload is None: + return + callback = config.get(CONF, {}).get(CONFIG_KEY_TIMED_ATTEMPT_OBSERVER) + if callback is None: + return + finish: _TimedAttemptPayload = { + **payload, + "event": "finish", + "finished_at": datetime.now(timezone.utc), + "status": "error" if error is not None else "success", + "error_type": type(error).__name__ if error is not None else None, + "error_message": str(error) if error is not None else None, + } + _dispatch_observer(callback, finish) + + +def _dispatch_observer( + callback: Callable[[_TimedAttemptPayload], None], payload: _TimedAttemptPayload +) -> None: + try: + callback(payload) + except Exception: + logger.warning("Timed attempt observer failed", exc_info=True) + + +async def _arun_with_timeout( + task: PregelExecutableTask, + config: RunnableConfig, + timeout_s: float, + *, + stream: bool, +) -> Any: + scope = _TimedAttemptScope() + scoped_config = scope.wrap_config(config) + start = time.monotonic() + if stream: + + async def run() -> Any: + async for _ in task.proc.astream(task.input, scoped_config): + pass + + else: + + async def run() -> Any: + return await task.proc.ainvoke(task.input, scoped_config) + + bg = _create_task_with_config_context(run, scoped_config) + try: + return await asyncio.wait_for(asyncio.shield(bg), timeout=timeout_s) + except asyncio.TimeoutError as exc: + elapsed = time.monotonic() - start + scope.close() + task.writes.clear() + bg.cancel() + bg.add_done_callback(_drain_cancelled) + raise NodeTimeoutError(task.name, timeout_s, elapsed) from exc + except asyncio.CancelledError: + scope.close() + bg.cancel() + bg.add_done_callback(_drain_cancelled) + raise + finally: + scope.close() + + def _ensure_execution_info( runtime: Runtime, config: RunnableConfig, task: PregelExecutableTask ) -> Runtime: @@ -90,6 +273,11 @@ def run_with_retry( ) -> None: """Run a task with retries.""" retry_policy = task.retry_policy or retry_policy + if task.timeout is not None: + # `validate_timeout_supported` catches sync nodes at compile time; + # this is a runtime safety net for paths (e.g. distributed runtime) + # that may bypass that validation. + raise sync_timeout_unsupported(task.name) attempts = 0 node_first_attempt_time = time.time() config = task.config @@ -195,6 +383,7 @@ async def arun_with_retry( ) -> None: """Run a task asynchronously with retries.""" retry_policy = task.retry_policy or retry_policy + timeout_s = task.timeout attempts = 0 node_first_attempt_time = time.time() config = task.config @@ -229,35 +418,47 @@ async def arun_with_retry( ) }, ) + attempt_payload = ( + _start_timed_attempt(task, config, timeout_s) + if timeout_s is not None + else None + ) try: - # clear any writes from previous attempts task.writes.clear() - # run the task - if stream: - async for _ in task.proc.astream(task.input, config): - pass - # if successful, end - break - else: + if timeout_s is None: + if stream: + async for _ in task.proc.astream(task.input, config): + pass + break return await task.proc.ainvoke(task.input, config) + result = await _arun_with_timeout(task, config, timeout_s, stream=stream) + _finish_timed_attempt(config, attempt_payload) + if stream: + break + return result except ParentCommand as exc: ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS] cmd = exc.args[0] # strip task_ids from namespace for comparison (ns format: "node1|node2:task_id") if cmd.graph in (ns, recast_checkpoint_ns(ns), task.name): - # this command is for the current graph, handle it - for w in task.writers: - w.invoke(cmd, config) + try: + for w in task.writers: + w.invoke(cmd, config) + except Exception as writer_exc: + _finish_timed_attempt(config, attempt_payload, writer_exc) + raise + _finish_timed_attempt(config, attempt_payload) break elif cmd.graph == Command.PARENT: # this command is for the parent graph, assign it to the parent. exc.args = (replace(cmd, graph=_checkpoint_ns_for_parent_command(ns)),) - # bubble up + _finish_timed_attempt(config, attempt_payload) raise except GraphBubbleUp: - # if interrupted, end + _finish_timed_attempt(config, attempt_payload) raise except Exception as exc: + _finish_timed_attempt(config, attempt_payload, exc) if SUPPORTS_EXC_NOTES: exc.add_note(f"During task with name '{task.name}' and id '{task.id}'") if not retry_policy: diff --git a/libs/langgraph/langgraph/pregel/_runner.py b/libs/langgraph/langgraph/pregel/_runner.py index fea4a7272..e7a352c4e 100644 --- a/libs/langgraph/langgraph/pregel/_runner.py +++ b/libs/langgraph/langgraph/pregel/_runner.py @@ -14,6 +14,7 @@ from collections.abc import ( Iterator, Sequence, ) +from datetime import timedelta from functools import partial from typing import ( Any, @@ -537,6 +538,7 @@ def _call( *, retry_policy: Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy | None = None, + timeout: float | timedelta | None = None, callbacks: Callbacks = None, futures: weakref.ref[FuturesDict], schedule_task: Callable[ @@ -560,6 +562,7 @@ def _call( retry_policy=retry_policy, cache_policy=cache_policy, callbacks=callbacks, + timeout=timeout, ), ): if fut := next( @@ -624,6 +627,7 @@ def _acall( *, retry_policy: Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy | None = None, + timeout: float | timedelta | None = None, callbacks: Callbacks = None, # injected dependencies futures: weakref.ref[FuturesDict], @@ -657,6 +661,7 @@ def _acall( input, retry_policy=retry_policy, cache_policy=cache_policy, + timeout=timeout, callbacks=callbacks, futures=futures, schedule_task=schedule_task, @@ -678,6 +683,7 @@ async def _acall_impl( *, retry_policy: Sequence[RetryPolicy] | None = None, cache_policy: CachePolicy | None = None, + timeout: float | timedelta | None = None, callbacks: Callbacks = None, # injected dependencies futures: weakref.ref[FuturesDict[asyncio.Future, asyncio.Event]], @@ -703,6 +709,7 @@ async def _acall_impl( retry_policy=retry_policy, cache_policy=cache_policy, callbacks=callbacks, + timeout=timeout, ), ): if fut := next( diff --git a/libs/langgraph/langgraph/pregel/_utils.py b/libs/langgraph/langgraph/pregel/_utils.py index 0c8a14eec..d4d7c103e 100644 --- a/libs/langgraph/langgraph/pregel/_utils.py +++ b/libs/langgraph/langgraph/pregel/_utils.py @@ -4,16 +4,21 @@ import ast import inspect import re import textwrap -from collections.abc import Callable +from collections.abc import Callable, Sequence +from functools import partial from typing import Any from langchain_core.runnables import Runnable, RunnableLambda, RunnableSequence +from langchain_core.runnables.config import run_in_executor from langgraph.checkpoint.base import ChannelVersions from typing_extensions import override from langgraph._internal._runnable import RunnableCallable, RunnableSeq +from langgraph._internal._timeout import sync_timeout_unsupported from langgraph.pregel.protocol import PregelProtocol +_SEQUENCE_TYPES = (RunnableSeq, RunnableSequence) + def get_new_channel_versions( previous_versions: ChannelVersions, current_versions: ChannelVersions @@ -64,6 +69,47 @@ def find_subgraph_pregel(candidate: Runnable) -> PregelProtocol | None: return None +def _sequence_steps(runnable: Runnable) -> Sequence[Runnable] | None: + if isinstance(runnable, _SEQUENCE_TYPES): + return runnable.steps + return None + + +def _has_method_override(runnable: Runnable, method_name: str) -> bool: + method = getattr(type(runnable), method_name, None) + return method is not None and method is not getattr(Runnable, method_name) + + +def _is_executor_backed_afunc(afunc: Callable[..., Any] | None) -> bool: + return isinstance(afunc, partial) and afunc.func is run_in_executor + + +def _has_native_async(runnable: Runnable) -> bool: + if isinstance(runnable, RunnableCallable): + return runnable.afunc is not None and not _is_executor_backed_afunc( + runnable.afunc + ) + if isinstance(runnable, RunnableLambda): + return bool(getattr(runnable, "afunc", False)) + return _has_method_override(runnable, "ainvoke") + + +def _runnable_has_native_async(runnable: Runnable) -> bool: + """Return whether a runnable can be timed without running sync code.""" + + if (steps := _sequence_steps(runnable)) is not None: + for step in steps: + if not _runnable_has_native_async(step): + return False + return True + return _has_native_async(runnable) + + +def validate_timeout_supported(runnable: Runnable, *, name: str) -> None: + if not _runnable_has_native_async(runnable): + raise sync_timeout_unsupported(name) + + def get_function_nonlocals(func: Callable) -> list[Any]: """Get the nonlocal variables accessed by a function. diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index c440e77f9..0b255085c 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -17,6 +17,7 @@ from collections.abc import ( Sequence, ) from dataclasses import is_dataclass, replace +from datetime import timedelta from functools import partial from inspect import isclass from typing import ( @@ -95,6 +96,7 @@ from langgraph._internal._runnable import ( RunnableSeq, coerce_to_runnable, ) +from langgraph._internal._timeout import coerce_timeout from langgraph._internal._typing import MISSING, DeprecatedKwargs from langgraph.callbacks import ( GraphInterruptEvent, @@ -137,7 +139,7 @@ from langgraph.pregel._messages import StreamMessagesHandler from langgraph.pregel._read import DEFAULT_BOUND, PregelNode from langgraph.pregel._retry import RetryPolicy from langgraph.pregel._runner import PregelRunner -from langgraph.pregel._utils import get_new_channel_versions +from langgraph.pregel._utils import get_new_channel_versions, validate_timeout_supported from langgraph.pregel._validate import validate_graph, validate_keys from langgraph.pregel._write import ChannelWrite, ChannelWriteEntry from langgraph.pregel.debug import get_bolded_text, get_colored_text, tasks_w_writes @@ -186,6 +188,7 @@ class NodeBuilder: "_bound", "_retry_policy", "_cache_policy", + "_timeout", ) _channels: str | list[str] @@ -196,6 +199,7 @@ class NodeBuilder: _bound: Runnable _retry_policy: list[RetryPolicy] _cache_policy: CachePolicy | None + _timeout: float | None def __init__( self, @@ -208,6 +212,7 @@ class NodeBuilder: self._bound = DEFAULT_BOUND self._retry_policy = [] self._cache_policy = None + self._timeout = None def subscribe_only( self, @@ -326,6 +331,11 @@ class NodeBuilder: self._cache_policy = policy return self + def set_timeout(self, timeout: float | timedelta | None) -> Self: + """Set the per-attempt timeout for this node.""" + self._timeout = coerce_timeout(timeout) + return self + def build(self) -> PregelNode: """Builds the node.""" return PregelNode( @@ -337,6 +347,7 @@ class NodeBuilder: bound=self._bound, retry_policy=self._retry_policy, cache_policy=self._cache_policy, + timeout=self._timeout, ) @@ -817,6 +828,9 @@ class Pregel( ) def validate(self) -> Self: + for name, node in self.nodes.items(): + if node.timeout is not None: + validate_timeout_supported(node.bound, name=name) validate_graph( self.nodes, {k: v for k, v in self.channels.items() if isinstance(v, BaseChannel)}, diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index d04d82da7..d2c4e0f64 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -548,6 +548,7 @@ class PregelExecutableTask: path: tuple[str | int | tuple, ...] writers: Sequence[Runnable] = () subgraphs: Sequence[PregelProtocol] = () + timeout: float | None = None class StateSnapshot(NamedTuple): diff --git a/libs/langgraph/tests/test_retry.py b/libs/langgraph/tests/test_retry.py index 3156f7599..4ba7fb163 100644 --- a/libs/langgraph/tests/test_retry.py +++ b/libs/langgraph/tests/test_retry.py @@ -1,7 +1,12 @@ +import asyncio +import threading +import time from collections import deque +from datetime import datetime, timedelta from unittest.mock import Mock, patch import pytest +from langchain_core.runnables import RunnableLambda from langgraph.checkpoint.memory import MemorySaver from typing_extensions import TypedDict @@ -10,18 +15,29 @@ from langgraph._internal._constants import ( CONFIG_KEY_CHECKPOINT_ID, CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_RUNTIME, + CONFIG_KEY_SEND, CONFIG_KEY_TASK_ID, CONFIG_KEY_THREAD_ID, + CONFIG_KEY_TIMED_ATTEMPT_OBSERVER, ) -from langgraph.graph import START, StateGraph +from langgraph._internal._runnable import RunnableCallable +from langgraph._internal._timeout import coerce_timeout +from langgraph.channels.ephemeral_value import EphemeralValue +from langgraph.channels.last_value import LastValue +from langgraph.errors import GraphInterrupt, NodeTimeoutError, ParentCommand +from langgraph.func import entrypoint, task +from langgraph.graph import END, START, StateGraph +from langgraph.pregel import NodeBuilder, Pregel +from langgraph.pregel._read import PregelNode from langgraph.pregel._retry import ( _checkpoint_ns_for_parent_command, _ensure_execution_info, _should_retry_on, + arun_with_retry, run_with_retry, ) from langgraph.runtime import DEFAULT_RUNTIME, ExecutionInfo, Runtime -from langgraph.types import PregelExecutableTask, RetryPolicy +from langgraph.types import Command, PregelExecutableTask, RetryPolicy def test_should_retry_on_single_exception(): @@ -567,3 +583,643 @@ def test_run_with_retry_creates_execution_info_when_missing(): assert info.run_id == "run-abc" assert info.node_attempt == 1 assert info.node_first_attempt_time is not None + + +def _make_task( + proc, *, timeout=None, retry_policy=(), name="timed", task_id="tid", writers=() +): + runtime = DEFAULT_RUNTIME.override(execution_info=None) + writes = deque() + config = { + "run_id": "run-x", + CONF: { + CONFIG_KEY_RUNTIME: runtime, + CONFIG_KEY_CHECKPOINT_ID: "cp", + CONFIG_KEY_CHECKPOINT_NS: f"{name}:{task_id}", + CONFIG_KEY_SEND: writes.extend, + CONFIG_KEY_TASK_ID: task_id, + CONFIG_KEY_THREAD_ID: "thr", + }, + } + return PregelExecutableTask( + name=name, + input=None, + proc=proc, + writes=writes, + config=config, + triggers=[name], + retry_policy=retry_policy, + cache_key=None, + id=task_id, + path=("__pregel_pull", name), + writers=writers, + timeout=coerce_timeout(timeout), + ) + + +def test_coerce_timeout(): + assert coerce_timeout(None) is None + assert coerce_timeout(1.5) == 1.5 + assert coerce_timeout(2) == 2.0 + assert coerce_timeout(timedelta(milliseconds=250)) == 0.25 + with pytest.raises(ValueError, match="greater than 0"): + coerce_timeout(0) + with pytest.raises(ValueError, match="greater than 0"): + coerce_timeout(timedelta()) + + +def test_run_with_retry_rejects_sync_timeout_without_starting_proc(): + started = False + + class Proc: + def invoke(self, input, config): + nonlocal started + started = True + return input + + task = _make_task(Proc(), timeout=0.05, name="sync") + + with pytest.raises(ValueError, match="only supported for async nodes"): + run_with_retry(task, retry_policy=None) + assert not started + + +def test_run_with_retry_without_timeout_runs_sync_directly(): + class FastProc: + def invoke(self, input, config): + return "ok" + + task = _make_task(FastProc(), timeout=None) + assert run_with_retry(task, retry_policy=None) == "ok" + + +def test_arun_with_retry_timeout_ok_when_fast(): + class FastProc: + async def ainvoke(self, input, config): + return "ok" + + task = _make_task(FastProc(), timeout=1.0) + + async def _run() -> None: + assert await arun_with_retry(task, retry_policy=None) == "ok" + + asyncio.run(_run()) + + +def test_arun_with_retry_timeout_retries_when_retry_on_timeout(): + calls: list[float] = [] + + class FlakyProc: + async def ainvoke(self, input, config): + calls.append(time.monotonic()) + if len(calls) < 2: + await asyncio.sleep(0.5) + return "late" + return "ok" + + policy = RetryPolicy( + max_attempts=3, + initial_interval=0.0, + jitter=False, + retry_on=NodeTimeoutError, + ) + task = _make_task(FlakyProc(), timeout=0.05, retry_policy=(policy,)) + + async def _run() -> None: + assert await arun_with_retry(task, retry_policy=None) == "ok" + assert len(calls) == 2 + + asyncio.run(_run()) + + +def test_entrypoint_timeout_allows_pre_timeout_child_task_to_run(): + child_started = threading.Event() + + @task() + def child(value: int) -> int: + child_started.set() + return value + 1 + + @entrypoint(timeout=0.05) + async def parent(value: int) -> int: + child(value) + await asyncio.sleep(0.2) + return value + + async def _run() -> None: + with pytest.raises(NodeTimeoutError): + await parent.ainvoke(1) + + asyncio.run(_run()) + assert child_started.wait(timeout=1.0) + + +def test_arun_with_retry_timeout_accepts_timedelta(): + class SlowProc: + async def ainvoke(self, input, config): + await asyncio.sleep(0.5) + return input + + task = _make_task(SlowProc(), timeout=timedelta(milliseconds=50)) + + async def _run() -> None: + with pytest.raises(NodeTimeoutError): + await arun_with_retry(task, retry_policy=None) + + asyncio.run(_run()) + + +def test_arun_with_retry_timeout_fires_async(): + class SlowProc: + async def ainvoke(self, input, config): + await asyncio.sleep(1.0) + return input + + task = _make_task(SlowProc(), timeout=0.05, name="aslow") + + async def _run(): + with pytest.raises(NodeTimeoutError) as excinfo: + await arun_with_retry(task, retry_policy=None) + assert excinfo.value.node == "aslow" + assert excinfo.value.timeout == 0.05 + + asyncio.run(_run()) + + +def test_arun_with_retry_timeout_discards_stale_executor_writes(): + release_first_attempt = threading.Event() + + class FlakyAsyncProc: + def __init__(self) -> None: + self.calls = 0 + + async def ainvoke(self, input, config): + self.calls += 1 + if self.calls == 1: + + def late_write() -> str: + release_first_attempt.wait(timeout=1.0) + config[CONF][CONFIG_KEY_SEND]([("value", "stale")]) + return "late" + + return await asyncio.to_thread(late_write) + release_first_attempt.set() + config[CONF][CONFIG_KEY_SEND]([("value", "fresh")]) + return "ok" + + policy = RetryPolicy( + max_attempts=2, + initial_interval=0.0, + jitter=False, + retry_on=NodeTimeoutError, + ) + task = _make_task(FlakyAsyncProc(), timeout=0.05, retry_policy=(policy,)) + + async def _run() -> None: + assert await arun_with_retry(task, retry_policy=None) == "ok" + await asyncio.sleep(0.05) + assert task.writes == deque([("value", "fresh")]) + + asyncio.run(_run()) + + +def test_arun_with_retry_timeout_discards_pre_timeout_writes(): + class SlowAsyncWriterProc: + async def ainvoke(self, input, config): + config[CONF][CONFIG_KEY_SEND]([("value", "stale-before-timeout")]) + await asyncio.sleep(0.2) + return "late" + + task = _make_task(SlowAsyncWriterProc(), timeout=0.05, name="aslow-writer") + + async def _run() -> None: + with pytest.raises(NodeTimeoutError): + await arun_with_retry(task, retry_policy=None) + assert task.writes == deque() + + asyncio.run(_run()) + + +def test_astream_with_retry_timeout_discards_pre_timeout_writes(): + class SlowStreamWriterProc: + async def astream(self, input, config): + config[CONF][CONFIG_KEY_SEND]([("value", "stale-before-timeout")]) + await asyncio.sleep(0.2) + if False: + yield None + + task = _make_task(SlowStreamWriterProc(), timeout=0.05, name="astream-writer") + + async def _run() -> None: + with pytest.raises(NodeTimeoutError): + await arun_with_retry(task, retry_policy=None, stream=True) + assert task.writes == deque() + + asyncio.run(_run()) + + +def test_arun_with_retry_timeout_cannot_be_swallowed(): + class StubbornProc: + async def ainvoke(self, input, config): + try: + await asyncio.sleep(1.0) + except asyncio.CancelledError: + config[CONF][CONFIG_KEY_SEND]([("value", "stale")]) + await asyncio.sleep(0) + return "late" + return "ok" + + task = _make_task(StubbornProc(), timeout=0.05, name="stubborn") + + async def _run() -> None: + with pytest.raises(NodeTimeoutError) as excinfo: + await arun_with_retry(task, retry_policy=None) + assert excinfo.value.node == "stubborn" + await asyncio.sleep(0.05) + assert task.writes == deque() + + asyncio.run(_run()) + + +def test_astream_with_retry_timeout_cannot_be_swallowed(): + class StubbornStreamProc: + async def astream(self, input, config): + try: + await asyncio.sleep(1.0) + except asyncio.CancelledError: + config[CONF][CONFIG_KEY_SEND]([("value", "stale")]) + await asyncio.sleep(0) + if False: + yield None + return + yield "ok" + + task = _make_task(StubbornStreamProc(), timeout=0.05, name="stubborn-stream") + + async def _run() -> None: + with pytest.raises(NodeTimeoutError) as excinfo: + await arun_with_retry(task, retry_policy=None, stream=True) + assert excinfo.value.node == "stubborn-stream" + await asyncio.sleep(0.05) + assert task.writes == deque() + + asyncio.run(_run()) + + +class _TimeoutState(TypedDict): + x: int + + +def test_timeout_validation_is_eager_across_apis(): + with pytest.raises(ValueError, match="greater than 0"): + task(timeout=0) + + with pytest.raises(ValueError, match="greater than 0"): + entrypoint(timeout=0) + + with pytest.raises(ValueError, match="greater than 0"): + NodeBuilder().set_timeout(0) + + with pytest.raises(ValueError, match="greater than 0"): + PregelNode(channels="x", triggers=["x"], timeout=0) + + builder = StateGraph(_TimeoutState) + with pytest.raises(ValueError, match="greater than 0"): + builder.add_node("slow", lambda state: state, timeout=0) + + +def test_timeout_rejects_sync_functional_apis_at_declaration_time(): + with pytest.raises(ValueError, match="only supported for async nodes"): + + @task(timeout=0.05) + def sync_task(value: int) -> int: + return value + + with pytest.raises(ValueError, match="only supported for async nodes"): + + @entrypoint(timeout=0.05) + def sync_entrypoint(value: int) -> int: + return value + + +def test_state_graph_compile_rejects_sync_node_timeout(): + def slow(state: _TimeoutState) -> _TimeoutState: + return {"x": state["x"] + 1} + + builder = StateGraph(_TimeoutState) + builder.add_node("slow", slow, timeout=0.05) + builder.add_edge(START, "slow") + builder.add_edge("slow", END) + + with pytest.raises(ValueError, match="only supported for async nodes"): + builder.compile() + + +def test_pregel_validate_rejects_sync_node_timeout(): + def slow(value: int) -> int: + return value + 1 + + with pytest.raises(ValueError, match="only supported for async nodes"): + Pregel( + nodes={ + "slow": ( + NodeBuilder() + .subscribe_only("input") + .do(slow) + .set_timeout(0.05) + .write_to("output") + ) + }, + channels={ + "input": EphemeralValue(int), + "output": LastValue(int), + }, + input_channels="input", + output_channels="output", + ) + + +def test_pregel_validate_accepts_async_runnable_lambda_timeout(): + async def slow(value: int) -> int: + await asyncio.sleep(0.2) + return value + 1 + + graph = Pregel( + nodes={ + "slow": ( + NodeBuilder() + .subscribe_only("input") + .do(RunnableLambda(slow)) + .set_timeout(0.05) + .write_to("output") + ) + }, + channels={ + "input": EphemeralValue(int), + "output": LastValue(int), + }, + input_channels="input", + output_channels="output", + ) + + async def _run() -> None: + with pytest.raises(NodeTimeoutError): + await graph.ainvoke(1) + + asyncio.run(_run()) + + +def test_pregel_validate_accepts_runnable_callable_with_sync_and_async_timeout(): + def sync(value: int) -> int: + return value + 1 + + async def async_(value: int) -> int: + await asyncio.sleep(0.2) + return value + 1 + + graph = Pregel( + nodes={ + "slow": ( + NodeBuilder() + .subscribe_only("input") + .do(RunnableCallable(sync, async_)) + .set_timeout(0.05) + .write_to("output") + ) + }, + channels={ + "input": EphemeralValue(int), + "output": LastValue(int), + }, + input_channels="input", + output_channels="output", + ) + + async def _run() -> None: + with pytest.raises(NodeTimeoutError): + await graph.ainvoke(1) + + asyncio.run(_run()) + + +def test_state_graph_add_node_timeout_e2e(): + async def slow(state: _TimeoutState) -> _TimeoutState: + await asyncio.sleep(1.0) + return {"x": state["x"] + 1} + + builder = StateGraph(_TimeoutState) + builder.add_node("slow", slow, timeout=0.05) + builder.add_edge(START, "slow") + builder.add_edge("slow", END) + graph = builder.compile() + + async def _run() -> None: + with pytest.raises(NodeTimeoutError): + await graph.ainvoke({"x": 1}) + + asyncio.run(_run()) + + +def test_state_graph_add_node_timeout_composes_with_retry(): + """add_node(..., timeout=...) + retry_policy retries then succeeds.""" + + attempts: list[int] = [] + + async def flaky(state: _TimeoutState) -> _TimeoutState: + attempts.append(len(attempts)) + if len(attempts) < 2: + await asyncio.sleep(0.5) + return {"x": state["x"] + 1} + + builder = StateGraph(_TimeoutState) + builder.add_node( + "flaky", + flaky, + timeout=0.1, + retry_policy=RetryPolicy( + max_attempts=3, + initial_interval=0.0, + jitter=False, + retry_on=NodeTimeoutError, + ), + ) + builder.add_edge(START, "flaky") + builder.add_edge("flaky", END) + graph = builder.compile() + + async def _run() -> None: + result = await graph.ainvoke({"x": 0}) + assert result == {"x": 1} + assert len(attempts) == 2 + + asyncio.run(_run()) + + +def test_task_decorator_timeout_e2e(): + @task(timeout=0.05) + async def slow_task(x: int) -> int: + await asyncio.sleep(0.2) + return x + 1 + + @entrypoint() + async def workflow(x: int) -> int: + return await slow_task(x) + + async def _run() -> None: + with pytest.raises(NodeTimeoutError): + await workflow.ainvoke(1) + + asyncio.run(_run()) + + +def test_entrypoint_timeout_e2e(): + @entrypoint(timeout=0.05) + async def slow_workflow(x: int) -> int: + await asyncio.sleep(0.2) + return x + + async def _run() -> None: + with pytest.raises(NodeTimeoutError): + await slow_workflow.ainvoke(1) + + asyncio.run(_run()) + + +def test_node_builder_timeout_e2e(): + async def slow(value: int) -> int: + await asyncio.sleep(0.2) + return value + 1 + + graph = Pregel( + nodes={ + "slow": ( + NodeBuilder() + .subscribe_only("input") + .do(slow) + .set_timeout(0.05) + .write_to("output") + ) + }, + channels={ + "input": EphemeralValue(int), + "output": LastValue(int), + }, + input_channels="input", + output_channels="output", + ) + + async def _run() -> None: + with pytest.raises(NodeTimeoutError): + await graph.ainvoke(1) + + asyncio.run(_run()) + + +def test_arun_with_retry_timeout_observer_tracks_attempts(): + events: list[dict] = [] + + class FlakyProc: + async def ainvoke(self, input, config): + runtime = config[CONF][CONFIG_KEY_RUNTIME] + if runtime.execution_info.node_attempt == 1: + await asyncio.sleep(0.2) + return "ok" + + policy = RetryPolicy( + max_attempts=2, + initial_interval=0.0, + jitter=False, + retry_on=NodeTimeoutError, + ) + task = _make_task(FlakyProc(), timeout=0.05, retry_policy=(policy,), name="flaky") + task.config[CONF][CONFIG_KEY_TIMED_ATTEMPT_OBSERVER] = events.append + + async def _run() -> None: + assert await arun_with_retry(task, retry_policy=None) == "ok" + + asyncio.run(_run()) + + starts = [payload for payload in events if payload["event"] == "start"] + finishes = [payload for payload in events if payload["event"] == "finish"] + assert [payload["attempt"] for payload in starts] == [1, 2] + assert [payload["attempt"] for payload in finishes] == [1, 2] + assert [payload["status"] for payload in finishes] == ["error", "success"] + assert starts[0]["execution_id"] != starts[1]["execution_id"] + assert starts[0]["timeout_secs"] == 0.05 + assert starts[0]["task_name"] == "flaky" + assert isinstance(starts[0]["started_at"], datetime) + assert isinstance(starts[0]["deadline_at"], datetime) + assert isinstance(finishes[0]["finished_at"], datetime) + assert starts[0]["deadline_at"] > starts[0]["started_at"] + + +def test_arun_with_retry_timeout_observer_treats_parent_command_as_non_error(): + events: list[dict] = [] + + class ParentProc: + async def ainvoke(self, input, config): + raise ParentCommand(Command(graph=Command.PARENT)) + + task = _make_task(ParentProc(), timeout=0.05, name="parent") + task.config[CONF][CONFIG_KEY_TIMED_ATTEMPT_OBSERVER] = events.append + + async def _run() -> None: + with pytest.raises(ParentCommand): + await arun_with_retry(task, retry_policy=None) + + asyncio.run(_run()) + + finish = next(payload for payload in events if payload["event"] == "finish") + assert finish["status"] == "success" + assert finish["error_type"] is None + assert finish["error_message"] is None + + +def test_arun_with_retry_timeout_observer_finishes_when_parent_writer_errors(): + events: list[dict] = [] + + class ParentProc: + async def ainvoke(self, input, config): + raise ParentCommand(Command(graph="parent", update={"value": "updated"})) + + class FailingWriter: + def invoke(self, input, config): + raise ValueError("writer failed") + + task = _make_task( + ParentProc(), timeout=0.05, name="parent", writers=(FailingWriter(),) + ) + task.config[CONF][CONFIG_KEY_TIMED_ATTEMPT_OBSERVER] = events.append + + async def _run() -> None: + with pytest.raises(ValueError, match="writer failed"): + await arun_with_retry(task, retry_policy=None) + + asyncio.run(_run()) + + finish = next(payload for payload in events if payload["event"] == "finish") + assert finish["status"] == "error" + assert finish["error_type"] == "ValueError" + assert finish["error_message"] == "writer failed" + + +def test_arun_with_retry_timeout_observer_treats_bubble_up_as_non_error(): + events: list[dict] = [] + + class BubbleProc: + async def ainvoke(self, input, config): + raise GraphInterrupt(()) + + task = _make_task(BubbleProc(), timeout=0.05, name="bubble") + task.config[CONF][CONFIG_KEY_TIMED_ATTEMPT_OBSERVER] = events.append + + async def _run() -> None: + with pytest.raises(GraphInterrupt): + await arun_with_retry(task, retry_policy=None) + + asyncio.run(_run()) + + finish = next(payload for payload in events if payload["event"] == "finish") + assert finish["status"] == "success" + assert finish["error_type"] is None + assert finish["error_message"] is None