feat: add upstream first-token timeout, stream idle timeout and provider cooldown

This commit is contained in:
9qeklajc
2026-09-27 03:10:10 +02:00
parent eb0f4a2cf9
commit ca4e93a819
6 changed files with 440 additions and 4 deletions
+15
View File
@@ -40,6 +40,21 @@ class Settings(BaseSettings):
upstream_5xx_retry_attempts: int = Field( upstream_5xx_retry_attempts: int = Field(
default=1, ge=0, env="UPSTREAM_5XX_RETRY_ATTEMPTS" default=1, ge=0, env="UPSTREAM_5XX_RETRY_ATTEMPTS"
) )
# Streaming guards, both disabled by 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.
upstream_first_token_timeout_seconds: float = Field(
default=60.0, ge=0, env="UPSTREAM_FIRST_TOKEN_TIMEOUT_SECONDS"
)
upstream_stream_idle_timeout_seconds: float = Field(
default=120.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 # Node info
name: str = Field(default="ARoutstrNode", env="NAME") name: str = Field(default="ARoutstrNode", env="NAME")
+21
View File
@@ -40,6 +40,7 @@ from .payment.helpers import (
) )
from .payment.models import Model from .payment.models import Model
from .upstream import BaseUpstreamProvider from .upstream import BaseUpstreamProvider
from .upstream.cooldown import is_cooling_down, record_failure
from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request from .upstream.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request
from .upstream.helpers import init_upstreams from .upstream.helpers import init_upstreams
from .upstream.model_paths import ( from .upstream.model_paths import (
@@ -405,6 +406,11 @@ _RETRYABLE_UPSTREAM_5XX = frozenset({502, 503, 504})
_UPSTREAM_5XX_RETRY_BACKOFF_SECONDS = 0.5 _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
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) @proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
async def proxy( async def proxy(
request: Request, path: str, session: AsyncSession = Depends(get_session) request: Request, path: str, session: AsyncSession = Depends(get_session)
@@ -621,6 +627,17 @@ async def _proxy(
request=request, 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(candidate[1].base_url, model_id)
]
if healthy:
candidates = healthy
# Reserve/max-cost checks use the best-ranked candidate; the failover loop # Reserve/max-cost checks use the best-ranked candidate; the failover loop
# below rebinds (model_obj, upstream) per candidate so forwarding and # below rebinds (model_obj, upstream) per candidate so forwarding and
# settlement always use the model of the provider actually being tried. # settlement always use the model of the provider actually being tried.
@@ -950,6 +967,8 @@ async def _proxy(
break break
if response.status_code != 200: if response.status_code != 200:
if _counts_toward_cooldown(response.status_code):
record_failure(upstream.base_url, model_id)
# 424 is an upstream failure re-reported by error_scope. # 424 is an upstream failure re-reported by error_scope.
# 502/503 are upstream errors, 429 rate limits. # 502/503 are upstream errors, 429 rate limits.
should_retry = response.status_code in [ should_retry = response.status_code in [
@@ -1037,6 +1056,8 @@ async def _proxy(
raise raise
except UpstreamError as e: except UpstreamError as e:
if _counts_toward_cooldown(e.status_code):
record_failure(upstream.base_url, model_id)
logger.warning( logger.warning(
"Upstream %s failed for model=%s: %s", "Upstream %s failed for model=%s: %s",
upstream.provider_type, upstream.provider_type,
+13 -4
View File
@@ -85,6 +85,7 @@ from .stream_ownership import (
close_upstream_exchange, close_upstream_exchange,
finalize_and_close_stream, finalize_and_close_stream,
) )
from .stream_timeout import open_guarded_stream
if typing.TYPE_CHECKING: if typing.TYPE_CHECKING:
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
@@ -1147,6 +1148,8 @@ class BaseUpstreamProvider:
Returns: Returns:
StreamingResponse with cost data injected at the end StreamingResponse with cost data injected at the end
""" """
guarded_chunks = await open_guarded_stream(response, self.provider_type)
if reservation_snapshot is None: if reservation_snapshot is None:
async with create_session() as snapshot_session: async with create_session() as snapshot_session:
snapshot_key = await snapshot_session.get(key.__class__, key.hashed_key) snapshot_key = await snapshot_session.get(key.__class__, key.hashed_key)
@@ -1348,7 +1351,7 @@ class BaseUpstreamProvider:
# multiple events can arrive together; buffering makes parsing # multiple events can arrive together; buffering makes parsing
# boundary-independent for every provider. # boundary-independent for every provider.
splitter = SSEEventSplitter() splitter = SSEEventSplitter()
async for chunk in response.aiter_bytes(): async for chunk in guarded_chunks:
for raw_event in splitter.feed(chunk): for raw_event in splitter.feed(chunk):
for out in _process_event(raw_event): for out in _process_event(raw_event):
yield out yield out
@@ -1639,6 +1642,8 @@ class BaseUpstreamProvider:
Returns: Returns:
StreamingResponse with cost data injected at the end StreamingResponse with cost data injected at the end
""" """
guarded_chunks = await open_guarded_stream(response, self.provider_type)
usage_estimator = MissingUsageEstimator(request_body, model_obj) usage_estimator = MissingUsageEstimator(request_body, model_obj)
logger.debug( logger.debug(
@@ -1790,7 +1795,7 @@ class BaseUpstreamProvider:
# Buffer across network chunks; dispatch only on the SSE event # Buffer across network chunks; dispatch only on the SSE event
# delimiter so parsing is independent of byte boundaries. # delimiter so parsing is independent of byte boundaries.
splitter = SSEEventSplitter() splitter = SSEEventSplitter()
async for chunk in response.aiter_bytes(): async for chunk in guarded_chunks:
for raw_event in splitter.feed(chunk): for raw_event in splitter.feed(chunk):
for out in _process_event(raw_event): for out in _process_event(raw_event):
yield out yield out
@@ -2137,7 +2142,9 @@ class BaseUpstreamProvider:
) )
) )
try: try:
async for chunk in response.aiter_bytes(): # This generator is already the response body, so a first-chunk
# timeout here can only abort the stream, never fail over.
async for chunk in await open_guarded_stream(response, self.provider_type):
yield chunk yield chunk
finally: finally:
await finalizer.run() await finalizer.run()
@@ -2192,6 +2199,8 @@ class BaseUpstreamProvider:
reservation_snapshot: ReservationSnapshot | None = None, reservation_snapshot: ReservationSnapshot | None = None,
request_body: bytes | None = None, request_body: bytes | None = None,
) -> StreamingResponse: ) -> StreamingResponse:
guarded_chunks = await open_guarded_stream(response, self.provider_type)
usage_estimator = MissingUsageEstimator(request_body, model_obj) usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_finalized = False usage_finalized = False
last_model_seen: str | None = None last_model_seen: str | None = None
@@ -2284,7 +2293,7 @@ class BaseUpstreamProvider:
total_cost = max(total_cost, _coerce_usd(usage_or_root.get(field))) total_cost = max(total_cost, _coerce_usd(usage_or_root.get(field)))
try: try:
async for chunk in response.aiter_bytes(): async for chunk in guarded_chunks:
stored_chunks.append(chunk) stored_chunks.append(chunk)
try: try:
decoded_chunk = chunk.decode("utf-8", errors="ignore") decoded_chunk = chunk.decode("utf-8", errors="ignore")
+60
View File
@@ -0,0 +1,60 @@
"""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 ..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 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()
+69
View File
@@ -0,0 +1,69 @@
"""Timeout guards for upstream streaming responses."""
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator
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
logger = get_logger(__name__)
async def open_guarded_stream(
response: httpx.Response, provider_type: str
) -> AsyncIterator[bytes]:
"""Await the upstream's first chunk, then hand back the whole stream.
Awaiting the first chunk before any ``StreamingResponse`` exists is what
makes a slow-starting provider recoverable: the proxy's candidate loop only
sees errors raised while it still owns the request, and no byte has reached
the client yet. A stall after that chunk cannot fail over, so the returned
iterator simply ends and the caller's finalizer settles actual usage.
"""
chunks = response.aiter_bytes().__aiter__()
timeout = settings.upstream_first_token_timeout_seconds
try:
first = await _next_chunk(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 _resume(first, chunks, provider_type)
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
async def _resume(
first: bytes | None, chunks: AsyncIterator[bytes], provider_type: str
) -> AsyncIterator[bytes]:
idle_timeout = settings.upstream_stream_idle_timeout_seconds
chunk = first
while chunk is not None:
yield chunk
try:
chunk = await _next_chunk(chunks, idle_timeout)
except TimeoutError:
logger.warning(
"Upstream stream stalled; aborting and billing actual usage",
extra={
"provider": provider_type,
"idle_timeout_seconds": idle_timeout,
},
)
return
+262
View File
@@ -0,0 +1,262 @@
"""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.exceptions import UpstreamError
from routstr.core.settings import settings
from routstr.upstream.cooldown import is_cooling_down, record_failure, reset_cooldowns
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"
@pytest.fixture(autouse=True)
def _clean_cooldowns() -> Any:
reset_cooldowns()
yield
reset_cooldowns()
@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_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"]
@pytest.mark.asyncio
async def test_idle_timeout_ends_the_stream_without_raising(
fast_timeouts: None,
) -> None:
stream = await open_guarded_stream(_response(_stalls_after_first()), "test")
# The stalled stream ends after the delivered bytes; the caller's finalizer
# then settles actual usage instead of the request hanging.
assert [chunk async for chunk in stream] == [b"first"]
@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.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,
) -> 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),
):
return await proxy_module.proxy(
_proxy_request(), "v1/chat/completions", session=MagicMock()
)
@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("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("https://only.example", "test-model")
assert await _run_proxy([(MagicMock(), only)], AsyncMock()) is only_response