mirror of
https://github.com/usestrix/strix.git
synced 2026-09-30 01:52:18 +00:00
379 lines
14 KiB
Python
379 lines
14 KiB
Python
"""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)
|