Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
67 changes: 34 additions & 33 deletions src/art_inference/append_only.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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 = [
Expand Down
5 changes: 4 additions & 1 deletion src/art_inference/vllm.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
aligned_values,
chat_response_prefixes,
merge_chat_delta,
openai_tool_arguments,
output_prefix_observations,
patch_deepseek_renderer,
patch_harmony,
Expand Down Expand Up @@ -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),
)

Expand Down
73 changes: 72 additions & 1 deletion tests/unit/test_inference_history.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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())
Loading