From 833b1faba2c51d55bf7887054ec2f3fbd250d740 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Thu, 18 Jul 2024 10:18:29 -0700 Subject: [PATCH] Enable configuring retry policy for each node of StateGraph --- libs/langgraph/langgraph/graph/state.py | 8 +++++++- libs/langgraph/langgraph/pregel/types.py | 8 ++++---- libs/langgraph/tests/test_pregel.py | 12 ++++++------ 3 files changed, 17 insertions(+), 11 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 2c9c729ac..abde9d0e9 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -42,7 +42,7 @@ from langgraph.graph.graph import ( ) from langgraph.managed.base import ManagedValue, is_managed_value from langgraph.pregel.read import ChannelRead, PregelNode -from langgraph.pregel.types import All +from langgraph.pregel.types import All, RetryPolicy from langgraph.pregel.write import SKIP_WRITE, ChannelWrite, ChannelWriteEntry from langgraph.utils import RunnableCallable, coerce_to_runnable @@ -65,6 +65,7 @@ class StateNodeSpec(NamedTuple): runnable: Runnable metadata: dict[str, Any] input: Type[Any] + retry_policy: Optional[RetryPolicy] class StateGraph(Graph): @@ -195,6 +196,7 @@ class StateGraph(Graph): *, metadata: Optional[dict[str, Any]] = None, input: Optional[Type[Any]] = None, + retry: Optional[RetryPolicy] = None, ) -> None: """Adds a new node to the state graph. Will take the name of the function/runnable as the node name. @@ -218,6 +220,7 @@ class StateGraph(Graph): *, metadata: Optional[dict[str, Any]] = None, input: Optional[Type[Any]] = None, + retry: Optional[RetryPolicy] = None, ) -> None: """Adds a new node to the state graph. @@ -240,6 +243,7 @@ class StateGraph(Graph): *, metadata: Optional[dict[str, Any]] = None, input: Optional[Type[Any]] = None, + retry: Optional[RetryPolicy] = None, ) -> None: """Adds a new node to the state graph. @@ -322,6 +326,7 @@ class StateGraph(Graph): coerce_to_runnable(action, name=node, trace=False), metadata, input=input or self.schema, + retry_policy=retry, ) def add_edge(self, start_key: Union[str, list[str]], end_key: str) -> None: @@ -552,6 +557,7 @@ class CompiledStateGraph(CompiledGraph): ), ], metadata=node.metadata, + retry_policy=node.retry_policy, ).pipe(node.runnable) def attach_edge(self, starts: Union[str, Sequence[str]], end: str) -> None: diff --git a/libs/langgraph/langgraph/pregel/types.py b/libs/langgraph/langgraph/pregel/types.py index 33d79bd9e..9a1afb596 100644 --- a/libs/langgraph/langgraph/pregel/types.py +++ b/libs/langgraph/langgraph/pregel/types.py @@ -1,5 +1,5 @@ from collections import deque -from typing import Any, Callable, Literal, NamedTuple, Optional, Union +from typing import Any, Callable, Literal, NamedTuple, Optional, Type, Union from langchain_core.runnables import Runnable, RunnableConfig @@ -46,14 +46,14 @@ class RetryPolicy(NamedTuple): """Multiplier by which the interval increases after each retry.""" max_interval: float = 128.0 """Maximum amount of time that may elapse between retries. In seconds.""" - max_attempts: int = 10 + max_attempts: int = 3 """Maximum number of attempts to make before giving up, including the first.""" jitter: bool = True """Whether to add random jitter to the interval between retries.""" retry_on: Union[ - tuple[Exception, ...], Callable[[Exception], bool] + Type[Exception], tuple[Type[Exception], ...], Callable[[Exception], bool] ] = default_retry_on - """List of exceptions that should trigger a retry, or a callable that returns True for exceptions that should trigger a retry.""" + """List of exception classes that should trigger a retry, or a callable that returns True for exceptions that should trigger a retry.""" class PregelTaskDescription(NamedTuple): diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index fc5cddb87..df6ba34ca 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -1072,21 +1072,21 @@ def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None: self.calls = 0 one = AwhileMaker(0.2, {"value": 2}) - two = AwhileMaker(0.6, ValueError("I'm not good")) + two = AwhileMaker(0.6, ConnectionError("I'm not good")) builder = StateGraph(State) builder.add_node("one", one) - builder.add_node("two", two) + builder.add_node("two", two, retry=RetryPolicy(max_attempts=2)) builder.add_edge(START, "one") builder.add_edge(START, "two") graph = builder.compile(checkpointer=checkpointer) thread1: RunnableConfig = {"configurable": {"thread_id": 1}} - with pytest.raises(ValueError, match="I'm not good"): + with pytest.raises(ConnectionError, match="I'm not good"): graph.invoke({"value": 1}, thread1) # both nodes should have been called once assert one.calls == 1 - assert two.calls == 1 + assert two.calls == 2 # two attempts # latest checkpoint should be before nodes "one", "two" state = graph.get_state(thread1) @@ -1105,13 +1105,13 @@ def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None: assert checkpoint.pending_writes[0][0] == checkpoint.pending_writes[1][0] # resume execution - with pytest.raises(ValueError, match="I'm not good"): + with pytest.raises(ConnectionError, match="I'm not good"): graph.invoke(None, thread1) # node "one" succeeded previously, so shouldn't be called again assert one.calls == 1 # node "two" should have been called once again - assert two.calls == 2 + assert two.calls == 4 # two attempts before + two attempts now # confirm no new checkpoints saved state_two = graph.get_state(thread1)