mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-02 14:35:18 +02:00
Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a55365f3f9 | ||
|
|
2cd7ecc81e | ||
|
|
f6746b39cc |
Generated
+1
-1
@@ -279,7 +279,7 @@ wheels = [
|
||||
|
||||
[[package]]
|
||||
name = "langgraph-checkpoint"
|
||||
version = "4.1.0"
|
||||
version = "4.1.0a4"
|
||||
source = { editable = "../checkpoint" }
|
||||
dependencies = [
|
||||
{ name = "langchain-core" },
|
||||
|
||||
@@ -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",
|
||||
|
||||
Generated
+2
-2
@@ -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" },
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
Generated
+2
-2
@@ -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" },
|
||||
|
||||
@@ -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"
|
||||
|
||||
Generated
+7
-7
@@ -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 @@
|
||||
__version__ = "0.4.26"
|
||||
__version__ = "0.4.25"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
Generated
+5
-18
@@ -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" },
|
||||
]
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
Generated
+4
-17
@@ -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" },
|
||||
]
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -628,6 +628,7 @@ class PregelExecutableTask:
|
||||
writers: Sequence[Runnable] = ()
|
||||
subgraphs: Sequence[PregelProtocol] = ()
|
||||
timeout: TimeoutPolicy | None = None
|
||||
error_handler: Runnable | None = None
|
||||
|
||||
|
||||
class StateSnapshot(NamedTuple):
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Generated
+5
-5
@@ -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" },
|
||||
|
||||
@@ -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"
|
||||
|
||||
Generated
+5
-5
@@ -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" },
|
||||
|
||||
Generated
+3
-3
@@ -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" },
|
||||
|
||||
Reference in New Issue
Block a user