mirror of
https://github.com/usestrix/strix.git
synced 2026-10-08 03:08:08 +00:00
feat(llm): bound the first stream event and the whole stream (#1487)
This commit is contained in:
parent
928f053610
commit
4695290682
4 changed files with 159 additions and 28 deletions
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue