feat(llm): bound the first stream event and the whole stream (#1487)
Some checks are pending
CI / python (push) Waiting to run
CI / tui (push) Waiting to run
CI / viewer (push) Waiting to run
CI / package (push) Waiting to run
CI / workflows (push) Waiting to run
CI / ci-passed (push) Blocked by required conditions

This commit is contained in:
ian-at-strix 2026-10-07 21:35:45 -04:00 • committed by GitHub
parent 928f053610
commit 4695290682
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 159 additions and 28 deletions

View file

@ -263,6 +263,8 @@ class _TurnGuardModel(Model):
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.
``LLM_STREAM_FIRST_EVENT_TIMEOUT`` and ``LLM_STREAM_TOTAL_TIMEOUT`` bound
the first event and the whole stream.
"""
def __init__(
@ -270,11 +272,15 @@ class _TurnGuardModel(Model):
inner: Model,
*,
max_tool_calls_per_turn: int = 0,
stream_idle_timeout: float = 0.0,
stream_idle_timeout: float | None = None,
first_event_timeout: float | None = None,
total_timeout: float | None = None,
) -> None:
self._inner = inner
self._max_tool_calls_per_turn = max_tool_calls_per_turn
self._stream_idle_timeout = stream_idle_timeout
self._first_event_timeout = first_event_timeout
self._total_timeout = total_timeout
def _limiter(self) -> TurnToolCallLimiter:
return TurnToolCallLimiter(self._max_tool_calls_per_turn)
@ -361,7 +367,12 @@ class _TurnGuardModel(Model):
conversation_id=conversation_id,
prompt=prompt,
)
async for event in _with_idle_timeout(stream, self._stream_idle_timeout):
async for event in _with_timeouts(
stream,
idle=self._stream_idle_timeout,
first_event=self._first_event_timeout,
total=self._total_timeout,
):
guarded = _guard_event(event, rewriter, limiter)
if guarded is not None:
yield guarded
@ -395,28 +406,68 @@ async def _aclose(stream: AsyncIterator[TResponseStreamEvent]) -> None:
await stream.aclose()
async def _with_idle_timeout(
stream: AsyncIterator[TResponseStreamEvent], timeout: float
async def _with_timeouts(
stream: AsyncIterator[TResponseStreamEvent],
*,
idle: float | None = None,
first_event: float | None = None,
total: float | None = None,
) -> AsyncIterator[TResponseStreamEvent]:
if timeout <= 0:
async for event in stream:
yield event
return
"""Bound the first event, the gap between events and the whole stream; None is no bound."""
iterator = stream.__aiter__()
started = time.monotonic()
while True:
limits: list[tuple[float, float, str]] = []
# Wait until the soonest limit expires: first-event (else idle) until the
# first event arrives, then idle; plus whatever is left of the total.
if first_event is not None:
limits.append((first_event, first_event, "stream_first_event_timeout"))
elif idle is not None:
limits.append((idle, idle, "stream_idle_timeout"))
if total is not None:
left = max(0.0, started + total - time.monotonic())
limits.append((left, total, "stream_total_timeout"))
try:
event = await asyncio.wait_for(iterator.__anext__(), timeout)
if limits:
event = await _next_event(iterator, *min(limits))
else:
event = await iterator.__anext__()
except StopAsyncIteration:
return
except TimeoutError:
except TimeoutError as exc:
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
logger.warning("%s; abandoning the turn", exc)
raise
first_event = None
yield event
async def _next_event(
iterator: AsyncIterator[TResponseStreamEvent], wait: float, limit: float, reason: str
) -> TResponseStreamEvent:
# asyncio.timeout() cancels without a message, so the request log could not
# tell this from a shutdown; this is asyncio.timeout() with a message.
task = asyncio.current_task()
assert task is not None
cancelling = task.cancelling()
expired = False
def expire() -> None:
nonlocal expired
expired = True
task.cancel(msg=f"{request_log.CANCEL_REASON_PREFIX}{reason}")
handle = asyncio.get_running_loop().call_later(wait, expire)
try:
return await iterator.__anext__()
except BaseException as exc:
if expired and task.uncancel() <= cancelling and isinstance(exc, asyncio.CancelledError):
raise TimeoutError(f"model stream hit {reason} ({limit:.0f}s)") from None
raise
finally:
handle.cancel()
def _guard_event(
event: TResponseStreamEvent, rewriter: TurnCallIdRewriter, limiter: TurnToolCallLimiter
) -> TResponseStreamEvent | None:
@ -551,7 +602,10 @@ 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)
# The settings use 0 for no bound.
idle_timeout = float(llm.stream_idle_timeout) or None
first_event_timeout = float(llm.stream_first_event_timeout) or None
total_timeout = float(llm.stream_total_timeout) or None
if slug:
# The ChatGPT subscription backend is always streamed; it has no
# non-streaming mode to fall back to, so LLM_DISABLE_STREAMING
@ -592,11 +646,13 @@ class StrixProvider(MultiProvider):
# 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
idle_timeout = first_event_timeout = total_timeout = None
return _TurnGuardModel(
model,
max_tool_calls_per_turn=llm.max_tool_calls_per_turn,
stream_idle_timeout=idle_timeout,
first_event_timeout=first_event_timeout,
total_timeout=total_timeout,
)

View file

@ -79,6 +79,10 @@ class LlmSettings(BaseSettings):
timeout: int = Field(default=300, alias="LLM_TIMEOUT")
preflight_timeout: int = Field(default=30, ge=1, alias="LLM_PREFLIGHT_TIMEOUT")
stream_idle_timeout: int = Field(default=300, ge=0, alias="LLM_STREAM_IDLE_TIMEOUT")
stream_first_event_timeout: int = Field(
default=120, ge=0, alias="LLM_STREAM_FIRST_EVENT_TIMEOUT"
)
stream_total_timeout: int = Field(default=600, ge=0, alias="LLM_STREAM_TOTAL_TIMEOUT")
max_tool_calls_per_turn: int = Field(
default=32,
ge=0,

View file

@ -1110,7 +1110,7 @@ class RequestLoggingModel(Model):
outcome="error",
status_code=status,
provider_request_id=request_id or base.provider_request_id,
error_type=type(exc).__name__,
error_type=_cancel_reason(exc) or type(exc).__name__,
error_message=_abandonment_message(exc) or clean_error_message(exc),
response_bytes=_exception_body_size(exc, status),
response_headers=headers_from_response(error_headers) or base.response_headers,
@ -1347,9 +1347,23 @@ def _is_abandonment(exc: BaseException | None) -> bool:
return isinstance(exc, asyncio.CancelledError | GeneratorExit)
# Prefixes the reason a stream timeout gives when it cancels an attempt.
CANCEL_REASON_PREFIX = "strix:"
def _cancel_reason(exc: BaseException) -> str | None:
"""The stream timeout that cancelled the attempt, e.g. ``stream_idle_timeout``."""
if isinstance(exc, asyncio.CancelledError) and exc.args:
message = str(exc.args[0])
if message.startswith(CANCEL_REASON_PREFIX):
return message.removeprefix(CANCEL_REASON_PREFIX)
return None
def _abandonment_message(exc: BaseException) -> str | None:
if isinstance(exc, asyncio.CancelledError):
return "attempt cancelled before the reply was consumed (stream idle timeout or shutdown)"
reason = f": {exc.args[0]}" if exc.args else ""
return f"attempt cancelled before the reply was consumed{reason}"
if isinstance(exc, GeneratorExit):
return "stream closed by the caller before it finished"
return None

View file

@ -23,7 +23,8 @@ 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
from strix.config.models import StrixProvider, _TurnGuardModel, _with_timeouts
from strix.llm import request_log
if TYPE_CHECKING:
@ -78,9 +79,21 @@ def stalling_gateway() -> Iterator[str]:
server.server_close()
def _stream(base_url: str, *, idle_timeout: float) -> AsyncIterator[Any]:
class _ServedByRelace(OpenAIChatCompletionsModel):
async def stream_response(self, *args: Any, **kwargs: Any) -> AsyncIterator[Any]:
request_log.record_upstream_provider("Relace")
async for event in super().stream_response(*args, **kwargs):
yield event
def _stream(base_url: str, *, idle_timeout: float | None) -> 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)
inner: Model = request_log.RequestLoggingModel(
_ServedByRelace(model="gw-model", openai_client=client),
model_name="gw-model",
provider="openai",
base_url=base_url,
)
guarded = _TurnGuardModel(inner, stream_idle_timeout=idle_timeout)
return guarded.stream_response(
None,
@ -96,7 +109,7 @@ def _stream(base_url: str, *, idle_timeout: float) -> AsyncIterator[Any]:
)
async def _drain(base_url: str, *, idle_timeout: float) -> list[Any]:
async def _drain(base_url: str, *, idle_timeout: float | None) -> list[Any]:
return [event async for event in _stream(base_url, idle_timeout=idle_timeout)]
@ -105,16 +118,31 @@ async def test_stalled_stream_hangs_without_the_watchdog(stalling_gateway: str)
# 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)
await asyncio.wait_for(_drain(stalling_gateway, idle_timeout=None), timeout=2)
@pytest.mark.asyncio
async def test_stalled_stream_is_abandoned_by_the_watchdog(stalling_gateway: str) -> None:
logged: list[request_log.LlmRequestEvent] = []
request_log.register_sink(logged.append)
started = time.monotonic()
with pytest.raises(TimeoutError, match="produced no event"):
await _drain(stalling_gateway, idle_timeout=1)
try:
with pytest.raises(TimeoutError, match="stream_idle_timeout"):
await _drain(stalling_gateway, idle_timeout=1)
finally:
request_log.unregister_sink(logged.append)
assert time.monotonic() - started < _STALL_SECONDS
# The request log can tell our timeout from a shutdown.
assert [
(e.error_type, (e.details or {}).get("upstream_provider"), e.error_message) for e in logged
] == [
(
"stream_idle_timeout",
"Relace",
"attempt cancelled before the reply was consumed: strix:stream_idle_timeout",
)
]
@pytest.mark.asyncio
@ -124,14 +152,39 @@ async def test_events_keep_flowing_while_the_stream_is_alive() -> None:
await asyncio.sleep(0.05)
yield f"event-{i}"
seen: list[Any] = [event async for event in _with_idle_timeout(_live(), 1.0)]
seen: list[Any] = [event async for event in _with_timeouts(_live(), idle=1.0)]
assert seen == [f"event-{i}" for i in range(5)]
@pytest.mark.asyncio
async def test_first_event_and_whole_stream_are_bounded() -> None:
async def _slow_start() -> AsyncIterator[Any]:
await asyncio.sleep(_STALL_SECONDS)
yield "late"
async def _endless() -> AsyncIterator[Any]:
while True:
await asyncio.sleep(0.05)
yield "event"
with pytest.raises(TimeoutError, match="stream_first_event_timeout"):
[event async for event in _with_timeouts(_slow_start(), idle=5, first_event=0.2)]
with pytest.raises(TimeoutError, match="stream_idle_timeout"):
[event async for event in _with_timeouts(_slow_start(), idle=0.2)]
with pytest.raises(TimeoutError, match="stream_total_timeout"):
[event async for event in _with_timeouts(_endless(), idle=5, total=0.3)]
@pytest.fixture
def _reset_settings(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
for key in ("STRIX_LLM", "LLM_DISABLE_STREAMING", "LLM_STREAM_IDLE_TIMEOUT"):
for key in (
"STRIX_LLM",
"LLM_DISABLE_STREAMING",
"LLM_STREAM_IDLE_TIMEOUT",
"LLM_STREAM_FIRST_EVENT_TIMEOUT",
"LLM_STREAM_TOTAL_TIMEOUT",
):
monkeypatch.delenv(key, raising=False)
monkeypatch.setattr(loader, "_cached", None)
monkeypatch.setattr(loader, "_override", None)
@ -151,11 +204,15 @@ def test_idle_timeout_is_configurable(
) -> None:
monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: _DummyModel())
monkeypatch.setenv("LLM_STREAM_IDLE_TIMEOUT", "45")
monkeypatch.setenv("LLM_STREAM_FIRST_EVENT_TIMEOUT", "30")
monkeypatch.setenv("LLM_STREAM_TOTAL_TIMEOUT", "400")
load_settings()
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert isinstance(model, _TurnGuardModel)
assert model._stream_idle_timeout == 45
assert model._first_event_timeout == 30
assert model._total_timeout == 400
def test_idle_timeout_is_off_without_streaming(
@ -170,4 +227,4 @@ def test_idle_timeout_is_off_without_streaming(
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert isinstance(model, _TurnGuardModel)
assert model._stream_idle_timeout == 0
assert model._stream_idle_timeout is model._first_event_timeout is model._total_timeout is None