mirror of
https://github.com/usestrix/strix.git
synced 2026-08-16 09:26:39 +02:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
60b1d72642 |
@@ -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
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user