mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-01 04:39:01 +02:00
feat(langgraph): durable error-handler resume across host crashes (#7773)
## 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.
This commit is contained in:
@@ -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
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Reference in New Issue
Block a user