mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-08 02:37:52 +02:00
feat: add execution info to runtime (#7143)
- [x] **Add tests and docs**: If you're adding a new integration, you must include: 1. A test for the integration, preferably unit tests that do not rely on network access, 2. An example notebook showing its use. It lives in `docs/docs/integrations` directory. - [x] **Lint and test**: Run `make format`, `make lint` and `make test` from the root of the package(s) you've modified. We will not consider a PR unless these three are passing in CI. See [contribution guidelines](https://docs.langchain.com/oss/python/contributing/overview) for more. Additional guidelines: - Make sure optional dependencies are imported within a function. - Please do not add dependencies to `pyproject.toml` files (even optional ones) unless they are **required** for unit tests. - Most PRs should not touch more than one package. - Changes should be backwards compatible.
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user