mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-06 12:22:11 +00:00
feat: add error normalization middleware for tool and model calls
This commit is contained in:
parent
9f78cf0b90
commit
7413befa79
4 changed files with 127 additions and 3 deletions
122
apps/agent/agent/middleware.py
Normal file
122
apps/agent/agent/middleware.py
Normal file
|
|
@ -0,0 +1,122 @@
|
||||||
|
"""Error normalization middleware for tool and model calls.
|
||||||
|
|
||||||
|
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,
|
||||||
|
ModelRequest,
|
||||||
|
ModelResponse,
|
||||||
|
)
|
||||||
|
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 ErrorNormalizationMiddleware(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
|
||||||
|
|
||||||
|
def wrap_tool_call(
|
||||||
|
self,
|
||||||
|
request: ToolCallRequest,
|
||||||
|
handler: Callable[[ToolCallRequest], ToolMessage | Command],
|
||||||
|
) -> ToolMessage | Command:
|
||||||
|
try:
|
||||||
|
return handler(request)
|
||||||
|
except Exception as e: # noqa: BLE001
|
||||||
|
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: # noqa: BLE001
|
||||||
|
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_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
|
||||||
|
|
@ -26,13 +26,14 @@ warnings.filterwarnings("ignore", message=".*Pydantic V1.*", category=UserWarnin
|
||||||
|
|
||||||
# Now safe to import agent (which imports LangChain modules)
|
# Now safe to import agent (which imports LangChain modules)
|
||||||
from deepagents import create_deep_agent
|
from deepagents import create_deep_agent
|
||||||
from .protocol import SandboxBackendProtocol
|
|
||||||
|
|
||||||
# Local import for encryption
|
# Local import for encryption
|
||||||
from langchain_anthropic import ChatAnthropic
|
from langchain_anthropic import ChatAnthropic
|
||||||
|
|
||||||
from .encryption import decrypt_token
|
from .encryption import decrypt_token
|
||||||
|
from .middleware import ErrorNormalizationMiddleware
|
||||||
from .prompt import construct_system_prompt
|
from .prompt import construct_system_prompt
|
||||||
|
from .protocol import SandboxBackendProtocol
|
||||||
from .tools import commit_and_open_pr, fetch_url, http_request
|
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],
|
tools=[http_request, fetch_url, commit_and_open_pr],
|
||||||
backend=sandbox_backend,
|
backend=sandbox_backend,
|
||||||
middleware=[
|
middleware=[
|
||||||
|
ErrorNormalizationMiddleware(),
|
||||||
check_message_queue_before_model,
|
check_message_queue_before_model,
|
||||||
post_to_linear_after_model,
|
post_to_linear_after_model,
|
||||||
open_pr_if_needed,
|
open_pr_if_needed,
|
||||||
|
|
|
||||||
|
|
@ -71,7 +71,7 @@ def get_service_jwt_token_for_user(
|
||||||
|
|
||||||
LINEAR_TEAM_TO_REPO: dict[str, dict[str, str]] = {
|
LINEAR_TEAM_TO_REPO: dict[str, dict[str, str]] = {
|
||||||
"Brace's test workspace": {"owner": "langchain-ai", "name": "open-swe"},
|
"Brace's test workspace": {"owner": "langchain-ai", "name": "open-swe"},
|
||||||
"Yogesh-dev": {"owner": "aran-yogesh", "name": "nimedge"},
|
"Yogesh-dev": {"owner": "aran-yogesh", "name": "TalkBack"},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue