mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-29 03:09:45 +02:00
Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
372d54dc4f | ||
|
|
f4aee546ad | ||
|
|
85cd64ed69 | ||
|
|
53a9806e65 | ||
|
|
219fbbe8d0 | ||
|
|
aeff9549c2 |
Generated
+1
-1
@@ -259,7 +259,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.0.2"
|
||||
version = "4.0.3"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
Generated
+1
-1
@@ -268,7 +268,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.0.2"
|
||||
version = "4.0.3"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -55,6 +55,19 @@ _warned_unregistered_types: set[tuple[str, str]] = set()
|
||||
_warned_blocked_types: set[tuple[str, str]] = set()
|
||||
|
||||
|
||||
def _is_safe_json_type(id_list: list[str]) -> bool:
|
||||
"""Return True if an lc=2 id refers to a type in SAFE_MSGPACK_TYPES.
|
||||
|
||||
Safe types bypass the ``allowed_json_modules`` gate so that old "json" format
|
||||
checkpoints (written before the msgpack migration) can be resumed without
|
||||
requiring users to configure an explicit allowlist.
|
||||
"""
|
||||
if len(id_list) < 2:
|
||||
return False
|
||||
module_name = ".".join(id_list[:-1])
|
||||
return (module_name, id_list[-1]) in _lg_msgpack.SAFE_MSGPACK_TYPES
|
||||
|
||||
|
||||
def _warn_once(
|
||||
seen: set[tuple[str, str]], key: tuple[str, str], msg: str, *args: object
|
||||
) -> None:
|
||||
@@ -164,19 +177,23 @@ class JsonPlusSerializer(SerializerProtocol):
|
||||
return out
|
||||
|
||||
def _reviver(self, value: dict[str, Any]) -> Any:
|
||||
if self._allowed_json_modules and (
|
||||
if (
|
||||
value.get("lc", None) == 2
|
||||
and value.get("type", None) == "constructor"
|
||||
and value.get("id", None) is not None
|
||||
):
|
||||
try:
|
||||
return self._revive_lc2(value)
|
||||
except InvalidModuleError as e:
|
||||
logger.warning(
|
||||
"Object %s is not in the deserialization allowlist.\n%s",
|
||||
value["id"],
|
||||
e.message,
|
||||
)
|
||||
id_list = value["id"]
|
||||
is_safe = _is_safe_json_type(id_list)
|
||||
if self._allowed_json_modules or is_safe:
|
||||
try:
|
||||
return self._revive_lc2(value)
|
||||
except InvalidModuleError as e:
|
||||
if not is_safe:
|
||||
logger.warning(
|
||||
"Object %s is not in the deserialization allowlist.\n%s",
|
||||
value["id"],
|
||||
e.message,
|
||||
)
|
||||
|
||||
return LC_REVIVER(value)
|
||||
|
||||
@@ -224,6 +241,13 @@ class JsonPlusSerializer(SerializerProtocol):
|
||||
method_display = "<init>"
|
||||
|
||||
dotted = ".".join(needed)
|
||||
# Safe types (the same set already allowed for msgpack deserialization) are
|
||||
# permitted without an explicit allowlist — they are known-safe LangGraph and
|
||||
# LangChain types. This restores backwards-compat for old "json" checkpoints
|
||||
# that pre-date the msgpack migration without reopening the broader security gate.
|
||||
if _is_safe_json_type(list(needed)):
|
||||
return
|
||||
|
||||
if not self._allowed_json_modules:
|
||||
raise InvalidModuleError(
|
||||
f"Refused to deserialize JSON constructor: {dotted} (method: {method_display}). "
|
||||
|
||||
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.0.2"
|
||||
version = "4.0.3"
|
||||
description = "Library with base interfaces for LangGraph checkpoint savers."
|
||||
authors = []
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -333,6 +333,57 @@ def test_serde_jsonplus_bytes() -> None:
|
||||
assert serde.loads_typed(dumped) == some_bytes
|
||||
|
||||
|
||||
def test_lc2_json_safe_type_revives_without_allowlist() -> None:
|
||||
"""Old 'json' blobs with lc=2 for safe types must revive without an explicit allowlist.
|
||||
|
||||
Regression test for: https://github.com/langchain-ai/langgraph/issues/7498
|
||||
Threads checkpointed before v1.0.1 (pre-msgpack) stored messages as lc=2 JSON
|
||||
constructor dicts. Resuming those threads must reconstruct proper BaseMessage objects
|
||||
rather than returning raw dicts that cause MESSAGE_COERCION_FAILURE in add_messages.
|
||||
"""
|
||||
from langchain_core.messages import AIMessage
|
||||
|
||||
serde = JsonPlusSerializer() # default: _allowed_json_modules=None
|
||||
|
||||
human_blob = {
|
||||
"lc": 2,
|
||||
"type": "constructor",
|
||||
"id": ["langchain_core", "messages", "human", "HumanMessage"],
|
||||
"kwargs": {"content": "hello", "type": "human"},
|
||||
}
|
||||
ai_blob = {
|
||||
"lc": 2,
|
||||
"type": "constructor",
|
||||
"id": ["langchain_core", "messages", "ai", "AIMessage"],
|
||||
"kwargs": {"content": "hi there", "type": "ai"},
|
||||
}
|
||||
result = serde.loads_typed(("json", json.dumps([human_blob, ai_blob]).encode()))
|
||||
|
||||
assert len(result) == 2
|
||||
assert isinstance(result[0], HumanMessage), (
|
||||
f"Expected HumanMessage, got {type(result[0])}: {result[0]!r}\n"
|
||||
"lc=2 JSON blobs for safe types must deserialize without an explicit allowlist"
|
||||
)
|
||||
assert result[0].content == "hello"
|
||||
assert isinstance(result[1], AIMessage)
|
||||
assert result[1].content == "hi there"
|
||||
|
||||
|
||||
def test_lc2_json_unknown_type_stays_blocked_without_allowlist() -> None:
|
||||
"""lc=2 JSON blobs for types NOT in SAFE_MSGPACK_TYPES still require an allowlist."""
|
||||
serde = JsonPlusSerializer()
|
||||
load = {
|
||||
"lc": 2,
|
||||
"type": "constructor",
|
||||
"id": ["pprint", "pprint"],
|
||||
"kwargs": {"object": "HELLO"},
|
||||
}
|
||||
# No allowlist configured → raw dict returned (not raised, not reconstructed)
|
||||
result = serde.loads_typed(("json", json.dumps(load).encode()))
|
||||
assert isinstance(result, dict), "Unknown lc=2 type must stay as raw dict"
|
||||
assert result.get("lc") == 2
|
||||
|
||||
|
||||
def test_deserde_invalid_module() -> None:
|
||||
serde = JsonPlusSerializer()
|
||||
load = {
|
||||
|
||||
Generated
+1
-1
@@ -286,7 +286,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.0.2"
|
||||
version = "4.0.3"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.")
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
[
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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)},
|
||||
|
||||
@@ -548,6 +548,7 @@ class PregelExecutableTask:
|
||||
path: tuple[str | int | tuple, ...]
|
||||
writers: Sequence[Runnable] = ()
|
||||
subgraphs: Sequence[PregelProtocol] = ()
|
||||
timeout: float | None = None
|
||||
|
||||
|
||||
class StateSnapshot(NamedTuple):
|
||||
|
||||
@@ -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
|
||||
|
||||
Generated
+8
-8
@@ -1,5 +1,5 @@
|
||||
version = 1
|
||||
revision = 2
|
||||
revision = 3
|
||||
requires-python = ">=3.10"
|
||||
resolution-markers = [
|
||||
"python_full_version >= '3.14'",
|
||||
@@ -1548,7 +1548,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.0.2"
|
||||
version = "4.0.3"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
@@ -2140,7 +2140,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "nbconvert"
|
||||
version = "7.17.0"
|
||||
version = "7.17.1"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "beautifulsoup4" },
|
||||
@@ -2158,9 +2158,9 @@ dependencies = [
|
||||
{ name = "pygments" },
|
||||
{ name = "traitlets" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/38/47/81f886b699450d0569f7bc551df2b1673d18df7ff25cc0c21ca36ed8a5ff/nbconvert-7.17.0.tar.gz", hash = "sha256:1b2696f1b5be12309f6c7d707c24af604b87dfaf6d950794c7b07acab96dda78", size = 862855, upload-time = "2026-01-29T16:37:48.478Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/01/b1/708e53fe2e429c103c6e6e159106bcf0357ac41aa4c28772bd8402339051/nbconvert-7.17.1.tar.gz", hash = "sha256:34d0d0a7e73ce3cbab6c5aae8f4f468797280b01fd8bd2ca746da8569eddd7d2", size = 865311, upload-time = "2026-04-08T00:44:14.914Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/0d/4b/8d5f796a792f8a25f6925a96032f098789f448571eb92011df1ae59e8ea8/nbconvert-7.17.0-py3-none-any.whl", hash = "sha256:4f99a63b337b9a23504347afdab24a11faa7d86b405e5c8f9881cd313336d518", size = 261510, upload-time = "2026-01-29T16:37:46.322Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/67/f8/bb0a9d5f46819c821dc1f004aa2cc29b1d91453297dbf5ff20470f00f193/nbconvert-7.17.1-py3-none-any.whl", hash = "sha256:aa85c087b435e7bf1ffd03319f658e285f2b89eccab33bc1ba7025495ab3e7c8", size = 261927, upload-time = "2026-04-08T00:44:12.845Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -3018,11 +3018,11 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "python-dotenv"
|
||||
version = "1.2.1"
|
||||
version = "1.2.2"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/f0/26/19cadc79a718c5edbec86fd4919a6b6d3f681039a2f6d66d14be94e75fb9/python_dotenv-1.2.1.tar.gz", hash = "sha256:42667e897e16ab0d66954af0e60a9caa94f0fd4ecf3aaf6d2d260eec1aa36ad6", size = 44221, upload-time = "2025-10-26T15:12:10.434Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/82/ed/0301aeeac3e5353ef3d94b6ec08bbcabd04a72018415dcb29e588514bba8/python_dotenv-1.2.2.tar.gz", hash = "sha256:2c371a91fbd7ba082c2c1dc1f8bf89ca22564a087c2c287cd9b662adde799cf3", size = 50135, upload-time = "2026-03-01T16:00:26.196Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/14/1b/a298b06749107c305e1fe0f814c6c74aea7b2f1e10989cb30f544a1b3253/python_dotenv-1.2.1-py3-none-any.whl", hash = "sha256:b81ee9561e9ca4004139c6cbba3a238c32b03e4894671e181b671e8cb8425d61", size = 21230, upload-time = "2025-10-26T15:12:09.109Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/0b/d7/1959b9648791274998a9c3526f6d0ec8fd2233e4d4acce81bbae76b44b2a/python_dotenv-1.2.2-py3-none-any.whl", hash = "sha256:1d8214789a24de455a8b8bd8ae6fe3c6b69a5e3d64aa8a8e5d68e694bbcb285a", size = 22101, upload-time = "2026-03-01T16:00:25.09Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
||||
@@ -82,6 +82,7 @@ from langchain_core.tools.base import (
|
||||
_is_injected_arg_type,
|
||||
get_all_basemodel_annotations,
|
||||
)
|
||||
from langgraph._internal._constants import CONF, CONFIG_KEY_READ
|
||||
from langgraph._internal._runnable import RunnableCallable
|
||||
from langgraph.errors import GraphBubbleUp
|
||||
from langgraph.graph.message import REMOVE_ALL_MESSAGES
|
||||
@@ -800,7 +801,7 @@ class ToolNode(RunnableCallable):
|
||||
# Construct ToolRuntime instances at the top level for each tool call
|
||||
tool_runtimes = []
|
||||
for call, cfg in zip(tool_calls, config_list, strict=False):
|
||||
state = self._extract_state(input)
|
||||
state = self._extract_state(input, cfg)
|
||||
tool_runtime = ToolRuntime(
|
||||
state=state,
|
||||
tool_call_id=call["id"],
|
||||
@@ -835,7 +836,7 @@ class ToolNode(RunnableCallable):
|
||||
# Construct ToolRuntime instances at the top level for each tool call
|
||||
tool_runtimes = []
|
||||
for call, cfg in zip(tool_calls, config_list, strict=False):
|
||||
state = self._extract_state(input)
|
||||
state = self._extract_state(input, cfg)
|
||||
tool_runtime = ToolRuntime(
|
||||
state=state,
|
||||
tool_call_id=call["id"],
|
||||
@@ -1277,18 +1278,37 @@ class ToolNode(RunnableCallable):
|
||||
return None
|
||||
|
||||
def _extract_state(
|
||||
self, input: list[AnyMessage] | dict[str, Any] | BaseModel
|
||||
self,
|
||||
input: list[AnyMessage] | dict[str, Any] | BaseModel,
|
||||
config: RunnableConfig,
|
||||
) -> list[AnyMessage] | dict[str, Any] | BaseModel:
|
||||
"""Extract state from input, handling ToolCallWithContext if present.
|
||||
"""Extract state from input.
|
||||
|
||||
Args:
|
||||
input: The input which may be raw state or ToolCallWithContext.
|
||||
Three input shapes:
|
||||
|
||||
Returns:
|
||||
The actual state to pass to wrap_tool_call wrappers.
|
||||
- `ToolCallWithContext` dict — legacy Send payload carrying an inlined
|
||||
state snapshot; return `input["state"]`.
|
||||
- list of `ToolCall` dicts — new Send payload with no inlined state;
|
||||
hydrate state from channels via `CONFIG_KEY_READ`.
|
||||
- regular graph state (dict/list/BaseModel) — return `input` as-is.
|
||||
"""
|
||||
if isinstance(input, dict) and input.get("__type") == "tool_call_with_context":
|
||||
return input["state"]
|
||||
if (
|
||||
isinstance(input, list)
|
||||
and input
|
||||
and isinstance(input[-1], dict)
|
||||
and input[-1].get("type") == "tool_call"
|
||||
):
|
||||
read = config.get(CONF, {}).get(CONFIG_KEY_READ)
|
||||
if read is None:
|
||||
return {}
|
||||
# Pregel installs CONFIG_KEY_READ as
|
||||
# `functools.partial(local_read, scratchpad, channels, managed, task)`.
|
||||
# Match the previous inlined-state contract by reading channels only;
|
||||
# managed values have their own injection path (`ToolRuntime.context`).
|
||||
channels = read.args[1]
|
||||
return cast("dict[str, Any]", read(list(channels), True))
|
||||
return input
|
||||
|
||||
def _inject_tool_args(
|
||||
|
||||
@@ -1320,6 +1320,98 @@ async def test_state_extraction_with_tool_call_with_context_async() -> None:
|
||||
assert "tool_call" not in state_seen[0]
|
||||
|
||||
|
||||
def _config_with_channel_read(
|
||||
channel_values: dict[str, object],
|
||||
store: BaseStore | None = None,
|
||||
) -> RunnableConfig:
|
||||
"""Build a config that mimics `CONFIG_KEY_READ` as Pregel installs it.
|
||||
|
||||
Pregel always installs a `functools.partial(local_read, scratchpad,
|
||||
channels, managed, task)`, and `ToolNode` introspects that partial to
|
||||
learn channel names. The stub matches the shape: partial whose second and
|
||||
third positional args are `channels` and `managed` mappings.
|
||||
"""
|
||||
import functools
|
||||
|
||||
channels_stub = {k: None for k in channel_values}
|
||||
managed_stub: dict[str, object] = {}
|
||||
|
||||
# Shape matches pregel's real partial:
|
||||
# functools.partial(local_read, scratchpad, channels, managed, task)
|
||||
def _read(scratchpad, channels, managed, task, select, fresh): # noqa: ARG001
|
||||
if isinstance(select, str):
|
||||
return channel_values[select]
|
||||
return {k: channel_values[k] for k in select if k in channel_values}
|
||||
|
||||
read = functools.partial(_read, None, channels_stub, managed_stub, None)
|
||||
cfg = _create_config_with_runtime(store)
|
||||
cfg["configurable"]["__pregel_read"] = read
|
||||
return cfg
|
||||
|
||||
|
||||
def test_list_form_send_hydrates_state_from_channel_read() -> None:
|
||||
"""Send('tools', [tool_call]) with no inlined state should hydrate
|
||||
ToolRuntime.state from CONFIG_KEY_READ (full state read)."""
|
||||
state_seen = []
|
||||
|
||||
def state_inspector_handler(
|
||||
request: ToolCallRequest,
|
||||
execute: Callable[[ToolCallRequest], ToolMessage | Command],
|
||||
) -> ToolMessage | Command:
|
||||
state_seen.append(request.state)
|
||||
return execute(request)
|
||||
|
||||
channel_values = {
|
||||
"messages": [AIMessage("from channels")],
|
||||
"files": {"/a.md": "body"},
|
||||
}
|
||||
|
||||
tool_node = ToolNode([add], wrap_tool_call=state_inspector_handler)
|
||||
|
||||
tool_call: ToolCall = {
|
||||
"name": "add",
|
||||
"args": {"a": 1, "b": 2},
|
||||
"id": "call_1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
|
||||
tool_node.invoke([tool_call], config=_config_with_channel_read(channel_values))
|
||||
|
||||
assert len(state_seen) == 1
|
||||
got = state_seen[0]
|
||||
assert got == channel_values
|
||||
assert "messages" in got and "files" in got
|
||||
|
||||
|
||||
async def test_list_form_send_hydrates_state_async() -> None:
|
||||
state_seen = []
|
||||
|
||||
def state_inspector_handler(
|
||||
request: ToolCallRequest,
|
||||
execute: Callable[[ToolCallRequest], ToolMessage | Command],
|
||||
) -> ToolMessage | Command:
|
||||
state_seen.append(request.state)
|
||||
return execute(request)
|
||||
|
||||
channel_values = {"messages": [AIMessage("from channels")], "files": {}}
|
||||
|
||||
tool_node = ToolNode([add], wrap_tool_call=state_inspector_handler)
|
||||
|
||||
tool_call: ToolCall = {
|
||||
"name": "add",
|
||||
"args": {"a": 1, "b": 2},
|
||||
"id": "call_1",
|
||||
"type": "tool_call",
|
||||
}
|
||||
|
||||
await tool_node.ainvoke(
|
||||
[tool_call], config=_config_with_channel_read(channel_values)
|
||||
)
|
||||
|
||||
assert len(state_seen) == 1
|
||||
assert state_seen[0] == channel_values
|
||||
|
||||
|
||||
def test_tool_call_request_is_frozen() -> None:
|
||||
"""Test that ToolCallRequest raises deprecation warnings on direct attribute reassignment."""
|
||||
tool_call: ToolCall = {"name": "add", "args": {"a": 1, "b": 2}, "id": "call_1"}
|
||||
|
||||
Generated
+2
-2
@@ -1,5 +1,5 @@
|
||||
version = 1
|
||||
revision = 2
|
||||
revision = 3
|
||||
requires-python = ">=3.10"
|
||||
|
||||
[[package]]
|
||||
@@ -352,7 +352,7 @@ test = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.0.2"
|
||||
version = "4.0.3"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
Generated
+2
-2
@@ -1,5 +1,5 @@
|
||||
version = 1
|
||||
revision = 2
|
||||
revision = 3
|
||||
requires-python = ">=3.10"
|
||||
|
||||
[[package]]
|
||||
@@ -365,7 +365,7 @@ test = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.0.2"
|
||||
version = "4.0.3"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
Reference in New Issue
Block a user