diff --git a/agent/middleware/__init__.py b/agent/middleware/__init__.py index 18f11fa3..14b9bf59 100644 --- a/agent/middleware/__init__.py +++ b/agent/middleware/__init__.py @@ -2,9 +2,11 @@ from .check_message_queue import check_message_queue_before_model from .ensure_no_empty_msg import ensure_no_empty_msg from .notify_step_limit import notify_step_limit_reached from .open_pr import open_pr_if_needed +from .sanitize_tool_inputs import SanitizeToolInputsMiddleware from .tool_error_handler import ToolErrorMiddleware __all__ = [ + "SanitizeToolInputsMiddleware", "ToolErrorMiddleware", "check_message_queue_before_model", "ensure_no_empty_msg", diff --git a/agent/middleware/sanitize_tool_inputs.py b/agent/middleware/sanitize_tool_inputs.py new file mode 100644 index 00000000..64d8f2b4 --- /dev/null +++ b/agent/middleware/sanitize_tool_inputs.py @@ -0,0 +1,88 @@ +"""Sanitize tool input middleware. + +Coerces malformed integer fields in read_file calls before they reach Pydantic +validation. The LLM occasionally generates strings like ``'1, 80'`` or +``'170, "limit": 60'`` for integer parameters; we extract the leading digit +sequence so the call succeeds instead of burning an LLM turn on a retry. +""" + +from __future__ import annotations + +import logging +import re +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__) + +_READ_FILE_INT_FIELDS = ("offset", "limit") + + +def _coerce_int(value: object) -> int | None: + """Extract the first integer from *value* if it is a non-integer string. + + Returns the parsed integer, or ``None`` if no leading digits are found. + If *value* is already an ``int`` (or ``None``), returns it unchanged. + """ + if value is None or isinstance(value, int): + return value + if isinstance(value, str): + match = re.match(r"\s*(\d+)", value) + if match: + return int(match.group(1)) + return None + return None + + +def _sanitize_read_file_args(args: dict) -> dict: + """Return a copy of *args* with integer fields coerced where needed.""" + sanitized = dict(args) + for field in _READ_FILE_INT_FIELDS: + if field in sanitized: + original = sanitized[field] + coerced = _coerce_int(original) + if coerced is not None and coerced != original: + logger.warning("Coercing read_file.%s from %r to %d", field, original, coerced) + sanitized[field] = coerced + return sanitized + + +class SanitizeToolInputsMiddleware(AgentMiddleware): + """Intercept read_file calls and coerce malformed integer parameters. + + When the LLM produces a string value for an integer field (e.g. + ``offset='1, 80'``), this middleware extracts the leading integer so that + Pydantic validation passes rather than raising a ``ValidationError`` and + forcing an unnecessary retry. + """ + + state_schema = AgentState + + def _sanitize_request(self, request: ToolCallRequest) -> ToolCallRequest: + tool_call = request.tool_call + if not isinstance(tool_call, dict) or tool_call.get("name") != "read_file": + return request + args = tool_call.get("args", {}) + sanitized_args = _sanitize_read_file_args(args) + if sanitized_args is args: + return request + new_tool_call = {**tool_call, "args": sanitized_args} + return request.override(tool_call=new_tool_call) + + def wrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], ToolMessage | Command], + ) -> ToolMessage | Command: + return handler(self._sanitize_request(request)) + + async def awrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]], + ) -> ToolMessage | Command: + return await handler(self._sanitize_request(request)) diff --git a/agent/server.py b/agent/server.py index 9c6ca69d..680c9e67 100644 --- a/agent/server.py +++ b/agent/server.py @@ -28,6 +28,7 @@ from langsmith.sandbox import SandboxClientError from .integrations.langsmith import _configure_github_proxy from .middleware import ( + SanitizeToolInputsMiddleware, ToolErrorMiddleware, check_message_queue_before_model, ensure_no_empty_msg, @@ -324,6 +325,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: ], backend=sandbox_backend, middleware=[ + SanitizeToolInputsMiddleware(), ModelCallLimitMiddleware(run_limit=60, exit_behavior="end"), ToolErrorMiddleware(), check_message_queue_before_model, diff --git a/tests/test_sanitize_tool_inputs.py b/tests/test_sanitize_tool_inputs.py new file mode 100644 index 00000000..56510d72 --- /dev/null +++ b/tests/test_sanitize_tool_inputs.py @@ -0,0 +1,76 @@ +"""Unit tests for SanitizeToolInputsMiddleware. + +Guards against the regression where the LLM generates a string value for an +integer field in read_file (e.g. offset='1, 80'), causing a Pydantic +ValidationError and an unnecessary retry. +""" + +from __future__ import annotations + +from agent.middleware.sanitize_tool_inputs import _coerce_int, _sanitize_read_file_args + + +class TestCoerceInt: + def test_already_int_passes_through(self) -> None: + assert _coerce_int(1) == 1 + + def test_none_passes_through(self) -> None: + assert _coerce_int(None) is None + + def test_extracts_leading_integer_from_comma_string(self) -> None: + # Production trace 1: offset='1, 80' + assert _coerce_int("1, 80") == 1 + + def test_extracts_leading_integer_from_embedded_json(self) -> None: + # Production trace 2: offset='170, "limit": 60' + assert _coerce_int('170, "limit": 60') == 170 + + def test_extracts_leading_integer_from_trailing_comma(self) -> None: + # Production trace 3: offset='1504, ' + assert _coerce_int("1504, ") == 1504 + + def test_returns_none_when_no_digits(self) -> None: + assert _coerce_int("abc") is None + + def test_handles_leading_whitespace(self) -> None: + assert _coerce_int(" 42, extra") == 42 + + +class TestSanitizeReadFileArgs: + def test_coerces_offset_string_to_int(self) -> None: + args = {"file_path": "foo.ts", "offset": "1, 80", "limit": 80} + result = _sanitize_read_file_args(args) + assert result["offset"] == 1 + assert result["limit"] == 80 + assert result["file_path"] == "foo.ts" + + def test_coerces_offset_with_embedded_json(self) -> None: + args = {"file_path": "bar.tsx", "offset": '170, "limit": 60', "limit": 60} + result = _sanitize_read_file_args(args) + assert result["offset"] == 170 + + def test_coerces_offset_with_trailing_comma(self) -> None: + args = {"file_path": "baz.go", "offset": "1504, ", "limit": 200} + result = _sanitize_read_file_args(args) + assert result["offset"] == 1504 + + def test_int_offset_unchanged(self) -> None: + args = {"file_path": "foo.ts", "offset": 42, "limit": 80} + result = _sanitize_read_file_args(args) + assert result["offset"] == 42 + + def test_missing_offset_unchanged(self) -> None: + args = {"file_path": "foo.ts"} + result = _sanitize_read_file_args(args) + assert "offset" not in result + + def test_uncoercible_offset_passed_through_unchanged(self) -> None: + # If no digits at all, we leave the value alone so ToolErrorMiddleware handles it. + args = {"file_path": "foo.ts", "offset": "bad"} + result = _sanitize_read_file_args(args) + assert result["offset"] == "bad" + + def test_does_not_mutate_original_dict(self) -> None: + args = {"file_path": "foo.ts", "offset": "1, 80"} + _ = _sanitize_read_file_args(args) + assert args["offset"] == "1, 80"