fix(models): replay malformed tool-call arguments as valid JSON

A model that emits a tool call with truncated or otherwise invalid JSON
arguments fails that one call, but the raw string is recorded in the
session. Strict OpenAI-compatible servers (vLLM, SGLang) validate every
assistant tool call in the request, so each later turn is rejected with
"Assistant tool call function.arguments must be valid JSON" and the agent
never recovers. The turn guard now rewrites such arguments on the way out
to a JSON object carrying the original text, alongside the existing
duplicate call-id repair.
This commit is contained in:
Alex Schapiro
2026-09-17 03:53:28 +00:00
parent 910c1ea4bb
commit b6668d13b3
3 changed files with 323 additions and 2 deletions

View File

@@ -38,6 +38,7 @@ from openai.types.shared import Reasoning
from strix.config import codex from strix.config import codex
from strix.config.loader import load_settings from strix.config.loader import load_settings
from strix.config.tool_call_arguments import repair_input
from strix.config.tool_call_ids import TurnCallIdRewriter, dedupe_input from strix.config.tool_call_ids import TurnCallIdRewriter, dedupe_input
from strix.config.tool_call_limits import TurnToolCallLimiter from strix.config.tool_call_limits import TurnToolCallLimiter
@@ -251,6 +252,12 @@ class _TurnGuardModel(Model):
Ids that collide with the history are rewritten before the turn is Ids that collide with the history are rewritten before the turn is
recorded, and already-corrupted histories are repaired on the way out. recorded, and already-corrupted histories are repaired on the way out.
Tool-call arguments: a turn whose ``arguments`` are not valid JSON fails
that one call with a parse error, but the raw string is recorded and
strict providers then reject every request that replays it. Such
arguments are rewritten to a valid JSON object on the way out (see
:mod:`strix.config.tool_call_arguments`).
Tool-call volume: a degenerate response can queue hundreds of calls that Tool-call volume: a degenerate response can queue hundreds of calls that
the run loop then honours one by one. Only the first the run loop then honours one by one. Only the first
``LLM_MAX_TOOL_CALLS_PER_TURN`` calls of a response are kept. ``LLM_MAX_TOOL_CALLS_PER_TURN`` calls of a response are kept.
@@ -303,7 +310,7 @@ class _TurnGuardModel(Model):
conversation_id: str | None, conversation_id: str | None,
prompt: ResponsePromptParam | None, prompt: ResponsePromptParam | None,
) -> ModelResponse: ) -> ModelResponse:
sanitized = dedupe_input(input) sanitized = _sanitize_input(input)
rewriter = TurnCallIdRewriter(sanitized) rewriter = TurnCallIdRewriter(sanitized)
response = await self._inner.get_response( response = await self._inner.get_response(
system_instructions, system_instructions,
@@ -336,7 +343,7 @@ class _TurnGuardModel(Model):
conversation_id: str | None, conversation_id: str | None,
prompt: ResponsePromptParam | None, prompt: ResponsePromptParam | None,
) -> AsyncIterator[TResponseStreamEvent]: ) -> AsyncIterator[TResponseStreamEvent]:
sanitized = dedupe_input(input) sanitized = _sanitize_input(input)
rewriter = TurnCallIdRewriter(sanitized) rewriter = TurnCallIdRewriter(sanitized)
limiter = self._limiter() limiter = self._limiter()
stream = self._inner.stream_response( stream = self._inner.stream_response(
@@ -358,6 +365,10 @@ class _TurnGuardModel(Model):
self._log_dropped(limiter) self._log_dropped(limiter)
def _sanitize_input(model_input: str | list[TResponseInputItem]) -> str | list[Any]:
return repair_input(dedupe_input(model_input))
async def _aclose(stream: AsyncIterator[TResponseStreamEvent]) -> None: async def _aclose(stream: AsyncIterator[TResponseStreamEvent]) -> None:
if isinstance(stream, AsyncGenerator): if isinstance(stream, AsyncGenerator):
with contextlib.suppress(Exception): with contextlib.suppress(Exception):

View File

@@ -0,0 +1,64 @@
"""Keep replayed tool-call arguments valid JSON.
A model occasionally emits a tool call whose ``arguments`` string is not a
JSON object (truncated, unbalanced, empty). The call itself fails safely: the
tool reports the parse error and the run continues. But the raw string is
recorded in the session as-is, and strict OpenAI-compatible servers (vLLM,
SGLang, ...) validate every assistant tool call in the request, so from then
on each turn is rejected with ``Assistant tool call function.arguments must
be valid JSON`` and the agent can never recover. Rewriting the replayed
arguments to a JSON object that carries the original text keeps the history
valid while the model still sees what it sent.
"""
from __future__ import annotations
import json
from typing import Any
from openai.types.responses import ResponseFunctionToolCall
MALFORMED_ARGUMENTS_KEY = "malformed_arguments"
def repair_arguments(arguments: object) -> str | None:
"""Return replacement arguments a strict server accepts, or ``None`` if already valid."""
if not isinstance(arguments, str) or not arguments.strip():
return "{}"
try:
parsed = json.loads(arguments)
except ValueError:
parsed = None
if isinstance(parsed, dict):
return None
return json.dumps({MALFORMED_ARGUMENTS_KEY: arguments}, ensure_ascii=False)
def repair_history_arguments(items: list[Any]) -> tuple[list[Any], bool]:
"""Rewrite function calls in a conversation history whose arguments are not a JSON object."""
rebuilt: list[Any] = []
changed = False
for item in items:
if isinstance(item, dict):
if item.get("type") == "function_call":
repaired = repair_arguments(item.get("arguments"))
if repaired is not None:
item = {**item, "arguments": repaired} # noqa: PLW2901
changed = True
elif isinstance(item, ResponseFunctionToolCall):
repaired = repair_arguments(item.arguments)
if repaired is not None:
item = item.model_copy(update={"arguments": repaired}) # noqa: PLW2901
changed = True
rebuilt.append(item)
return rebuilt, changed
def repair_input(model_input: str | list[Any]) -> str | list[Any]:
if isinstance(model_input, str):
return model_input
rebuilt, changed = repair_history_arguments(model_input)
return rebuilt if changed else model_input

View File

@@ -0,0 +1,246 @@
"""Tests for keeping replayed tool-call arguments valid JSON.
A model that emits a tool call with malformed ``arguments`` fails that one
call, but the raw string is recorded in the session. Strict OpenAI-compatible
servers validate every assistant tool call in the request, so each later turn
is rejected and the agent can never recover. A gateway that validates
arguments the way those servers do proves both the failure and the fix.
"""
from __future__ import annotations
import json
import threading
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import TYPE_CHECKING, Any
import pytest
from agents import Agent, Runner, function_tool
from agents.models.interface import Model, ModelProvider
from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel
from agents.run import RunConfig
from openai import AsyncOpenAI, BadRequestError
from openai.types.responses import ResponseFunctionToolCall
from strix.config.models import _NonStreamingModel, _TurnGuardModel
from strix.config.tool_call_arguments import (
MALFORMED_ARGUMENTS_KEY,
repair_arguments,
repair_history_arguments,
repair_input,
)
if TYPE_CHECKING:
from collections.abc import Iterator
TRUNCATED = '{"n": 1'
def _tool_call_completion(arguments: str) -> dict[str, Any]:
return {
"id": "chatcmpl-1",
"object": "chat.completion",
"created": 0,
"model": "gw-model",
"choices": [
{
"index": 0,
"finish_reason": "tool_calls",
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "do_thing", "arguments": arguments},
}
],
},
}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7},
}
def _text_completion(text: str) -> dict[str, Any]:
return {
"id": "chatcmpl-2",
"object": "chat.completion",
"created": 0,
"model": "gw-model",
"choices": [
{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": text}}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8},
}
_REQUESTS: list[list[dict[str, Any]]] = []
def _assistant_arguments(messages: list[dict[str, Any]]) -> list[str]:
return [
str(call["function"]["arguments"])
for message in messages
for call in message.get("tool_calls") or []
]
def _tool_results(messages: list[dict[str, Any]]) -> list[str]:
return [str(m.get("content")) for m in messages if m.get("role") == "tool"]
class _StrictHandler(BaseHTTPRequestHandler):
"""Gateway that rejects malformed assistant tool-call arguments, like vLLM/SGLang do."""
def log_message(self, *args: Any) -> None:
pass
def do_POST(self) -> None:
length = int(self.headers.get("Content-Length", 0))
body = json.loads(self.rfile.read(length) or b"{}")
messages = body.get("messages", [])
_REQUESTS.append(messages)
for arguments in _assistant_arguments(messages):
try:
json.loads(arguments)
except ValueError:
self._respond(
400,
{
"object": "error",
"message": "Assistant tool call function.arguments must be valid JSON.",
"type": "BadRequest",
"param": None,
"code": 400,
},
)
return
if len(_REQUESTS) == 1:
self._respond(200, _tool_call_completion(TRUNCATED))
else:
self._respond(200, _text_completion("all done"))
def _respond(self, status: int, payload: dict[str, Any]) -> None:
encoded = json.dumps(payload).encode()
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(encoded)))
self.end_headers()
self.wfile.write(encoded)
@pytest.fixture
def strict_gateway() -> Iterator[str]:
_REQUESTS.clear()
server = HTTPServer(("127.0.0.1", 0), _StrictHandler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_address[1]}/v1"
finally:
server.shutdown()
server.server_close()
def _model(base_url: str) -> Model:
client = AsyncOpenAI(api_key="tok", base_url=base_url, max_retries=0)
return _NonStreamingModel(OpenAIChatCompletionsModel(model="gw-model", openai_client=client))
async def _run_agent(base_url: str, *, wrap: bool) -> Any:
@function_tool
def do_thing(n: int) -> str:
return f"did {n}"
class _Provider(ModelProvider):
def get_model(self, model_name: str | None) -> Model: # noqa: ARG002
model = _model(base_url)
return _TurnGuardModel(model) if wrap else model
agent = Agent(name="t", instructions="use the tool", tools=[do_thing], model="gw-model")
result = Runner.run_streamed(
agent, input="please", run_config=RunConfig(model_provider=_Provider())
)
async for _ in result.stream_events():
pass
return result
@pytest.mark.asyncio
async def test_malformed_arguments_poison_every_later_turn_without_the_wrapper(
strict_gateway: str,
) -> None:
# Repro: the model truncates a tool call's arguments once. The tool fails
# that call gracefully, but the next request replays the raw string and
# the provider rejects the whole conversation from then on.
with pytest.raises(BadRequestError, match="must be valid JSON"):
await _run_agent(strict_gateway, wrap=False)
assert _assistant_arguments(_REQUESTS[-1]) == [TRUNCATED]
@pytest.mark.asyncio
async def test_malformed_arguments_are_replayed_as_valid_json(strict_gateway: str) -> None:
result = await _run_agent(strict_gateway, wrap=True)
assert result.final_output == "all done"
(arguments,) = _assistant_arguments(_REQUESTS[-1])
assert json.loads(arguments) == {MALFORMED_ARGUMENTS_KEY: TRUNCATED}
(tool_result,) = _tool_results(_REQUESTS[-1])
assert "JSON" in tool_result
@pytest.mark.parametrize(
("arguments", "expected"),
[
('{"cmd": "ls', json.dumps({MALFORMED_ARGUMENTS_KEY: '{"cmd": "ls'})),
("[1, 2]", json.dumps({MALFORMED_ARGUMENTS_KEY: "[1, 2]"})),
("", "{}"),
(" ", "{}"),
(None, "{}"),
],
)
def test_repair_arguments_rewrites_anything_but_a_json_object(
arguments: str | None, expected: str
) -> None:
assert repair_arguments(arguments) == expected
@pytest.mark.parametrize("arguments", ["{}", '{"cmd": "ls -la"}', '{"nested": {"a": [1]}}'])
def test_repair_arguments_leaves_json_objects_alone(arguments: str) -> None:
assert repair_arguments(arguments) is None
def test_history_repair_rewrites_only_malformed_calls() -> None:
items = [
{"role": "user", "content": "go"},
{"type": "function_call", "call_id": "a", "name": "x", "arguments": '{"cmd": "ls'},
{"type": "function_call_output", "call_id": "a", "output": "invalid JSON"},
{"type": "function_call", "call_id": "b", "name": "y", "arguments": '{"ok": true}'},
ResponseFunctionToolCall(call_id="c", name="z", arguments="", type="function_call"),
]
rebuilt, changed = repair_history_arguments(items)
assert changed
assert rebuilt[0] is items[0]
assert json.loads(rebuilt[1]["arguments"]) == {MALFORMED_ARGUMENTS_KEY: '{"cmd": "ls'}
assert rebuilt[1]["call_id"] == "a"
assert rebuilt[2] is items[2]
assert rebuilt[3] is items[3]
assert isinstance(rebuilt[4], ResponseFunctionToolCall)
assert rebuilt[4].arguments == "{}"
assert items[1]["arguments"] == '{"cmd": "ls'
def test_repair_input_returns_same_object_when_nothing_changes() -> None:
items = [{"type": "function_call", "call_id": "a", "name": "x", "arguments": "{}"}]
assert repair_input(items) is items
assert repair_input("plain prompt") == "plain prompt"