diff --git a/src/art_inference/append_only.py b/src/art_inference/append_only.py index 989c1c61a..7eeb54c33 100644 --- a/src/art_inference/append_only.py +++ b/src/art_inference/append_only.py @@ -425,7 +425,48 @@ 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 = [ + openai_tool_arguments(message) if isinstance(message, Mapping) else message + for message in messages + ] + async def complete(message: Mapping[str, Any]) -> list[int] | None: + message = openai_tool_arguments(message) return await render( type(request).model_validate({**payload, "messages": [*messages, message]}) ) diff --git a/tests/unit/test_append_only.py b/tests/unit/test_append_only.py index 99bff1ad5..4c761020d 100644 --- a/tests/unit/test_append_only.py +++ b/tests/unit/test_append_only.py @@ -1,4 +1,6 @@ import asyncio +from copy import deepcopy +import json from types import SimpleNamespace from pydantic import BaseModel @@ -140,6 +142,104 @@ async def render(_): ) +class StrictFunction(BaseModel): + name: str + arguments: str + + +class StrictToolCall(BaseModel): + id: str + type: str + function: StrictFunction + + +class StrictMessage(BaseModel): + role: str + content: str | None = None + tool_calls: list[StrictToolCall] | None = None + function_call: StrictFunction | None = None + + +class StrictRequest(BaseModel): + messages: list[StrictMessage] + + +@pytest.mark.filterwarnings("ignore:Pydantic serializer warnings") +@pytest.mark.parametrize( + ("arguments", "legacy", "historical"), + [ + ({"id": 3}, False, False), + ('{ "id": 3 }', False, False), + ({"id": 3}, True, False), + ({"previous": True}, False, True), + ], + ids=("mapping", "string", "legacy-mapping", "historical-mapping"), +) +def test_response_observation_accepts_structured_tool_arguments( + arguments, legacy, historical +): + tokenizer = Tokenizer() + messages = [StrictMessage(role="user", content="question")] + if historical: + function = StrictFunction(name="previous", arguments="{}") + object.__setattr__(function, "arguments", arguments) + messages.extend( + [ + StrictMessage( + role="assistant", + tool_calls=[ + StrictToolCall( + id="previous", type="function", function=function + ) + ], + ), + StrictMessage(role="user", content="next"), + ] + ) + request = StrictRequest(messages=messages) + message = {"role": "assistant", "content": "answer"} + if not historical: + function = {"name": "lookup", "arguments": arguments} + if legacy: + message["function_call"] = function + else: + message["tool_calls"] = [ + {"id": "call", "type": "function", "function": function} + ] + original_request = request.model_dump(mode="python", warnings=False) + original_message = deepcopy(message) + rendered_arguments: list[list[str]] = [] + + async def render(value: StrictRequest) -> list[int]: + if len(value.messages) == len(request.messages): + return tokenizer.encode("prompt:") + completed: list[str] = [] + for rendered_message in value.messages: + completed.extend( + call.function.arguments for call in rendered_message.tool_calls or [] + ) + if rendered_message.function_call: + completed.append(rendered_message.function_call.arguments) + rendered_arguments.append(completed) + return tokenizer.encode("prompt:answerEND") + + observations = asyncio.run( + chat_response_prefixes( + tokenizer, + request, + tokenizer.encode("prompt:"), + [(message, tokenizer.encode("answerEND"), True)], + render, + ) + ) + + expected = json.dumps(arguments) if isinstance(arguments, dict) else arguments + assert rendered_arguments == [[expected]] + assert observations + assert request.model_dump(mode="python", warnings=False) == original_request + assert message == original_message + + def test_custom_stop_does_not_delete_the_template_terminator(): encode = Tokenizer().encode assert (