Compare commits

...
4 changed files with 229 additions and 3 deletions
+1
View File
@@ -238,6 +238,7 @@ ignore = [
"tests/test_disable_streaming.py" = ["N802"]
"tests/test_tool_call_ids.py" = ["N802"]
"tests/test_tool_call_limits.py" = ["N802", "SLF001"]
"tests/test_stream_idle_timeout.py" = ["N802", "SLF001"]
"tests/test_unknown_tool_recovery.py" = ["N802"]
"tests/test_report_pdf.py" = ["S105", "S106"]
# Stdlib HTTP handler overrides (do_GET/do_POST) and lazy imports that avoid a
+54 -3
View File
@@ -2,11 +2,13 @@
from __future__ import annotations
import asyncio
import contextlib
import inspect
import logging
import os
import time
from collections.abc import AsyncGenerator
from typing import TYPE_CHECKING, Any, cast
from agents import (
@@ -252,11 +254,23 @@ class _TurnGuardModel(Model):
Tool-call volume: a degenerate response can queue hundreds of calls that
the run loop then honours one by one. Only the first
``LLM_MAX_TOOL_CALLS_PER_TURN`` calls of a response are kept.
Stalled streams: a turn that emits a few tokens and then goes silent is
not covered by the request timeout, which resets on any byte (keepalives
included). ``LLM_STREAM_IDLE_TIMEOUT`` bounds the gap between events so the
turn fails instead of hanging, and the existing retry path replays it.
"""
def __init__(self, inner: Model, *, max_tool_calls_per_turn: int = 0) -> None:
def __init__(
self,
inner: Model,
*,
max_tool_calls_per_turn: int = 0,
stream_idle_timeout: float = 0.0,
) -> None:
self._inner = inner
self._max_tool_calls_per_turn = max_tool_calls_per_turn
self._stream_idle_timeout = stream_idle_timeout
def _limiter(self) -> TurnToolCallLimiter:
return TurnToolCallLimiter(self._max_tool_calls_per_turn)
@@ -337,13 +351,41 @@ class _TurnGuardModel(Model):
conversation_id=conversation_id,
prompt=prompt,
)
async for event in stream:
async for event in _with_idle_timeout(stream, self._stream_idle_timeout):
guarded = _guard_event(event, rewriter, limiter)
if guarded is not None:
yield guarded
self._log_dropped(limiter)
async def _aclose(stream: AsyncIterator[TResponseStreamEvent]) -> None:
if isinstance(stream, AsyncGenerator):
with contextlib.suppress(Exception):
await stream.aclose()
async def _with_idle_timeout(
stream: AsyncIterator[TResponseStreamEvent], timeout: float
) -> AsyncIterator[TResponseStreamEvent]:
if timeout <= 0:
async for event in stream:
yield event
return
iterator = stream.__aiter__()
while True:
try:
event = await asyncio.wait_for(iterator.__anext__(), timeout)
except StopAsyncIteration:
return
except TimeoutError:
await _aclose(stream)
message = f"model stream produced no event for {timeout:.0f}s"
logger.warning("%s; abandoning the turn", message)
raise TimeoutError(message) from None
yield event
def _guard_event(
event: TResponseStreamEvent, rewriter: TurnCallIdRewriter, limiter: TurnToolCallLimiter
) -> TResponseStreamEvent | None:
@@ -429,6 +471,7 @@ class StrixProvider(MultiProvider):
def get_model(self, model_name: str | None) -> Model:
llm = load_settings().llm
slug = codex.subscription_model(model_name)
idle_timeout = float(llm.stream_idle_timeout)
if slug:
# The ChatGPT subscription backend is always streamed; it has no
# non-streaming mode to fall back to, so LLM_DISABLE_STREAMING
@@ -442,7 +485,15 @@ class StrixProvider(MultiProvider):
model = super().get_model(model_name)
if llm.disable_streaming:
model = _NonStreamingModel(model)
return _TurnGuardModel(model, max_tool_calls_per_turn=llm.max_tool_calls_per_turn)
# The wrapper emits its single event only once the whole request
# is done, so an idle gap is meaningless here; the request
# timeout bounds it instead.
idle_timeout = 0.0
return _TurnGuardModel(
model,
max_tool_calls_per_turn=llm.max_tool_calls_per_turn,
stream_idle_timeout=idle_timeout,
)
DEFAULT_MODEL_RETRY = ModelRetrySettings(
+1
View File
@@ -57,6 +57,7 @@ class LlmSettings(BaseSettings):
alias="LLM_DISABLE_STREAMING",
)
timeout: int = Field(default=300, alias="LLM_TIMEOUT")
stream_idle_timeout: int = Field(default=300, ge=0, alias="LLM_STREAM_IDLE_TIMEOUT")
max_tool_calls_per_turn: int = Field(
default=32,
ge=0,
+173
View File
@@ -0,0 +1,173 @@
"""Tests for the model-stream idle watchdog.
A turn that streams a few tokens and then goes silent is not covered by the
request timeout: the read timeout resets on every byte, keepalives included.
The watchdog bounds the gap between events so the turn fails and can be
retried instead of parking the agent forever.
"""
from __future__ import annotations
import asyncio
import json
import threading
import time
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import TYPE_CHECKING, Any
import pytest
from agents.model_settings import ModelSettings
from agents.models.interface import Model, ModelTracing
from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel
from openai import AsyncOpenAI
from strix.config import loader
from strix.config.loader import load_settings
from strix.config.models import StrixProvider, _TurnGuardModel, _with_idle_timeout
if TYPE_CHECKING:
from collections.abc import AsyncIterator, Iterator
_STALL_SECONDS = 30.0
def _chunk(text: str) -> bytes:
payload = {
"id": "chatcmpl-1",
"object": "chat.completion.chunk",
"created": 0,
"model": "gw-model",
"choices": [{"index": 0, "delta": {"content": text}, "finish_reason": None}],
}
return b"data: " + json.dumps(payload).encode() + b"\n\n"
class _StallingHandler(BaseHTTPRequestHandler):
"""Streams a couple of tokens, then stops producing anything."""
stop = threading.Event()
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)
self.send_response(200)
self.send_header("Content-Type", "text/event-stream")
self.end_headers()
self.wfile.write(_chunk("Now"))
self.wfile.write(_chunk(" spawning"))
self.wfile.flush()
self.stop.wait(_STALL_SECONDS)
@pytest.fixture
def stalling_gateway() -> Iterator[str]:
_StallingHandler.stop.clear()
server = HTTPServer(("127.0.0.1", 0), _StallingHandler)
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:
_StallingHandler.stop.set()
server.shutdown()
server.server_close()
def _stream(base_url: str, *, idle_timeout: float) -> AsyncIterator[Any]:
client = AsyncOpenAI(api_key="tok", base_url=base_url, max_retries=0, timeout=_STALL_SECONDS)
inner: Model = OpenAIChatCompletionsModel(model="gw-model", openai_client=client)
guarded = _TurnGuardModel(inner, stream_idle_timeout=idle_timeout)
return guarded.stream_response(
None,
"go",
ModelSettings(),
[],
None,
[],
ModelTracing.DISABLED,
previous_response_id=None,
conversation_id=None,
prompt=None,
)
async def _drain(base_url: str, *, idle_timeout: float) -> list[Any]:
return [event async for event in _stream(base_url, idle_timeout=idle_timeout)]
@pytest.mark.asyncio
async def test_stalled_stream_hangs_without_the_watchdog(stalling_gateway: str) -> None:
# Repro: tokens arrive, then nothing. Un-watched, the turn just sits there;
# the request timeout is far away and would reset on any keepalive byte.
with pytest.raises(TimeoutError):
await asyncio.wait_for(_drain(stalling_gateway, idle_timeout=0), timeout=2)
@pytest.mark.asyncio
async def test_stalled_stream_is_abandoned_by_the_watchdog(stalling_gateway: str) -> None:
started = time.monotonic()
with pytest.raises(TimeoutError, match="produced no event"):
await _drain(stalling_gateway, idle_timeout=1)
assert time.monotonic() - started < _STALL_SECONDS
@pytest.mark.asyncio
async def test_events_keep_flowing_while_the_stream_is_alive() -> None:
async def _live() -> AsyncIterator[Any]:
for i in range(5):
await asyncio.sleep(0.05)
yield f"event-{i}"
seen: list[Any] = [event async for event in _with_idle_timeout(_live(), 1.0)]
assert seen == [f"event-{i}" for i in range(5)]
@pytest.fixture
def _reset_settings(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
for key in ("STRIX_LLM", "LLM_DISABLE_STREAMING", "LLM_STREAM_IDLE_TIMEOUT"):
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_idle_timeout_is_configurable(
monkeypatch: pytest.MonkeyPatch, _reset_settings: None
) -> None:
monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: _DummyModel())
monkeypatch.setenv("LLM_STREAM_IDLE_TIMEOUT", "45")
load_settings()
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert isinstance(model, _TurnGuardModel)
assert model._stream_idle_timeout == 45
def test_idle_timeout_is_off_without_streaming(
monkeypatch: pytest.MonkeyPatch, _reset_settings: None
) -> None:
# LLM_DISABLE_STREAMING turns the whole request into one event, so an idle
# gap would just be the request duration — the request timeout bounds that.
monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: _DummyModel())
monkeypatch.setenv("LLM_STREAM_IDLE_TIMEOUT", "45")
monkeypatch.setenv("LLM_DISABLE_STREAMING", "true")
load_settings()
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert isinstance(model, _TurnGuardModel)
assert model._stream_idle_timeout == 0