From 6dff3b3bce67386e18d7e5dcbc3f5b569b8821e7 Mon Sep 17 00:00:00 2001 From: Quanzheng Long Date: Mon, 11 May 2026 18:39:03 -0700 Subject: [PATCH] feat(langgraph): durable error-handler resume across host crashes (#7773) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Summary - **Consolidate error writes:** When a node fails and has an error handler, `commit()` now appends both `ERROR` and `ERROR_SOURCE_NODE` in a single `put_writes` call, eliminating the redundant overwrite that `schedule_error_handler` used to do. - **Ensure durability before handler execution:** Reuses the `_delta_write_futs` pattern — a new `_error_handler_write_futs` list collects the persistence future from `put_writes` when `ERROR_SOURCE_NODE` is written, and `schedule_error_handler` / `aschedule_error_handler` drain it (sync: `concurrent.futures.wait`, async: `asyncio.gather`) before preparing the handler task. - **Resume directly to error handler:** Adds `_resume_error_handlers_if_applicable()` to `PregelLoop`, called from `tick()` after `_reapply_writes_to_succeeded_nodes()`. On resume, it detects `ERROR_SOURCE_NODE` markers in `checkpoint_pending_writes`, marks the original task as done (so the runner skips it), and schedules a fresh handler task. - **Rename internal methods for clarity:** `_match_writes` → `_reapply_writes_to_succeeded_nodes` (makes it clear that failed/interrupted tasks are skipped); `_resume_error_handlers` → `_resume_error_handlers_if_applicable`. ## Test plan - [x] `test_error_handler_resumes_after_crash`: single node fails, handler crashes, resume re-runs the handler (not the original node). Verifies `NodeError.node` and error content survive checkpoint round-trip. - [x] `test_error_handler_resumes_after_crash_multiple_nodes`: two nodes fail concurrently in the same superstep, each with its own handler. Verifies error handler starts while other nodes are still in-flight (via `threading.Event`), and on resume both handlers re-run with correct `NodeError.node` and error content. - [x] All 101 tests in `test_retry.py` pass (including 19 error-handler tests). - [x] `make format` + `make lint` clean. --- libs/langgraph/langgraph/pregel/_loop.py | 121 +++++++++++++++--- libs/langgraph/langgraph/pregel/_runner.py | 3 +- libs/langgraph/tests/test_retry.py | 135 +++++++++++++++++++++ 3 files changed, 241 insertions(+), 18 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 9949a62d9..9b98fc91c 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -203,6 +203,12 @@ class PregelLoop: # `__enter__`; stays `None` only when no checkpointer. _delta_write_futs: list[Any] | None = None + # Same pattern as `_delta_write_futs` but for error-handler writes. + # When `put_writes` persists an ERROR_SOURCE_NODE marker, the future is + # appended here. `schedule_error_handler` / `aschedule_error_handler` + # drain this list so the write is durable before the handler starts. + _error_handler_write_futs: list[Any] | None = None + # Exit-mode accumulator: every delta-channel write produced during this # run (input writes from `_first` + per-superstep writes captured in # `after_tick`). At exit, `_put_exit_delta_writes` filters out channels @@ -474,6 +480,13 @@ class PregelLoop: isinstance(self.specs.get(c), DeltaChannel) for c, _ in writes_to_save ): self._delta_write_futs.append(fut) + # ERROR_SOURCE_NODE is only appended by commit() when the task + # has an error handler (_should_route_to_error_handler), so this + # check naturally limits future collection to those tasks. + if self._error_handler_write_futs is not None and any( + c == ERROR_SOURCE_NODE for c, _ in writes + ): + self._error_handler_write_futs.append(fut) # output writes if hasattr(self, "tasks"): self.output_writes(task_id, writes) @@ -553,7 +566,7 @@ class PregelLoop: self.tasks[pushed.id] = pushed # match any pending writes to the new task if not self.is_replaying: - self._match_writes({pushed.id: pushed}) + self._reapply_writes_to_succeeded_nodes({pushed.id: pushed}) # return the new task, to be started if not run before return pushed @@ -631,7 +644,8 @@ class PregelLoop: # if there are pending writes from a previous loop, apply them if not self.is_replaying and self.checkpoint_pending_writes: - self._match_writes(self.tasks) + self._reapply_writes_to_succeeded_nodes(self.tasks) + self._resume_error_handlers_if_applicable() # before execution, check if we should interrupt if self.interrupt_before and should_interrupt( @@ -698,13 +712,88 @@ class PregelLoop: # private - def _match_writes(self, tasks: Mapping[str, PregelExecutableTask]) -> None: + def _reapply_writes_to_succeeded_nodes( + self, tasks: Mapping[str, PregelExecutableTask] + ) -> None: + """Restore successful channel writes from checkpoint to in-memory tasks. + + Skips control signals (ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME) + so that failed/interrupted tasks remain with empty writes and will be + re-executed (or routed to error handlers) by the runner. + """ for tid, k, v in self.checkpoint_pending_writes: if k in (ERROR, ERROR_SOURCE_NODE, INTERRUPT, RESUME): continue if task := tasks.get(tid): task.writes.append((k, v)) + def _resume_error_handlers_if_applicable(self) -> None: + """On resume, schedule error handlers for tasks that failed in a prior run. + + Called right after ``_reapply_writes_to_succeeded_nodes`` during ``tick()``. + At that point, ``_reapply_writes_to_succeeded_nodes`` has already skipped + ERROR / ERROR_SOURCE_NODE writes, so a previously-failed task still has + empty ``writes``. Without intervention the runner (which executes only + tasks where ``not t.writes``) would re-run the original node. + + This method prevents that re-execution for nodes that have an error + handler: + + 1. Scan ``checkpoint_pending_writes`` for ERROR_SOURCE_NODE markers + persisted by a prior ``commit()``. Each marker means "this task + already failed and was routed to an error handler". + 2. For each such task, write ``(ERROR, error)`` into ``task.writes`` + so the task is no longer empty — the runner will skip it. + 3. Prepare a fresh error-handler task and add it to ``self.tasks``. + Because the handler task starts with empty ``writes``, the runner + will pick it up and execute it. + """ + # Phase 1: collect task-ids that have ERROR_SOURCE_NODE + ERROR pairs. + failed: dict[str, BaseException] = {} + for tid, chan, val in self.checkpoint_pending_writes: + if chan == ERROR_SOURCE_NODE: + error = next( + ( + v + for t, c, v in self.checkpoint_pending_writes + if t == tid and c == ERROR + ), + None, + ) + if error is not None: + failed[tid] = error + # Phase 2: mark originals as done, schedule handler tasks. + for task_id, error in failed.items(): + task = self.tasks.get(task_id) + if task is None: + continue + handler_node = self.nodes[task.name].error_handler_node + if not handler_node: + continue + # Non-empty writes → runner's `not t.writes` filter skips this task. + task.writes.append((ERROR, error)) + # The handler task starts with empty writes → runner will execute it. + handler_task = prepare_node_error_handler_task( + task, + handler_node_name=handler_node, + failed_error=error, + checkpoint=self.checkpoint, + pending_writes=self.checkpoint_pending_writes, + processes=self.nodes, + channels=self.channels, + managed=self.managed, + config=task.config, + step=self.step, + stop=self.stop, + store=self.store, + checkpointer=self.checkpointer, + manager=self.manager, + retry_policy=self.retry_policy, + cache_policy=self.cache_policy, + ) + if handler_task is not None: + self.tasks[handler_task.id] = handler_task + def _pending_interrupts(self) -> set[str]: """Return the set of interrupt ids that are pending without corresponding resume values.""" # mapping of task ids to interrupt ids @@ -1454,12 +1543,10 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): handler_node = self.nodes[failed_task.name].error_handler_node if not handler_node: return None - writes = list(failed_task.writes) - writes.append((ERROR_SOURCE_NODE, failed_task.name)) - self.put_writes( - failed_task.id, - writes, - ) + # ensure error + ERROR_SOURCE_NODE writes are durable before handler runs + if self._error_handler_write_futs: + futs, self._error_handler_write_futs = self._error_handler_write_futs, [] + concurrent.futures.wait(futs) handler_task = prepare_node_error_handler_task( failed_task, handler_node_name=handler_node, @@ -1482,7 +1569,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): return None self.tasks[handler_task.id] = handler_task if not self.is_replaying: - self._match_writes({handler_task.id: handler_task}) + self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task}) for task in self.match_cached_writes(): self.output_writes(task.id, task.writes, cached=True) return handler_task @@ -1564,6 +1651,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): else [] ) self._delta_write_futs = [] + self._error_handler_write_futs = [] self._exit_delta_writes = ( [] if self.durability == "exit" and self.checkpointer is not None else None ) @@ -1709,12 +1797,10 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): handler_node = self.nodes[failed_task.name].error_handler_node if not handler_node: return None - writes = list(failed_task.writes) - writes.append((ERROR_SOURCE_NODE, failed_task.name)) - self.put_writes( - failed_task.id, - writes, - ) + # ensure error + ERROR_SOURCE_NODE writes are durable before handler runs + if self._error_handler_write_futs: + futs, self._error_handler_write_futs = self._error_handler_write_futs, [] + await asyncio.gather(*futs) handler_task = prepare_node_error_handler_task( failed_task, handler_node_name=handler_node, @@ -1737,7 +1823,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): return None self.tasks[handler_task.id] = handler_task if not self.is_replaying: - self._match_writes({handler_task.id: handler_task}) + self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task}) for task in await self.amatch_cached_writes(): self.output_writes(task.id, task.writes, cached=True) return handler_task @@ -1822,6 +1908,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): else [] ) self._delta_write_futs = [] + self._error_handler_write_futs = [] self._exit_delta_writes = ( [] if self.durability == "exit" and self.checkpointer is not None else None ) diff --git a/libs/langgraph/langgraph/pregel/_runner.py b/libs/langgraph/langgraph/pregel/_runner.py index 979935c9a..4d53b9f9d 100644 --- a/libs/langgraph/langgraph/pregel/_runner.py +++ b/libs/langgraph/langgraph/pregel/_runner.py @@ -31,6 +31,7 @@ from langgraph._internal._constants import ( CONFIG_KEY_CALL, CONFIG_KEY_SCRATCHPAD, ERROR, + ERROR_SOURCE_NODE, INTERRUPT, NO_WRITES, RESUME, @@ -597,7 +598,7 @@ class PregelRunner: if self._should_route_to_error_handler(task) and not isinstance( exception, GraphBubbleUp ): - # Mark early in commit path; loop-side routing may happen later. + task.writes.append((ERROR_SOURCE_NODE, task.name)) self._handled_exception_ids.add(id(exception)) self.put_writes()(task.id, task.writes) # type: ignore[misc] else: diff --git a/libs/langgraph/tests/test_retry.py b/libs/langgraph/tests/test_retry.py index be57e9360..ee54f4b39 100644 --- a/libs/langgraph/tests/test_retry.py +++ b/libs/langgraph/tests/test_retry.py @@ -2667,3 +2667,138 @@ def test_set_node_defaults_combined_retry_and_error_handler(): assert result["foo"] == "handled" assert attempts == 2 assert captured["error"] == "Always fails" + + +def test_error_handler_resumes_after_crash(): + """If the error handler crashes, resuming should re-schedule the handler + (not re-execute the original failed node).""" + + class State(TypedDict): + foo: str + + call_count = {"node": 0, "handler": 0} + captured_errors: list[NodeError] = [] + + def failing_node(state: State) -> State: + call_count["node"] += 1 + raise RuntimeError("boom") + + handler_should_fail = [True] + + def handler(state: State, error: NodeError) -> State: + call_count["handler"] += 1 + captured_errors.append(error) + if handler_should_fail[0]: + raise RuntimeError("handler crash") + return {"foo": "recovered"} + + checkpointer = MemorySaver() + graph = ( + StateGraph(State) + .set_node_defaults(error_handler=handler) + .add_node("fail", failing_node) + .add_edge(START, "fail") + .compile(checkpointer=checkpointer) + ) + + config = {"configurable": {"thread_id": "t1"}} + + # First invoke: node fails -> handler runs -> handler crashes -> run fails + with pytest.raises(RuntimeError, match="handler crash"): + graph.invoke({"foo": ""}, config) + + assert call_count["node"] == 1 + assert call_count["handler"] == 1 + assert captured_errors[0].node == "fail" + assert isinstance(captured_errors[0].error, RuntimeError) + assert str(captured_errors[0].error) == "boom" + + # Resume: handler should run again, NOT the original node + handler_should_fail[0] = False + result = graph.invoke(None, config) + + assert result["foo"] == "recovered" + assert call_count["node"] == 1 # NOT re-executed + assert call_count["handler"] == 2 # ran again on resume + # on resume the error was round-tripped through the checkpointer, so it + # may be deserialized as a string representation rather than the original + # exception type — verify the node name and that the error content matches. + assert captured_errors[1].node == "fail" + assert "boom" in str(captured_errors[1].error) + + +def test_error_handler_resumes_after_crash_multiple_nodes(): + """When multiple nodes fail in the same superstep and all have error handlers: + - error handlers start running while other nodes may still be in-flight + - resuming re-schedules each handler (not re-executes the original nodes) + """ + + class State(TypedDict): + results: Annotated[list[str], operator.add] + + call_count = {"a": 0, "b": 0, "handler_a": 0, "handler_b": 0} + handler_a_started = threading.Event() + + def node_a(state: State) -> State: + call_count["a"] += 1 + raise RuntimeError("a failed") + + def node_b(state: State) -> State: + call_count["b"] += 1 + # Block until handler_a has started — proves the error handler runs + # concurrently with in-flight nodes in the same superstep. + assert handler_a_started.wait(timeout=5), "handler_a never started" + raise RuntimeError("b failed") + + handler_should_fail = [True] + + def handler_a(state: State, error: NodeError) -> State: + call_count["handler_a"] += 1 + assert error.node == "a" + assert "a failed" in str(error.error) + handler_a_started.set() + if handler_should_fail[0]: + raise RuntimeError("handler_a crash") + return {"results": [f"recovered_a:{error.node}"]} + + def handler_b(state: State, error: NodeError) -> State: + call_count["handler_b"] += 1 + assert error.node == "b" + assert "b failed" in str(error.error) + if handler_should_fail[0]: + raise RuntimeError("handler_b crash") + return {"results": [f"recovered_b:{error.node}"]} + + checkpointer = MemorySaver() + graph = ( + StateGraph(State) + .add_node("a", node_a, error_handler=handler_a) + .add_node("b", node_b, error_handler=handler_b) + .add_edge(START, "a") + .add_edge(START, "b") + .compile(checkpointer=checkpointer) + ) + + config = {"configurable": {"thread_id": "t1"}} + + # First invoke: node_a fails immediately -> handler_a starts (sets event) -> + # node_b unblocks and fails -> handler_b starts -> both handlers crash + with pytest.raises(RuntimeError): + graph.invoke({"results": []}, config) + + assert call_count["a"] == 1 + assert call_count["b"] == 1 + assert call_count["handler_a"] == 1 + assert call_count["handler_b"] == 1 + + # Resume: both handlers should run again, NOT the original nodes + handler_should_fail[0] = False + handler_a_started.clear() + result = graph.invoke(None, config) + + assert call_count["a"] == 1 # NOT re-executed + assert call_count["b"] == 1 # NOT re-executed + assert call_count["handler_a"] == 2 # ran again on resume + assert call_count["handler_b"] == 2 # ran again on resume + assert "recovered_a:a" in result["results"] + assert "recovered_b:b" in result["results"]