From ee2f4b1a4c6517eee94ef63c4e207cd2e37bb6a5 Mon Sep 17 00:00:00 2001 From: PrimalPimmy Date: Mon, 15 Jun 2026 10:41:52 +0530 Subject: [PATCH] If tool call is empty, then return error to the server Signed-off-by: PrimalPimmy --- server/backend/commons.py | 14 +++++--- server/backend/linux.py | 50 ++++++++++++++++------------ server/tests/test_commons.py | 29 +++++++++++++++- server/tests/test_linux_streaming.py | 22 +++++++++++- 4 files changed, 87 insertions(+), 28 deletions(-) diff --git a/server/backend/commons.py b/server/backend/commons.py index 104fd4a..e688ff9 100644 --- a/server/backend/commons.py +++ b/server/backend/commons.py @@ -508,6 +508,8 @@ def _process_stop_tool_call_events( event_name = "response.function_call_arguments.done" has_recipient = bool(tool_name) tool_name = tool_name or _find_tool(request.tools, text) # pyright: ignore + if not tool_name: + raise ValueError("tool call is missing a tool name and could not be inferred") try: arguments_map = json.loads(text) @@ -552,11 +554,14 @@ def _process_stop_tool_call_events( ) -def _find_tool(tools: list, arguments_str: str) -> str: +def _find_tool(tools: list | None, arguments_str: str) -> str | None: try: arguments_map = json.loads(arguments_str) except json.JSONDecodeError as e: - arguments_map = {} + return None + + if not tools: + return None # To increase the accuracy of the selected tool, since we # check the required params is a subset of model responded @@ -583,9 +588,8 @@ def _find_tool(tools: list, arguments_str: str) -> str: break if tool_name == "": - return "read" - else: - return tool_name + return None + return tool_name def _is_correct_tool(required_params: list, model_argument_list: list) -> bool: diff --git a/server/backend/linux.py b/server/backend/linux.py index 694a999..431381a 100644 --- a/server/backend/linux.py +++ b/server/backend/linux.py @@ -424,29 +424,37 @@ async def generate_response_chat_stream( return # Emit the stop events current state - if state == "reasoning": - resp_str, sequence_number, output_index, item = _process_stop_reasoning_events( - reasoning_id, output_index, reasoning_text, sequence_number - ) - output_items.append(item) - yield resp_str - elif state == "toolcall": - resp_str, sequence_number, output_index, item = _process_stop_tool_call_events( - tool_id, - output_index, - tool_call_text, - sequence_number, - request, - tool_name, - ) - output_items.append(item) - yield resp_str - elif state == "answer": - resp_str, sequence_number, output_index, item = _process_output_item_done( - "message", message_id, answer_text, output_index, sequence_number + try: + if state == "reasoning": + resp_str, sequence_number, output_index, item = _process_stop_reasoning_events( + reasoning_id, output_index, reasoning_text, sequence_number + ) + output_items.append(item) + yield resp_str + elif state == "toolcall": + resp_str, sequence_number, output_index, item = _process_stop_tool_call_events( + tool_id, + output_index, + tool_call_text, + sequence_number, + request, + tool_name, + ) + output_items.append(item) + yield resp_str + elif state == "answer": + resp_str, sequence_number, output_index, item = _process_output_item_done( + "message", message_id, answer_text, output_index, sequence_number + ) + output_items.append(item) + yield resp_str + except Exception as e: + traceback.print_exc() + resp_str, sequence_number = _process_error_event( + str(e), response_id, request, created, sequence_number ) - output_items.append(item) yield resp_str + return ## Envelope, response.completed if generation_metrics is not None: diff --git a/server/tests/test_commons.py b/server/tests/test_commons.py index 24f2c05..f513460 100644 --- a/server/tests/test_commons.py +++ b/server/tests/test_commons.py @@ -1,6 +1,10 @@ from openai_harmony import HarmonyEncodingName, ReasoningEffort, Role, load_harmony_encoding -from server.backend.commons import build_harmony_conversation, normalize_harmony_tool_name +from server.backend.commons import ( + _find_tool, + build_harmony_conversation, + normalize_harmony_tool_name, +) from server.schemas import ResponsesRequest @@ -114,3 +118,26 @@ def test_normalize_harmony_tool_name_only_strips_known_function_namespace(): normalize_harmony_tool_name("functions.unknown", request.tools) == "functions.unknown" ) + + +def test_find_tool_does_not_fabricate_read_when_arguments_do_not_match(): + request = ResponsesRequest.model_validate( + { + "model": "unsloth/gpt-oss-20b-GGUF", + "input": "hello", + "tools": [ + { + "type": "function", + "name": "read", + "parameters": { + "type": "object", + "required": ["path"], + "properties": {"path": {"type": "string"}}, + }, + } + ], + } + ) + + assert _find_tool(request.tools, "{}") is None + assert _find_tool(request.tools, "{") is None diff --git a/server/tests/test_linux_streaming.py b/server/tests/test_linux_streaming.py index 12594bc..e7f5ca3 100644 --- a/server/tests/test_linux_streaming.py +++ b/server/tests/test_linux_streaming.py @@ -15,6 +15,7 @@ class FakeModel: class FakeRunner: model = FakeModel() tool_name = "read" + arguments = '{"path":"changelog.md"}' def cleanup(self): pass @@ -32,7 +33,7 @@ class FakeRunner: yield "**[Reasoning]**\n\n" yield "Need to read file." yield ToolCallStart(self.tool_name) - yield '{"path":"changelog.md"}' + yield self.arguments # TODO, make a better test for this as this one takes a lot of time # def test_get_or_load_model_without_cache_path_reuses_loaded_runner(): @@ -162,3 +163,22 @@ async def test_gpt_streaming_infers_tool_when_commentary_has_no_recipient(): assert function_call_added["item"]["name"] == "read" assert "event: response.function_call_arguments.delta\n" in stream assert "event: response.failed\n" not in stream + + +@pytest.mark.asyncio +async def test_gpt_streaming_fails_when_tool_cannot_be_inferred(): + request = make_request() + runner = FakeRunner() + runner.tool_name = "" + runner.arguments = "{}" + + with patch.object(linux, "get_or_load_model", return_value=runner): + chunks = [ + chunk async for chunk in linux.generate_response_chat_stream(request) + ] + + stream = "".join(chunks) + + assert "event: response.failed\n" in stream + assert "tool call is missing a tool name and could not be inferred" in stream + assert "event: response.output_item.done\n" not in stream -- 2.51.2