diff --git a/.env.example b/.env.example index e79b7770..a063da56 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/routstr/core/settings.py b/routstr/core/settings.py index d9324f13..d77d6fc1 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -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") diff --git a/routstr/proxy.py b/routstr/proxy.py index c585188f..681a3aeb 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -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, diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 2f65ebd3..d94610d0 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -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, diff --git a/routstr/upstream/cooldown.py b/routstr/upstream/cooldown.py new file mode 100644 index 00000000..7031817c --- /dev/null +++ b/routstr/upstream/cooldown.py @@ -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() diff --git a/routstr/upstream/stream_timeout.py b/routstr/upstream/stream_timeout.py new file mode 100644 index 00000000..d6256da4 --- /dev/null +++ b/routstr/upstream/stream_timeout.py @@ -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 diff --git a/tests/conftest.py b/tests/conftest.py index d1bfa919..0e0ccbb3 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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() diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index f8babb17..badaf549 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -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 diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index 10c5a3e8..eef0c44f 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -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, diff --git a/tests/unit/test_upstream_stream_timeout.py b/tests/unit/test_upstream_stream_timeout.py new file mode 100644 index 00000000..7942337c --- /dev/null +++ b/tests/unit/test_upstream_stream_timeout.py @@ -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")