diff --git a/libs/langgraph/langgraph/_internal/_constants.py b/libs/langgraph/langgraph/_internal/_constants.py index c51c5c77b..68cb48fe8 100644 --- a/libs/langgraph/langgraph/_internal/_constants.py +++ b/libs/langgraph/langgraph/_internal/_constants.py @@ -56,8 +56,6 @@ 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") @@ -108,7 +106,6 @@ 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 eebb02183..63e03f544 100644 --- a/libs/langgraph/langgraph/_internal/_runnable.py +++ b/libs/langgraph/langgraph/_internal/_runnable.py @@ -706,12 +706,7 @@ class RunnableSeq(Runnable): step.ainvoke(input, config, **kwargs), context=context ) else: - with set_config_context(config) as context: - input = await context.run( - lambda: asyncio.create_task( - step.ainvoke(input, config, **kwargs) - ) - ) + input = await 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 deleted file mode 100644 index 3eac4d589..000000000 --- a/libs/langgraph/langgraph/_internal/_timeout.py +++ /dev/null @@ -1,26 +0,0 @@ -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 53a3660b7..fb648879e 100644 --- a/libs/langgraph/langgraph/errors.py +++ b/libs/langgraph/langgraph/errors.py @@ -20,7 +20,6 @@ __all__ = ( "GraphBubbleUp", "GraphInterrupt", "NodeInterrupt", - "NodeTimeoutError", "ParentCommand", "EmptyInputError", "TaskNotFound", @@ -126,25 +125,3 @@ 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 e0b0a955a..c7443e0a2 100644 --- a/libs/langgraph/langgraph/func/__init__.py +++ b/libs/langgraph/langgraph/func/__init__.py @@ -5,7 +5,6 @@ import inspect import warnings from collections.abc import Awaitable, Callable, Sequence from dataclasses import dataclass -from datetime import timedelta from typing import ( Any, Generic, @@ -23,8 +22,6 @@ 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 @@ -54,7 +51,6 @@ 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: @@ -71,7 +67,6 @@ 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]: @@ -79,7 +74,6 @@ class _TaskFunction(Generic[P, T]): self.func, retry_policy=self.retry_policy, cache_policy=self.cache_policy, - timeout=self.timeout, *args, **kwargs, ) @@ -104,7 +98,6 @@ 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]], @@ -126,7 +119,6 @@ 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]] @@ -150,9 +142,6 @@ 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. @@ -207,7 +196,6 @@ def task( ) if retry_policy is None: retry_policy = retry # type: ignore[assignment] - timeout_s = coerce_timeout(timeout) retry_policies: Sequence[RetryPolicy] = ( () @@ -220,15 +208,8 @@ 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, - timeout=timeout_s, - name=name, + func, retry_policy=retry_policies, cache_policy=cache_policy, name=name ) if __func_or_none__ is not None: @@ -419,7 +400,6 @@ 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.""" @@ -446,7 +426,6 @@ 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) @@ -556,7 +535,6 @@ 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 b5c4d4a3e..cadf097d9 100644 --- a/libs/langgraph/langgraph/graph/_node.py +++ b/libs/langgraph/langgraph/graph/_node.py @@ -90,4 +90,3 @@ 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 8d96508ff..b1c24de2b 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -7,7 +7,6 @@ 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 @@ -46,7 +45,6 @@ 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 @@ -302,7 +300,6 @@ 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. @@ -370,7 +367,6 @@ 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. @@ -443,7 +439,6 @@ 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. @@ -511,7 +506,6 @@ 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. @@ -586,7 +580,6 @@ 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`. @@ -616,12 +609,6 @@ 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 @@ -675,7 +662,6 @@ 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 @@ -771,7 +757,6 @@ 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( @@ -782,7 +767,6 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): cache_policy=cache_policy, ends=ends, defer=defer, - timeout=timeout, ) else: self.nodes[node] = StateNodeSpec[StateT, ContextT]( @@ -793,7 +777,6 @@ 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 @@ -1349,7 +1332,6 @@ 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 d8f83d519..d7157e239 100644 --- a/libs/langgraph/langgraph/pregel/_algo.py +++ b/libs/langgraph/langgraph/pregel/_algo.py @@ -7,7 +7,6 @@ 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 ( @@ -62,7 +61,6 @@ 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 @@ -116,21 +114,13 @@ class PregelTaskWrites(NamedTuple): class Call: - __slots__ = ( - "func", - "input", - "retry_policy", - "cache_policy", - "callbacks", - "timeout", - ) + __slots__ = ("func", "input", "retry_policy", "cache_policy", "callbacks") 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, @@ -140,14 +130,12 @@ 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( @@ -745,7 +733,6 @@ 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]) @@ -883,7 +870,6 @@ 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) @@ -1055,7 +1041,6 @@ 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 0b2b5af4c..0cd007042 100644 --- a/libs/langgraph/langgraph/pregel/_call.py +++ b/libs/langgraph/langgraph/pregel/_call.py @@ -8,7 +8,6 @@ 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 @@ -21,7 +20,6 @@ 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 @@ -257,13 +255,8 @@ 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( @@ -272,6 +265,5 @@ 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 2c53f4485..8d4c21135 100644 --- a/libs/langgraph/langgraph/pregel/_read.py +++ b/libs/langgraph/langgraph/pregel/_read.py @@ -1,7 +1,6 @@ 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, @@ -12,7 +11,6 @@ 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 @@ -125,11 +123,6 @@ 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.""" @@ -152,7 +145,6 @@ 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) @@ -164,7 +156,6 @@ 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 9772ad62f..538d6b915 100644 --- a/libs/langgraph/langgraph/pregel/_retry.py +++ b/libs/langgraph/langgraph/pregel/_retry.py @@ -4,16 +4,12 @@ import asyncio import logging import random import sys -import threading import time -from collections.abc import Awaitable, Callable, Coroutine, Sequence -from contextlib import suppress +from collections.abc import Awaitable, Callable, Sequence from dataclasses import replace -from datetime import datetime, timedelta, timezone -from typing import Any, Literal +from typing import Any 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 ( @@ -22,14 +18,11 @@ 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._internal._timeout import sync_timeout_unsupported -from langgraph.errors import GraphBubbleUp, NodeTimeoutError, ParentCommand +from langgraph.errors import GraphBubbleUp, ParentCommand from langgraph.runtime import ExecutionInfo, Runtime from langgraph.types import Command, PregelExecutableTask, RetryPolicy @@ -37,182 +30,6 @@ 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: @@ -273,11 +90,6 @@ 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 @@ -383,7 +195,6 @@ 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 @@ -418,47 +229,35 @@ 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() - 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) + # run the task if stream: + async for _ in task.proc.astream(task.input, config): + pass + # if successful, end break - return result + else: + return await task.proc.ainvoke(task.input, config) 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): - 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) + # this command is for the current graph, handle it + for w in task.writers: + w.invoke(cmd, config) 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)),) - _finish_timed_attempt(config, attempt_payload) + # bubble up raise except GraphBubbleUp: - _finish_timed_attempt(config, attempt_payload) + # if interrupted, end 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 e7a352c4e..fea4a7272 100644 --- a/libs/langgraph/langgraph/pregel/_runner.py +++ b/libs/langgraph/langgraph/pregel/_runner.py @@ -14,7 +14,6 @@ from collections.abc import ( Iterator, Sequence, ) -from datetime import timedelta from functools import partial from typing import ( Any, @@ -538,7 +537,6 @@ 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[ @@ -562,7 +560,6 @@ def _call( retry_policy=retry_policy, cache_policy=cache_policy, callbacks=callbacks, - timeout=timeout, ), ): if fut := next( @@ -627,7 +624,6 @@ 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], @@ -661,7 +657,6 @@ def _acall( input, retry_policy=retry_policy, cache_policy=cache_policy, - timeout=timeout, callbacks=callbacks, futures=futures, schedule_task=schedule_task, @@ -683,7 +678,6 @@ 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]], @@ -709,7 +703,6 @@ 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 d4d7c103e..0c8a14eec 100644 --- a/libs/langgraph/langgraph/pregel/_utils.py +++ b/libs/langgraph/langgraph/pregel/_utils.py @@ -4,21 +4,16 @@ import ast import inspect import re import textwrap -from collections.abc import Callable, Sequence -from functools import partial +from collections.abc import Callable 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 @@ -69,47 +64,6 @@ 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 0b255085c..c440e77f9 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -17,7 +17,6 @@ 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 ( @@ -96,7 +95,6 @@ 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, @@ -139,7 +137,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, validate_timeout_supported +from langgraph.pregel._utils import get_new_channel_versions 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 @@ -188,7 +186,6 @@ class NodeBuilder: "_bound", "_retry_policy", "_cache_policy", - "_timeout", ) _channels: str | list[str] @@ -199,7 +196,6 @@ class NodeBuilder: _bound: Runnable _retry_policy: list[RetryPolicy] _cache_policy: CachePolicy | None - _timeout: float | None def __init__( self, @@ -212,7 +208,6 @@ class NodeBuilder: self._bound = DEFAULT_BOUND self._retry_policy = [] self._cache_policy = None - self._timeout = None def subscribe_only( self, @@ -331,11 +326,6 @@ 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( @@ -347,7 +337,6 @@ class NodeBuilder: bound=self._bound, retry_policy=self._retry_policy, cache_policy=self._cache_policy, - timeout=self._timeout, ) @@ -828,9 +817,6 @@ 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 d2c4e0f64..d04d82da7 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -548,7 +548,6 @@ 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 4ba7fb163..3156f7599 100644 --- a/libs/langgraph/tests/test_retry.py +++ b/libs/langgraph/tests/test_retry.py @@ -1,12 +1,7 @@ -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 @@ -15,29 +10,18 @@ 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._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.graph import START, StateGraph 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 Command, PregelExecutableTask, RetryPolicy +from langgraph.types import PregelExecutableTask, RetryPolicy def test_should_retry_on_single_exception(): @@ -583,643 +567,3 @@ 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