Refactor middleware into package and update repo config

This commit is contained in:
aran-yogesh 2026-02-09 12:53:34 -08:00
parent 03744fbf2d
commit 810bf412e9
3 changed files with 7 additions and 20 deletions

View file

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

View file

@ -1,4 +1,4 @@
"""Error normalization middleware for tool and model calls.
"""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.
@ -13,8 +13,6 @@ from collections.abc import Awaitable, Callable
from langchain.agents.middleware.types import (
AgentMiddleware,
AgentState,
ModelRequest,
ModelResponse,
)
from langchain_core.messages import ToolMessage
from langgraph.prebuilt.tool_node import ToolCallRequest
@ -63,15 +61,12 @@ def _get_tool_call_id(request: ToolCallRequest) -> str | None:
return None
class ErrorNormalizationMiddleware(AgentMiddleware):
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.
Model call errors (invalid API key, rate limit, etc.) are logged but
re-raised so they surface to the caller.
"""
state_schema = AgentState
@ -107,14 +102,3 @@ class ErrorNormalizationMiddleware(AgentMiddleware):
tool_call_id=_get_tool_call_id(request),
status="error",
)
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
try:
return await handler(request)
except Exception:
logger.exception("Error during model invocation")
raise

View file

@ -31,7 +31,7 @@ from deepagents import create_deep_agent
from langchain_anthropic import ChatAnthropic
from .encryption import decrypt_token
from .middleware import ErrorNormalizationMiddleware
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
@ -960,7 +960,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: # noqa: PLR0915
tools=[http_request, fetch_url, commit_and_open_pr],
backend=sandbox_backend,
middleware=[
ErrorNormalizationMiddleware(),
ToolErrorMiddleware(),
check_message_queue_before_model,
post_to_linear_after_model,
open_pr_if_needed,