Compare commits

..
Author SHA1 Message Date
Sydney Runkle a55365f3f9 simplify: drop _should_route_to_error_handler, remove functional API error handler, clean up prepare_node_error_handler_task signature 2026-05-11 16:22:10 -07:00
Sydney Runkle 2cd7ecc81e simplify: handlers always receive state as first arg, drop takes_input flag 2026-05-11 16:04:47 -07:00
Sydney Runkle f6746b39cc refactor(langgraph): implement error_handler as a policy field, not a hidden node
Replace the hidden __error_handler__<name> node approach with a
callable field on PregelNode/PregelExecutableTask, matching how
retry_policy and cache_policy work:

- error_handler: Runnable | None lives on StateNodeSpec, PregelNode,
  and PregelExecutableTask — no separate node registration at compile time
- ErrorHandlerNode type alias added to _node.py covering the common
  handler signatures (state + NodeError, runtime + NodeError, etc.)
- Handler task is synthesized at runtime in schedule_error_handler using
  the failed task's write pipeline (failed_task.writers), so no separate
  PregelNode is needed
- compile(error_handler=...) adds a graph-level fallback handler;
  per-node handlers take precedence
- @task(error_handler=...) in the functional API uses inline try/except
  wrapping so the caller's future always resolves to a value
- interrupt_before/after ["__error_handler__<node>"] still works via
  updated validate() and _validate.py checks
- RunnableCallable gains takes_input flag so handlers whose only
  parameters are injected kwargs (e.g. def h(error: NodeError) -> T)
  are called without a spurious positional input arg
