diff --git a/libs/langgraph/langgraph/pregel/_retry.py b/libs/langgraph/langgraph/pregel/_retry.py index ba62bda6e..3c60f478e 100644 --- a/libs/langgraph/langgraph/pregel/_retry.py +++ b/libs/langgraph/langgraph/pregel/_retry.py @@ -14,9 +14,11 @@ from langgraph._internal._constants import ( CONF, CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_RESUMING, + CONFIG_KEY_RUNTIME, NS_SEP, ) from langgraph.errors import GraphBubbleUp, ParentCommand +from langgraph.runtime import Runtime from langgraph.types import Command, PregelExecutableTask, RetryPolicy logger = logging.getLogger(__name__) @@ -60,10 +62,33 @@ def run_with_retry( """Run a task with retries.""" retry_policy = task.retry_policy or retry_policy attempts = 0 + node_first_attempt_time = time.time() config = task.config if configurable is not None: config = patch_configurable(config, configurable) + runtime = config.get(CONF, {}).get(CONFIG_KEY_RUNTIME) + if isinstance(runtime, Runtime): + config = patch_configurable( + config, + { + CONFIG_KEY_RUNTIME: runtime.patch_execution_info( + node_first_attempt_time=node_first_attempt_time, + ) + }, + ) while True: + runtime = config.get(CONF, {}).get(CONFIG_KEY_RUNTIME) + if isinstance(runtime, Runtime): + config = patch_configurable( + config, + { + CONFIG_KEY_RUNTIME: runtime.patch_execution_info( + # node_attempt is execution count (1-indexed): 1 on first run, + # then 2, 3, ... on subsequent retries. + node_attempt=attempts + 1, + ) + }, + ) try: # clear any writes from previous attempts task.writes.clear() @@ -102,7 +127,7 @@ def run_with_retry( if not matching_policy: raise - # increment attempts + # attempts tracks failed tries only; it increments after a failure. attempts += 1 # check if we should give up if attempts >= matching_policy.max_attempts: @@ -141,15 +166,38 @@ async def arun_with_retry( """Run a task asynchronously with retries.""" retry_policy = task.retry_policy or retry_policy attempts = 0 + node_first_attempt_time = time.time() config = task.config if configurable is not None: config = patch_configurable(config, configurable) + runtime = config.get(CONF, {}).get(CONFIG_KEY_RUNTIME) + if isinstance(runtime, Runtime): + config = patch_configurable( + config, + { + CONFIG_KEY_RUNTIME: runtime.patch_execution_info( + node_first_attempt_time=node_first_attempt_time, + ) + }, + ) if match_cached_writes is not None and task.cache_key is not None: for t in await match_cached_writes(): if t is task: # if the task is already cached, return return while True: + runtime = config.get(CONF, {}).get(CONFIG_KEY_RUNTIME) + if isinstance(runtime, Runtime): + config = patch_configurable( + config, + { + CONFIG_KEY_RUNTIME: runtime.patch_execution_info( + # node_attempt is execution count (1-indexed): 1 on first run, + # then 2, 3, ... on subsequent retries. + node_attempt=attempts + 1, + ) + }, + ) try: # clear any writes from previous attempts task.writes.clear() @@ -194,7 +242,8 @@ async def arun_with_retry( if not matching_policy: raise - # increment attempts + # attempts tracks failed tries only; it increments after a failure. + # The next execution's node_attempt is derived as attempts + 1. attempts += 1 # check if we should give up if attempts >= matching_policy.max_attempts: diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index 3ecefc1dc..a09575827 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -136,7 +136,7 @@ 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 from langgraph.pregel.protocol import PregelProtocol, StreamChunk, StreamProtocol -from langgraph.runtime import DEFAULT_RUNTIME, Runtime +from langgraph.runtime import DEFAULT_RUNTIME, ExecutionInfo, Runtime from langgraph.types import ( All, CachePolicy, @@ -2648,6 +2648,7 @@ class Pregel( store=store, stream_writer=stream_writer, previous=None, + execution_info=ExecutionInfo(), ) parent_runtime = config[CONF].get(CONFIG_KEY_RUNTIME, DEFAULT_RUNTIME) runtime = parent_runtime.merge(runtime) @@ -3014,6 +3015,7 @@ class Pregel( store=store, stream_writer=stream_writer, previous=None, + execution_info=ExecutionInfo(), ) parent_runtime = config[CONF].get(CONFIG_KEY_RUNTIME, DEFAULT_RUNTIME) runtime = parent_runtime.merge(runtime) diff --git a/libs/langgraph/langgraph/runtime.py b/libs/langgraph/langgraph/runtime.py index 872b9987c..b4c0afac5 100644 --- a/libs/langgraph/langgraph/runtime.py +++ b/libs/langgraph/langgraph/runtime.py @@ -1,7 +1,7 @@ from __future__ import annotations from dataclasses import dataclass, field, replace -from typing import Any, Generic, cast +from typing import Any, Generic, NamedTuple, cast from langgraph.store.base import BaseStore from typing_extensions import TypedDict, Unpack @@ -11,7 +11,21 @@ from langgraph.config import get_config from langgraph.types import _DC_KWARGS, StreamWriter from langgraph.typing import ContextT -__all__ = ("Runtime", "get_runtime") +__all__ = ("ExecutionInfo", "Runtime", "get_runtime") + + +class ExecutionInfo(NamedTuple): + """Read-only execution info/metadata for the execution of current thread/run/node.""" + + node_attempt: int = 1 + """Current node execution attempt number (1-indexed).""" + + node_first_attempt_time: float | None = None + """Unix timestamp (seconds) for when the first attempt started.""" + + def patch(self, **overrides: Any) -> ExecutionInfo: + """Return a new execution info object with selected fields replaced.""" + return self._replace(**overrides) def _no_op_stream_writer(_: Any) -> None: ... @@ -22,6 +36,7 @@ class _RuntimeOverrides(TypedDict, Generic[ContextT], total=False): store: BaseStore | None stream_writer: StreamWriter previous: Any + execution_info: ExecutionInfo @dataclass(**_DC_KWARGS) @@ -29,7 +44,7 @@ class Runtime(Generic[ContextT]): """Convenience class that bundles run-scoped context and other runtime utilities. This class is injected into graph nodes and middleware. It provides access to - `context`, `store`, `stream_writer`, and `previous`. + `context`, `store`, `stream_writer`, `previous`, and `execution_info`. !!! note "Accessing `config`" @@ -115,6 +130,9 @@ class Runtime(Generic[ContextT]): Only available with the functional API when a checkpointer is provided. """ + execution_info: ExecutionInfo = field(default_factory=ExecutionInfo) + """Read-only execution information/metadata for the current node run.""" + def merge(self, other: Runtime[ContextT]) -> Runtime[ContextT]: """Merge two runtimes together. @@ -127,6 +145,7 @@ class Runtime(Generic[ContextT]): if other.stream_writer is not _no_op_stream_writer else self.stream_writer, previous=self.previous if other.previous is None else other.previous, + execution_info=other.execution_info, ) def override( @@ -135,12 +154,20 @@ class Runtime(Generic[ContextT]): """Replace the runtime with a new runtime with the given overrides.""" return replace(self, **overrides) + def patch_execution_info(self, **overrides: Any) -> Runtime[ContextT]: + """Return a new runtime with selected execution_info fields replaced.""" + return replace( + self, + execution_info=self.execution_info.patch(**overrides), + ) + DEFAULT_RUNTIME = Runtime( context=None, store=None, stream_writer=_no_op_stream_writer, previous=None, + execution_info=ExecutionInfo(), ) diff --git a/libs/langgraph/tests/test_retry.py b/libs/langgraph/tests/test_retry.py index 864affe2f..a7fda9740 100644 --- a/libs/langgraph/tests/test_retry.py +++ b/libs/langgraph/tests/test_retry.py @@ -5,6 +5,7 @@ from typing_extensions import TypedDict from langgraph.graph import START, StateGraph from langgraph.pregel._retry import _checkpoint_ns_for_parent_command, _should_retry_on +from langgraph.runtime import Runtime from langgraph.types import RetryPolicy @@ -168,10 +169,15 @@ def test_graph_with_single_retry_policy(): foo: str attempt_count = 0 + attempt_numbers: list[int] = [] + first_attempt_times: list[float | None] = [] - def failing_node(state: State): + def failing_node(state: State, runtime: Runtime): nonlocal attempt_count attempt_count += 1 + assert runtime.execution_info.node_attempt == attempt_count + attempt_numbers.append(runtime.execution_info.node_attempt) + first_attempt_times.append(runtime.execution_info.node_first_attempt_time) if attempt_count < 3: # Fail the first two attempts raise ValueError("Intentional failure") return {"foo": "success"} @@ -203,6 +209,11 @@ def test_graph_with_single_retry_policy(): # Verify retry behavior assert attempt_count == 3 # The node should have been tried 3 times + assert attempt_numbers == [1, 2, 3] + assert len(first_attempt_times) == 3 + assert first_attempt_times[0] is not None + assert first_attempt_times[1] == first_attempt_times[0] + assert first_attempt_times[2] == first_attempt_times[0] assert result["foo"] == "other_node" # Final result should be from other_node # Verify the sleep intervals @@ -210,6 +221,30 @@ def test_graph_with_single_retry_policy(): assert call_args_list == [0.01, 0.02] +def test_runtime_execution_info_defaults_without_retry(): + """Test execution_info defaults when no retry and no config are provided.""" + + class State(TypedDict): + foo: str + + captured = {} + + def node(state: State, runtime: Runtime): + captured["node_attempt"] = runtime.execution_info.node_attempt + captured["node_first_attempt_time"] = ( + runtime.execution_info.node_first_attempt_time + ) + return {"foo": "ok"} + + graph = StateGraph(State).add_node("node", node).add_edge(START, "node").compile() + + result = graph.invoke({"foo": ""}) + + assert result["foo"] == "ok" + assert captured["node_attempt"] == 1 + assert isinstance(captured["node_first_attempt_time"], float) + + def test_graph_with_jitter_retry_policy(): """Test a graph with a RetryPolicy that uses jitter."""