Files
strix/tests/test_tool_call_limits.py
T
8bd6c8e87a fix(llm): cap the tool calls one assistant response may queue (#977)
* fix(llm): cap the tool calls one assistant response may queue

* fix(llm): cap the subscription backend's responses too

---------

Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-08-06 00:06:59 +03:00

190 lines
5.9 KiB
Python

"""Tests for the per-response tool-call cap.
A degenerate generation can emit hundreds of tool calls in one assistant
response — a wait/poll loop the model writes out ahead of time. The run loop
honours every one of them, so the agent stops reacting for hours. The cap
keeps the first N calls of a response and drops the tail.
"""
from __future__ import annotations
import json
import threading
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import TYPE_CHECKING, Any
import pytest
from agents import Agent, Runner, function_tool
from agents.models.interface import Model, ModelProvider
from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel
from agents.run import RunConfig
from openai import AsyncOpenAI
from strix.config import loader
from strix.config.loader import load_settings
from strix.config.models import StrixProvider, _NonStreamingModel, _TurnGuardModel
if TYPE_CHECKING:
from collections.abc import Iterator
_RUNAWAY_CALLS = 200
_CAP = 32
def _runaway_completion() -> dict[str, Any]:
return {
"id": "chatcmpl-1",
"object": "chat.completion",
"created": 0,
"model": "gw-model",
"choices": [
{
"index": 0,
"finish_reason": "tool_calls",
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": f"call_{i}",
"type": "function",
"function": {"name": "wait_for_message", "arguments": "{}"},
}
for i in range(_RUNAWAY_CALLS)
],
},
}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7},
}
def _text_completion() -> dict[str, Any]:
return {
"id": "chatcmpl-2",
"object": "chat.completion",
"created": 0,
"model": "gw-model",
"choices": [
{
"index": 0,
"finish_reason": "stop",
"message": {"role": "assistant", "content": "done"},
}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8},
}
_TURNS: list[int] = []
class _RunawayHandler(BaseHTTPRequestHandler):
"""First turn queues a huge poll loop; the next turn ends the run."""
def log_message(self, *args: Any) -> None:
pass
def do_POST(self) -> None:
length = int(self.headers.get("Content-Length", 0))
self.rfile.read(length)
_TURNS.append(1)
payload = _runaway_completion() if len(_TURNS) == 1 else _text_completion()
encoded = json.dumps(payload).encode()
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(encoded)))
self.end_headers()
self.wfile.write(encoded)
@pytest.fixture
def runaway_gateway() -> Iterator[str]:
_TURNS.clear()
server = HTTPServer(("127.0.0.1", 0), _RunawayHandler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_address[1]}/v1"
finally:
server.shutdown()
server.server_close()
def _model(base_url: str) -> Model:
client = AsyncOpenAI(api_key="tok", base_url=base_url, max_retries=0)
return _NonStreamingModel(OpenAIChatCompletionsModel(model="gw-model", openai_client=client))
async def _run_agent(base_url: str, *, cap: int) -> list[int]:
executed: list[int] = []
@function_tool
def wait_for_message() -> str:
executed.append(1)
return "nothing new"
class _Provider(ModelProvider):
def get_model(self, model_name: str | None) -> Model: # noqa: ARG002
return _TurnGuardModel(_model(base_url), max_tool_calls_per_turn=cap)
agent = Agent(name="t", instructions="orchestrate", tools=[wait_for_message], model="gw-model")
result = Runner.run_streamed(
agent, input="go", run_config=RunConfig(model_provider=_Provider())
)
async for _ in result.stream_events():
pass
assert result.final_output == "done"
return executed
@pytest.mark.asyncio
async def test_runaway_response_runs_every_queued_call_when_uncapped(runaway_gateway: str) -> None:
# Repro: one response queues 200 calls and the run loop honours all of them.
executed = await _run_agent(runaway_gateway, cap=0)
assert len(executed) == _RUNAWAY_CALLS
@pytest.mark.asyncio
async def test_runaway_response_is_capped(runaway_gateway: str) -> None:
executed = await _run_agent(runaway_gateway, cap=_CAP)
assert len(executed) == _CAP
@pytest.mark.asyncio
async def test_response_below_the_cap_is_untouched(runaway_gateway: str) -> None:
executed = await _run_agent(runaway_gateway, cap=_RUNAWAY_CALLS + 1)
assert len(executed) == _RUNAWAY_CALLS
@pytest.fixture
def _reset_settings(monkeypatch: pytest.MonkeyPatch) -> Iterator[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)
yield
class _DummyModel(Model):
async def get_response(self, *args: Any, **kwargs: Any) -> Any:
raise NotImplementedError
def stream_response(self, *args: Any, **kwargs: Any) -> Any:
raise NotImplementedError
def test_cap_is_configurable(monkeypatch: pytest.MonkeyPatch, _reset_settings: None) -> None:
monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: _DummyModel())
monkeypatch.setenv("LLM_MAX_TOOL_CALLS_PER_TURN", "7")
load_settings()
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert isinstance(model, _TurnGuardModel)
assert model._max_tool_calls_per_turn == 7