fix: guard meaningful stream events and scope cooldown failures

This commit is contained in:
9qeklajc
2026-09-29 02:58:40 +02:00
parent cd2fe38ae6
commit a1071f4a98
6 changed files with 455 additions and 61 deletions
+47 -11
View File
@@ -1,6 +1,7 @@
import asyncio import asyncio
import inspect import inspect
import json import json
import re
from typing import Any from typing import Any
from fastapi import APIRouter, HTTPException, Request from fastapi import APIRouter, HTTPException, Request
@@ -23,6 +24,7 @@ from .core.db import (
create_session, create_session,
) )
from .core.error_scope import ( from .core.error_scope import (
ERROR_SCOPE_HEADER,
ERROR_SCOPE_UPSTREAM, ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS, UPSTREAM_ERROR_STATUS,
UPSTREAM_UNAVAILABLE, UPSTREAM_UNAVAILABLE,
@@ -39,7 +41,12 @@ 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.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.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 (
@@ -116,8 +123,6 @@ def get_candidates(
if candidates := _provider_map.get(model_id_lower): if candidates := _provider_map.get(model_id_lower):
return candidates return candidates
import re
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower) base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
if base_model_id != model_id_lower: if base_model_id != model_id_lower:
if candidates := _provider_map.get(base_model_id): if candidates := _provider_map.get(base_model_id):
@@ -410,6 +415,13 @@ def _counts_toward_cooldown(status_code: int) -> bool:
return status_code >= 500 or status_code == UPSTREAM_ERROR_STATUS 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( def _attribute_request(
request: Request, model_obj: Model, upstream: BaseUpstreamProvider request: Request, model_obj: Model, upstream: BaseUpstreamProvider
) -> None: ) -> None:
@@ -705,7 +717,10 @@ async def _proxy(
healthy = [ healthy = [
candidate candidate
for candidate in candidates for candidate in candidates
if not is_cooling_down(candidate[1].base_url, model_id) if not is_cooling_down(
provider_identity(candidate[1]),
candidate_model_identity(candidate[0], model_id),
)
] ]
if healthy: if healthy:
candidates = healthy candidates = healthy
@@ -737,7 +752,7 @@ async def _proxy(
model_id, model_id,
) )
continue continue
return await forward_ehbp_x_cashu_request( response = await forward_ehbp_x_cashu_request(
request=request, request=request,
x_cashu_token=x_cashu, x_cashu_token=x_cashu,
path=path, path=path,
@@ -746,7 +761,7 @@ async def _proxy(
upstream=upstream, upstream=upstream,
) )
elif is_responses_api: elif is_responses_api:
return await upstream.handle_x_cashu_responses( response = await upstream.handle_x_cashu_responses(
request, request,
x_cashu, x_cashu,
path, path,
@@ -755,7 +770,7 @@ async def _proxy(
request_body=request_body, request_body=request_body,
) )
else: else:
return await upstream.handle_x_cashu( response = await upstream.handle_x_cashu(
request, request,
x_cashu, x_cashu,
path, path,
@@ -763,6 +778,12 @@ async def _proxy(
model_obj, model_obj,
request_body=request_body, 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: except UpstreamError as e:
logger.warning( logger.warning(
"Upstream %s failed (x-cashu) for model=%s: %s", "Upstream %s failed (x-cashu) for model=%s: %s",
@@ -775,6 +796,13 @@ async def _proxy(
"status_code": e.status_code, "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: if i == len(candidates) - 1:
last_error = e last_error = e
continue continue
@@ -1044,8 +1072,11 @@ async def _proxy(
break break
if response.status_code != 200: if response.status_code != 200:
if _counts_toward_cooldown(response.status_code): if _upstream_response_failure(response):
record_failure(upstream.base_url, model_id) record_failure(
provider_identity(upstream),
candidate_model_identity(model_obj, 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 [
@@ -1133,8 +1164,13 @@ async def _proxy(
raise raise
except UpstreamError as e: except UpstreamError as e:
if _counts_toward_cooldown(e.status_code): if e.scope == ERROR_SCOPE_UPSTREAM and _counts_toward_cooldown(
record_failure(upstream.base_url, model_id) e.status_code
):
record_failure(
provider_identity(upstream),
candidate_model_identity(model_obj, 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,
+56 -13
View File
@@ -34,6 +34,7 @@ from ..core.error_scope import (
ERROR_SCOPE_HEADER, ERROR_SCOPE_HEADER,
ERROR_SCOPE_NODE, ERROR_SCOPE_NODE,
ERROR_SCOPE_UPSTREAM, ERROR_SCOPE_UPSTREAM,
UPSTREAM_ERROR_STATUS,
client_code_for_upstream_error, client_code_for_upstream_error,
client_status_for_upstream_error, client_status_for_upstream_error,
upstream_status_details, upstream_status_details,
@@ -68,6 +69,7 @@ from .cache_breakpoints import (
inject_anthropic_cache_breakpoints, inject_anthropic_cache_breakpoints,
is_explicit_cache_model, is_explicit_cache_model,
) )
from .cooldown import model_identity, provider_identity, record_failure
from .count_tokens import MissingUsageEstimator, count_tokens_locally from .count_tokens import MissingUsageEstimator, count_tokens_locally
from .http_client import acquire_upstream_http_client, build_x_cashu_client from .http_client import acquire_upstream_http_client, build_x_cashu_client
from .litellm_routing import detect_litellm_prefix from .litellm_routing import detect_litellm_prefix
@@ -85,7 +87,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 from .stream_timeout import GuardedStream, open_guarded_stream
if typing.TYPE_CHECKING: if typing.TYPE_CHECKING:
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
@@ -1127,6 +1129,17 @@ class BaseUpstreamProvider:
) )
return True 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( async def handle_streaming_chat_completion(
self, self,
response: httpx.Response, response: httpx.Response,
@@ -1148,7 +1161,7 @@ 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) guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
if reservation_snapshot is None: if reservation_snapshot is None:
async with create_session() as snapshot_session: async with create_session() as snapshot_session:
@@ -1436,7 +1449,9 @@ class BaseUpstreamProvider:
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() 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" yield b"data: [DONE]\n\n"
except httpx.RemoteProtocolError as stream_error: except httpx.RemoteProtocolError as stream_error:
@@ -1642,7 +1657,7 @@ 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) guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
usage_estimator = MissingUsageEstimator(request_body, model_obj) usage_estimator = MissingUsageEstimator(request_body, model_obj)
@@ -1844,7 +1859,9 @@ class BaseUpstreamProvider:
if usage_chunk_data is None: if usage_chunk_data is None:
usage_chunk_data = { usage_chunk_data = {
"type": "response.completed", "type": "response.failed"
if guarded_chunks.timed_out
else "response.completed",
"response": { "response": {
"model": last_model_seen or "unknown", "model": last_model_seen or "unknown",
"usage": { "usage": {
@@ -1865,6 +1882,14 @@ class BaseUpstreamProvider:
+ cost_data.get("output_tokens", 0), + 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: try:
self.inject_cost_metadata( self.inject_cost_metadata(
@@ -1880,7 +1905,12 @@ class BaseUpstreamProvider:
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() 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" yield b"data: [DONE]\n\n"
except httpx.RemoteProtocolError as stream_error: except httpx.RemoteProtocolError as stream_error:
@@ -2125,6 +2155,7 @@ class BaseUpstreamProvider:
provider_fee: float | None, provider_fee: float | None,
reservation_snapshot: ReservationSnapshot, reservation_snapshot: ReservationSnapshot,
finalizer: PersistentStreamFinalizer | None = None, finalizer: PersistentStreamFinalizer | None = None,
guarded_chunks: GuardedStream | None = None,
) -> AsyncGenerator[bytes, None]: ) -> AsyncGenerator[bytes, None]:
"""Relay an opaque stream and settle it even if the caller disconnects.""" """Relay an opaque stream and settle it even if the caller disconnects."""
if finalizer is None: if finalizer is None:
@@ -2142,14 +2173,22 @@ class BaseUpstreamProvider:
) )
) )
try: try:
# This generator is already the response body, so a first-chunk if guarded_chunks is None:
# timeout here can only abort the stream, never fail over. guarded_chunks = await self._guard_stream(
async for chunk in await open_guarded_stream(response, self.provider_type): response, model_obj, sse=False
)
async for chunk in guarded_chunks:
yield chunk yield chunk
if guarded_chunks.timed_out:
raise UpstreamError(
"Upstream stream stalled",
status_code=UPSTREAM_ERROR_STATUS,
code="UPSTREAM_TIMEOUT",
)
finally: finally:
await finalizer.run() await finalizer.run()
def _generic_streaming_response( async def _generic_streaming_response(
self, self,
response: httpx.Response, response: httpx.Response,
key_hash: str, key_hash: str,
@@ -2159,6 +2198,7 @@ class BaseUpstreamProvider:
provider_fee: float | None, provider_fee: float | None,
reservation_snapshot: ReservationSnapshot, reservation_snapshot: ReservationSnapshot,
) -> ClosingStreamingResponse: ) -> ClosingStreamingResponse:
guarded_chunks = await self._guard_stream(response, model_obj, sse=False)
finalizer = PersistentStreamFinalizer( finalizer = PersistentStreamFinalizer(
lambda: finalize_and_close_stream( lambda: finalize_and_close_stream(
lambda: self._finalize_generic_streaming_payment( lambda: self._finalize_generic_streaming_payment(
@@ -2181,6 +2221,7 @@ class BaseUpstreamProvider:
provider_fee, provider_fee,
reservation_snapshot, reservation_snapshot,
finalizer, finalizer,
guarded_chunks,
) )
return ClosingStreamingResponse( return ClosingStreamingResponse(
stream, stream,
@@ -2199,7 +2240,7 @@ 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) guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
usage_estimator = MissingUsageEstimator(request_body, model_obj) usage_estimator = MissingUsageEstimator(request_body, model_obj)
usage_finalized = False usage_finalized = False
@@ -2469,6 +2510,8 @@ class BaseUpstreamProvider:
maybe_cost_event = await finalize_without_usage() maybe_cost_event = await finalize_without_usage()
if maybe_cost_event is not None: if maybe_cost_event is not None:
yield maybe_cost_event 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: except httpx.ReadError:
if not usage_finalized: if not usage_finalized:
@@ -3421,7 +3464,7 @@ class BaseUpstreamProvider:
}, },
) )
result = self._generic_streaming_response( result = await self._generic_streaming_response(
response, response,
key.hashed_key, key.hashed_key,
max_cost_for_model, max_cost_for_model,
@@ -3706,7 +3749,7 @@ class BaseUpstreamProvider:
}, },
) )
result = self._generic_streaming_response( result = await self._generic_streaming_response(
response, response,
key.hashed_key, key.hashed_key,
max_cost_for_model, max_cost_for_model,
+19
View File
@@ -7,6 +7,7 @@ cooldown that outlives a restart would hide a provider that has recovered.
from __future__ import annotations from __future__ import annotations
import time import time
from typing import Any
from ..core import get_logger from ..core import get_logger
from ..core.settings import settings from ..core.settings import settings
@@ -19,6 +20,24 @@ _failures: dict[tuple[str, str], list[float]] = {}
_cooling_until: dict[tuple[str, str], 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: 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.""" """Count a timeout or 5xx, opening a cooldown once too many land in a minute."""
if settings.upstream_cooldown_seconds <= 0: if settings.upstream_cooldown_seconds <= 0:
+80 -32
View File
@@ -3,7 +3,7 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
from collections.abc import AsyncIterator from collections.abc import AsyncIterator, Callable
import httpx import httpx
@@ -11,25 +11,93 @@ from ..core import get_logger
from ..core.error_scope import UPSTREAM_ERROR_STATUS from ..core.error_scope import UPSTREAM_ERROR_STATUS
from ..core.exceptions import UpstreamError from ..core.exceptions import UpstreamError
from ..core.settings import settings from ..core.settings import settings
from .sse_splitter import SSEEventSplitter
logger = get_logger(__name__) logger = get_logger(__name__)
async def open_guarded_stream( class GuardedStream(AsyncIterator[bytes]):
response: httpx.Response, provider_type: str def __init__(
) -> AsyncIterator[bytes]: self,
"""Await the upstream's first chunk, then hand back the whole stream. 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)
Awaiting the first chunk before any ``StreamingResponse`` exists is what def __aiter__(self) -> GuardedStream:
makes a slow-starting provider recoverable: the proxy's candidate loop only return self
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 async def __anext__(self) -> bytes:
iterator simply ends and the caller's finalizer settles actual usage. 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__() chunks = response.aiter_bytes().__aiter__()
guarded_chunks = _sse_events(chunks) if sse else chunks
timeout = settings.upstream_first_token_timeout_seconds timeout = settings.upstream_first_token_timeout_seconds
try: try:
first = await _next_chunk(chunks, timeout) first = await _next_chunk(guarded_chunks, timeout)
except TimeoutError: except TimeoutError:
await response.aclose() await response.aclose()
raise UpstreamError( raise UpstreamError(
@@ -37,7 +105,7 @@ async def open_guarded_stream(
status_code=UPSTREAM_ERROR_STATUS, status_code=UPSTREAM_ERROR_STATUS,
code="UPSTREAM_TIMEOUT", code="UPSTREAM_TIMEOUT",
) from None ) from None
return _resume(first, chunks, provider_type) return GuardedStream(first, guarded_chunks, provider_type, on_idle_timeout)
async def _next_chunk(chunks: AsyncIterator[bytes], timeout: float) -> bytes | None: async def _next_chunk(chunks: AsyncIterator[bytes], timeout: float) -> bytes | None:
@@ -47,23 +115,3 @@ async def _next_chunk(chunks: AsyncIterator[bytes], timeout: float) -> bytes | N
return await (asyncio.wait_for(step, timeout) if timeout > 0 else step) return await (asyncio.wait_for(step, timeout) if timeout > 0 else step)
except StopAsyncIteration: except StopAsyncIteration:
return None 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
@@ -316,7 +316,7 @@ async def test_streaming_response_closes_iterator_when_downstream_send_is_cancel
reservation = MagicMock(spec=ReservationSnapshot) reservation = MagicMock(spec=ReservationSnapshot)
upstream_response.status_code = 201 upstream_response.status_code = 201
upstream_response.headers = {"x-upstream": "preserved"} upstream_response.headers = {"x-upstream": "preserved"}
response = provider._generic_streaming_response( response = await provider._generic_streaming_response(
upstream_response, upstream_response,
"key-hash", "key-hash",
500, 500,
@@ -377,7 +377,7 @@ async def test_generic_stream_settles_when_response_start_fails() -> None:
upstream_response.status_code = 201 upstream_response.status_code = 201
upstream_response.headers = {"x-upstream": "preserved"} upstream_response.headers = {"x-upstream": "preserved"}
reservation = MagicMock(spec=ReservationSnapshot) reservation = MagicMock(spec=ReservationSnapshot)
response = provider._generic_streaming_response( response = await provider._generic_streaming_response(
upstream_response, upstream_response,
"key-hash", "key-hash",
500, 500,
+251 -3
View File
@@ -9,8 +9,14 @@ from unittest.mock import AsyncMock, MagicMock, patch
import httpx import httpx
import pytest 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.exceptions import UpstreamError
from routstr.core.settings import settings from routstr.core.settings import settings
from routstr.upstream.base import BaseUpstreamProvider
from routstr.upstream.cooldown import is_cooling_down, record_failure from routstr.upstream.cooldown import is_cooling_down, record_failure
from routstr.upstream.stream_timeout import open_guarded_stream from routstr.upstream.stream_timeout import open_guarded_stream
@@ -33,6 +39,12 @@ async def _stalls_after_first() -> AsyncIterator[bytes]:
yield b"never delivered" 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 @pytest.fixture
def fast_timeouts(monkeypatch: pytest.MonkeyPatch) -> None: def fast_timeouts(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(settings, "upstream_first_token_timeout_seconds", 0.01) monkeypatch.setattr(settings, "upstream_first_token_timeout_seconds", 0.01)
@@ -53,6 +65,74 @@ async def test_first_token_timeout_closes_response_and_raises(
response.aclose.assert_awaited_once() 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 @pytest.mark.asyncio
async def test_zero_first_token_timeout_disables_the_guard( async def test_zero_first_token_timeout_disables_the_guard(
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
@@ -80,6 +160,67 @@ async def test_idle_timeout_ends_the_stream_without_raising(
assert [chunk async for chunk in stream] == [b"first"] assert [chunk async for chunk in stream] == [b"first"]
@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 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 @pytest.mark.asyncio
async def test_guarded_stream_passes_every_chunk_through() -> None: async def test_guarded_stream_passes_every_chunk_through() -> None:
async def _chunks() -> AsyncIterator[bytes]: async def _chunks() -> AsyncIterator[bytes]:
@@ -135,6 +276,7 @@ def _upstream(base_url: str, forward: AsyncMock) -> MagicMock:
upstream = MagicMock() upstream = MagicMock()
upstream.provider_type = "test" upstream.provider_type = "test"
upstream.base_url = base_url upstream.base_url = base_url
upstream.db_id = None
upstream.prepare_headers = MagicMock(side_effect=lambda h: h) upstream.prepare_headers = MagicMock(side_effect=lambda h: h)
upstream.forward_request = forward upstream.forward_request = forward
return upstream return upstream
@@ -143,6 +285,7 @@ def _upstream(base_url: str, forward: AsyncMock) -> MagicMock:
async def _run_proxy( async def _run_proxy(
candidates: list[tuple[MagicMock, MagicMock]], candidates: list[tuple[MagicMock, MagicMock]],
revert_mock: AsyncMock, revert_mock: AsyncMock,
request: MagicMock | None = None,
) -> Any: ) -> Any:
from routstr import proxy as proxy_module from routstr import proxy as proxy_module
from routstr.auth import ReservationSnapshot from routstr.auth import ReservationSnapshot
@@ -173,7 +316,7 @@ async def _run_proxy(
), ),
patch.object(proxy_module, "revert_pay_for_request", revert_mock), patch.object(proxy_module, "revert_pay_for_request", revert_mock),
): ):
request = _proxy_request() request = request or _proxy_request()
return await proxy_module._proxy( return await proxy_module._proxy(
request, "v1/chat/completions", MagicMock(), await request.body() request, "v1/chat/completions", MagicMock(), await request.body()
) )
@@ -232,7 +375,7 @@ async def test_cooling_down_candidate_is_skipped_then_recovers(
healthy = _upstream("https://ok.example", AsyncMock(return_value=healthy_response)) healthy = _upstream("https://ok.example", AsyncMock(return_value=healthy_response))
candidates = [(MagicMock(), sick), (MagicMock(), healthy)] candidates = [(MagicMock(), sick), (MagicMock(), healthy)]
record_failure("https://sick.example", "test-model") record_failure("test|https://sick.example", "test-model")
assert await _run_proxy(candidates, AsyncMock()) is healthy_response assert await _run_proxy(candidates, AsyncMock()) is healthy_response
sick.forward_request.assert_not_awaited() sick.forward_request.assert_not_awaited()
@@ -251,6 +394,111 @@ async def test_cooldown_never_empties_the_candidate_list(
only_response.status_code = 200 only_response.status_code = 200
only = _upstream("https://only.example", AsyncMock(return_value=only_response)) only = _upstream("https://only.example", AsyncMock(return_value=only_response))
record_failure("https://only.example", "test-model") record_failure("test|https://only.example", "test-model")
assert await _run_proxy([(MagicMock(), only)], AsyncMock()) is only_response 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")