mirror of
https://github.com/usestrix/strix.git
synced 2026-09-30 01:52:18 +00:00
Fill in blank tool-call ids so strict providers accept the history (#1355)
This commit is contained in:
parent
4c1be22150
commit
ae38fe70cd
4 changed files with 578 additions and 27 deletions
|
|
@ -391,7 +391,7 @@ def _guard_event(
|
|||
event: TResponseStreamEvent, rewriter: TurnCallIdRewriter, limiter: TurnToolCallLimiter
|
||||
) -> TResponseStreamEvent | None:
|
||||
if isinstance(event, ResponseOutputItemAddedEvent | ResponseOutputItemDoneEvent):
|
||||
rewritten = rewriter.rewrite_item(event.item)
|
||||
rewritten = rewriter.rewrite_item(event.item, event.output_index)
|
||||
if not limiter.allow(rewritten):
|
||||
return None
|
||||
if rewritten is not event.item:
|
||||
|
|
|
|||
|
|
@ -1,12 +1,17 @@
|
|||
"""Keep tool-call ids unique within a conversation.
|
||||
"""Keep tool-call ids present and unique within a conversation.
|
||||
|
||||
Some providers return per-turn tool-call ids (``exec_command:0``,
|
||||
``exec_command:1``, ...) whose counter restarts on every turn. Once the same
|
||||
id appears twice in one conversation, the request payload has two assistant
|
||||
tool calls sharing an id and strict providers reject the whole turn, which
|
||||
permanently kills the agent because the malformed history is replayed on
|
||||
every retry. Rewriting duplicates to fresh unique ids keeps the history
|
||||
valid for any provider.
|
||||
every retry.
|
||||
|
||||
Others omit the id altogether, or return it as an empty string. That turns
|
||||
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.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -15,6 +20,7 @@ from collections import defaultdict, deque
|
|||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
from agents.models.fake_id import FAKE_RESPONSES_ID
|
||||
from openai.types.responses import ResponseFunctionToolCall
|
||||
|
||||
|
||||
|
|
@ -22,23 +28,29 @@ def new_call_id() -> str:
|
|||
return f"call_{uuid4().hex}"
|
||||
|
||||
|
||||
def _pairing_key(call_id: Any) -> str:
|
||||
"""Bucket an id for call/output pairing; all blank ids share one bucket."""
|
||||
return call_id if isinstance(call_id, str) else ""
|
||||
|
||||
|
||||
def collect_call_ids(items: list[Any]) -> set[str]:
|
||||
used: set[str] = set()
|
||||
for item in items:
|
||||
if isinstance(item, dict):
|
||||
call_id = item.get("call_id")
|
||||
if isinstance(call_id, str):
|
||||
if isinstance(call_id, str) and call_id:
|
||||
used.add(call_id)
|
||||
elif isinstance(item, ResponseFunctionToolCall):
|
||||
elif isinstance(item, ResponseFunctionToolCall) and item.call_id:
|
||||
used.add(item.call_id)
|
||||
return used
|
||||
|
||||
|
||||
def dedupe_history_call_ids(items: list[Any]) -> tuple[list[Any], bool]:
|
||||
"""Rewrite duplicate call ids in a conversation history.
|
||||
"""Rewrite blank and duplicate call ids in a conversation history.
|
||||
|
||||
Outputs are paired with their call by order, so parallel calls that share
|
||||
an id keep answering the right call after the rewrite.
|
||||
an id — or are all missing one — keep answering the right call after the
|
||||
rewrite.
|
||||
"""
|
||||
used: set[str] = set()
|
||||
pending: dict[str, deque[str]] = defaultdict(deque)
|
||||
|
|
@ -49,22 +61,24 @@ def dedupe_history_call_ids(items: list[Any]) -> tuple[list[Any], bool]:
|
|||
if not isinstance(item, dict):
|
||||
rebuilt.append(item)
|
||||
continue
|
||||
call_id = item.get("call_id")
|
||||
if not isinstance(call_id, str):
|
||||
|
||||
kind = item.get("type")
|
||||
if kind not in ("function_call", "function_call_output"):
|
||||
rebuilt.append(item)
|
||||
continue
|
||||
|
||||
kind = item.get("type")
|
||||
call_id = item.get("call_id")
|
||||
key = _pairing_key(call_id)
|
||||
if kind == "function_call":
|
||||
effective = call_id
|
||||
if call_id in used:
|
||||
effective = key
|
||||
if not effective or effective in used:
|
||||
effective = new_call_id()
|
||||
item = {**item, "call_id": effective} # noqa: PLW2901
|
||||
changed = True
|
||||
used.add(effective)
|
||||
pending[call_id].append(effective)
|
||||
elif kind == "function_call_output":
|
||||
queue = pending.get(call_id)
|
||||
pending[key].append(effective)
|
||||
else:
|
||||
queue = pending.get(key)
|
||||
if queue:
|
||||
effective = queue.popleft()
|
||||
if effective != call_id:
|
||||
|
|
@ -83,7 +97,7 @@ def dedupe_input(model_input: str | list[Any]) -> str | list[Any]:
|
|||
|
||||
|
||||
class TurnCallIdRewriter:
|
||||
"""Rewrite a single turn's tool-call ids that collide with the history.
|
||||
"""Rewrite a single turn's tool-call ids that are blank or collide with the history.
|
||||
|
||||
A turn's items surface several times (streamed item events, then the
|
||||
completed response), so the same original id must always map to the same
|
||||
|
|
@ -93,12 +107,37 @@ class TurnCallIdRewriter:
|
|||
def __init__(self, model_input: str | list[Any]) -> None:
|
||||
self._used = set() if isinstance(model_input, str) else collect_call_ids(model_input)
|
||||
self._remap: dict[str, str] = {}
|
||||
self._blank_remap: dict[str, str] = {}
|
||||
self._settled: set[str] = set()
|
||||
|
||||
def rewrite_item(self, item: Any) -> Any:
|
||||
def _rewrite_blank(
|
||||
self, item: ResponseFunctionToolCall, position: int | None
|
||||
) -> ResponseFunctionToolCall:
|
||||
"""Give a call with no id one that stays the same on every sighting.
|
||||
|
||||
A real item id tells parallel calls apart. Chat Completions routes give
|
||||
every item the same placeholder id, so there the call's position in
|
||||
the turn's output tells them apart instead.
|
||||
"""
|
||||
if item.id and item.id != FAKE_RESPONSES_ID:
|
||||
key = item.id
|
||||
else:
|
||||
key = f"{FAKE_RESPONSES_ID}#{position}"
|
||||
replacement = self._blank_remap.get(key)
|
||||
if replacement is None:
|
||||
replacement = new_call_id()
|
||||
self._blank_remap[key] = replacement
|
||||
self._used.add(replacement)
|
||||
self._settled.add(replacement)
|
||||
return item.model_copy(update={"call_id": replacement})
|
||||
|
||||
def rewrite_item(self, item: Any, position: int | None = None) -> Any:
|
||||
"""Rewrite one item; ``position`` is its index in the turn's output."""
|
||||
if not isinstance(item, ResponseFunctionToolCall):
|
||||
return item
|
||||
original = item.call_id
|
||||
if not original:
|
||||
return self._rewrite_blank(item, position)
|
||||
if original in self._settled:
|
||||
return item
|
||||
replacement = self._remap.get(original)
|
||||
|
|
@ -114,4 +153,4 @@ class TurnCallIdRewriter:
|
|||
return item.model_copy(update={"call_id": replacement})
|
||||
|
||||
def rewrite_items(self, items: list[Any]) -> list[Any]:
|
||||
return [self.rewrite_item(item) for item in items]
|
||||
return [self.rewrite_item(item, position) for position, item in enumerate(items)]
|
||||
|
|
|
|||
|
|
@ -1,11 +1,12 @@
|
|||
"""Tests for tool-call id uniqueness.
|
||||
"""Tests for tool-call id presence and uniqueness.
|
||||
|
||||
Providers that number tool calls per turn (``exec_command:0``, ``:1``, ...)
|
||||
restart the counter on every turn, so the same id eventually appears twice in
|
||||
one conversation. Strict providers then reject the whole request, and because
|
||||
the history is replayed on every retry the agent can never recover. A gateway
|
||||
that validates id uniqueness the way those providers do proves both the
|
||||
failure and the fix.
|
||||
one conversation. Others hand back a tool call with no id at all, leaving the
|
||||
paired tool message with a blank ``tool_call_id``. Strict providers reject the
|
||||
whole request either way, and because the history is replayed on every retry
|
||||
the agent can never recover. Gateways that validate ids the way those
|
||||
providers do prove both failures and their fixes.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -125,10 +126,51 @@ class _StrictHandler(BaseHTTPRequestHandler):
|
|||
self.wfile.write(encoded)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def strict_gateway() -> Iterator[str]:
|
||||
class _BlankIdHandler(BaseHTTPRequestHandler):
|
||||
"""Gateway that hands out an id-less tool call and rejects blank ids, 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 index, message in enumerate(messages):
|
||||
if message.get("role") == "tool" and not message.get("tool_call_id"):
|
||||
self._respond(
|
||||
400,
|
||||
{
|
||||
"error": {
|
||||
"message": (
|
||||
f"messages[{index}]: tool messages must include "
|
||||
"a non-empty string tool_call_id"
|
||||
),
|
||||
"code": 400,
|
||||
}
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
if len(_REQUESTS) == 1:
|
||||
self._respond(200, _tool_call_completion(""))
|
||||
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), _StrictHandler)
|
||||
server = HTTPServer(("127.0.0.1", 0), handler)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
|
|
@ -138,6 +180,16 @@ def strict_gateway() -> Iterator[str]:
|
|||
server.server_close()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def strict_gateway() -> Iterator[str]:
|
||||
yield from _serve(_StrictHandler)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def blank_id_gateway() -> Iterator[str]:
|
||||
yield from _serve(_BlankIdHandler)
|
||||
|
||||
|
||||
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.
|
||||
|
|
@ -222,6 +274,61 @@ def test_history_dedupe_pairs_parallel_calls_by_order() -> None:
|
|||
assert rebuilt[1]["call_id"] != "dup"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blank_call_id_is_rejected_by_the_provider_without_the_wrapper(
|
||||
blank_id_gateway: str,
|
||||
) -> None:
|
||||
# Repro: the model answers with a tool call carrying no id, so the paired
|
||||
# tool result goes back as a ``tool`` message with a blank ``tool_call_id``
|
||||
# and the provider rejects the whole replayed history with a 400.
|
||||
with pytest.raises(Exception, match="non-empty string tool_call_id"):
|
||||
await _run_agent(blank_id_gateway, wrap=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blank_call_id_is_filled_in_so_the_history_stays_valid(
|
||||
blank_id_gateway: str,
|
||||
) -> None:
|
||||
result = await _run_agent(blank_id_gateway, wrap=True)
|
||||
|
||||
assert result.final_output == "all done"
|
||||
call_ids = _assistant_call_ids(_REQUESTS[-1])
|
||||
assert len(call_ids) == 1
|
||||
assert call_ids[0].startswith("call_")
|
||||
assert _tool_results(_REQUESTS[-1]) == ["did 1"]
|
||||
|
||||
|
||||
def test_history_dedupe_fills_in_blank_ids_and_keeps_outputs_paired() -> None:
|
||||
items = [
|
||||
{"type": "function_call", "call_id": "", "name": "a", "arguments": "{}"},
|
||||
{"type": "function_call", "call_id": "", "name": "b", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "", "output": "for-a"},
|
||||
{"type": "function_call_output", "call_id": "", "output": "for-b"},
|
||||
]
|
||||
|
||||
rebuilt, changed = dedupe_history_call_ids(items)
|
||||
|
||||
assert changed
|
||||
ids = [item["call_id"] for item in rebuilt]
|
||||
assert all(call_id.startswith("call_") for call_id in ids)
|
||||
assert ids[0] == ids[2]
|
||||
assert ids[1] == ids[3]
|
||||
assert ids[0] != ids[1]
|
||||
|
||||
|
||||
def test_history_dedupe_fills_in_a_missing_call_id_key() -> None:
|
||||
items = [
|
||||
{"type": "function_call", "name": "a", "arguments": "{}"},
|
||||
{"type": "function_call_output", "output": "x"},
|
||||
]
|
||||
|
||||
rebuilt, changed = dedupe_history_call_ids(items)
|
||||
|
||||
assert changed
|
||||
assert rebuilt[0]["call_id"] == rebuilt[1]["call_id"]
|
||||
assert rebuilt[0]["call_id"].startswith("call_")
|
||||
|
||||
|
||||
def test_history_dedupe_leaves_unique_ids_alone() -> None:
|
||||
items = [
|
||||
{"type": "function_call", "call_id": "call_a", "name": "a", "arguments": "{}"},
|
||||
|
|
@ -247,3 +354,29 @@ def test_turn_rewriter_is_stable_across_repeated_sightings() -> None:
|
|||
|
||||
assert first.call_id != "exec_command:0"
|
||||
assert second.call_id == first.call_id
|
||||
|
||||
|
||||
def test_turn_rewriter_fills_in_a_blank_id_stably() -> None:
|
||||
rewriter = TurnCallIdRewriter([])
|
||||
call = ResponseFunctionToolCall(
|
||||
id="fc_1", call_id="", name="a", arguments="{}", type="function_call"
|
||||
)
|
||||
|
||||
first = rewriter.rewrite_item(call)
|
||||
second = rewriter.rewrite_item(call)
|
||||
|
||||
assert first.call_id.startswith("call_")
|
||||
assert second.call_id == first.call_id
|
||||
assert rewriter.rewrite_item(first).call_id == first.call_id
|
||||
|
||||
|
||||
def test_turn_rewriter_gives_parallel_blank_calls_distinct_ids() -> None:
|
||||
rewriter = TurnCallIdRewriter([])
|
||||
a = ResponseFunctionToolCall(
|
||||
id="fc_1", call_id="", name="a", arguments="{}", type="function_call"
|
||||
)
|
||||
b = ResponseFunctionToolCall(
|
||||
id="fc_2", call_id="", name="b", arguments="{}", type="function_call"
|
||||
)
|
||||
|
||||
assert rewriter.rewrite_item(a).call_id != rewriter.rewrite_item(b).call_id
|
||||
|
|
|
|||
379
tests/test_tool_call_ids_providers.py
Normal file
379
tests/test_tool_call_ids_providers.py
Normal file
|
|
@ -0,0 +1,379 @@
|
|||
"""Tool-call id repair across provider routes, models, and streaming modes.
|
||||
|
||||
Every model Strix resolves goes through ``StrixProvider``, which picks the
|
||||
OpenAI SDK for ``openai/...`` and LiteLLM for every other prefix, then wraps the
|
||||
result in the turn guard that repairs tool-call ids. Each route parses a
|
||||
provider's tool calls on its own code path, so a blank or missing id — or a
|
||||
per-turn counter that repeats — has to be repaired no matter which route
|
||||
carried it, streamed or not.
|
||||
|
||||
A local OpenAI-compatible gateway stands in for the provider. It hands out
|
||||
tool calls with whatever ids the scenario calls for and rejects a history the
|
||||
way strict providers do: any ``tool`` message with an empty ``tool_call_id``,
|
||||
or two assistant tool calls sharing one id.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import threading
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from agents import Agent, ModelSettings, Runner, function_tool
|
||||
from agents.run import RunConfig
|
||||
|
||||
from strix.config import loader
|
||||
from strix.config import models as strix_models
|
||||
from strix.config.models import StrixProvider
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Iterator
|
||||
|
||||
|
||||
_MISSING = object()
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _Scenario:
|
||||
"""Tool-call ids the gateway hands out, one list per tool-calling turn."""
|
||||
|
||||
turns: tuple[tuple[Any, ...], ...]
|
||||
|
||||
@property
|
||||
def calls(self) -> int:
|
||||
return sum(len(turn) for turn in self.turns)
|
||||
|
||||
@property
|
||||
def has_null_id(self) -> bool:
|
||||
return any(
|
||||
call_id is None or call_id is _MISSING for turn in self.turns for call_id in turn
|
||||
)
|
||||
|
||||
|
||||
SCENARIOS = {
|
||||
"empty-id": _Scenario(turns=(("",),)),
|
||||
"null-id": _Scenario(turns=((None,),)),
|
||||
"missing-id": _Scenario(turns=((_MISSING,),)),
|
||||
"parallel-empty-ids": _Scenario(turns=(("", ""),)),
|
||||
"empty-id-every-turn": _Scenario(turns=(("",), ("",), ("",))),
|
||||
"recycled-counter": _Scenario(turns=(("exec_command:0",), ("exec_command:0",))),
|
||||
"parallel-recycled-counter": _Scenario(
|
||||
turns=(("exec_command:0", "exec_command:1"), ("exec_command:0", "exec_command:1"))
|
||||
),
|
||||
"mixed-blank-and-recycled": _Scenario(
|
||||
turns=(("exec_command:0", ""), ("exec_command:0", None), (_MISSING,))
|
||||
),
|
||||
"valid-ids": _Scenario(turns=(("call_a",), ("call_b", "call_c"))),
|
||||
}
|
||||
|
||||
# One model per provider family Strix recommends or documents, each on the
|
||||
# route ``StrixProvider`` gives it: the OpenAI SDK for ``openai/``, LiteLLM's
|
||||
# OpenAI-compatible adapters for the rest, and a generic OpenAI-compatible
|
||||
# endpoint through ``litellm/openai/``.
|
||||
MODELS = [
|
||||
"openai/gpt-5.4",
|
||||
"openrouter/z-ai/glm-5.3",
|
||||
"openrouter/anthropic/claude-sonnet-4.6",
|
||||
"openrouter/google/gemini-3-pro-preview",
|
||||
"openrouter/qwen/qwen3-coder",
|
||||
"zai/glm-5.3",
|
||||
"zai/glm-5.3-flash",
|
||||
"deepseek/deepseek-chat",
|
||||
"moonshot/kimi-k2.5",
|
||||
"xai/grok-4",
|
||||
"mistral/mistral-large-latest",
|
||||
"together_ai/Qwen/Qwen3-235B-A22B",
|
||||
"fireworks_ai/accounts/fireworks/models/kimi-k2",
|
||||
"dashscope/qwen3-max",
|
||||
"deepinfra/Qwen/Qwen3-32B",
|
||||
"nebius/Qwen/Qwen3-32B",
|
||||
"hosted_vllm/Qwen/Qwen3-32B",
|
||||
"litellm/openai/gw-model",
|
||||
]
|
||||
|
||||
|
||||
class _Gateway(BaseHTTPRequestHandler):
|
||||
scenario: ClassVar[_Scenario]
|
||||
requests: ClassVar[list[dict[str, Any]]]
|
||||
lock: ClassVar[threading.Lock]
|
||||
|
||||
def log_message(self, *args: Any) -> None:
|
||||
pass
|
||||
|
||||
def do_POST(self) -> None: # noqa: N802
|
||||
length = int(self.headers.get("Content-Length", 0))
|
||||
body = json.loads(self.rfile.read(length) or b"{}")
|
||||
messages = body.get("messages", [])
|
||||
with self.lock:
|
||||
self.requests.append(body)
|
||||
turn = len(self.requests) - 1
|
||||
|
||||
error = _history_error(messages)
|
||||
if error:
|
||||
self._send_json(400, {"error": {"message": error, "code": 400}})
|
||||
return
|
||||
|
||||
if turn < len(self.scenario.turns):
|
||||
ids = self.scenario.turns[turn]
|
||||
message = _tool_call_message(ids, first_n=turn * 10)
|
||||
finish = "tool_calls"
|
||||
else:
|
||||
message = {"role": "assistant", "content": "all done"}
|
||||
finish = "stop"
|
||||
|
||||
if body.get("stream"):
|
||||
self._send_stream(message, finish)
|
||||
else:
|
||||
self._send_json(
|
||||
200,
|
||||
{
|
||||
"id": f"chatcmpl-{turn}",
|
||||
"object": "chat.completion",
|
||||
"created": 0,
|
||||
"model": body.get("model", "gw-model"),
|
||||
"choices": [{"index": 0, "finish_reason": finish, "message": message}],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7},
|
||||
},
|
||||
)
|
||||
|
||||
def _send_json(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 _send_stream(self, message: dict[str, Any], finish: str) -> None:
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "text/event-stream")
|
||||
self.end_headers()
|
||||
|
||||
def chunk(delta: dict[str, Any], finish_reason: str | None = None) -> None:
|
||||
payload = {
|
||||
"id": "chatcmpl-s",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 0,
|
||||
"model": "gw-model",
|
||||
"choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}],
|
||||
}
|
||||
self.wfile.write(f"data: {json.dumps(payload)}\n\n".encode())
|
||||
|
||||
chunk({"role": "assistant", "content": ""})
|
||||
if message.get("content"):
|
||||
chunk({"content": message["content"]})
|
||||
for index, call in enumerate(message.get("tool_calls") or []):
|
||||
head: dict[str, Any] = {
|
||||
"index": index,
|
||||
"type": "function",
|
||||
"function": {"name": call["function"]["name"], "arguments": ""},
|
||||
}
|
||||
if "id" in call:
|
||||
head["id"] = call["id"]
|
||||
chunk({"tool_calls": [head]})
|
||||
arguments = call["function"]["arguments"]
|
||||
middle = len(arguments) // 2
|
||||
for part in (arguments[:middle], arguments[middle:]):
|
||||
chunk({"tool_calls": [{"index": index, "function": {"arguments": part}}]})
|
||||
chunk({}, finish)
|
||||
usage = {
|
||||
"id": "chatcmpl-s",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 0,
|
||||
"model": "gw-model",
|
||||
"choices": [],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7},
|
||||
}
|
||||
self.wfile.write(f"data: {json.dumps(usage)}\n\n".encode())
|
||||
self.wfile.write(b"data: [DONE]\n\n")
|
||||
self.wfile.flush()
|
||||
|
||||
|
||||
def _tool_call_message(ids: tuple[Any, ...], *, first_n: int) -> dict[str, Any]:
|
||||
calls = []
|
||||
for offset, call_id in enumerate(ids):
|
||||
call: dict[str, Any] = {
|
||||
"type": "function",
|
||||
"function": {"name": "do_thing", "arguments": json.dumps({"n": first_n + offset})},
|
||||
}
|
||||
if call_id is not _MISSING:
|
||||
call["id"] = call_id
|
||||
calls.append(call)
|
||||
return {"role": "assistant", "content": None, "tool_calls": calls}
|
||||
|
||||
|
||||
def _history_error(messages: list[dict[str, Any]]) -> str | None:
|
||||
seen: set[str] = set()
|
||||
for index, message in enumerate(messages):
|
||||
if message.get("role") == "tool":
|
||||
call_id = message.get("tool_call_id")
|
||||
if not isinstance(call_id, str) or not call_id:
|
||||
return (
|
||||
f"messages[{index}]: tool messages must include a non-empty string tool_call_id"
|
||||
)
|
||||
for call in message.get("tool_calls") or []:
|
||||
call_id = call.get("id")
|
||||
if not isinstance(call_id, str) or not call_id:
|
||||
return f"messages[{index}]: assistant tool_calls must include a non-empty id"
|
||||
if call_id in seen:
|
||||
return f"messages[{index}]: duplicate tool_call id {call_id!r}"
|
||||
seen.add(call_id)
|
||||
return None
|
||||
|
||||
|
||||
def _serve(scenario: _Scenario) -> tuple[HTTPServer, type[_Gateway]]:
|
||||
handler = type(
|
||||
"_ScenarioGateway",
|
||||
(_Gateway,),
|
||||
{"scenario": scenario, "requests": [], "lock": threading.Lock()},
|
||||
)
|
||||
server = HTTPServer(("127.0.0.1", 0), handler)
|
||||
threading.Thread(target=server.serve_forever, daemon=True).start()
|
||||
return server, handler
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_settings(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
for key in ("STRIX_LLM", "LLM_DISABLE_STREAMING", "LLM_MAX_TOOL_CALLS_PER_TURN"):
|
||||
monkeypatch.delenv(key, raising=False)
|
||||
monkeypatch.setattr(loader, "_cached", None)
|
||||
monkeypatch.setattr(loader, "_override", None)
|
||||
# As ``configure_sdk_model_defaults`` sets it, so routes like ``zai/`` that
|
||||
# reject ``parallel_tool_calls`` still get a request out.
|
||||
monkeypatch.setattr(litellm, "drop_params", True)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _gateway(scenario: _Scenario) -> Iterator[tuple[str, type[_Gateway]]]:
|
||||
server, handler = _serve(scenario)
|
||||
try:
|
||||
yield f"http://127.0.0.1:{server.server_address[1]}/v1", handler
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
|
||||
|
||||
async def _run(
|
||||
monkeypatch: pytest.MonkeyPatch, model: str, base_url: str, *, stream: bool, parallel: bool
|
||||
) -> tuple[Any, list[int]]:
|
||||
monkeypatch.setenv("LLM_DISABLE_STREAMING", "false" if stream else "true")
|
||||
monkeypatch.setattr(loader, "_cached", None)
|
||||
# Binding the provider to the gateway sends every route there, the way a
|
||||
# custom endpoint would, while each prefix keeps its own adapter.
|
||||
provider = StrixProvider(api_key="tok", base_url=base_url)
|
||||
ran: list[int] = []
|
||||
|
||||
@function_tool
|
||||
def do_thing(n: int) -> str:
|
||||
ran.append(n)
|
||||
return f"did {n}"
|
||||
|
||||
agent = Agent(
|
||||
name="t",
|
||||
instructions="use the tool",
|
||||
tools=[do_thing],
|
||||
model=model,
|
||||
model_settings=ModelSettings(parallel_tool_calls=parallel),
|
||||
)
|
||||
result = Runner.run_streamed(
|
||||
agent, input="please", max_turns=10, run_config=RunConfig(model_provider=provider)
|
||||
)
|
||||
async for _ in result.stream_events():
|
||||
pass
|
||||
return result, ran
|
||||
|
||||
|
||||
def _final_history(handler: type[_Gateway]) -> list[dict[str, Any]]:
|
||||
messages: list[dict[str, Any]] = handler.requests[-1]["messages"]
|
||||
return messages
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("stream", [True, False], ids=["streamed", "non-streamed"])
|
||||
@pytest.mark.parametrize("scenario_name", list(SCENARIOS))
|
||||
@pytest.mark.parametrize("model", MODELS)
|
||||
async def test_tool_call_ids_are_repaired_on_every_route(
|
||||
request: pytest.FixtureRequest,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
model: str,
|
||||
scenario_name: str,
|
||||
*,
|
||||
stream: bool,
|
||||
) -> None:
|
||||
scenario = SCENARIOS[scenario_name]
|
||||
if model.startswith("openai/") and not stream and scenario.has_null_id:
|
||||
request.applymarker(
|
||||
pytest.mark.xfail(
|
||||
strict=True,
|
||||
reason=(
|
||||
"the Agents SDK's Chat Completions converter rejects a null tool-call id "
|
||||
"in a non-streamed response before the turn guard sees it"
|
||||
),
|
||||
)
|
||||
)
|
||||
with _gateway(scenario) as (base_url, handler):
|
||||
result, ran = await _run(
|
||||
monkeypatch,
|
||||
model,
|
||||
base_url,
|
||||
stream=stream,
|
||||
parallel=any(len(t) > 1 for t in scenario.turns),
|
||||
)
|
||||
|
||||
assert result.final_output == "all done"
|
||||
assert len(handler.requests) == len(scenario.turns) + 1
|
||||
assert all(bool(body.get("stream")) is stream for body in handler.requests)
|
||||
|
||||
history = _final_history(handler)
|
||||
call_ids = [call["id"] for m in history for call in m.get("tool_calls") or []]
|
||||
assert len(call_ids) == scenario.calls
|
||||
assert all(isinstance(call_id, str) and call_id for call_id in call_ids)
|
||||
assert len(set(call_ids)) == len(call_ids)
|
||||
|
||||
outputs = {m["tool_call_id"]: m["content"] for m in history if m.get("role") == "tool"}
|
||||
assert set(outputs) == set(call_ids)
|
||||
|
||||
originals = {
|
||||
turn * 10 + offset: call_id
|
||||
for turn, ids in enumerate(scenario.turns)
|
||||
for offset, call_id in enumerate(ids)
|
||||
}
|
||||
assert sorted(ran) == sorted(originals)
|
||||
|
||||
# Each output still answers the call that produced it.
|
||||
id_by_n = {
|
||||
json.loads(call["function"]["arguments"])["n"]: call["id"]
|
||||
for m in history
|
||||
for call in m.get("tool_calls") or []
|
||||
}
|
||||
assert {id_by_n[n]: f"did {n}" for n in originals} == outputs
|
||||
|
||||
# A usable id is kept the first time it appears; only blanks and repeats change.
|
||||
seen: set[str] = set()
|
||||
for n, original in sorted(originals.items()):
|
||||
if isinstance(original, str) and original and original not in seen:
|
||||
assert id_by_n[n] == original
|
||||
seen.add(original)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("stream", [True, False], ids=["streamed", "non-streamed"])
|
||||
@pytest.mark.parametrize("model", MODELS)
|
||||
async def test_blank_call_id_is_rejected_on_every_route_without_the_guard(
|
||||
monkeypatch: pytest.MonkeyPatch, model: str, *, stream: bool
|
||||
) -> None:
|
||||
# Repro: with the turn guard removed, an empty id reaches the provider on
|
||||
# every route, so the matrix above is exercising the repair and not a
|
||||
# route that happens to fill ids in on its own.
|
||||
monkeypatch.setattr(strix_models, "_TurnGuardModel", lambda model, **_: model)
|
||||
with (
|
||||
_gateway(SCENARIOS["empty-id"]) as (base_url, _handler),
|
||||
pytest.raises(Exception, match="must include a non-empty"),
|
||||
):
|
||||
await _run(monkeypatch, model, base_url, stream=stream, parallel=False)
|
||||
Loading…
Add table
Reference in a new issue