From 4695290682900eeedf4819a0336976d56dc028e0 Mon Sep 17 00:00:00 2001 From: ian-at-strix Date: Wed, 7 Oct 2026 21:35:45 -0400 Subject: [PATCH] feat(llm): bound the first stream event and the whole stream (#1487) --- strix/config/models.py | 88 +++++++++++++++++++++++++------ strix/config/settings.py | 4 ++ strix/llm/request_log.py | 18 ++++++- tests/test_stream_idle_timeout.py | 77 +++++++++++++++++++++++---- 4 files changed, 159 insertions(+), 28 deletions(-) diff --git a/strix/config/models.py b/strix/config/models.py index 09ff493e..f36c1d9e 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -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, ) diff --git a/strix/config/settings.py b/strix/config/settings.py index 54243b74..5aa6da06 100644 --- a/strix/config/settings.py +++ b/strix/config/settings.py @@ -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, diff --git a/strix/llm/request_log.py b/strix/llm/request_log.py index 8e55a736..eda0d222 100644 --- a/strix/llm/request_log.py +++ b/strix/llm/request_log.py @@ -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 diff --git a/tests/test_stream_idle_timeout.py b/tests/test_stream_idle_timeout.py index 9d618978..f45035bf 100644 --- a/tests/test_stream_idle_timeout.py +++ b/tests/test_stream_idle_timeout.py @@ -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