From a1071f4a98d2813a15f46a88ef445e29c96e4d4e Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 29 Sep 2026 02:58:40 +0200 Subject: [PATCH] fix: guard meaningful stream events and scope cooldown failures --- routstr/proxy.py | 58 +++- routstr/upstream/base.py | 69 ++++- routstr/upstream/cooldown.py | 19 ++ routstr/upstream/stream_timeout.py | 112 +++++--- .../test_streaming_billing_finalization.py | 4 +- tests/unit/test_upstream_stream_timeout.py | 254 +++++++++++++++++- 6 files changed, 455 insertions(+), 61 deletions(-) diff --git a/routstr/proxy.py b/routstr/proxy.py index 94ff74a9..c6dd9628 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, @@ -39,7 +41,12 @@ from .payment.helpers import ( ) from .payment.models import Model 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.helpers import init_upstreams from .upstream.model_paths import ( @@ -116,8 +123,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): @@ -410,6 +415,13 @@ def _counts_toward_cooldown(status_code: int) -> bool: 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: @@ -705,7 +717,10 @@ async def _proxy( healthy = [ candidate 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: candidates = healthy @@ -737,7 +752,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, @@ -746,7 +761,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, @@ -755,7 +770,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, @@ -763,6 +778,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", @@ -775,6 +796,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 @@ -1044,8 +1072,11 @@ async def _proxy( break if response.status_code != 200: - if _counts_toward_cooldown(response.status_code): - record_failure(upstream.base_url, model_id) + 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 [ @@ -1133,8 +1164,13 @@ async def _proxy( raise except UpstreamError as e: - if _counts_toward_cooldown(e.status_code): - record_failure(upstream.base_url, model_id) + 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 984ca53e..7e5b67ef 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,7 +87,7 @@ from .stream_ownership import ( close_upstream_exchange, finalize_and_close_stream, ) -from .stream_timeout import open_guarded_stream +from .stream_timeout import GuardedStream, open_guarded_stream if typing.TYPE_CHECKING: from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget @@ -1127,6 +1129,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, @@ -1148,7 +1161,7 @@ class BaseUpstreamProvider: Returns: 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: async with create_session() as snapshot_session: @@ -1436,7 +1449,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: @@ -1642,7 +1657,7 @@ class BaseUpstreamProvider: Returns: 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) @@ -1844,7 +1859,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", "response": { "model": last_model_seen or "unknown", "usage": { @@ -1865,6 +1882,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( @@ -1880,7 +1905,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: @@ -2125,6 +2155,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: @@ -2142,14 +2173,22 @@ class BaseUpstreamProvider: ) ) try: - # 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): + 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, @@ -2159,6 +2198,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( @@ -2181,6 +2221,7 @@ class BaseUpstreamProvider: provider_fee, reservation_snapshot, finalizer, + guarded_chunks, ) return ClosingStreamingResponse( stream, @@ -2199,7 +2240,7 @@ class BaseUpstreamProvider: reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, ) -> 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_finalized = False @@ -2469,6 +2510,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: @@ -3421,7 +3464,7 @@ class BaseUpstreamProvider: }, ) - result = self._generic_streaming_response( + result = await self._generic_streaming_response( response, key.hashed_key, max_cost_for_model, @@ -3706,7 +3749,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 index bb4984fc..7031817c 100644 --- a/routstr/upstream/cooldown.py +++ b/routstr/upstream/cooldown.py @@ -7,6 +7,7 @@ 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 @@ -19,6 +20,24 @@ _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: diff --git a/routstr/upstream/stream_timeout.py b/routstr/upstream/stream_timeout.py index ece38c32..d6256da4 100644 --- a/routstr/upstream/stream_timeout.py +++ b/routstr/upstream/stream_timeout.py @@ -3,7 +3,7 @@ from __future__ import annotations import asyncio -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Callable import httpx @@ -11,25 +11,93 @@ 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__) -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. +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) - 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. + 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(chunks, timeout) + first = await _next_chunk(guarded_chunks, timeout) except TimeoutError: await response.aclose() raise UpstreamError( @@ -37,7 +105,7 @@ async def open_guarded_stream( status_code=UPSTREAM_ERROR_STATUS, code="UPSTREAM_TIMEOUT", ) 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: @@ -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) 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_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 index 83a1ccbb..7c42df41 100644 --- a/tests/unit/test_upstream_stream_timeout.py +++ b/tests/unit/test_upstream_stream_timeout.py @@ -9,8 +9,14 @@ 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 +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 @@ -33,6 +39,12 @@ async def _stalls_after_first() -> AsyncIterator[bytes]: 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) @@ -53,6 +65,74 @@ async def test_first_token_timeout_closes_response_and_raises( 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, @@ -80,6 +160,67 @@ async def test_idle_timeout_ends_the_stream_without_raising( 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 async def test_guarded_stream_passes_every_chunk_through() -> None: async def _chunks() -> AsyncIterator[bytes]: @@ -135,6 +276,7 @@ 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 @@ -143,6 +285,7 @@ def _upstream(base_url: str, forward: AsyncMock) -> MagicMock: 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 @@ -173,7 +316,7 @@ async def _run_proxy( ), patch.object(proxy_module, "revert_pay_for_request", revert_mock), ): - request = _proxy_request() + request = request or _proxy_request() return await proxy_module._proxy( 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)) 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 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 = _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 + + +@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")