diff --git a/routstr/core/settings.py b/routstr/core/settings.py index da503a50..2f57d7a1 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -40,6 +40,21 @@ class Settings(BaseSettings): upstream_5xx_retry_attempts: int = Field( 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 name: str = Field(default="ARoutstrNode", env="NAME") diff --git a/routstr/proxy.py b/routstr/proxy.py index 26bf20e7..5be15e62 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -40,6 +40,7 @@ from .payment.helpers import ( ) from .payment.models import Model 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.helpers import init_upstreams from .upstream.model_paths import ( @@ -405,6 +406,11 @@ _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 + + @proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None) async def proxy( request: Request, path: str, session: AsyncSession = Depends(get_session) @@ -621,6 +627,17 @@ 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(candidate[1].base_url, 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. @@ -950,6 +967,8 @@ async def _proxy( break 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. # 502/503 are upstream errors, 429 rate limits. should_retry = response.status_code in [ @@ -1037,6 +1056,8 @@ async def _proxy( raise except UpstreamError as e: + if _counts_toward_cooldown(e.status_code): + record_failure(upstream.base_url, 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 5b0e5c0f..984ca53e 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -85,6 +85,7 @@ from .stream_ownership import ( close_upstream_exchange, finalize_and_close_stream, ) +from .stream_timeout import open_guarded_stream if typing.TYPE_CHECKING: from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget @@ -1147,6 +1148,8 @@ class BaseUpstreamProvider: Returns: StreamingResponse with cost data injected at the end """ + guarded_chunks = await open_guarded_stream(response, self.provider_type) + if reservation_snapshot is None: async with create_session() as snapshot_session: 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 # 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 @@ -1639,6 +1642,8 @@ class BaseUpstreamProvider: Returns: 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) logger.debug( @@ -1790,7 +1795,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 @@ -2137,7 +2142,9 @@ class BaseUpstreamProvider: ) ) 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 finally: await finalizer.run() @@ -2192,6 +2199,8 @@ class BaseUpstreamProvider: reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, ) -> StreamingResponse: + guarded_chunks = await open_guarded_stream(response, self.provider_type) + usage_estimator = MissingUsageEstimator(request_body, model_obj) usage_finalized = False last_model_seen: str | None = None @@ -2284,7 +2293,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") diff --git a/routstr/upstream/cooldown.py b/routstr/upstream/cooldown.py new file mode 100644 index 00000000..bb4984fc --- /dev/null +++ b/routstr/upstream/cooldown.py @@ -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() diff --git a/routstr/upstream/stream_timeout.py b/routstr/upstream/stream_timeout.py new file mode 100644 index 00000000..ece38c32 --- /dev/null +++ b/routstr/upstream/stream_timeout.py @@ -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 diff --git a/tests/unit/test_upstream_stream_timeout.py b/tests/unit/test_upstream_stream_timeout.py new file mode 100644 index 00000000..2dd25310 --- /dev/null +++ b/tests/unit/test_upstream_stream_timeout.py @@ -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