2026-05-11 15:59:22 -07:00
29 changed files with 279 additions and 938 deletions
+1 -1
View File
@@ -279,7 +279,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.1.0"
version = "4.1.0a4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
+2 -2
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-checkpoint-postgres"
version = "3.1.0"
version = "3.1.0a4"
description = "Library with a Postgres implementation of LangGraph checkpoint saver."
authors = []
requires-python = ">=3.10"
@@ -12,7 +12,7 @@ readme = "README.md"
license = "MIT"
license-files = ['LICENSE']
dependencies = [
"langgraph-checkpoint>=4.1.0,<5.0.0",
"langgraph-checkpoint>=4.1.0a4,<5.0.0",
"orjson>=3.11.5",
"psycopg>=3.2.0",
"psycopg-pool>=3.2.0",
+2 -2
View File
@@ -276,7 +276,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.1.0"
version = "4.1.0a4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -324,7 +324,7 @@ test = [
[[package]]
name = "langgraph-checkpoint-postgres"
version = "3.1.0"
version = "3.1.0a4"
source = { editable = "." }
dependencies = [
{ name = "langgraph-checkpoint" },
+2 -2
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-checkpoint-sqlite"
version = "3.1.0"
version = "3.1.0a1"
description = "Library with a SQLite implementation of LangGraph checkpoint saver."
authors = []
requires-python = ">=3.10"
@@ -12,7 +12,7 @@ readme = "README.md"
license = "MIT"
license-files = ['LICENSE']
dependencies = [
"langgraph-checkpoint>=4.1.0,<5.0.0",
"langgraph-checkpoint>=4.1.0a4,<5.0.0",
"aiosqlite>=0.20",
"sqlite-vec>=0.1.6",
]
+2 -2
View File
@@ -285,7 +285,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.1.0"
version = "4.1.0a4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -333,7 +333,7 @@ test = [
[[package]]
name = "langgraph-checkpoint-sqlite"
version = "3.1.0"
version = "3.1.0a1"
source = { editable = "." }
dependencies = [
{ name = "aiosqlite" },
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-checkpoint"
version = "4.1.0"
version = "4.1.0a4"
description = "Library with base interfaces for LangGraph checkpoint savers."
authors = []
requires-python = ">=3.10"
+7 -7
View File
@@ -268,7 +268,7 @@ wheels = [
[[package]]
name = "langchain-core"
version = "1.3.3"
version = "1.3.2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jsonpatch" },
@@ -281,9 +281,9 @@ dependencies = [
{ name = "typing-extensions" },
{ name = "uuid-utils" },
]
sdist = { url = "https://files.pythonhosted.org/packages/d3/ae/8b74458fc3850ec3d150eb9f45e857db129dafa801fb5cf173dfc9f8bbf3/langchain_core-1.3.3.tar.gz", hash = "sha256:fa510a5db8efdc0c6ff41c0939fb5c00a0183c11f6b84233e892e3227ff69182", size = 915041, upload-time = "2026-05-05T19:02:36.612Z" }
sdist = { url = "https://files.pythonhosted.org/packages/a8/03/7219502e8ca728d65eb44d7a3eb60239230742a70dbfc9241b9bfd61c4ab/langchain_core-1.3.2.tar.gz", hash = "sha256:fd7a50b2f28ba561fd9d7f5d2760bc9e06cf00cdf820a3ccafe88a94ffa8d5b7", size = 911813, upload-time = "2026-04-24T15:49:23.699Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/1f/01/4771b7ab2af1d1aba5b710bd8f13d9225c609425214b357590a17b01be77/langchain_core-1.3.3-py3-none-any.whl", hash = "sha256:18aae8506f37da7f74398492279a7d6efcee4f8e23c4c41c7af080eeb7ef7bd1", size = 543857, upload-time = "2026-05-05T19:02:34.52Z" },
{ url = "https://files.pythonhosted.org/packages/7d/d5/8fa4431007cbb7cfed7590f4d6a5dea3ad724f4174d248f6642ef5ce7d05/langchain_core-1.3.2-py3-none-any.whl", hash = "sha256:d44a66127f9f8db735bdfd0ab9661bccb47a97113cfd3f2d89c74864422b7274", size = 542390, upload-time = "2026-04-24T15:49:21.991Z" },
]
[[package]]
@@ -300,7 +300,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.1.0"
version = "4.1.0a4"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
@@ -1463,11 +1463,11 @@ wheels = [
[[package]]
name = "urllib3"
version = "2.7.0"
version = "2.6.3"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/53/0c/06f8b233b8fd13b9e5ee11424ef85419ba0d8ba0b3138bf360be2ff56953/urllib3-2.7.0.tar.gz", hash = "sha256:231e0ec3b63ceb14667c67be60f2f2c40a518cb38b03af60abc813da26505f4c", size = 433602, upload-time = "2026-05-07T16:13:18.596Z" }
sdist = { url = "https://files.pythonhosted.org/packages/c7/24/5f1b3bdffd70275f6661c76461e25f024d5a38a46f04aaca912426a2b1d3/urllib3-2.6.3.tar.gz", hash = "sha256:1b62b6884944a57dbe321509ab94fd4d3b307075e0c2eae991ac71ee15ad38ed", size = 435556, upload-time = "2026-01-07T16:24:43.925Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/7f/3e/5db95bcf282c52709639744ca2a8b149baccf648e39c8cc87553df9eae0c/urllib3-2.7.0-py3-none-any.whl", hash = "sha256:9fb4c81ebbb1ce9531cce37674bbc6f1360472bc18ca9a553ede278ef7276897", size = 131087, upload-time = "2026-05-07T16:13:17.151Z" },
{ url = "https://files.pythonhosted.org/packages/39/08/aaaad47bc4e9dc8c725e68f9d04865dbcb2052843ff09c97b08904852d84/urllib3-2.6.3-py3-none-any.whl", hash = "sha256:bf272323e553dfb2e87d9bfd225ca7b0f467b919d7bbd355436d3fd37cb0acd4", size = 131584, upload-time = "2026-01-07T16:24:42.685Z" },
]
[[package]]
+1 -1
View File
@@ -1 +1 @@
__version__ = "0.4.26"
__version__ = "0.4.25"
+2 -20
View File
@@ -35,12 +35,6 @@ DISALLOWED_BUILD_COMMAND_CHARS = [
# This blocks background execution (cmd &) while allowing command
# chaining (cmd1 && cmd2) which is common in build commands.
_SINGLE_AMPERSAND_RE = re.compile(r"(?<!&)&(?:&&)*(?!&)")
_API_VERSION_PATTERN = re.compile(
r"^(?P<major>\d+)"
r"(?:\.(?P<minor>\d+))?"
r"(?:\.(?P<patch>\d+))?"
r"(?:(?:\.|)(?:[A-Za-z][0-9A-Za-z]*))?$"
)
def has_disallowed_build_command_content(command: str) -> bool:
@@ -129,18 +123,6 @@ def _parse_node_version(version_str: str) -> int:
) from None
def _parse_api_version_parts(version_str: str) -> tuple[int, ...]:
"""Parse an API version into numeric components.
Supports optional prerelease suffixes, e.g. `0.9.0rc1`.
"""
version_core = version_str.split("-", 1)[0]
match = _API_VERSION_PATTERN.fullmatch(version_core)
if not match:
raise ValueError("Version must be major or major.minor or major.minor.patch.")
return tuple(int(part) for part in match.groups() if part is not None)
def _is_node_graph(spec: str | dict) -> bool:
"""Check if a graph is a Node.js graph based on the file extension."""
if isinstance(spec, dict):
@@ -194,12 +176,12 @@ def validate_config(config: Config) -> Config:
)
if api_version:
try:
parts = _parse_api_version_parts(api_version)
parts = tuple(map(int, api_version.split("-")[0].split(".")))
if len(parts) > 3:
raise ValueError(
"Version must be major or major.minor or major.minor.patch."
)
except (TypeError, ValueError):
except TypeError:
raise click.UsageError(
f"Invalid version format: {api_version}.\n\n"
"Pin to a minor version, e.g.:\n"
-17
View File
@@ -2944,23 +2944,6 @@ def test_docker_tag_with_api_version(in_config: bool):
assert tag == f"langchain/langgraph-server:{version}-py3.11"
@pytest.mark.parametrize("in_config", [False, True])
@pytest.mark.parametrize("version", ["0.9.0rc1", "0.9.0.dev1"])
def test_docker_tag_with_prerelease_api_version(version: str, in_config: bool):
"""Test docker_tag with prerelease and dev api_version values."""
config = validate_config(
{
"python_version": "3.11",
"dependencies": ["."],
"graphs": {"agent": "./agent.py:graph"},
"api_version": version if in_config else None,
}
)
tag = docker_tag(config, api_version=version if not in_config else None)
assert tag == f"langchain/langgraph-api:{version}-py3.11"
def test_config_to_docker_with_api_version():
"""Test config_to_docker function with api_version parameter."""
+1 -1
View File
@@ -5,7 +5,7 @@ description = "uv workspace monorepo example for LangGraph CLI integration test"
requires-python = ">=3.11"
dependencies = [
"langgraph>=0.6.0,<2",
"langchain-core>=1.3.3",
"langchain-core>=0.2.14",
]
[tool.uv.workspace]
+5 -18
View File
@@ -21,7 +21,7 @@ dependencies = [
[package.metadata]
requires-dist = [
{ name = "langchain-core", specifier = ">=1.3.3" },
{ name = "langchain-core", specifier = ">=0.2.14" },
{ name = "langgraph", specifier = ">=0.6.0,<2" },
{ name = "shared", editable = "libs/shared" },
]
@@ -215,11 +215,10 @@ wheels = [
[[package]]
name = "langchain-core"
version = "1.3.3"
version = "1.2.28"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jsonpatch" },
{ name = "langchain-protocol" },
{ name = "langsmith" },
{ name = "packaging" },
{ name = "pydantic" },
@@ -228,21 +227,9 @@ dependencies = [
{ name = "typing-extensions" },
{ name = "uuid-utils" },
]
sdist = { url = "https://files.pythonhosted.org/packages/d3/ae/8b74458fc3850ec3d150eb9f45e857db129dafa801fb5cf173dfc9f8bbf3/langchain_core-1.3.3.tar.gz", hash = "sha256:fa510a5db8efdc0c6ff41c0939fb5c00a0183c11f6b84233e892e3227ff69182", size = 915041, upload-time = "2026-05-05T19:02:36.612Z" }
sdist = { url = "https://files.pythonhosted.org/packages/f8/a4/317a1a3ac1df33a64adb3670bf88bbe3b3d5baa274db6863a979db472897/langchain_core-1.2.28.tar.gz", hash = "sha256:271a3d8bd618f795fdeba112b0753980457fc90537c46a0c11998516a74dc2cb", size = 846119, upload-time = "2026-04-08T18:19:34.867Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/1f/01/4771b7ab2af1d1aba5b710bd8f13d9225c609425214b357590a17b01be77/langchain_core-1.3.3-py3-none-any.whl", hash = "sha256:18aae8506f37da7f74398492279a7d6efcee4f8e23c4c41c7af080eeb7ef7bd1", size = 543857, upload-time = "2026-05-05T19:02:34.52Z" },
]
[[package]]
name = "langchain-protocol"
version = "0.0.15"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/4f/24/9777489d6fbbee64af0c8f96d4f840239c408cf694f3394672807dafc490/langchain_protocol-0.0.15.tar.gz", hash = "sha256:9ab2d11ee73944754f10e037e717098d3a6796f0e58afa9cadda6154e7655ade", size = 5862, upload-time = "2026-05-01T22:30:04.748Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/1d/7a/9c97a7b9cbe4c5dc6a44cdb1545450c28f0c8ce89b9c1f0ee7fbad896263/langchain_protocol-0.0.15-py3-none-any.whl", hash = "sha256:461eb794358f83d5e42635a5797799ffec7b4702314e34edf73ac21e75d3ef79", size = 6982, upload-time = "2026-05-01T22:30:03.877Z" },
{ url = "https://files.pythonhosted.org/packages/a8/92/32f785f077c7e898da97064f113c73fbd9ad55d1e2169cf3a391b183dedb/langchain_core-1.2.28-py3-none-any.whl", hash = "sha256:80764232581eaf8057bcefa71dbf8adc1f6a28d257ebd8b95ba9b8b452e8c6ac", size = 508727, upload-time = "2026-04-08T18:19:32.823Z" },
]
[[package]]
@@ -724,7 +711,7 @@ dependencies = [
[package.metadata]
requires-dist = [
{ name = "langchain-core", specifier = ">=1.3.3" },
{ name = "langchain-core", specifier = ">=0.2.14" },
{ name = "langgraph", specifier = ">=0.6.0,<2" },
]
+1 -1
View File
@@ -5,7 +5,7 @@ description = "Simple single-package uv example for LangGraph CLI integration te
requires-python = ">=3.11"
dependencies = [
"langgraph>=0.6.0,<2",
"langchain-core>=1.3.3",
"langchain-core>=0.2.14",
]
[build-system]
+4 -17
View File
@@ -191,11 +191,10 @@ wheels = [
[[package]]
name = "langchain-core"
version = "1.3.3"
version = "1.2.28"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jsonpatch" },
{ name = "langchain-protocol" },
{ name = "langsmith" },
{ name = "packaging" },
{ name = "pydantic" },
@@ -204,21 +203,9 @@ dependencies = [
{ name = "typing-extensions" },
{ name = "uuid-utils" },
]
sdist = { url = "https://files.pythonhosted.org/packages/d3/ae/8b74458fc3850ec3d150eb9f45e857db129dafa801fb5cf173dfc9f8bbf3/langchain_core-1.3.3.tar.gz", hash = "sha256:fa510a5db8efdc0c6ff41c0939fb5c00a0183c11f6b84233e892e3227ff69182", size = 915041, upload-time = "2026-05-05T19:02:36.612Z" }
sdist = { url = "https://files.pythonhosted.org/packages/f8/a4/317a1a3ac1df33a64adb3670bf88bbe3b3d5baa274db6863a979db472897/langchain_core-1.2.28.tar.gz", hash = "sha256:271a3d8bd618f795fdeba112b0753980457fc90537c46a0c11998516a74dc2cb", size = 846119, upload-time = "2026-04-08T18:19:34.867Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/1f/01/4771b7ab2af1d1aba5b710bd8f13d9225c609425214b357590a17b01be77/langchain_core-1.3.3-py3-none-any.whl", hash = "sha256:18aae8506f37da7f74398492279a7d6efcee4f8e23c4c41c7af080eeb7ef7bd1", size = 543857, upload-time = "2026-05-05T19:02:34.52Z" },
]
[[package]]
name = "langchain-protocol"
version = "0.0.15"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions" },
]
sdist = { url = "https://files.pythonhosted.org/packages/4f/24/9777489d6fbbee64af0c8f96d4f840239c408cf694f3394672807dafc490/langchain_protocol-0.0.15.tar.gz", hash = "sha256:9ab2d11ee73944754f10e037e717098d3a6796f0e58afa9cadda6154e7655ade", size = 5862, upload-time = "2026-05-01T22:30:04.748Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/1d/7a/9c97a7b9cbe4c5dc6a44cdb1545450c28f0c8ce89b9c1f0ee7fbad896263/langchain_protocol-0.0.15-py3-none-any.whl", hash = "sha256:461eb794358f83d5e42635a5797799ffec7b4702314e34edf73ac21e75d3ef79", size = 6982, upload-time = "2026-05-01T22:30:03.877Z" },
{ url = "https://files.pythonhosted.org/packages/a8/92/32f785f077c7e898da97064f113c73fbd9ad55d1e2169cf3a391b183dedb/langchain_core-1.2.28-py3-none-any.whl", hash = "sha256:80764232581eaf8057bcefa71dbf8adc1f6a28d257ebd8b95ba9b8b452e8c6ac", size = 508727, upload-time = "2026-04-08T18:19:32.823Z" },
]
[[package]]
@@ -627,7 +614,7 @@ dependencies = [
[package.metadata]
requires-dist = [
{ name = "langchain-core", specifier = ">=1.3.3" },
{ name = "langchain-core", specifier = ">=0.2.14" },
{ name = "langgraph", specifier = ">=0.6.0,<2" },
]
+25 -2
View File
@@ -8,6 +8,7 @@ from langchain_core.runnables import Runnable, RunnableConfig
from langgraph.store.base import BaseStore
from langgraph._internal._typing import EMPTY_SEQ
from langgraph.errors import NodeError
from langgraph.runtime import Runtime
from langgraph.types import CachePolicy, RetryPolicy, StreamWriter, TimeoutPolicy
from langgraph.typing import ContextT, NodeInputT, NodeInputT_contra
@@ -64,6 +65,22 @@ class _NodeWithRuntime(Protocol[NodeInputT_contra, ContextT]):
) -> Any: ...
class _NodeWithNodeError(Protocol[NodeInputT_contra]):
def __call__(self, state: NodeInputT_contra, *, error: NodeError) -> Any: ...
class _NodeWithConfigNodeError(Protocol[NodeInputT_contra]):
def __call__(
self, state: NodeInputT_contra, *, config: RunnableConfig, error: NodeError
) -> Any: ...
class _NodeWithRuntimeNodeError(Protocol[NodeInputT_contra, ContextT]):
def __call__(
self, state: NodeInputT_contra, *, runtime: Runtime[ContextT], error: NodeError
) -> Any: ...
# TODO: we probably don't want to explicitly support the config / store signatures once
# we move to adding a context arg. Maybe what we do is we add support for kwargs with param spec
# this is purely for typing purposes though, so can easily change in the coming weeks.
@@ -80,6 +97,13 @@ StateNode: TypeAlias = (
| Runnable[NodeInputT, Any]
)
ErrorHandlerNode: TypeAlias = (
StateNode[NodeInputT, ContextT]
| _NodeWithNodeError[NodeInputT]
| _NodeWithConfigNodeError[NodeInputT]
| _NodeWithRuntimeNodeError[NodeInputT, ContextT]
)
@dataclass(slots=True)
class StateNodeSpec(Generic[NodeInputT, ContextT]):
@@ -88,8 +112,7 @@ class StateNodeSpec(Generic[NodeInputT, ContextT]):
input_schema: type[NodeInputT]
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None
cache_policy: CachePolicy | None
is_error_handler: bool = False
error_handler_node: str | None = None
error_handler: Runnable[Any, Any] | None = None
ends: tuple[str, ...] | dict[str, str] | None = EMPTY_SEQ
defer: bool = False
timeout: TimeoutPolicy | None = None
+35 -161
View File
@@ -6,7 +6,7 @@ import typing
import warnings
from collections import defaultdict
from collections.abc import Awaitable, Callable, Hashable, Sequence
from dataclasses import dataclass, is_dataclass
from dataclasses import is_dataclass
from datetime import timedelta
from functools import partial
from inspect import isclass, isfunction, ismethod, signature
@@ -65,7 +65,7 @@ from langgraph.errors import (
create_error_message,
)
from langgraph.graph._branch import BranchSpec
from langgraph.graph._node import StateNode, StateNodeSpec
from langgraph.graph._node import ErrorHandlerNode, StateNode, StateNodeSpec
from langgraph.managed.base import (
ManagedValueSpec,
is_managed_value,
@@ -95,17 +95,6 @@ __all__ = ("StateGraph", "CompiledStateGraph")
logger = logging.getLogger(__name__)
_CHANNEL_BRANCH_TO = "branch:to:{}"
_DEFAULT_ERROR_HANDLER_NODE = "__default_error_handler__"
@dataclass(slots=True)
class _NodeDefaults:
"""Default node policies applied to every node at compile time."""
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None
cache_policy: CachePolicy | None = None
error_handler: StateNode[Any, Any] | None = None
timeout: TimeoutPolicy | None = None
def _warn_invalid_state_schema(schema: type[Any] | Any) -> None:
@@ -262,77 +251,10 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
self.output_schema = cast(type[OutputT], output_schema or state_schema)
self.context_schema = context_schema
self._node_defaults: _NodeDefaults = _NodeDefaults()
self._add_schema(self.state_schema)
self._add_schema(self.input_schema, allow_managed=False)
self._add_schema(self.output_schema, allow_managed=False)
def set_node_defaults(
self,
*,
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
error_handler: StateNode[Any, ContextT] | None = None,
timeout: float | timedelta | TimeoutPolicy | None = None,
) -> Self:
"""Set default node policies that apply to every node in this graph.
Per-node values passed to `add_node` always take precedence over these
defaults. Defaults are applied at `compile()` time. Policies set here
are **not** inherited by subgraphs.
`retry_policy` and `timeout` defaults apply to **all** nodes,
including error-handler nodes. `cache_policy` and `error_handler`
defaults only apply to regular nodes -- caching error-handler results
is unsafe, and handlers must never catch themselves.
Args:
retry_policy: Default retry policy for nodes that don't specify
their own via `add_node(..., retry_policy=...)`. Also applies
to error-handler nodes.
cache_policy: Default cache policy for nodes that don't specify
their own via `add_node(..., cache_policy=...)`. Does **not**
apply to error-handler nodes.
error_handler: Default error handler invoked when any regular node
raises and does not have its own `error_handler` set via
`add_node`. The handler is **not** invoked when an
error-handler node itself raises -- handler failures fail the
run.
timeout: Default timeout policy for nodes that don't specify their
own via `add_node(..., timeout=...)`. Also applies to
error-handler nodes. Accepts a `TimeoutPolicy`, a number of
seconds (`float`), or a `timedelta`.
Returns:
Self: The builder instance, for chaining.
Example:
```python
graph = (
StateGraph(State)
.set_node_defaults(
retry_policy=RetryPolicy(max_attempts=3),
error_handler=my_fallback_handler,
)
.add_node("a", node_a)
.add_node("b", node_b, retry_policy=custom_retry) # overrides default
.add_edge(START, "a")
.compile()
)
```
"""
defaults = self._node_defaults
if retry_policy is not None:
defaults.retry_policy = retry_policy
if cache_policy is not None:
defaults.cache_policy = cache_policy
if error_handler is not None:
defaults.error_handler = error_handler
if timeout is not None:
defaults.timeout = coerce_timeout_policy(timeout)
return self
@property
def _all_edges(self) -> set[tuple[str, str]]:
return self.edges | {
@@ -850,24 +772,15 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
if destinations is not None:
ends = destinations
resolved_input_schema: type[Any] = (
input_schema or inferred_input_schema or self.state_schema
)
handler_node_name: str | None = None
if error_handler is not None:
handler_node_name = f"__error_handler__{node}"
if handler_node_name in self.nodes:
raise ValueError(
f"Auto-generated error handler node `{handler_node_name}` already exists."
)
self.nodes[handler_node_name] = StateNodeSpec[Any, ContextT](
coerce_to_runnable(error_handler, name=handler_node_name, trace=False), # type: ignore[arg-type]
metadata=None,
input_schema=resolved_input_schema,
retry_policy=None,
cache_policy=None,
is_error_handler=True,
coerced_error_handler: Runnable[Any, Any] | None = (
coerce_to_runnable( # type: ignore[arg-type]
error_handler,
name=f"__error_handler__{node}",
trace=False,
)
if error_handler is not None
else None
)
if input_schema is not None:
self.nodes[node] = StateNodeSpec[NodeInputT, ContextT](
@@ -876,7 +789,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema=input_schema,
retry_policy=retry_policy,
cache_policy=cache_policy,
error_handler_node=handler_node_name,
error_handler=coerced_error_handler,
ends=ends,
defer=defer,
timeout=timeout,
@@ -888,7 +801,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema=inferred_input_schema,
retry_policy=retry_policy,
cache_policy=cache_policy,
error_handler_node=handler_node_name,
error_handler=coerced_error_handler,
ends=ends,
defer=defer,
timeout=timeout,
@@ -900,7 +813,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
input_schema=self.state_schema,
retry_policy=retry_policy,
cache_policy=cache_policy,
error_handler_node=handler_node_name,
error_handler=coerced_error_handler,
ends=ends,
defer=defer,
timeout=timeout,
@@ -1157,7 +1070,17 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
if interrupt:
for node in interrupt:
if node not in self.nodes:
raise ValueError(f"Interrupt node `{node}` not found")
# __error_handler__<name> is a valid virtual task name when the
# base node has an error_handler configured.
if node.startswith("__error_handler__"):
base = node[len("__error_handler__"):]
if (
base not in self.nodes
or self.nodes[base].error_handler is None
):
raise ValueError(f"Interrupt node `{node}` not found")
else:
raise ValueError(f"Interrupt node `{node}` not found")
self.compiled = True
return self
@@ -1172,6 +1095,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
debug: bool = False,
name: str | None = None,
transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
error_handler: ErrorHandlerNode[Any, ContextT] | None = None,
) -> CompiledStateGraph[StateT, ContextT, InputT, OutputT]:
"""Compiles the `StateGraph` into a `CompiledStateGraph` object.
@@ -1271,64 +1195,15 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
key for key, val in self.channels.items() if not is_managed_value(val)
]
)
# Apply builder defaults to node specs. Per-node values always win.
# Error-handler routing and cache_policy are only assigned to regular
# nodes. Retry and timeout defaults also apply to error-handler nodes.
defaults = self._node_defaults
default_handler_name: str | None = None
if defaults.error_handler is not None:
if _DEFAULT_ERROR_HANDLER_NODE in self.nodes:
raise ValueError(
f"Auto-generated default error handler node "
f"`{_DEFAULT_ERROR_HANDLER_NODE}` already exists."
)
default_handler_name = _DEFAULT_ERROR_HANDLER_NODE
self.nodes[default_handler_name] = StateNodeSpec[Any, ContextT](
coerce_to_runnable(
defaults.error_handler, # type: ignore[arg-type]
name=default_handler_name,
trace=False,
),
metadata=None,
input_schema=self.state_schema,
retry_policy=None,
cache_policy=None,
is_error_handler=True,
error_handler: Runnable[Any, Any] | None = (
coerce_to_runnable( # type: ignore[arg-type]
error_handler,
name="__graph_error_handler__",
trace=False,
)
# Apply builder defaults to node specs. Per-node values always win.
for spec in self.nodes.values():
# error_handler: regular nodes only — handlers must never
# catch themselves or other handlers.
if (
not spec.is_error_handler
and default_handler_name is not None
and spec.error_handler_node is None
):
spec.error_handler_node = default_handler_name
# retry: all nodes — handlers should be retried on transient
# failures just like regular nodes.
if defaults.retry_policy is not None and spec.retry_policy is None:
spec.retry_policy = defaults.retry_policy
# cache: regular nodes only — caching an error-handler result
# is unsafe because the input (failed-node state) may differ
# across failures even when the cache key matches.
if (
not spec.is_error_handler
and defaults.cache_policy is not None
and spec.cache_policy is None
):
spec.cache_policy = defaults.cache_policy
# timeout: all nodes — a stuck handler should be cancelled the
# same way a stuck regular node would be.
if defaults.timeout is not None and spec.timeout is None:
spec.timeout = defaults.timeout
node_error_handler_map = {
node_name: spec.error_handler_node
for node_name, spec in self.nodes.items()
if not spec.is_error_handler and spec.error_handler_node is not None
}
if error_handler is not None
else None
)
compiled = CompiledStateGraph[StateT, ContextT, InputT, OutputT](
builder=self,
@@ -1351,7 +1226,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
debug=debug,
store=store,
cache=cache,
node_error_handler_map=node_error_handler_map,
error_handler=error_handler,
name=name or "LangGraph",
stream_transformers=transformers,
)
@@ -1526,8 +1401,7 @@ class CompiledStateGraph(
metadata=node.metadata,
retry_policy=node.retry_policy,
cache_policy=node.cache_policy,
is_error_handler=node.is_error_handler,
error_handler_node=node.error_handler_node,
error_handler=node.error_handler,
bound=node.runnable, # type: ignore[arg-type]
timeout=node.timeout,
)
+29 -39
View File
@@ -20,6 +20,7 @@ from typing import (
from langchain_core.callbacks import Callbacks
from langchain_core.callbacks.manager import AsyncParentRunManager, ParentRunManager
from langchain_core.runnables import Runnable
from langchain_core.runnables.config import RunnableConfig
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
@@ -32,6 +33,7 @@ from langgraph.store.base import BaseStore
from xxhash import xxh3_128_hexdigest
from langgraph._internal._config import merge_configs, patch_config
from langgraph._internal._runnable import RunnableSeq
from langgraph._internal._constants import (
CACHE_NS_WRITES,
CONF,
@@ -71,6 +73,7 @@ from langgraph.constants import TAG_HIDDEN
from langgraph.errors import NodeError
from langgraph.managed.base import ManagedValueMapping
from langgraph.pregel._call import get_runnable_for_task, identifier
from langgraph.pregel._write import ChannelWrite, ChannelWriteEntry
from langgraph.pregel._io import read_channels
from langgraph.pregel._log import logger
from langgraph.pregel._read import INPUT_CACHE_KEY_TYPE, PregelNode
@@ -407,6 +410,7 @@ def prepare_next_tasks(
updated_channels: set[str] | None = None,
retry_policy: Sequence[RetryPolicy] = (),
cache_policy: CachePolicy | None = None,
error_handler: Runnable[Any, Any] | None = None,
) -> dict[str, PregelTask] | dict[str, PregelExecutableTask]:
"""Prepare the set of tasks that will make up the next Pregel step.
@@ -462,6 +466,7 @@ def prepare_next_tasks(
input_cache=input_cache,
cache_policy=cache_policy,
retry_policy=retry_policy,
error_handler=error_handler,
):
tasks.append(task)
@@ -508,6 +513,7 @@ def prepare_next_tasks(
input_cache=input_cache,
cache_policy=cache_policy,
retry_policy=retry_policy,
error_handler=error_handler,
):
tasks.append(task)
return {t.id: t for t in tasks}
@@ -542,6 +548,7 @@ def prepare_single_task(
input_cache: dict[INPUT_CACHE_KEY_TYPE, Any] | None = None,
cache_policy: CachePolicy | None = None,
retry_policy: Sequence[RetryPolicy] = (),
error_handler: Runnable[Any, Any] | None = None,
) -> None | PregelTask | PregelExecutableTask:
"""Prepares a single task for the next Pregel step, given a task path, which
uniquely identifies a PUSH or PULL task within the graph."""
@@ -756,6 +763,7 @@ def prepare_single_task(
writers=proc.flat_writers,
subgraphs=proc.subgraphs,
timeout=proc.timeout,
error_handler=proc.error_handler or error_handler,
)
else:
return PregelTask(task_id, name, task_path[:3])
@@ -1110,11 +1118,10 @@ def prepare_push_task_send(
def prepare_node_error_handler_task(
failed_task: PregelExecutableTask,
*,
handler_node_name: str,
handler: Runnable,
failed_error: BaseException,
checkpoint: Checkpoint,
pending_writes: list[PendingWrite],
processes: Mapping[str, PregelNode],
channels: Mapping[str, BaseChannel],
managed: ManagedValueMapping,
config: RunnableConfig,
@@ -1123,17 +1130,14 @@ def prepare_node_error_handler_task(
store: BaseStore | None = None,
checkpointer: BaseCheckpointSaver | None = None,
manager: None | ParentRunManager | AsyncParentRunManager = None,
cache_policy: CachePolicy | None = None,
retry_policy: Sequence[RetryPolicy] = (),
) -> PregelExecutableTask | None:
"""Prepare an immediate node-level error handler task for a failed task."""
if handler_node_name not in processes:
return None
proc = processes[handler_node_name]
proc_node = proc.node
if proc_node is None:
return None
) -> PregelExecutableTask:
"""Prepare an error handler task for a failed task.
The handler borrows the failed task's write pipeline (same state channels),
so no separate node registration is needed.
"""
handler_node_name = f"__error_handler__{failed_task.name}"
checkpoint_id_bytes = binascii.unhexlify(checkpoint["id"].replace("-", ""))
task_id_func = _xxhash_str if checkpoint["v"] > 1 else _uuid5_str
configurable = config.get(CONF, {})
@@ -1159,27 +1163,17 @@ def prepare_node_error_handler_task(
"langgraph_path": translated_task_path,
"langgraph_checkpoint_ns": task_checkpoint_ns,
}
if proc.metadata:
metadata.update(proc.metadata)
writes: deque[tuple[str, Any]] = deque()
effective_retry_policy = proc.retry_policy or retry_policy
effective_cache_policy = proc.cache_policy or cache_policy
if effective_cache_policy:
args_key = effective_cache_policy.key_func(failed_task.input)
cache_key = CacheKey(
(
CACHE_NS_WRITES,
(identifier(proc) or "__dynamic__"),
handler_node_name,
),
xxh3_128_hexdigest(
args_key.encode() if isinstance(args_key, str) else args_key
),
effective_cache_policy.ttl,
)
# Mirror how regular node procs are built: combine handler with a write pipeline
# so run_with_retry invokes the full pipeline in one shot.
# - PULL node tasks: writers are in failed_task.writers → reuse them
# - PUSH functional tasks: writers are embedded in proc (empty writers list) →
# add a RETURN write so the handler's result becomes the future's value.
handler_writers = failed_task.writers
if handler_writers:
handler_proc: Runnable = RunnableSeq(handler, *handler_writers)
else:
cache_key = None
handler_proc = RunnableSeq(handler, ChannelWrite([ChannelWriteEntry(RETURN)]))
scratchpad = _scratchpad(
config[CONF].get(CONFIG_KEY_SCRATCHPAD),
@@ -1194,14 +1188,11 @@ def prepare_node_error_handler_task(
runtime = runtime.override(
store=store, previous=checkpoint["channel_values"].get(PREVIOUS, None)
)
additional_config: RunnableConfig = {
"metadata": metadata,
"tags": proc.tags,
}
additional_config: RunnableConfig = {"metadata": metadata}
return PregelExecutableTask(
handler_node_name,
failed_task.input,
proc_node,
handler_proc,
writes,
patch_config(
merge_configs(config, additional_config),
@@ -1239,12 +1230,11 @@ def prepare_node_error_handler_task(
},
),
PUSH_TRIGGER,
effective_retry_policy,
cache_key,
retry_policy,
None, # handlers don't cache
task_id,
translated_task_path,
writers=proc.flat_writers,
subgraphs=proc.subgraphs,
writers=handler_writers, # for ParentCommand / subgraph routing
)
+29 -119
View File
@@ -22,7 +22,8 @@ from typing import (
)
from langchain_core.callbacks import AsyncParentRunManager, ParentRunManager
from langchain_core.runnables import RunnableConfig
from langchain_core.runnables import Runnable, RunnableConfig
from langgraph.cache.base import BaseCache
from langgraph.checkpoint.base import (
WRITES_IDX_MAP,
@@ -203,12 +204,6 @@ 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
@@ -285,6 +280,7 @@ class PregelLoop:
migrate_checkpoint: Callable[[Checkpoint], None] | None = None,
retry_policy: Sequence[RetryPolicy] = (),
cache_policy: CachePolicy | None = None,
error_handler: Runnable[Any, Any] | None = None,
has_graph_lifecycle_callbacks: bool = False,
) -> None:
self.stream = stream
@@ -309,6 +305,7 @@ class PregelLoop:
self.trigger_to_nodes = trigger_to_nodes
self.retry_policy = retry_policy
self.cache_policy = cache_policy
self.error_handler = error_handler
self.durability = durability
self._has_graph_lifecycle_callbacks = has_graph_lifecycle_callbacks
self._graph_lifecycle_events = deque()
@@ -480,13 +477,6 @@ 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)
@@ -558,6 +548,7 @@ class PregelLoop:
manager=self.manager,
retry_policy=self.retry_policy,
cache_policy=self.cache_policy,
error_handler=self.error_handler,
),
):
# produce debug output
@@ -566,7 +557,7 @@ class PregelLoop:
self.tasks[pushed.id] = pushed
# match any pending writes to the new task
if not self.is_replaying:
self._reapply_writes_to_succeeded_nodes({pushed.id: pushed})
self._match_writes({pushed.id: pushed})
# return the new task, to be started if not run before
return pushed
@@ -610,6 +601,7 @@ class PregelLoop:
updated_channels=self.updated_channels,
retry_policy=self.retry_policy,
cache_policy=self.cache_policy,
error_handler=self.error_handler,
)
# produce debug output
@@ -644,8 +636,7 @@ class PregelLoop:
# if there are pending writes from a previous loop, apply them
if not self.is_replaying and self.checkpoint_pending_writes:
self._reapply_writes_to_succeeded_nodes(self.tasks)
self._resume_error_handlers_if_applicable()
self._match_writes(self.tasks)
# before execution, check if we should interrupt
if self.interrupt_before and should_interrupt(
@@ -712,88 +703,13 @@ class PregelLoop:
# private
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.
"""
def _match_writes(self, tasks: Mapping[str, PregelExecutableTask]) -> None:
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
@@ -1457,6 +1373,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
migrate_checkpoint: Callable[[Checkpoint], None] | None = None,
retry_policy: Sequence[RetryPolicy] = (),
cache_policy: CachePolicy | None = None,
error_handler: Runnable[Any, Any] | None = None,
has_graph_lifecycle_callbacks: bool = False,
) -> None:
super().__init__(
@@ -1478,6 +1395,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
trigger_to_nodes=trigger_to_nodes,
retry_policy=retry_policy,
cache_policy=cache_policy,
error_handler=error_handler,
durability=durability,
has_graph_lifecycle_callbacks=has_graph_lifecycle_callbacks,
)
@@ -1540,20 +1458,18 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
def schedule_error_handler(
self, failed_task: PregelExecutableTask, error: BaseException
) -> PregelExecutableTask | None:
handler_node = self.nodes[failed_task.name].error_handler_node
if not handler_node:
handler = failed_task.error_handler or self.error_handler
if handler is None:
return None
# 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)
writes = list(failed_task.writes)
writes.append((ERROR_SOURCE_NODE, failed_task.name))
self.put_writes(failed_task.id, writes)
handler_task = prepare_node_error_handler_task(
failed_task,
handler_node_name=handler_node,
handler=handler,
failed_error=error,
checkpoint=self.checkpoint,
pending_writes=self.checkpoint_pending_writes,
processes=self.nodes,
channels=self.channels,
managed=self.managed,
config=failed_task.config,
@@ -1563,17 +1479,16 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
checkpointer=self.checkpointer,
manager=self.manager,
retry_policy=self.retry_policy,
cache_policy=self.cache_policy,
)
if handler_task is None:
return None
self.tasks[handler_task.id] = handler_task
if not self.is_replaying:
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
self._match_writes({handler_task.id: handler_task})
for task in self.match_cached_writes():
self.output_writes(task.id, task.writes, cached=True)
return handler_task
def put_writes(self, task_id: str, writes: WritesT) -> None:
"""Put writes for a task, to be read by the next tick."""
super().put_writes(task_id, writes)
@@ -1651,7 +1566,6 @@ 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,6 +1623,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
migrate_checkpoint: Callable[[Checkpoint], None] | None = None,
retry_policy: Sequence[RetryPolicy] = (),
cache_policy: CachePolicy | None = None,
error_handler: Runnable[Any, Any] | None = None,
has_graph_lifecycle_callbacks: bool = False,
) -> None:
super().__init__(
@@ -1730,6 +1645,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
trigger_to_nodes=trigger_to_nodes,
retry_policy=retry_policy,
cache_policy=cache_policy,
error_handler=error_handler,
durability=durability,
has_graph_lifecycle_callbacks=has_graph_lifecycle_callbacks,
)
@@ -1794,20 +1710,18 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
async def aschedule_error_handler(
self, failed_task: PregelExecutableTask, error: BaseException
) -> PregelExecutableTask | None:
handler_node = self.nodes[failed_task.name].error_handler_node
if not handler_node:
handler = failed_task.error_handler or self.error_handler
if handler is None:
return None
# 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)
writes = list(failed_task.writes)
writes.append((ERROR_SOURCE_NODE, failed_task.name))
self.put_writes(failed_task.id, writes)
handler_task = prepare_node_error_handler_task(
failed_task,
handler_node_name=handler_node,
handler=handler,
failed_error=error,
checkpoint=self.checkpoint,
pending_writes=self.checkpoint_pending_writes,
processes=self.nodes,
channels=self.channels,
managed=self.managed,
config=failed_task.config,
@@ -1817,13 +1731,10 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
checkpointer=self.checkpointer,
manager=self.manager,
retry_policy=self.retry_policy,
cache_policy=self.cache_policy,
)
if handler_task is None:
return None
self.tasks[handler_task.id] = handler_task
if not self.is_replaying:
self._reapply_writes_to_succeeded_nodes({handler_task.id: handler_task})
self._match_writes({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
@@ -1908,7 +1819,6 @@ 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
)
+4 -9
View File
@@ -138,11 +138,8 @@ class PregelNode:
metadata: Mapping[str, Any] | None
"""Metadata to attach to the node for tracing."""
is_error_handler: bool
"""Whether this node is registered as an error handler node."""
error_handler_node: str | None
"""Optional handler node name for failures from this node."""
error_handler: Runnable[Any, Any] | None
"""Callable invoked after retries are exhausted; receives same input as the node."""
subgraphs: Sequence[PregelProtocol]
"""Subgraphs used by the node."""
@@ -159,8 +156,7 @@ class PregelNode:
bound: Runnable[Any, Any] | None = None,
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
is_error_handler: bool = False,
error_handler_node: str | None = None,
error_handler: Runnable[Any, Any] | None = None,
subgraphs: Sequence[PregelProtocol] | None = None,
timeout: float | timedelta | TimeoutPolicy | None = None,
) -> None:
@@ -177,8 +173,7 @@ class PregelNode:
self.timeout = coerce_timeout_policy(timeout)
self.tags = tags
self.metadata = metadata
self.is_error_handler = is_error_handler
self.error_handler_node = error_handler_node
self.error_handler = error_handler
if subgraphs is not None:
self.subgraphs = subgraphs
elif self.bound is not DEFAULT_BOUND:
+6 -19
View File
@@ -13,7 +13,6 @@ from collections.abc import (
Collection,
Iterable,
Iterator,
Mapping,
Sequence,
)
from functools import partial
@@ -31,7 +30,6 @@ from langgraph._internal._constants import (
CONFIG_KEY_CALL,
CONFIG_KEY_SCRATCHPAD,
ERROR,
ERROR_SOURCE_NODE,
INTERRUPT,
NO_WRITES,
RESUME,
@@ -144,7 +142,6 @@ class PregelRunner:
put_writes: weakref.ref[Callable[[str, Sequence[tuple[str, Any]]], None]],
use_astream: bool = False,
node_finished: Callable[[str], None] | None = None,
node_error_handler_map: Mapping[str, str] | None = None,
schedule_error_handler: Callable[
[PregelExecutableTask, BaseException], PregelExecutableTask | None
]
@@ -159,20 +156,10 @@ class PregelRunner:
self.put_writes = put_writes
self.use_astream = use_astream
self.node_finished = node_finished
self.node_error_handler_map = dict(node_error_handler_map or {})
self.error_handler_nodes = set(self.node_error_handler_map.values())
self.schedule_error_handler = schedule_error_handler
self.aschedule_error_handler = aschedule_error_handler
# Exception object ids that are already routed to graph-level error handler.
# These ids are consulted by stop/panic checks to avoid re-raising handled
# exceptions via the normal fatal path in the same run.
self._handled_exception_ids: set[int] = set()
def _should_route_to_error_handler(self, task: PregelExecutableTask) -> bool:
if task.name in self.error_handler_nodes:
return False
return task.name in self.node_error_handler_map
def tick(
self,
tasks: Iterable[PregelExecutableTask],
@@ -223,7 +210,7 @@ class PregelRunner:
self.commit(t, exc)
if (
not isinstance(exc, GraphBubbleUp)
and self._should_route_to_error_handler(t)
and t.error_handler is not None
and self.schedule_error_handler is not None
):
self._handled_exception_ids.add(id(exc))
@@ -296,7 +283,7 @@ class PregelRunner:
futures[get_waiter()] = None
elif (
(task_exc := _exception(fut))
and self._should_route_to_error_handler(task)
and task.error_handler is not None
and not isinstance(task_exc, GraphBubbleUp)
):
self._handled_exception_ids.add(id(task_exc))
@@ -415,7 +402,7 @@ class PregelRunner:
self.commit(t, exc)
if (
not isinstance(exc, GraphBubbleUp)
and self._should_route_to_error_handler(t)
and t.error_handler is not None
and self.aschedule_error_handler is not None
):
self._handled_exception_ids.add(id(exc))
@@ -495,7 +482,7 @@ class PregelRunner:
futures[get_waiter()] = None
elif (
(task_exc := _exception(fut))
and self._should_route_to_error_handler(task)
and task.error_handler is not None
and not isinstance(task_exc, GraphBubbleUp)
):
self._handled_exception_ids.add(id(task_exc))
@@ -595,10 +582,10 @@ class PregelRunner:
else:
# save error to checkpointer
task.writes.append((ERROR, exception))
if self._should_route_to_error_handler(task) and not isinstance(
if task.error_handler is not None and not isinstance(
exception, GraphBubbleUp
):
task.writes.append((ERROR_SOURCE_NODE, task.name))
# Mark early in commit path; loop-side routing may happen later.
self._handled_exception_ids.add(id(exception))
self.put_writes()(task.id, task.writes) # type: ignore[misc]
else:
+9 -2
View File
@@ -99,14 +99,21 @@ def validate_graph(
if interrupt_after_nodes != "*":
for n in interrupt_after_nodes:
if n not in nodes:
if n not in nodes and not _is_valid_error_handler_interrupt(n, nodes):
raise ValueError(f"Node {n} not in nodes")
if interrupt_before_nodes != "*":
for n in interrupt_before_nodes:
if n not in nodes:
if n not in nodes and not _is_valid_error_handler_interrupt(n, nodes):
raise ValueError(f"Node {n} not in nodes")
def _is_valid_error_handler_interrupt(name: str, nodes: Mapping[str, PregelNode]) -> bool:
if not name.startswith("__error_handler__"):
return False
base = name[len("__error_handler__"):]
return base in nodes and nodes[base].error_handler is not None
def validate_keys(
keys: str | Sequence[str] | None,
channels: Mapping[str, Any],
+6 -5
View File
@@ -33,6 +33,7 @@ from uuid import UUID, uuid5
from langchain_core._api import beta
from langchain_core.globals import get_debug
from langchain_core.runnables import (
Runnable,
RunnableSequence,
)
from langchain_core.runnables.base import Input, Output
@@ -751,7 +752,7 @@ class Pregel(
name: str = "LangGraph"
trigger_to_nodes: Mapping[str, Sequence[str]]
node_error_handler_map: Mapping[str, str]
error_handler: Runnable[Any, Any] | None
def __init__(
self,
@@ -776,7 +777,7 @@ class Pregel(
context_schema: type[ContextT] | None = None,
config: RunnableConfig | None = None,
trigger_to_nodes: Mapping[str, Sequence[str]] | None = None,
node_error_handler_map: Mapping[str, str] | None = None,
error_handler: Runnable[Any, Any] | None = None,
name: str = "LangGraph",
stream_transformers: Sequence[Callable[[tuple[str, ...]], Any]] | None = None,
**deprecated_kwargs: Unpack[DeprecatedKwargs],
@@ -824,7 +825,7 @@ class Pregel(
self.context_schema = context_schema
self.config = config
self.trigger_to_nodes = trigger_to_nodes or {}
self.node_error_handler_map = node_error_handler_map or {}
self.error_handler = error_handler
self.name = name
self.stream_transformers: tuple[Callable[[tuple[str, ...]], Any], ...] = tuple(
stream_transformers or ()
@@ -2885,6 +2886,7 @@ class Pregel(
migrate_checkpoint=self._migrate_checkpoint,
retry_policy=self.retry_policy,
cache_policy=self.cache_policy,
error_handler=self.error_handler,
has_graph_lifecycle_callbacks=bool(graph_callback_manager.handlers),
) as loop:
emit_graph_lifecycle_events(loop)
@@ -2895,7 +2897,6 @@ class Pregel(
),
put_writes=weakref.WeakMethod(loop.put_writes),
node_finished=config[CONF].get(CONFIG_KEY_NODE_FINISHED),
node_error_handler_map=self.node_error_handler_map,
schedule_error_handler=loop.schedule_error_handler,
)
# enable subgraph streaming
@@ -3337,6 +3338,7 @@ class Pregel(
migrate_checkpoint=self._migrate_checkpoint,
retry_policy=self.retry_policy,
cache_policy=self.cache_policy,
error_handler=self.error_handler,
has_graph_lifecycle_callbacks=bool(graph_callback_manager.handlers),
) as loop:
await aemit_graph_lifecycle_events(loop)
@@ -3348,7 +3350,6 @@ class Pregel(
put_writes=weakref.WeakMethod(loop.put_writes),
use_astream=do_stream,
node_finished=config[CONF].get(CONFIG_KEY_NODE_FINISHED),
node_error_handler_map=self.node_error_handler_map,
aschedule_error_handler=loop.aschedule_error_handler,
)
# enable subgraph streaming
+1
View File
@@ -628,6 +628,7 @@ class PregelExecutableTask:
writers: Sequence[Runnable] = ()
subgraphs: Sequence[PregelProtocol] = ()
timeout: TimeoutPolicy | None = None
error_handler: Runnable | None = None
class StateSnapshot(NamedTuple):
+3 -3
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph"
version = "1.2.0"
version = "1.2.0a7"
description = "Building stateful, multi-actor applications with LLMs"
authors = []
requires-python = ">=3.10"
@@ -25,9 +25,9 @@ classifiers = [
]
dependencies = [
"langchain-core>=1.4.0,<2",
"langgraph-checkpoint>=4.1.0,<5.0.0",
"langgraph-checkpoint>=4.1.0a4,<5.0.0",
"langgraph-sdk>=0.3.0,<0.4.0",
"langgraph-prebuilt>=1.1.0,<1.2.0",
"langgraph-prebuilt>=1.1.0a2,<1.2.0",
"xxhash>=3.5.0",
"pydantic>=2.7.4",
]
+87 -473
View File
@@ -15,7 +15,7 @@ from langchain_core.callbacks import AsyncCallbackManagerForLLMRun, BaseCallback
from langchain_core.language_models.fake_chat_models import GenericFakeChatModel
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, HumanMessage
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult
from langchain_core.runnables import RunnableConfig, RunnableLambda, RunnableParallel
from langchain_core.runnables import RunnableLambda, RunnableParallel
from langgraph.checkpoint.memory import InMemorySaver, MemorySaver
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
from typing_extensions import TypedDict
@@ -2283,522 +2283,136 @@ def test_node_without_error_handler_still_fails_run():
# ---------------------------------------------------------------------------
# set_node_defaults()
# Structural invariants from the policy-style refactor
# ---------------------------------------------------------------------------
def test_set_node_defaults_error_handler_catches_all_nodes():
class State(TypedDict):
route: str
foo: Annotated[list[str], operator.add]
def route_node(state: State) -> Command:
return Command(goto=state["route"])
def fail_a(state: State) -> State:
raise RuntimeError("a failed")
def fail_b(state: State) -> State:
raise RuntimeError("b failed")
captured: dict[str, list[str]] = {"nodes": []}
def default_handler(state: State, error: NodeError) -> State:
captured["nodes"].append(error.node)
return {"foo": [f"handled_{error.node}"]}
graph = (
StateGraph(State)
.set_node_defaults(error_handler=default_handler)
.add_node("route_node", route_node)
.add_node("fail_a", fail_a)
.add_node("fail_b", fail_b)
.add_edge(START, "route_node")
.add_conditional_edges(
"route_node", lambda s: s["route"], path_map=["fail_a", "fail_b"]
)
.compile()
)
result_a = graph.invoke({"route": "fail_a", "foo": []})
result_b = graph.invoke({"route": "fail_b", "foo": []})
assert result_a["foo"] == ["handled_fail_a"]
assert result_b["foo"] == ["handled_fail_b"]
assert "fail_a" in captured["nodes"]
assert "fail_b" in captured["nodes"]
def test_set_node_defaults_error_handler_overridden_by_node_handler():
class State(TypedDict):
route: str
foo: Annotated[list[str], operator.add]
def route_node(state: State) -> Command:
return Command(goto=state["route"])
def fail_a(state: State) -> State:
raise RuntimeError("a failed")
def fail_b(state: State) -> State:
raise RuntimeError("b failed")
captured: dict[str, list[str]] = {"handler": []}
def node_handler(state: State, error: NodeError) -> State:
captured["handler"].append(f"node:{error.node}")
return {"foo": [f"node_handled_{error.node}"]}
def default_handler(state: State, error: NodeError) -> State:
captured["handler"].append(f"default:{error.node}")
return {"foo": [f"default_handled_{error.node}"]}
graph = (
StateGraph(State)
.set_node_defaults(error_handler=default_handler)
.add_node("route_node", route_node)
.add_node("fail_a", fail_a, error_handler=node_handler)
.add_node("fail_b", fail_b)
.add_edge(START, "route_node")
.add_conditional_edges(
"route_node", lambda s: s["route"], path_map=["fail_a", "fail_b"]
)
.compile()
)
result_a = graph.invoke({"route": "fail_a", "foo": []})
assert result_a["foo"] == ["node_handled_fail_a"]
assert "node:fail_a" in captured["handler"]
assert "default:fail_a" not in captured["handler"]
result_b = graph.invoke({"route": "fail_b", "foo": []})
assert result_b["foo"] == ["default_handled_fail_b"]
assert "default:fail_b" in captured["handler"]
def test_set_node_defaults_error_handler_skips_per_node_handler_nodes():
"""If a per-node error handler itself raises, the default handler must NOT
catch it -- the run should fail."""
def test_error_handler_not_registered_as_node():
"""After compile, no hidden __error_handler__* nodes should exist in the graph."""
class State(TypedDict):
foo: str
def always_failing(state: State) -> State:
raise RuntimeError("node boom")
def broken_handler(state: State, error: NodeError) -> State:
raise RuntimeError("handler boom")
def default_handler(state: State, error: NodeError) -> State:
return {"foo": "default recovered"}
graph = (
StateGraph(State)
.set_node_defaults(error_handler=default_handler)
.add_node("always_failing", always_failing, error_handler=broken_handler)
.add_edge(START, "always_failing")
.compile()
)
with pytest.raises(RuntimeError, match="handler boom"):
graph.invoke({"foo": ""})
def test_set_node_defaults_error_handler_failure_fails_run():
"""When the default handler itself raises, the run fails (no infinite
recursion, no double-routing)."""
class State(TypedDict):
foo: str
def always_failing(state: State) -> State:
raise RuntimeError("node boom")
def broken_default_handler(state: State, error: NodeError) -> State:
raise RuntimeError("default handler boom")
graph = (
StateGraph(State)
.set_node_defaults(error_handler=broken_default_handler)
.add_node("always_failing", always_failing)
.add_edge(START, "always_failing")
.compile()
)
with pytest.raises(RuntimeError, match="default handler boom"):
graph.invoke({"foo": ""})
def test_set_node_defaults_error_handler_receives_runnable_config():
class State(TypedDict):
foo: str
def always_failing(state: State) -> State:
raise RuntimeError("boom")
captured: dict[str, Any] = {}
def default_handler(
state: State, error: NodeError, config: RunnableConfig
) -> State:
captured["thread_id"] = config["configurable"].get("thread_id")
return {"foo": "handled"}
checkpointer = MemorySaver()
graph = (
StateGraph(State)
.set_node_defaults(error_handler=default_handler)
.add_node("always_failing", always_failing)
.add_edge(START, "always_failing")
.compile(checkpointer=checkpointer)
)
thread_id = str(uuid4())
result = graph.invoke(
{"foo": ""}, config={"configurable": {"thread_id": thread_id}}
)
assert result["foo"] == "handled"
assert captured["thread_id"] == thread_id
def test_set_node_defaults_error_handler_collides_with_user_node():
class State(TypedDict):
foo: str
def default_handler(state: State, error: NodeError) -> State:
return {"foo": "handled"}
builder = (
StateGraph(State)
.set_node_defaults(error_handler=default_handler)
.add_node("__default_error_handler__", lambda s: s)
.add_edge(START, "__default_error_handler__")
)
with pytest.raises(ValueError, match="__default_error_handler__"):
builder.compile()
def test_set_node_defaults_retry_policy():
class State(TypedDict):
foo: str
attempts = 0
def flaky_node(state: State) -> State:
nonlocal attempts
attempts += 1
if attempts < 3:
raise ValueError("not yet")
return {"foo": "ok"}
graph = (
StateGraph(State)
.set_node_defaults(
retry_policy=RetryPolicy(
max_attempts=3, initial_interval=0.01, jitter=False, retry_on=ValueError
)
)
.add_node("flaky", flaky_node)
.add_edge(START, "flaky")
.compile()
)
with patch("time.sleep"):
result = graph.invoke({"foo": ""})
assert result["foo"] == "ok"
assert attempts == 3
def test_set_node_defaults_retry_policy_per_node_wins():
class State(TypedDict):
foo: str
attempts = 0
def flaky_node(state: State) -> State:
nonlocal attempts
attempts += 1
if attempts < 2:
raise ValueError("not yet")
return {"foo": "ok"}
graph = (
StateGraph(State)
.set_node_defaults(
retry_policy=RetryPolicy(
max_attempts=1, initial_interval=0.01, jitter=False, retry_on=ValueError
)
)
.add_node(
"flaky",
flaky_node,
retry_policy=RetryPolicy(
max_attempts=3,
initial_interval=0.01,
jitter=False,
retry_on=ValueError,
),
)
.add_edge(START, "flaky")
.compile()
)
with patch("time.sleep"):
result = graph.invoke({"foo": ""})
assert result["foo"] == "ok"
assert attempts == 2
@pytest.mark.anyio
async def test_set_node_defaults_timeout():
class State(TypedDict):
foo: str
async def slow_node(state: State) -> State:
await asyncio.sleep(10)
return {"foo": "should-not-happen"}
graph = (
StateGraph(State)
.set_node_defaults(timeout=TimeoutPolicy(run_timeout=0.05))
.add_node("slow", slow_node)
.add_edge(START, "slow")
.compile()
)
from langgraph.errors import NodeTimeoutError
with pytest.raises(NodeTimeoutError):
await graph.ainvoke({"foo": ""})
@pytest.mark.anyio
async def test_set_node_defaults_timeout_per_node_wins():
"""Per-node timeout overrides the default; a generous per-node timeout
allows a node to complete even when the builder default is very short."""
class State(TypedDict):
foo: str
async def quick_node(state: State) -> State:
await asyncio.sleep(0.05)
return {"foo": "done"}
graph = (
StateGraph(State)
.set_node_defaults(timeout=TimeoutPolicy(run_timeout=0.01))
.add_node("quick", quick_node, timeout=TimeoutPolicy(run_timeout=5.0))
.add_edge(START, "quick")
.compile()
)
result = await graph.ainvoke({"foo": ""})
assert result["foo"] == "done"
def test_set_node_defaults_chaining():
"""set_node_defaults() is chainable and can be called in any order relative to add_node."""
class State(TypedDict):
foo: str
def always_failing(state: State) -> State:
raise RuntimeError("boom")
def failing_node(state: State) -> State:
raise ValueError("boom")
def handler(state: State, error: NodeError) -> State:
return {"foo": "handled"}
graph = (
StateGraph(State)
.add_node("a", always_failing)
.add_edge(START, "a")
.set_node_defaults(
retry_policy=RetryPolicy(
max_attempts=1, initial_interval=0.01, jitter=False
),
error_handler=handler,
)
.add_node("failing_node", failing_node, error_handler=handler)
.add_edge(START, "failing_node")
.compile()
)
hidden = [k for k in graph.nodes if k.startswith("__error_handler__")]
assert hidden == [], f"unexpected hidden nodes: {hidden}"
def test_error_handler_stored_on_pregel_node():
"""The error_handler callable should be a Runnable field on PregelNode, not a name pointer."""
class State(TypedDict):
foo: str
def failing_node(state: State) -> State:
raise ValueError("boom")
def handler(state: State) -> State:
return {"foo": "handled"}
graph = (
StateGraph(State)
.add_node("failing_node", failing_node, error_handler=handler)
.add_edge(START, "failing_node")
.compile()
)
pregel_node = graph.nodes["failing_node"]
assert pregel_node.error_handler is not None, "error_handler should be set on PregelNode"
assert not hasattr(pregel_node, "error_handler_node"), "old string-pointer field should be gone"
assert not hasattr(pregel_node, "is_error_handler"), "is_error_handler flag should be gone"
def test_error_handler_dispatched_from_task_field():
"""error_handler on PregelExecutableTask drives dispatch — no node-map lookup needed."""
class State(TypedDict):
foo: str
def failing_node(state: State) -> State:
raise ValueError("boom")
def handler(state: State) -> State:
return {"foo": "handled"}
graph = (
StateGraph(State)
.add_node("failing_node", failing_node, error_handler=handler)
.add_edge(START, "failing_node")
.compile()
)
result = graph.invoke({"foo": ""})
assert result["foo"] == "handled"
def test_set_node_defaults_combined_retry_and_error_handler():
"""Retries are exhausted first, then the error handler runs."""
# ---------------------------------------------------------------------------
# Graph-level error handler
# ---------------------------------------------------------------------------
def test_graph_level_error_handler_used_when_no_per_node_handler():
"""compile(error_handler=fallback) should catch failures from nodes without their own handler."""
class State(TypedDict):
foo: str
attempts = 0
captured: dict[str, Any] = {}
def always_failing(state: State) -> State:
nonlocal attempts
attempts += 1
raise ValueError("Always fails")
def handler(state: State, error: NodeError) -> State:
captured["error"] = str(error.error)
return {"foo": "handled"}
graph = (
StateGraph(State)
.set_node_defaults(
retry_policy=RetryPolicy(
max_attempts=2,
initial_interval=0.01,
jitter=False,
retry_on=ValueError,
),
error_handler=handler,
)
.add_node("fail", always_failing)
.add_edge(START, "fail")
.compile()
)
with patch("time.sleep"):
result = graph.invoke({"foo": ""})
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")
raise RuntimeError("node failed")
handler_should_fail = [True]
def graph_handler(state: State, error: NodeError) -> State:
return {"foo": f"graph_handler_caught:{error.node}"}
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)
.add_node("failing_node", failing_node)
.add_edge(START, "failing_node")
.compile(error_handler=graph_handler)
)
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)
result = graph.invoke({"foo": ""})
assert result["foo"] == "graph_handler_caught:failing_node"
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)
"""
def test_per_node_handler_takes_precedence_over_graph_level():
"""When a node has its own error_handler, it should win over the graph-level fallback."""
class State(TypedDict):
results: Annotated[list[str], operator.add]
foo: str
call_count = {"a": 0, "b": 0, "handler_a": 0, "handler_b": 0}
handler_a_started = threading.Event()
def failing_node(state: State) -> State:
raise RuntimeError("node failed")
def node_a(state: State) -> State:
call_count["a"] += 1
raise RuntimeError("a failed")
def node_handler(state: State, error: NodeError) -> State:
return {"foo": "node_handler"}
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")
def graph_handler(state: State, error: NodeError) -> State:
return {"foo": "graph_handler"}
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)
.add_node("failing_node", failing_node, error_handler=node_handler)
.add_edge(START, "failing_node")
.compile(error_handler=graph_handler)
)
config = {"configurable": {"thread_id": "t1"}}
result = graph.invoke({"foo": ""})
assert result["foo"] == "node_handler"
# 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
# ---------------------------------------------------------------------------
# Functional API
# ---------------------------------------------------------------------------
# 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"]
+5 -5
View File
@@ -1382,7 +1382,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.2.0"
version = "1.2.0a7"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
@@ -1563,7 +1563,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.1.0"
version = "4.1.0a4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -1611,7 +1611,7 @@ test = [
[[package]]
name = "langgraph-checkpoint-postgres"
version = "3.1.0"
version = "3.1.0a4"
source = { editable = "../checkpoint-postgres" }
dependencies = [
{ name = "langgraph-checkpoint" },
@@ -1658,7 +1658,7 @@ test = [
[[package]]
name = "langgraph-checkpoint-sqlite"
version = "3.1.0"
version = "3.1.0a1"
source = { editable = "../checkpoint-sqlite" }
dependencies = [
{ name = "aiosqlite" },
@@ -1757,7 +1757,7 @@ test = [
[[package]]
name = "langgraph-prebuilt"
version = "1.1.0"
version = "1.1.0a2"
source = { editable = "../prebuilt" }
dependencies = [
{ name = "langchain-core" },
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-prebuilt"
version = "1.1.0"
version = "1.1.0a2"
description = "Library with high-level APIs for creating and executing LangGraph agents and tools."
authors = []
requires-python = ">=3.10"
+5 -5
View File
@@ -285,7 +285,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.2.0"
version = "1.2.0a7"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },
@@ -369,7 +369,7 @@ test = [
[[package]]
name = "langgraph-checkpoint"
version = "4.1.0"
version = "4.1.0a4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -417,7 +417,7 @@ test = [
[[package]]
name = "langgraph-checkpoint-postgres"
version = "3.1.0"
version = "3.1.0a4"
source = { editable = "../checkpoint-postgres" }
dependencies = [
{ name = "langgraph-checkpoint" },
@@ -464,7 +464,7 @@ test = [
[[package]]
name = "langgraph-checkpoint-sqlite"
version = "3.1.0"
version = "3.1.0a1"
source = { editable = "../checkpoint-sqlite" }
dependencies = [
{ name = "aiosqlite" },
@@ -507,7 +507,7 @@ test = [
[[package]]
name = "langgraph-prebuilt"
version = "1.1.0"
version = "1.1.0a2"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
+3 -3
View File
@@ -298,7 +298,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.2.0"
version = "1.2.0a7"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },
@@ -382,7 +382,7 @@ test = [
[[package]]
name = "langgraph-checkpoint"
version = "4.1.0"
version = "4.1.0a4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -430,7 +430,7 @@ test = [
[[package]]
name = "langgraph-prebuilt"
version = "1.1.0"
version = "1.1.0a2"
source = { editable = "../prebuilt" }
dependencies = [
{ name = "langchain-core" },