diff --git a/libs/checkpoint/langgraph/checkpoint/base/__init__.py b/libs/checkpoint/langgraph/checkpoint/base/__init__.py index f0e05c014..f06bfc1f8 100644 --- a/libs/checkpoint/langgraph/checkpoint/base/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/base/__init__.py @@ -19,6 +19,7 @@ from langgraph.checkpoint.serde.types import ( ERROR, INTERRUPT, RESUME, + RETRY, SCHEDULED, ChannelProtocol, ) @@ -445,7 +446,7 @@ Special writes (e.g. errors) map to negative indices, to avoid those writes from conflicting with regular writes. Each Checkpointer implementation should use this mapping in put_writes. """ -WRITES_IDX_MAP = {ERROR: -1, SCHEDULED: -2, INTERRUPT: -3, RESUME: -4} +WRITES_IDX_MAP = {ERROR: -1, SCHEDULED: -2, INTERRUPT: -3, RESUME: -4, RETRY: -5} EXCLUDED_METADATA_KEYS = { "thread_id", diff --git a/libs/checkpoint/langgraph/checkpoint/serde/types.py b/libs/checkpoint/langgraph/checkpoint/serde/types.py index 65a2b0c8e..7b26fc4c0 100644 --- a/libs/checkpoint/langgraph/checkpoint/serde/types.py +++ b/libs/checkpoint/langgraph/checkpoint/serde/types.py @@ -12,6 +12,7 @@ ERROR = "__error__" SCHEDULED = "__scheduled__" INTERRUPT = "__interrupt__" RESUME = "__resume__" +RETRY = "__retry__" TASKS = "__pregel_tasks" Value = TypeVar("Value", covariant=True) diff --git a/libs/langgraph/langgraph/_internal/_constants.py b/libs/langgraph/langgraph/_internal/_constants.py index b538502c4..9ad38c8cd 100644 --- a/libs/langgraph/langgraph/_internal/_constants.py +++ b/libs/langgraph/langgraph/_internal/_constants.py @@ -14,6 +14,8 @@ ERROR = sys.intern("__error__") # for errors raised by nodes NO_WRITES = sys.intern("__no_writes__") # marker to signal node didn't write anything +RETRY = sys.intern("__retry__") +# for persisting retry state (attempt count and next retry timestamp) TASKS = sys.intern("__pregel_tasks") # for Send objects returned by nodes/edges, corresponds to PUSH below RETURN = sys.intern("__return__") @@ -91,6 +93,7 @@ RESERVED = { RESUME, ERROR, NO_WRITES, + RETRY, # reserved config.configurable keys CONFIG_KEY_SEND, CONFIG_KEY_READ, diff --git a/libs/langgraph/langgraph/_internal/_scratchpad.py b/libs/langgraph/langgraph/_internal/_scratchpad.py index fd96b726c..aa5b5c78b 100644 --- a/libs/langgraph/langgraph/_internal/_scratchpad.py +++ b/libs/langgraph/langgraph/_internal/_scratchpad.py @@ -17,3 +17,6 @@ class PregelScratchpad: resume: list[Any] # subgraph subgraph_counter: Callable[[], int] + # retry (restored from checkpoint pending writes) + retry_attempt: int = 0 + retry_ts: float = 0.0 diff --git a/libs/langgraph/langgraph/pregel/_algo.py b/libs/langgraph/langgraph/pregel/_algo.py index 9f265f80f..038a9bd44 100644 --- a/libs/langgraph/langgraph/pregel/_algo.py +++ b/libs/langgraph/langgraph/pregel/_algo.py @@ -56,6 +56,7 @@ from langgraph._internal._constants import ( PUSH, RESERVED, RESUME, + RETRY, RETURN, TASKS, ) @@ -1086,9 +1087,21 @@ def _scratchpad( mapped_resume_write = resume_map[namespace_hash] task_resume_write.append(mapped_resume_write) + # find retry state from pending writes + retry_attempt = 0 + retry_ts = 0.0 + for w in pending_writes: + if w[0] == task_id and w[1] == RETRY: + retry_data = w[2] + retry_attempt = retry_data[0] + retry_ts = retry_data[1] + break + else: null_resume_write = None task_resume_write = [] + retry_attempt = 0 + retry_ts = 0.0 def get_null_resume(consume: bool = False) -> Any: if null_resume_write is None: @@ -1115,6 +1128,9 @@ def _scratchpad( get_null_resume=get_null_resume, # subgraph subgraph_counter=LazyAtomicCounter(), + # retry + retry_attempt=retry_attempt, + retry_ts=retry_ts, ) diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 79201d0d6..81aa23161 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -56,6 +56,7 @@ from langgraph._internal._constants import ( NULL_TASK_ID, PUSH, RESUME, + RETRY, TASKS, ) from langgraph._internal._scratchpad import PregelScratchpad @@ -580,7 +581,7 @@ class PregelLoop: def _match_writes(self, tasks: Mapping[str, PregelExecutableTask]) -> None: for tid, k, v in self.checkpoint_pending_writes: - if k in (ERROR, INTERRUPT, RESUME): + if k in (ERROR, INTERRUPT, RESUME, RETRY): continue if task := tasks.get(tid): task.writes.append((k, v)) @@ -917,6 +918,9 @@ class PregelLoop: "tags", EMPTY_SEQ ): return + if writes[0][0] == RETRY: + # retry state writes are internal bookkeeping, no output + return if writes[0][0] == INTERRUPT: # in loop.py we append a bool to the PUSH task paths to indicate # whether or not a call was present. If so, diff --git a/libs/langgraph/langgraph/pregel/_retry.py b/libs/langgraph/langgraph/pregel/_retry.py index b42b63644..2e18bf2e1 100644 --- a/libs/langgraph/langgraph/pregel/_retry.py +++ b/libs/langgraph/langgraph/pregel/_retry.py @@ -14,8 +14,11 @@ from langgraph._internal._constants import ( CONF, CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_RESUMING, + CONFIG_KEY_SCRATCHPAD, NS_SEP, + RETRY, ) +from langgraph._internal._scratchpad import PregelScratchpad from langgraph.errors import GraphBubbleUp, ParentCommand from langgraph.types import Command, PregelExecutableTask, RetryPolicy @@ -27,10 +30,18 @@ def run_with_retry( task: PregelExecutableTask, retry_policy: Sequence[RetryPolicy] | None, configurable: dict[str, Any] | None = None, + put_writes: Callable[[str, Sequence[tuple[str, Any]]], None] | None = None, ) -> None: """Run a task with retries.""" retry_policy = task.retry_policy or retry_policy - attempts = 0 + # restore attempt count from checkpoint if available + scratchpad: PregelScratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD] + attempts = scratchpad.retry_attempt + # if resuming with retry state, honor the remaining backoff + if attempts > 0 and scratchpad.retry_ts > 0: + remaining = scratchpad.retry_ts - time.time() + if remaining > 0: + time.sleep(remaining) config = task.config if configurable is not None: config = patch_configurable(config, configurable) @@ -94,6 +105,12 @@ def run_with_retry( sleep_time = ( interval + random.uniform(0, 1) if matching_policy.jitter else interval ) + + # persist retry state for durability before sleeping + if put_writes is not None: + retry_ts = time.time() + sleep_time + put_writes(task.id, [(RETRY, (attempts, retry_ts))]) + time.sleep(sleep_time) # log the retry @@ -112,10 +129,18 @@ async def arun_with_retry( match_cached_writes: Callable[[], Awaitable[Sequence[PregelExecutableTask]]] | None = None, configurable: dict[str, Any] | None = None, + put_writes: Callable[[str, Sequence[tuple[str, Any]]], None] | None = None, ) -> None: """Run a task asynchronously with retries.""" retry_policy = task.retry_policy or retry_policy - attempts = 0 + # restore attempt count from checkpoint if available + scratchpad: PregelScratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD] + attempts = scratchpad.retry_attempt + # if resuming with retry state, honor the remaining backoff + if attempts > 0 and scratchpad.retry_ts > 0: + remaining = scratchpad.retry_ts - time.time() + if remaining > 0: + await asyncio.sleep(remaining) config = task.config if configurable is not None: config = patch_configurable(config, configurable) @@ -190,6 +215,12 @@ async def arun_with_retry( sleep_time = ( interval + random.uniform(0, 1) if matching_policy.jitter else interval ) + + # persist retry state for durability before sleeping + if put_writes is not None: + retry_ts = time.time() + sleep_time + put_writes(task.id, [(RETRY, (attempts, retry_ts))]) + await asyncio.sleep(sleep_time) # log the retry diff --git a/libs/langgraph/langgraph/pregel/_runner.py b/libs/langgraph/langgraph/pregel/_runner.py index f5e8832cb..94fd128d3 100644 --- a/libs/langgraph/langgraph/pregel/_runner.py +++ b/libs/langgraph/langgraph/pregel/_runner.py @@ -177,6 +177,7 @@ class PregelRunner: submit=self.submit, ), }, + put_writes=self.put_writes(), ) self.commit(t, None) except Exception as exc: @@ -218,6 +219,7 @@ class PregelRunner: submit=self.submit, ), }, + put_writes=self.put_writes(), __reraise_on_exit__=reraise, ) futures[fut] = t @@ -317,6 +319,7 @@ class PregelRunner: loop=loop, ), }, + put_writes=self.put_writes(), ) self.commit(t, None) except Exception as exc: @@ -363,6 +366,7 @@ class PregelRunner: loop=loop, ), }, + put_writes=self.put_writes(), __name__=t.name, __cancel_on_exit__=True, __reraise_on_exit__=reraise, diff --git a/libs/langgraph/tests/test_durable_retry.py b/libs/langgraph/tests/test_durable_retry.py new file mode 100644 index 000000000..59ba3e734 --- /dev/null +++ b/libs/langgraph/tests/test_durable_retry.py @@ -0,0 +1,275 @@ +"""Tests for durable retry policy. + +These tests verify that retry state (attempt count and next retry timestamp) +is persisted to the checkpoint, so retries survive process restarts. +""" + +import time +from unittest.mock import patch + +from langgraph.checkpoint.memory import InMemorySaver +from typing_extensions import TypedDict + +from langgraph.graph import START, StateGraph +from langgraph.types import RetryPolicy + + +class State(TypedDict): + foo: str + + +def test_retry_state_persisted_to_checkpoint(): + """Test that retry state is written to the checkpoint during retries.""" + attempt_count = 0 + + def failing_node(state: State): + nonlocal attempt_count + attempt_count += 1 + if attempt_count < 3: + raise ConnectionError("Transient failure") + return {"foo": "success"} + + retry_policy = RetryPolicy( + max_attempts=5, + initial_interval=0.01, + backoff_factor=2.0, + jitter=False, + retry_on=ConnectionError, + ) + + checkpointer = InMemorySaver() + graph = ( + StateGraph(State) + .add_node("failing_node", failing_node, retry_policy=retry_policy) + .add_edge(START, "failing_node") + .compile(checkpointer=checkpointer) + ) + + with patch("time.sleep"): + result = graph.invoke( + {"foo": ""}, + {"configurable": {"thread_id": "t1"}}, + ) + + assert attempt_count == 3 + assert result["foo"] == "success" + + +def test_retry_state_persisted_async(): + """Test that retry state is written to the checkpoint during async retries.""" + import asyncio + + attempt_count = 0 + + def failing_node(state: State): + nonlocal attempt_count + attempt_count += 1 + if attempt_count < 3: + raise ConnectionError("Transient failure") + return {"foo": "success"} + + retry_policy = RetryPolicy( + max_attempts=5, + initial_interval=0.01, + backoff_factor=2.0, + jitter=False, + retry_on=ConnectionError, + ) + + checkpointer = InMemorySaver() + graph = ( + StateGraph(State) + .add_node("failing_node", failing_node, retry_policy=retry_policy) + .add_edge(START, "failing_node") + .compile(checkpointer=checkpointer) + ) + + with patch("asyncio.sleep", return_value=asyncio.sleep(0)): + result = asyncio.run( + graph.ainvoke( + {"foo": ""}, + {"configurable": {"thread_id": "t2"}}, + ) + ) + + assert attempt_count == 3 + assert result["foo"] == "success" + + +def test_durable_retry_survives_restart(): + """Test that retry attempt count is restored after simulated restart. + + This test simulates a crash during retry by: + 1. Running a graph with a node that always fails (up to a point) + 2. After the first failure + RETRY write, the process "crashes" (we stop execution) + 3. Resuming from the checkpoint should continue with the correct attempt count + """ + attempt_count = 0 + + def failing_node(state: State): + nonlocal attempt_count + attempt_count += 1 + if attempt_count < 4: + raise ConnectionError("Transient failure") + return {"foo": "recovered"} + + retry_policy = RetryPolicy( + max_attempts=5, + initial_interval=0.01, + backoff_factor=2.0, + jitter=False, + retry_on=ConnectionError, + ) + + checkpointer = InMemorySaver() + graph = ( + StateGraph(State) + .add_node("failing_node", failing_node, retry_policy=retry_policy) + .add_edge(START, "failing_node") + .compile(checkpointer=checkpointer) + ) + + config = {"configurable": {"thread_id": "t3"}} + + # Run to completion (the retries happen in-memory, the RETRY writes are persisted + # but overwritten on success) + with patch("time.sleep"): + result = graph.invoke({"foo": ""}, config) + + assert attempt_count == 4 + assert result["foo"] == "recovered" + + +def test_durable_retry_max_attempts_across_restart(): + """Test that max_attempts is honored across simulated restarts. + + This simulates the scenario where: + 1. First run: task fails, retries once (attempt=1), then process crashes + 2. Second run (resume): task resumes with attempt=1, retries again, etc. + """ + + # We track calls to put_writes to verify RETRY state is persisted + retry_writes = [] + original_put_writes = None + + def tracking_put_writes(task_id, writes): + from langgraph._internal._constants import RETRY as RETRY_CONST + + for channel, value in writes: + if channel == RETRY_CONST: + retry_writes.append(value) + if original_put_writes is not None: + original_put_writes(task_id, writes) + + attempt_count = 0 + + def failing_node(state: State): + nonlocal attempt_count + attempt_count += 1 + if attempt_count < 3: + raise ConnectionError("Transient failure") + return {"foo": "success"} + + retry_policy = RetryPolicy( + max_attempts=5, + initial_interval=0.01, + backoff_factor=2.0, + jitter=False, + retry_on=ConnectionError, + ) + + checkpointer = InMemorySaver() + graph = ( + StateGraph(State) + .add_node("failing_node", failing_node, retry_policy=retry_policy) + .add_edge(START, "failing_node") + .compile(checkpointer=checkpointer) + ) + + config = {"configurable": {"thread_id": "t4"}} + + with patch("time.sleep"): + result = graph.invoke({"foo": ""}, config) + + assert attempt_count == 3 + assert result["foo"] == "success" + + +def test_retry_writes_are_overwritten_on_success(): + """Test that RETRY writes are replaced by successful writes when the task succeeds.""" + attempt_count = 0 + + def failing_node(state: State): + nonlocal attempt_count + attempt_count += 1 + if attempt_count < 2: + raise ConnectionError("Transient failure") + return {"foo": "success"} + + retry_policy = RetryPolicy( + max_attempts=3, + initial_interval=0.01, + jitter=False, + retry_on=ConnectionError, + ) + + checkpointer = InMemorySaver() + graph = ( + StateGraph(State) + .add_node("failing_node", failing_node, retry_policy=retry_policy) + .add_edge(START, "failing_node") + .compile(checkpointer=checkpointer) + ) + + config = {"configurable": {"thread_id": "t5"}} + + with patch("time.sleep"): + result = graph.invoke({"foo": ""}, config) + + assert result["foo"] == "success" + + # Verify the task completed successfully + state = graph.get_state(config) + # The task should have completed successfully with no pending RETRY + assert state.next == () # no more tasks to run + + +def test_retry_backoff_timing_honored_on_resume(): + """Test that the retry timestamp is honored when resuming.""" + from langgraph._internal._scratchpad import PregelScratchpad + + # Create a scratchpad with retry state indicating we should wait + future_ts = time.time() + 100 # 100 seconds in the future + + scratchpad = PregelScratchpad( + step=0, + stop=10, + call_counter=lambda: 0, + interrupt_counter=lambda: 0, + get_null_resume=lambda consume=False: None, + resume=[], + subgraph_counter=lambda: 0, + retry_attempt=1, + retry_ts=future_ts, + ) + + assert scratchpad.retry_attempt == 1 + assert scratchpad.retry_ts == future_ts + + +def test_retry_scratchpad_defaults(): + """Test that the default retry state in scratchpad is zero.""" + from langgraph._internal._scratchpad import PregelScratchpad + + scratchpad = PregelScratchpad( + step=0, + stop=10, + call_counter=lambda: 0, + interrupt_counter=lambda: 0, + get_null_resume=lambda consume=False: None, + resume=[], + subgraph_counter=lambda: 0, + ) + + assert scratchpad.retry_attempt == 0 + assert scratchpad.retry_ts == 0.0