mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-01 08:33:14 +00:00
fix: coerce malformed integer strings in read_file offset/limit params (#1216)
- Root cause: LLM occasionally generates strings like '1, 80' or '170, "limit": 60' for integer fields, causing a Pydantic ValidationError and wasting an LLM turn - Change: add SanitizeToolInputsMiddleware in agent/middleware/sanitize_tool_inputs.py that extracts the leading integer from any string value in offset/limit before the call reaches Pydantic validation; registered before ToolErrorMiddleware in server.py - Verified: 14 unit tests covering all three production trace patterns pass Co-authored-by: LangSmith Forge <forge-agent@langsmith.ai> Co-authored-by: Johannes du Plessis <johannes@langchain.dev>
This commit is contained in:
parent
e5bc27a0ad
commit
48965a82a6
4 changed files with 168 additions and 0 deletions
|
|
@ -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",
|
||||
|
|
|
|||
88
agent/middleware/sanitize_tool_inputs.py
Normal file
88
agent/middleware/sanitize_tool_inputs.py
Normal file
|
|
@ -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))
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
76
tests/test_sanitize_tool_inputs.py
Normal file
76
tests/test_sanitize_tool_inputs.py
Normal file
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue