diff --git a/strix/config/tool_call_ids.py b/strix/config/tool_call_ids.py index c418643db..83c23cce5 100644 --- a/strix/config/tool_call_ids.py +++ b/strix/config/tool_call_ids.py @@ -12,6 +12,12 @@ the paired tool result into a ``tool`` message with an empty ``tool_call_id``, which strict providers reject the same way and with the same permanent outcome. Rewriting both blank and duplicate ids to fresh unique ones keeps the history valid for any provider. + +A tool call can also arrive with no function name, when the streamed deltas +never carry one. Providers reject an assistant message whose tool call has an +empty ``function.name`` just as permanently. Filling it in with a placeholder +name keeps the request valid; the run loop then treats the call like any other +hallucinated tool and hands the model a "tool not found" result to recover from. """ from __future__ import annotations @@ -24,6 +30,9 @@ from agents.models.fake_id import FAKE_RESPONSES_ID from openai.types.responses import ResponseFunctionToolCall +UNNAMED_TOOL = "unnamed_tool" + + def new_call_id() -> str: return f"call_{uuid4().hex}" @@ -70,6 +79,9 @@ def dedupe_history_call_ids(items: list[Any]) -> tuple[list[Any], bool]: call_id = item.get("call_id") key = _pairing_key(call_id) if kind == "function_call": + if not item.get("name"): + item = {**item, "name": UNNAMED_TOOL} # noqa: PLW2901 + changed = True effective = key if not effective or effective in used: effective = new_call_id() @@ -135,6 +147,8 @@ class TurnCallIdRewriter: """Rewrite one item; ``position`` is its index in the turn's output.""" if not isinstance(item, ResponseFunctionToolCall): return item + if not item.name: + item = item.model_copy(update={"name": UNNAMED_TOOL}) original = item.call_id if not original: return self._rewrite_blank(item, position) diff --git a/tests/test_tool_call_ids.py b/tests/test_tool_call_ids.py index 2647315f8..3726c7f26 100644 --- a/tests/test_tool_call_ids.py +++ b/tests/test_tool_call_ids.py @@ -25,14 +25,18 @@ from openai import AsyncOpenAI from openai.types.responses import ResponseFunctionToolCall from strix.config.models import _NonStreamingModel, _TurnGuardModel -from strix.config.tool_call_ids import TurnCallIdRewriter, dedupe_history_call_ids +from strix.config.tool_call_ids import ( + UNNAMED_TOOL, + TurnCallIdRewriter, + dedupe_history_call_ids, +) if TYPE_CHECKING: from collections.abc import Iterator -def _tool_call_completion(call_id: str, n: int = 1) -> dict[str, Any]: +def _tool_call_completion(call_id: str, n: int = 1, name: str = "do_thing") -> dict[str, Any]: return { "id": "chatcmpl-1", "object": "chat.completion", @@ -49,7 +53,7 @@ def _tool_call_completion(call_id: str, n: int = 1) -> dict[str, Any]: { "id": call_id, "type": "function", - "function": {"name": "do_thing", "arguments": json.dumps({"n": n})}, + "function": {"name": name, "arguments": json.dumps({"n": n})}, } ], }, @@ -168,6 +172,49 @@ class _BlankIdHandler(BaseHTTPRequestHandler): self.wfile.write(encoded) +class _BlankNameHandler(BaseHTTPRequestHandler): + """Gateway that hands out a nameless tool call and rejects blank names, like GLM does.""" + + 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 message in messages: + for index, call in enumerate(message.get("tool_calls") or []): + if not (call.get("function") or {}).get("name"): + self._respond( + 400, + { + "error": { + "message": ( + f"tool_calls[{index}].function.name must be a " + "non-empty string (got empty string)" + ), + "code": 400, + } + }, + ) + return + + if len(_REQUESTS) == 1: + self._respond(200, _tool_call_completion("call_1", name="")) + 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) + + def _serve(handler: type[BaseHTTPRequestHandler]) -> Iterator[str]: _REQUESTS.clear() server = HTTPServer(("127.0.0.1", 0), handler) @@ -190,6 +237,11 @@ def blank_id_gateway() -> Iterator[str]: yield from _serve(_BlankIdHandler) +@pytest.fixture +def blank_name_gateway() -> Iterator[str]: + yield from _serve(_BlankNameHandler) + + def _model(base_url: str) -> Model: # The gateway answers plain JSON, so the run loop's streamed turns are # served non-streamed; the ids on the wire are the same either way. @@ -209,7 +261,11 @@ async def _run_agent(base_url: str, *, wrap: bool) -> Any: 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()) + agent, + input="please", + run_config=RunConfig( + model_provider=_Provider(), tool_not_found_behavior="return_error_to_model" + ), ) async for _ in result.stream_events(): pass @@ -316,6 +372,60 @@ def test_history_dedupe_fills_in_blank_ids_and_keeps_outputs_paired() -> None: assert ids[0] != ids[1] +@pytest.mark.asyncio +async def test_blank_tool_name_is_rejected_by_the_provider_without_the_wrapper( + blank_name_gateway: str, +) -> None: + # Repro: the model answers with a tool call carrying no function name. The + # run loop reports it as an unknown tool, but the nameless call stays in + # the history and the provider rejects every replay of it with a 400. + with pytest.raises(Exception, match=r"function\.name must be a non-empty string"): + await _run_agent(blank_name_gateway, wrap=False) + + +@pytest.mark.asyncio +async def test_blank_tool_name_is_filled_in_so_the_history_stays_valid( + blank_name_gateway: str, +) -> None: + result = await _run_agent(blank_name_gateway, wrap=True) + + assert result.final_output == "all done" + names = [ + call["function"]["name"] + for message in _REQUESTS[-1] + for call in message.get("tool_calls") or [] + ] + assert names == [UNNAMED_TOOL] + # The model learns the call went nowhere and can pick a real tool next turn. + assert len(_tool_results(_REQUESTS[-1])) == 1 + assert UNNAMED_TOOL in _tool_results(_REQUESTS[-1])[0] + + +def test_history_dedupe_fills_in_a_blank_tool_name() -> None: + items = [ + {"type": "function_call", "call_id": "call_a", "name": "", "arguments": "{}"}, + {"type": "function_call_output", "call_id": "call_a", "output": "x"}, + {"type": "function_call", "call_id": "call_b", "arguments": "{}"}, + ] + + rebuilt, changed = dedupe_history_call_ids(items) + + assert changed + assert rebuilt[0]["name"] == rebuilt[2]["name"] == UNNAMED_TOOL + assert rebuilt[0]["call_id"] == rebuilt[1]["call_id"] == "call_a" + + +def test_turn_rewriter_fills_in_a_blank_tool_name() -> None: + rewriter = TurnCallIdRewriter([]) + call = ResponseFunctionToolCall(call_id="call_1", name="", arguments="{}", type="function_call") + + first = rewriter.rewrite_item(call) + second = rewriter.rewrite_item(call) + + assert first.name == second.name == UNNAMED_TOOL + assert first.call_id == "call_1" + + def test_history_dedupe_fills_in_a_missing_call_id_key() -> None: items = [ {"type": "function_call", "name": "a", "arguments": "{}"},