From 273a3d80b58b69ef5e2445936ba7ddc41325458b Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 06:33:32 +0000 Subject: [PATCH] Fix Responses history observation after tool argument normalization --- src/art_inference/append_only.py | 67 ++++++++++++------------- src/art_inference/vllm.py | 5 +- tests/unit/test_inference_history.py | 73 +++++++++++++++++++++++++++- 3 files changed, 110 insertions(+), 35 deletions(-) diff --git a/src/art_inference/append_only.py b/src/art_inference/append_only.py index 7eeb54c33..fb0184cbe 100644 --- a/src/art_inference/append_only.py +++ b/src/art_inference/append_only.py @@ -391,6 +391,40 @@ def chat_prefix_observations( return entries +def openai_tool_arguments(message: Mapping[str, Any]) -> Mapping[str, Any]: + function_call = message.get("function_call") + if isinstance(function_call, Mapping) and isinstance( + function_call.get("arguments"), Mapping + ): + message = { + **message, + "function_call": { + **function_call, + "arguments": json.dumps(function_call["arguments"]), + }, + } + calls = message.get("tool_calls") + if isinstance(calls, list): + return { + **message, + "tool_calls": [ + { + **call, + "function": { + **call["function"], + "arguments": json.dumps(call["function"]["arguments"]), + }, + } + if isinstance(call, Mapping) + and isinstance(call.get("function"), Mapping) + and isinstance(call["function"].get("arguments"), Mapping) + else call + for call in calls + ], + } + return message + + async def chat_response_prefixes( tokenizer: Any, request: Any, @@ -425,39 +459,6 @@ async def chat_response_prefixes( if name in fields ) - def openai_tool_arguments(message: Mapping[str, Any]) -> Mapping[str, Any]: - function_call = message.get("function_call") - if isinstance(function_call, Mapping) and isinstance( - function_call.get("arguments"), Mapping - ): - message = { - **message, - "function_call": { - **function_call, - "arguments": json.dumps(function_call["arguments"]), - }, - } - calls = message.get("tool_calls") - if isinstance(calls, list): - return { - **message, - "tool_calls": [ - { - **call, - "function": { - **call["function"], - "arguments": json.dumps(call["function"]["arguments"]), - }, - } - if isinstance(call, Mapping) - and isinstance(call.get("function"), Mapping) - and isinstance(call["function"].get("arguments"), Mapping) - else call - for call in calls - ], - } - return message - # vLLM renders tool-call arguments as mappings and shallow-copies their # containers, mutating historical request messages before this observer runs. messages = [ diff --git a/src/art_inference/vllm.py b/src/art_inference/vllm.py index ec007a576..f796f0546 100644 --- a/src/art_inference/vllm.py +++ b/src/art_inference/vllm.py @@ -18,6 +18,7 @@ aligned_values, chat_response_prefixes, merge_chat_delta, + openai_tool_arguments, output_prefix_observations, patch_deepseek_renderer, patch_harmony, @@ -352,7 +353,9 @@ async def observe_response(response): ) view = protocol.ChatCompletionRequest( model=request.model, - messages=conversation, + messages=[ + openai_tool_arguments(message) for message in conversation + ], chat_template_kwargs=self._effective_chat_template_kwargs(request), ) diff --git a/tests/unit/test_inference_history.py b/tests/unit/test_inference_history.py index a823a9bce..646e8d9b9 100644 --- a/tests/unit/test_inference_history.py +++ b/tests/unit/test_inference_history.py @@ -1,8 +1,10 @@ import asyncio +import copy import json from types import SimpleNamespace -from pydantic import BaseModel, model_validator +from openai.types.chat import ChatCompletionMessageToolCallParam +from pydantic import BaseModel, TypeAdapter, model_validator import pytest from art_inference import vllm @@ -399,3 +401,72 @@ async def run(): ) asyncio.run(run()) + + +@pytest.mark.parametrize("stream", [False, True]) +def test_responses_observe_rendered_tool_history_without_rejecting_result( + serving, stream +): + server, modules = serving + original_make = server._make_request + rendered = [] + + class StrictRequest(Request): + @model_validator(mode="after") + def validate_tool_arguments(self): + for message in self.messages: + for call in message.get("tool_calls", []): + TypeAdapter(ChatCompletionMessageToolCallParam).validate_python( + call + ) + return self + + async def make_request(request, previous): + conversation, inputs = await original_make(request, previous) + conversation = copy.deepcopy(conversation) + # vLLM's renderer converts historical JSON argument strings to mappings. + function = conversation[1]["tool_calls"][0]["function"] + function["arguments"] = json.loads(function["arguments"]) + rendered.append(conversation) + return conversation, inputs + + server._make_request = make_request + modules[ + "vllm.entrypoints.openai.chat_completion.protocol" + ].ChatCompletionRequest = StrictRequest + arguments = '{"query": "headphones", "max_price": 100}' + request = Request( + messages=[ + {"role": "user", "content": "question"}, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "search", "arguments": arguments}, + } + ], + }, + {"role": "tool", "tool_call_id": "call-1", "content": "[]"}, + ], + stream=stream, + ) + + async def run(): + response = await server.create_responses(request) + if stream: + event = await anext(response) + assert event.type == "response.completed" + await response.aclose() + response = event.response + assert response.output[0]["content"] == "action" + assert len(server.engine.prompts) == 1 + assert rendered[0][1]["tool_calls"][0]["function"]["arguments"] == json.loads( + arguments + ) + assert ( + request.messages[1]["tool_calls"][0]["function"]["arguments"] == arguments + ) + + asyncio.run(run())