Merge pull request #922 from langchain-ai/yogesh/add-error-normalization-middleware

feat: add error normalization middleware for tool and model calls [Closes #904]
This commit is contained in:
Aran Yogesh 2026-02-09 13:30:52 -08:00 • committed by GitHub
commit b8caf4c4a1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 111 additions and 2 deletions

View file

@ -0,0 +1,3 @@
from .tool_error_handler import ToolErrorMiddleware
__all__ = ["ToolErrorMiddleware"]

View file

@ -0,0 +1,104 @@
"""Tool error handling middleware.
Wraps all tool calls in try/except so that unhandled exceptions are
returned as error ToolMessages instead of crashing the agent run.
"""
from __future__ import annotations
import json
import logging
from collections.abc import Awaitable, Callable
from langchain.agents.middleware.types import (
AgentMiddleware,
AgentState,
)
from langchain_core.messages import ToolMessage
from langgraph.prebuilt.tool_node import ToolCallRequest
from langgraph.types import Command
logger = logging.getLogger(__name__)
def _get_name(candidate: object) -> str | None:
if not candidate:
return None
if isinstance(candidate, str):
return candidate
if isinstance(candidate, dict):
name = candidate.get("name")
else:
name = getattr(candidate, "name", None)
return name if isinstance(name, str) and name else None
def _extract_tool_name(request: ToolCallRequest | None) -> str | None:
if request is None:
return None
for attr in ("tool_call", "tool_name", "name"):
name = _get_name(getattr(request, attr, None))
if name:
return name
return None
def _to_error_payload(e: Exception, request: ToolCallRequest | None = None) -> dict[str, str]:
data: dict[str, str] = {
"error": str(e),
"error_type": e.__class__.__name__,
"status": "error",
}
tool_name = _extract_tool_name(request)
if tool_name:
data["name"] = tool_name
return data
def _get_tool_call_id(request: ToolCallRequest) -> str | None:
if isinstance(request.tool_call, dict):
return request.tool_call.get("id")
return None
class ToolErrorMiddleware(AgentMiddleware):
"""Normalize tool execution errors into predictable payloads.
Catches any exception thrown during a tool call and converts it into
a ToolMessage with status="error" so the LLM can see the failure and
self-correct, rather than crashing the entire agent run.
"""
state_schema = AgentState
def wrap_tool_call(
self,
request: ToolCallRequest,
handler: Callable[[ToolCallRequest], ToolMessage | Command],
) -> ToolMessage | Command:
try:
return handler(request)
except Exception as e:
logger.exception("Error during tool call handling; request=%r", request)
data = _to_error_payload(e, request)
return ToolMessage(
content=json.dumps(data),
tool_call_id=_get_tool_call_id(request),
status="error",
)
async def awrap_tool_call(
self,
request: ToolCallRequest,
handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]],
) -> ToolMessage | Command:
try:
return await handler(request)
except Exception as e:
logger.exception("Error during tool call handling; request=%r", request)
data = _to_error_payload(e, request)
return ToolMessage(
content=json.dumps(data),
tool_call_id=_get_tool_call_id(request),
status="error",
)

View file

@ -26,13 +26,14 @@ warnings.filterwarnings("ignore", message=".*Pydantic V1.*", category=UserWarnin
# Now safe to import agent (which imports LangChain modules)
from deepagents import create_deep_agent
from .protocol import SandboxBackendProtocol
# Local import for encryption
from langchain_anthropic import ChatAnthropic
from .encryption import decrypt_token
from .middleware import ToolErrorMiddleware
from .prompt import construct_system_prompt
from .protocol import SandboxBackendProtocol
from .tools import commit_and_open_pr, fetch_url, http_request
@ -959,6 +960,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: # noqa: PLR0915
tools=[http_request, fetch_url, commit_and_open_pr],
backend=sandbox_backend,
middleware=[
ToolErrorMiddleware(),
check_message_queue_before_model,
post_to_linear_after_model,
open_pr_if_needed,

View file

@ -71,7 +71,7 @@ def get_service_jwt_token_for_user(
LINEAR_TEAM_TO_REPO: dict[str, dict[str, str]] = {
"Brace's test workspace": {"owner": "langchain-ai", "name": "open-swe"},
"Yogesh-dev": {"owner": "aran-yogesh", "name": "nimedge"},
"Yogesh-dev": {"owner": "aran-yogesh", "name": "TalkBack"},
}