Merge pull request #776 from Routstr/feat/upstream-first-token-timeout

feat: upstream first-token timeout, stream idle timeout and provider cooldown
This commit is contained in:
9qeklajc
2026-10-01 13:26:25 +02:00
committed by GitHub
10 changed files with 865 additions and 17 deletions
+6
View File
@@ -72,6 +72,12 @@ ROUTSTR_SECRET_KEY=
# UPSTREAM_POOL_TIMEOUT=5
# UPSTREAM_READ_TIMEOUT=900
# Upstream Streaming Guards (0 disables; keep above reasoning models' think time)
# UPSTREAM_FIRST_TOKEN_TIMEOUT_SECONDS=0
# UPSTREAM_STREAM_IDLE_TIMEOUT_SECONDS=0
# UPSTREAM_ALLOWED_FAILS=3
# UPSTREAM_COOLDOWN_SECONDS=30
# Request and reservation lifetime limits (seconds)
# STALE_RESERVATION_TIMEOUT_SECONDS=300
# MAX_REQUEST_LIFETIME_SECONDS=1800
+16
View File
@@ -40,6 +40,22 @@ class Settings(BaseSettings):
upstream_5xx_retry_attempts: int = Field(
default=1, ge=0, env="UPSTREAM_5XX_RETRY_ATTEMPTS"
)
# Streaming guards, off by default (0). A stream that never produces a
# first chunk can still fail over; one that stalls later can only be
# aborted and billed for what it delivered. Reasoning models can stay
# silent for minutes, so set these above the longest expected think time.
upstream_first_token_timeout_seconds: float = Field(
default=0.0, ge=0, env="UPSTREAM_FIRST_TOKEN_TIMEOUT_SECONDS"
)
upstream_stream_idle_timeout_seconds: float = Field(
default=0.0, ge=0, env="UPSTREAM_STREAM_IDLE_TIMEOUT_SECONDS"
)
# Circuit breaker: timeouts/5xx per (provider, model) within a minute that
# take the pair out of candidate selection. 0 seconds disables it.
upstream_allowed_fails: int = Field(default=3, ge=1, env="UPSTREAM_ALLOWED_FAILS")
upstream_cooldown_seconds: float = Field(
default=30.0, ge=0, env="UPSTREAM_COOLDOWN_SECONDS"
)
# Node info
name: str = Field(default="ARoutstrNode", env="NAME")
+62 -5
View File
@@ -1,6 +1,7 @@
import asyncio
import inspect
import json
import re
from typing import Any
from fastapi import APIRouter, HTTPException, Request
@@ -23,6 +24,7 @@ from .core.db import (
create_session,
)
from .core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
UPSTREAM_UNAVAILABLE,
@@ -40,6 +42,12 @@ from .payment.helpers import (
)
from .payment.models import Model
from .upstream import BaseUpstreamProvider
from .upstream.cooldown import (
candidate_model_identity,
is_cooling_down,
provider_identity,
record_failure,
)
from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request
from .upstream.helpers import init_upstreams
from .upstream.model_paths import (
@@ -116,8 +124,6 @@ def get_candidates(
if candidates := _provider_map.get(model_id_lower):
return candidates
import re
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
if base_model_id != model_id_lower:
if candidates := _provider_map.get(base_model_id):
@@ -405,6 +411,18 @@ _RETRYABLE_UPSTREAM_5XX = frozenset({502, 503, 504})
_UPSTREAM_5XX_RETRY_BACKOFF_SECONDS = 0.5
def _counts_toward_cooldown(status_code: int) -> bool:
"""Provider faults and timeouts only — not client errors or rate limits."""
return status_code >= 500 or status_code == UPSTREAM_ERROR_STATUS
def _upstream_response_failure(response: Response) -> bool:
return (
_counts_toward_cooldown(response.status_code)
and response.headers.get(ERROR_SCOPE_HEADER) == ERROR_SCOPE_UPSTREAM
)
def _attribute_request(
request: Request, model_obj: Model, upstream: BaseUpstreamProvider
) -> None:
@@ -695,6 +713,20 @@ async def _proxy(
request=request,
)
# A provider that just failed this model repeatedly is skipped while some
# other candidate can serve it. An explicit route is never rerouted.
if selector is None:
healthy = [
candidate
for candidate in candidates
if not is_cooling_down(
provider_identity(candidate[1]),
candidate_model_identity(candidate[0], model_id),
)
]
if healthy:
candidates = healthy
# Reserve/max-cost checks use the best-ranked candidate; the failover loop
# below rebinds (model_obj, upstream) per candidate so forwarding and
# settlement always use the model of the provider actually being tried.
@@ -722,7 +754,7 @@ async def _proxy(
model_id,
)
continue
return await forward_ehbp_x_cashu_request(
response = await forward_ehbp_x_cashu_request(
request=request,
x_cashu_token=x_cashu,
path=path,
@@ -731,7 +763,7 @@ async def _proxy(
upstream=upstream,
)
elif is_responses_api:
return await upstream.handle_x_cashu_responses(
response = await upstream.handle_x_cashu_responses(
request,
x_cashu,
path,
@@ -740,7 +772,7 @@ async def _proxy(
request_body=request_body,
)
else:
return await upstream.handle_x_cashu(
response = await upstream.handle_x_cashu(
request,
x_cashu,
path,
@@ -748,6 +780,12 @@ async def _proxy(
model_obj,
request_body=request_body,
)
if _upstream_response_failure(response):
record_failure(
provider_identity(upstream),
candidate_model_identity(model_obj, model_id),
)
return response
except UpstreamError as e:
logger.warning(
"Upstream %s failed (x-cashu) for model=%s: %s",
@@ -760,6 +798,13 @@ async def _proxy(
"status_code": e.status_code,
},
)
if e.scope == ERROR_SCOPE_UPSTREAM and _counts_toward_cooldown(
e.status_code
):
record_failure(
provider_identity(upstream),
candidate_model_identity(model_obj, model_id),
)
if i == len(candidates) - 1:
last_error = e
continue
@@ -1030,6 +1075,11 @@ async def _proxy(
break
if response.status_code != 200:
if _upstream_response_failure(response):
record_failure(
provider_identity(upstream),
candidate_model_identity(model_obj, model_id),
)
# 424 is an upstream failure re-reported by error_scope.
# 502/503 are upstream errors, 429 rate limits.
should_retry = response.status_code in [
@@ -1117,6 +1167,13 @@ async def _proxy(
raise
except UpstreamError as e:
if e.scope == ERROR_SCOPE_UPSTREAM and _counts_toward_cooldown(
e.status_code
):
record_failure(
provider_identity(upstream),
candidate_model_identity(model_obj, model_id),
)
logger.warning(
"Upstream %s failed for model=%s: %s",
upstream.provider_type,
+62 -10
View File
@@ -34,6 +34,7 @@ from ..core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_NODE,
ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
client_code_for_upstream_error,
client_status_for_upstream_error,
upstream_status_details,
@@ -68,6 +69,7 @@ from .cache_breakpoints import (
inject_anthropic_cache_breakpoints,
is_explicit_cache_model,
)
from .cooldown import model_identity, provider_identity, record_failure
from .count_tokens import MissingUsageEstimator, count_tokens_locally
from .http_client import acquire_upstream_http_client, build_x_cashu_client
from .litellm_routing import detect_litellm_prefix
@@ -85,6 +87,7 @@ from .stream_ownership import (
close_upstream_exchange,
finalize_and_close_stream,
)
from .stream_timeout import GuardedStream, open_guarded_stream
if typing.TYPE_CHECKING:
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
@@ -1150,6 +1153,17 @@ class BaseUpstreamProvider:
)
return True
async def _guard_stream(
self, response: httpx.Response, model_obj: Model | None, *, sse: bool
) -> GuardedStream:
def on_idle() -> None:
if model_obj is not None and model_obj.id:
record_failure(provider_identity(self), model_identity(model_obj.id))
return await open_guarded_stream(
response, self.provider_type, sse=sse, on_idle_timeout=on_idle
)
async def handle_streaming_chat_completion(
self,
response: httpx.Response,
@@ -1171,6 +1185,8 @@ class BaseUpstreamProvider:
Returns:
StreamingResponse with cost data injected at the end
"""
guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
if reservation_snapshot is None:
async with create_session() as snapshot_session:
snapshot_key = await snapshot_session.get(key.__class__, key.hashed_key)
@@ -1374,7 +1390,7 @@ class BaseUpstreamProvider:
# multiple events can arrive together; buffering makes parsing
# boundary-independent for every provider.
splitter = SSEEventSplitter()
async for chunk in response.aiter_bytes():
async for chunk in guarded_chunks:
for raw_event in splitter.feed(chunk):
for out in _process_event(raw_event):
yield out
@@ -1460,7 +1476,9 @@ class BaseUpstreamProvider:
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
if done_seen:
if guarded_chunks.timed_out:
yield b'data: {"error":{"code":"UPSTREAM_TIMEOUT","message":"Upstream stream stalled"}}\n\n'
elif done_seen:
yield b"data: [DONE]\n\n"
except httpx.RemoteProtocolError as stream_error:
@@ -1666,6 +1684,8 @@ class BaseUpstreamProvider:
Returns:
StreamingResponse with cost data injected at the end
"""
guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
usage_estimator = MissingUsageEstimator(request_body, model_obj)
logger.debug(
@@ -1818,7 +1838,7 @@ class BaseUpstreamProvider:
# Buffer across network chunks; dispatch only on the SSE event
# delimiter so parsing is independent of byte boundaries.
splitter = SSEEventSplitter()
async for chunk in response.aiter_bytes():
async for chunk in guarded_chunks:
for raw_event in splitter.feed(chunk):
for out in _process_event(raw_event):
yield out
@@ -1867,7 +1887,9 @@ class BaseUpstreamProvider:
if usage_chunk_data is None:
usage_chunk_data = {
"type": "response.completed",
"type": "response.failed"
if guarded_chunks.timed_out
else "response.completed",
"provider": provider_seen,
"response": {
"model": last_model_seen or "unknown",
@@ -1889,6 +1911,14 @@ class BaseUpstreamProvider:
+ cost_data.get("output_tokens", 0),
},
}
if guarded_chunks.timed_out:
usage_chunk_data["type"] = "response.failed"
response_data = usage_chunk_data.get("response")
if isinstance(response_data, dict):
response_data["error"] = {
"code": "UPSTREAM_TIMEOUT",
"message": "Upstream stream stalled",
}
try:
self.inject_cost_metadata(
@@ -1904,7 +1934,12 @@ class BaseUpstreamProvider:
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
if done_seen:
if guarded_chunks.timed_out and (
usage_chunk_data is None
or usage_chunk_data.get("type") != "response.failed"
):
yield b'data: {"error":{"code":"UPSTREAM_TIMEOUT","message":"Upstream stream stalled"}}\n\n'
if done_seen and not guarded_chunks.timed_out:
yield b"data: [DONE]\n\n"
except httpx.RemoteProtocolError as stream_error:
@@ -2149,6 +2184,7 @@ class BaseUpstreamProvider:
provider_fee: float | None,
reservation_snapshot: ReservationSnapshot,
finalizer: PersistentStreamFinalizer | None = None,
guarded_chunks: GuardedStream | None = None,
) -> AsyncGenerator[bytes, None]:
"""Relay an opaque stream and settle it even if the caller disconnects."""
if finalizer is None:
@@ -2166,12 +2202,22 @@ class BaseUpstreamProvider:
)
)
try:
async for chunk in response.aiter_bytes():
if guarded_chunks is None:
guarded_chunks = await self._guard_stream(
response, model_obj, sse=False
)
async for chunk in guarded_chunks:
yield chunk
if guarded_chunks.timed_out:
raise UpstreamError(
"Upstream stream stalled",
status_code=UPSTREAM_ERROR_STATUS,
code="UPSTREAM_TIMEOUT",
)
finally:
await finalizer.run()
def _generic_streaming_response(
async def _generic_streaming_response(
self,
response: httpx.Response,
key_hash: str,
@@ -2181,6 +2227,7 @@ class BaseUpstreamProvider:
provider_fee: float | None,
reservation_snapshot: ReservationSnapshot,
) -> ClosingStreamingResponse:
guarded_chunks = await self._guard_stream(response, model_obj, sse=False)
finalizer = PersistentStreamFinalizer(
lambda: finalize_and_close_stream(
lambda: self._finalize_generic_streaming_payment(
@@ -2203,6 +2250,7 @@ class BaseUpstreamProvider:
provider_fee,
reservation_snapshot,
finalizer,
guarded_chunks,
)
return ClosingStreamingResponse(
stream,
@@ -2221,6 +2269,8 @@ class BaseUpstreamProvider:
reservation_snapshot: ReservationSnapshot | None = None,
request_body: bytes | None = None,
) -> StreamingResponse:
guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_finalized = False
last_model_seen: str | None = None
@@ -2314,7 +2364,7 @@ class BaseUpstreamProvider:
total_cost = max(total_cost, _coerce_usd(usage_or_root.get(field)))
try:
async for chunk in response.aiter_bytes():
async for chunk in guarded_chunks:
stored_chunks.append(chunk)
try:
decoded_chunk = chunk.decode("utf-8", errors="ignore")
@@ -2493,6 +2543,8 @@ class BaseUpstreamProvider:
maybe_cost_event = await finalize_without_usage()
if maybe_cost_event is not None:
yield maybe_cost_event
if guarded_chunks.timed_out:
yield b'event: error\ndata: {"error":{"code":"UPSTREAM_TIMEOUT","message":"Upstream stream stalled"}}\n\n'
except httpx.ReadError:
if not usage_finalized:
@@ -3451,7 +3503,7 @@ class BaseUpstreamProvider:
},
)
result = self._generic_streaming_response(
result = await self._generic_streaming_response(
response,
key.hashed_key,
max_cost_for_model,
@@ -3736,7 +3788,7 @@ class BaseUpstreamProvider:
},
)
result = self._generic_streaming_response(
result = await self._generic_streaming_response(
response,
key.hashed_key,
max_cost_for_model,
+79
View File
@@ -0,0 +1,79 @@
"""In-memory circuit breaker for a failing (provider, model) pair.
Process-local by design: each node observes its own upstream failures, and a
cooldown that outlives a restart would hide a provider that has recovered.
"""
from __future__ import annotations
import time
from typing import Any
from ..core import get_logger
from ..core.settings import settings
logger = get_logger(__name__)
_FAILURE_WINDOW_SECONDS = 60.0
_failures: dict[tuple[str, str], list[float]] = {}
_cooling_until: dict[tuple[str, str], float] = {}
def provider_identity(upstream: Any) -> str:
db_id = getattr(upstream, "db_id", None)
if isinstance(db_id, int):
return f"db:{db_id}"
return f"{upstream.provider_type.lower()}|{upstream.base_url.lower()}"
def model_identity(model_id: str) -> str:
return model_id.lower()
def candidate_model_identity(model: Any, requested_model_id: str) -> str:
model_id = getattr(model, "id", None)
return model_identity(
model_id if isinstance(model_id, str) and model_id else requested_model_id
)
def record_failure(provider_id: str, model_id: str) -> None:
"""Count a timeout or 5xx, opening a cooldown once too many land in a minute."""
if settings.upstream_cooldown_seconds <= 0:
return
pair = (provider_id, model_id)
now = time.monotonic()
recent = [t for t in _failures.get(pair, []) if now - t < _FAILURE_WINDOW_SECONDS]
recent.append(now)
if len(recent) >= settings.upstream_allowed_fails:
_failures.pop(pair, None)
_cooling_until[pair] = now + settings.upstream_cooldown_seconds
logger.warning(
"Upstream cooling down after repeated failures",
extra={
"provider": provider_id,
"model": model_id,
"cooldown_seconds": settings.upstream_cooldown_seconds,
},
)
else:
_failures[pair] = recent
def is_cooling_down(provider_id: str, model_id: str) -> bool:
pair = (provider_id, model_id)
until = _cooling_until.get(pair)
if until is None:
return False
if time.monotonic() >= until:
del _cooling_until[pair]
return False
return True
def reset_cooldowns() -> None:
_failures.clear()
_cooling_until.clear()
+117
View File
@@ -0,0 +1,117 @@
"""Timeout guards for upstream streaming responses."""
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator, Callable
import httpx
from ..core import get_logger
from ..core.error_scope import UPSTREAM_ERROR_STATUS
from ..core.exceptions import UpstreamError
from ..core.settings import settings
from .sse_splitter import SSEEventSplitter
logger = get_logger(__name__)
class GuardedStream(AsyncIterator[bytes]):
def __init__(
self,
first: bytes | None,
chunks: AsyncIterator[bytes],
provider_type: str,
on_idle_timeout: Callable[[], None] | None,
) -> None:
self.timed_out = False
self._chunks = self._resume(first, chunks, provider_type, on_idle_timeout)
def __aiter__(self) -> GuardedStream:
return self
async def __anext__(self) -> bytes:
return await anext(self._chunks)
async def _resume(
self,
first: bytes | None,
chunks: AsyncIterator[bytes],
provider_type: str,
on_idle_timeout: Callable[[], None] | None,
) -> AsyncIterator[bytes]:
chunk = first
while chunk is not None:
yield chunk
try:
chunk = await _next_chunk(
chunks, settings.upstream_stream_idle_timeout_seconds
)
except TimeoutError:
self.timed_out = True
logger.warning(
"Upstream stream stalled; aborting and billing actual usage",
extra={
"provider": provider_type,
"idle_timeout_seconds": settings.upstream_stream_idle_timeout_seconds,
},
)
if on_idle_timeout is not None:
on_idle_timeout()
return
def _has_data(event: bytes) -> bool:
return any(
line.startswith(b"data:") and line[5:].strip() for line in event.split(b"\n")
)
async def _sse_events(chunks: AsyncIterator[bytes]) -> AsyncIterator[bytes]:
"""Yield only deliverable SSE data events; comments cannot reset deadlines."""
splitter = SSEEventSplitter()
async for chunk in chunks:
for event in splitter.feed(chunk):
if _has_data(event):
yield event + b"\n\n"
tail = splitter.flush()
if _has_data(tail):
# Keep an unterminated tail unterminated: the caller's final flush must
# not mistake truncated JSON for a complete SSE frame.
yield tail
async def open_guarded_stream(
response: httpx.Response,
provider_type: str,
*,
sse: bool = False,
on_idle_timeout: Callable[[], None] | None = None,
) -> GuardedStream:
"""Prefetch a deliverable event before handing a response to the client.
Once the first event is sent, a stall cannot fail over; the stream ends and
the caller's finalizer settles usage observed before the interruption.
"""
chunks = response.aiter_bytes().__aiter__()
guarded_chunks = _sse_events(chunks) if sse else chunks
timeout = settings.upstream_first_token_timeout_seconds
try:
first = await _next_chunk(guarded_chunks, timeout)
except TimeoutError:
await response.aclose()
raise UpstreamError(
f"Upstream {provider_type} sent no first chunk within {timeout}s",
status_code=UPSTREAM_ERROR_STATUS,
code="UPSTREAM_TIMEOUT",
) from None
return GuardedStream(first, guarded_chunks, provider_type, on_idle_timeout)
async def _next_chunk(chunks: AsyncIterator[bytes], timeout: float) -> bytes | None:
"""Next chunk, or ``None`` at end of stream. ``timeout <= 0`` disables it."""
step = anext(chunks)
try:
return await (asyncio.wait_for(step, timeout) if timeout > 0 else step)
except StopAsyncIteration:
return None
+10
View File
@@ -31,3 +31,13 @@ def _isolate_redemption_negative_cache() -> Iterator[None]:
redemption_negative_cache.clear()
yield
redemption_negative_cache.clear()
@pytest.fixture(autouse=True)
def _isolate_upstream_cooldowns() -> Iterator[None]:
"""Clear process-wide upstream cooldowns so one test's failures can't skip providers in the next."""
from routstr.upstream.cooldown import reset_cooldowns
reset_cooldowns()
yield
reset_cooldowns()
+5
View File
@@ -365,6 +365,10 @@ async def test_database_url(tmp_path: Any) -> str:
@pytest_asyncio.fixture
async def integration_engine(test_database_url: str) -> AsyncGenerator[Any, None]:
"""Create an async engine for integration tests"""
from routstr.core.settings import settings
# Match the production engine's busy timeout; sqlite3's 5s default makes
# concurrency tests flake with "database is locked" on slow CI runners.
engine = create_async_engine(
test_database_url,
echo=False,
@@ -372,6 +376,7 @@ async def integration_engine(test_database_url: str) -> AsyncGenerator[Any, None
pool_pre_ping=True,
pool_size=5,
max_overflow=10,
connect_args={"timeout": settings.database_busy_timeout},
)
# Initialize database schema
@@ -316,7 +316,7 @@ async def test_streaming_response_closes_iterator_when_downstream_send_is_cancel
reservation = MagicMock(spec=ReservationSnapshot)
upstream_response.status_code = 201
upstream_response.headers = {"x-upstream": "preserved"}
response = provider._generic_streaming_response(
response = await provider._generic_streaming_response(
upstream_response,
"key-hash",
500,
@@ -377,7 +377,7 @@ async def test_generic_stream_settles_when_response_start_fails() -> None:
upstream_response.status_code = 201
upstream_response.headers = {"x-upstream": "preserved"}
reservation = MagicMock(spec=ReservationSnapshot)
response = provider._generic_streaming_response(
response = await provider._generic_streaming_response(
upstream_response,
"key-hash",
500,
+506
View File
@@ -0,0 +1,506 @@
"""First-token / idle stream guards and the per-(provider, model) cooldown."""
import asyncio
import json
from collections.abc import AsyncIterator
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from routstr.core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_NODE,
ERROR_SCOPE_UPSTREAM,
)
from routstr.core.exceptions import UpstreamError
from routstr.core.settings import Settings, settings
from routstr.upstream.base import BaseUpstreamProvider
from routstr.upstream.cooldown import is_cooling_down, record_failure
from routstr.upstream.stream_timeout import open_guarded_stream
def _response(chunks: AsyncIterator[bytes]) -> MagicMock:
response = MagicMock(spec=httpx.Response)
response.aiter_bytes = MagicMock(return_value=chunks)
response.aclose = AsyncMock()
return response
async def _never() -> AsyncIterator[bytes]:
await asyncio.sleep(10)
yield b"late"
async def _stalls_after_first() -> AsyncIterator[bytes]:
yield b"first"
await asyncio.sleep(10)
yield b"never delivered"
async def _heartbeat_only(frame: bytes = b": keepalive\n\n") -> AsyncIterator[bytes]:
while True:
yield frame
await asyncio.sleep(0.002)
@pytest.fixture
def fast_timeouts(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(settings, "upstream_first_token_timeout_seconds", 0.01)
monkeypatch.setattr(settings, "upstream_stream_idle_timeout_seconds", 0.01)
@pytest.mark.asyncio
async def test_first_token_timeout_closes_response_and_raises(
fast_timeouts: None,
) -> None:
response = _response(_never())
with pytest.raises(UpstreamError) as exc_info:
await open_guarded_stream(response, "test")
assert exc_info.value.code == "UPSTREAM_TIMEOUT"
assert exc_info.value.from_upstream_response is False
response.aclose.assert_awaited_once()
@pytest.mark.asyncio
async def test_generic_stream_times_out_before_response_is_handed_off(
fast_timeouts: None,
) -> None:
provider = BaseUpstreamProvider(base_url="https://slow.example", api_key="test")
response = _response(_never())
with pytest.raises(UpstreamError, match="no first chunk"):
await provider._generic_streaming_response(
response, "key-hash", 100, "audio/speech", None, None, MagicMock()
)
response.aclose.assert_awaited_once()
@pytest.mark.asyncio
async def test_generic_stream_idle_abort_settles_without_clean_completion(
fast_timeouts: None,
) -> None:
provider = BaseUpstreamProvider(base_url="https://slow.example", api_key="test")
finalize = AsyncMock()
provider._finalize_generic_streaming_payment = finalize # type: ignore[method-assign]
upstream = _response(_stalls_after_first())
upstream.status_code = 200
upstream.headers = {}
response = await provider._generic_streaming_response(
upstream, "key-hash", 100, "audio/speech", None, None, MagicMock()
)
chunks = []
with pytest.raises(UpstreamError, match="stream stalled"):
async for chunk in response.body_iterator:
chunks.append(chunk)
assert chunks == [b"first"]
finalize.assert_awaited_once()
upstream.aclose.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("frame", [b": keepalive\n\n", b"data: \n\n"])
async def test_sse_heartbeats_do_not_satisfy_first_token_timeout(
fast_timeouts: None, frame: bytes
) -> None:
response = _response(_heartbeat_only(frame))
with pytest.raises(UpstreamError, match="no first chunk"):
await open_guarded_stream(response, "test", sse=True)
response.aclose.assert_awaited_once()
@pytest.mark.asyncio
@pytest.mark.parametrize("frame", [b": keepalive\n\n", b"data: \n\n"])
async def test_sse_heartbeats_do_not_reset_idle_timeout(
fast_timeouts: None, frame: bytes
) -> None:
async def chunks() -> AsyncIterator[bytes]:
yield b'data: {"delta":"first"}\n\n'
async for chunk in _heartbeat_only(frame):
yield chunk
failures = MagicMock()
stream = await open_guarded_stream(
_response(chunks()), "test", sse=True, on_idle_timeout=failures
)
assert [chunk async for chunk in stream] == [b'data: {"delta":"first"}\n\n']
assert stream.timed_out is True
failures.assert_called_once()
@pytest.mark.asyncio
async def test_zero_first_token_timeout_disables_the_guard(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "upstream_first_token_timeout_seconds", 0)
monkeypatch.setattr(settings, "upstream_stream_idle_timeout_seconds", 0)
async def _slow() -> AsyncIterator[bytes]:
await asyncio.sleep(0.02)
yield b"first"
stream = await open_guarded_stream(_response(_slow()), "test")
assert [chunk async for chunk in stream] == [b"first"]
def test_stream_guards_are_off_by_default() -> None:
# Reasoning models can think silently for minutes; on by default, the
# guards would fail requests that succeed without them.
fields = Settings.__fields__
assert fields["upstream_first_token_timeout_seconds"].default == 0
assert fields["upstream_stream_idle_timeout_seconds"].default == 0
@pytest.mark.asyncio
async def test_idle_timeout_cools_down_the_serving_provider(
fast_timeouts: None, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(settings, "upstream_allowed_fails", 1)
provider = BaseUpstreamProvider(base_url="https://slow.example", api_key="test")
provider.db_id = 17
model = MagicMock(id="test-model")
guarded = await provider._guard_stream(
_response(_stalls_after_first()), model, sse=False
)
assert [chunk async for chunk in guarded] == [b"first"]
assert is_cooling_down("db:17", "test-model")
@pytest.mark.asyncio
@pytest.mark.parametrize("terminal_before_stall", [False, True])
async def test_responses_idle_timeout_does_not_emit_completed(
fast_timeouts: None, terminal_before_stall: bool
) -> None:
async def chunks() -> AsyncIterator[bytes]:
event = (
b'data: {"type":"response.completed","response":{"model":"test","usage":{"input_tokens":0,"output_tokens":1}}}\n\n'
if terminal_before_stall
else b'data: {"type":"response.created","response":{"model":"test"}}\n\n'
)
yield event
await asyncio.sleep(10)
response = _response(chunks())
response.status_code = 200
response.headers = {"content-type": "text/event-stream"}
key = MagicMock()
key.hashed_key = "test-key"
key.balance = 1000
session = MagicMock()
session.get = AsyncMock(return_value=key)
session_context = MagicMock()
session_context.__aenter__ = AsyncMock(return_value=session)
session_context.__aexit__ = AsyncMock(return_value=None)
provider = BaseUpstreamProvider(base_url="https://slow.example", api_key="test")
with (
patch("routstr.upstream.base.create_session", return_value=session_context),
patch(
"routstr.upstream.base.adjust_payment_for_tokens",
AsyncMock(return_value={"input_tokens": 0, "output_tokens": 1}),
),
):
result = await provider.handle_streaming_responses_completion(
response, key, 100, reservation_snapshot=MagicMock()
)
emitted = b"".join(
[
chunk.encode() if isinstance(chunk, str) else bytes(chunk)
async for chunk in result.body_iterator
]
)
assert b'"type": "response.failed"' in emitted
assert b'"code": "UPSTREAM_TIMEOUT"' in emitted
assert b'"type": "response.completed"' not in emitted
response.aclose.assert_awaited_once()
@pytest.mark.asyncio
async def test_guarded_stream_passes_every_chunk_through() -> None:
async def _chunks() -> AsyncIterator[bytes]:
yield b"a"
yield b"b"
yield b"c"
stream = await open_guarded_stream(_response(_chunks()), "test")
assert [chunk async for chunk in stream] == [b"a", b"b", b"c"]
def test_cooldown_opens_after_allowed_fails_and_expires(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "upstream_allowed_fails", 3)
monkeypatch.setattr(settings, "upstream_cooldown_seconds", 30)
for _ in range(2):
record_failure("https://a.example", "m")
assert is_cooling_down("https://a.example", "m") is False
record_failure("https://a.example", "m")
assert is_cooling_down("https://a.example", "m") is True
# Scoped to the exact pair.
assert is_cooling_down("https://b.example", "m") is False
assert is_cooling_down("https://a.example", "other") is False
with patch("routstr.upstream.cooldown.time.monotonic", return_value=1e6):
assert is_cooling_down("https://a.example", "m") is False
def test_zero_cooldown_disables_skipping(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(settings, "upstream_allowed_fails", 1)
monkeypatch.setattr(settings, "upstream_cooldown_seconds", 0)
record_failure("https://a.example", "m")
assert is_cooling_down("https://a.example", "m") is False
def _proxy_request() -> MagicMock:
request = MagicMock()
request.method = "POST"
request.headers = {"authorization": "Bearer sk-key"}
request.body = AsyncMock(return_value=b'{"model": "test-model", "stream": true}')
request.state = MagicMock()
request.state.request_id = "req-1"
return request
def _upstream(base_url: str, forward: AsyncMock) -> MagicMock:
upstream = MagicMock()
upstream.provider_type = "test"
upstream.base_url = base_url
upstream.db_id = None
upstream.prepare_headers = MagicMock(side_effect=lambda h: h)
upstream.forward_request = forward
return upstream
async def _run_proxy(
candidates: list[tuple[MagicMock, MagicMock]],
revert_mock: AsyncMock,
request: MagicMock | None = None,
) -> Any:
from routstr import proxy as proxy_module
from routstr.auth import ReservationSnapshot
from routstr.core.db import ApiKey
key = ApiKey(hashed_key="streamkey", balance=10_000)
reservation = ReservationSnapshot(
release_id="release",
key_hash=key.hashed_key,
billing_key_hash=key.hashed_key,
reserved_msats=1_000,
)
with (
patch.object(proxy_module, "get_candidates", return_value=candidates),
patch.object(
proxy_module, "get_max_cost_for_model", AsyncMock(return_value=1_000)
),
patch.object(
proxy_module,
"calculate_discounted_max_cost",
AsyncMock(return_value=1_000),
),
patch.object(proxy_module, "check_token_balance", MagicMock()),
patch.object(proxy_module, "get_bearer_token_key", AsyncMock(return_value=key)),
patch.object(
proxy_module, "pay_for_request", AsyncMock(return_value=reservation)
),
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
):
request = request or _proxy_request()
return await proxy_module._proxy(
request, "v1/chat/completions", MagicMock(), await request.body()
)
@pytest.mark.asyncio
async def test_first_token_timeout_fails_over_to_the_next_candidate(
fast_timeouts: None,
) -> None:
async def _timing_out(*args: Any, **kwargs: Any) -> Any:
return await open_guarded_stream(_response(_never()), "test")
served = MagicMock()
served.status_code = 200
slow = _upstream("https://slow.example", AsyncMock(side_effect=_timing_out))
fast = _upstream("https://fast.example", AsyncMock(return_value=served))
revert_mock = AsyncMock(return_value=True)
response = await _run_proxy([(MagicMock(), slow), (MagicMock(), fast)], revert_mock)
assert response is served
fast.forward_request.assert_awaited_once()
# The reservation carries over to the candidate that served the request.
revert_mock.assert_not_awaited()
@pytest.mark.asyncio
async def test_first_token_timeout_on_last_candidate_reverts_reservation(
fast_timeouts: None,
) -> None:
async def _timing_out(*args: Any, **kwargs: Any) -> Any:
return await open_guarded_stream(_response(_never()), "test")
slow = _upstream("https://slow.example", AsyncMock(side_effect=_timing_out))
revert_mock = AsyncMock(return_value=True)
response = await _run_proxy([(MagicMock(), slow)], revert_mock)
assert response.status_code == 424
assert json.loads(bytes(response.body))["error"]["code"] == "UPSTREAM_TIMEOUT"
revert_mock.assert_awaited_once()
@pytest.mark.asyncio
async def test_cooling_down_candidate_is_skipped_then_recovers(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "upstream_allowed_fails", 1)
monkeypatch.setattr(settings, "upstream_cooldown_seconds", 30)
sick_response = MagicMock()
sick_response.status_code = 200
healthy_response = MagicMock()
healthy_response.status_code = 200
sick = _upstream("https://sick.example", AsyncMock(return_value=sick_response))
healthy = _upstream("https://ok.example", AsyncMock(return_value=healthy_response))
candidates = [(MagicMock(), sick), (MagicMock(), healthy)]
record_failure("test|https://sick.example", "test-model")
assert await _run_proxy(candidates, AsyncMock()) is healthy_response
sick.forward_request.assert_not_awaited()
with patch("routstr.upstream.cooldown.time.monotonic", return_value=1e6):
assert await _run_proxy(candidates, AsyncMock()) is sick_response
@pytest.mark.asyncio
async def test_cooldown_never_empties_the_candidate_list(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "upstream_allowed_fails", 1)
monkeypatch.setattr(settings, "upstream_cooldown_seconds", 30)
only_response = MagicMock()
only_response.status_code = 200
only = _upstream("https://only.example", AsyncMock(return_value=only_response))
record_failure("test|https://only.example", "test-model")
assert await _run_proxy([(MagicMock(), only)], AsyncMock()) is only_response
@pytest.mark.asyncio
async def test_cooldown_distinguishes_credentials_at_same_url(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "upstream_allowed_fails", 1)
bad = _upstream("https://same.example", AsyncMock())
bad.db_id = 1
good_response = MagicMock(status_code=200)
good = _upstream("https://same.example", AsyncMock(return_value=good_response))
good.db_id = 2
other = _upstream("https://other.example", AsyncMock())
record_failure("db:1", "test-model")
assert (
await _run_proxy(
[(MagicMock(), bad), (MagicMock(), good), (MagicMock(), other)], AsyncMock()
)
is good_response
)
bad.forward_request.assert_not_awaited()
@pytest.mark.asyncio
async def test_cooldown_normalizes_model_spelling(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "upstream_allowed_fails", 1)
bad = _upstream("https://bad.example", AsyncMock())
good_response = MagicMock(status_code=200)
good = _upstream("https://good.example", AsyncMock(return_value=good_response))
record_failure("test|https://bad.example", "test-model")
request = _proxy_request()
request.body = AsyncMock(
return_value=b'{"model":"TEST-MODEL-20251222","stream":true}'
)
assert (
await _run_proxy(
[(MagicMock(id="test-model"), bad), (MagicMock(id="test-model"), good)],
AsyncMock(),
request,
)
is good_response
)
bad.forward_request.assert_not_awaited()
@pytest.mark.asyncio
async def test_x_cashu_upstream_failure_opens_cooldown(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "upstream_allowed_fails", 1)
upstream = _upstream("https://cashu.example", AsyncMock())
upstream.handle_x_cashu = AsyncMock(
return_value=MagicMock(
status_code=503, headers={ERROR_SCOPE_HEADER: ERROR_SCOPE_UPSTREAM}
)
)
request = _proxy_request()
request.headers = {"x-cashu": "token"}
response = await _run_proxy([(MagicMock(), upstream)], AsyncMock(), request)
assert response.status_code == 503
assert is_cooling_down("test|https://cashu.example", "test-model")
@pytest.mark.asyncio
async def test_x_cashu_local_mint_failure_does_not_cool_provider(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "upstream_allowed_fails", 1)
upstream = _upstream("https://cashu.example", AsyncMock())
upstream.handle_x_cashu = AsyncMock(
return_value=MagicMock(status_code=503, headers={})
)
request = _proxy_request()
request.headers = {"x-cashu": "token"}
response = await _run_proxy([(MagicMock(), upstream)], AsyncMock(), request)
assert response.status_code == 503
assert not is_cooling_down("test|https://cashu.example", "test-model")
@pytest.mark.asyncio
async def test_node_scoped_upstream_exception_does_not_cool_provider(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(settings, "upstream_allowed_fails", 1)
upstream = _upstream(
"https://healthy.example",
AsyncMock(
side_effect=UpstreamError(
"local fault", status_code=500, scope=ERROR_SCOPE_NODE
)
),
)
response = await _run_proxy([(MagicMock(), upstream)], AsyncMock())
assert response.status_code == 500
assert not is_cooling_down("test|https://healthy.example", "test-model")