From ca4e93a819061af2a52c515813afc48996dd4384 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 27 Sep 2026 03:10:10 +0200 Subject: [PATCH 01/12] feat: add upstream first-token timeout, stream idle timeout and provider cooldown --- routstr/core/settings.py | 15 ++ routstr/proxy.py | 21 ++ routstr/upstream/base.py | 17 +- routstr/upstream/cooldown.py | 60 +++++ routstr/upstream/stream_timeout.py | 69 ++++++ tests/unit/test_upstream_stream_timeout.py | 262 +++++++++++++++++++++ 6 files changed, 440 insertions(+), 4 deletions(-) create mode 100644 routstr/upstream/cooldown.py create mode 100644 routstr/upstream/stream_timeout.py create mode 100644 tests/unit/test_upstream_stream_timeout.py 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 From e9db5ffd302967b6316908f78cf2933226f707b0 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 27 Sep 2026 11:29:59 +0200 Subject: [PATCH 02/12] test: reset upstream cooldowns between all tests --- tests/conftest.py | 10 ++++++++++ tests/unit/test_upstream_stream_timeout.py | 9 +-------- 2 files changed, 11 insertions(+), 8 deletions(-) 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/unit/test_upstream_stream_timeout.py b/tests/unit/test_upstream_stream_timeout.py index 2dd25310..0c5a57cf 100644 --- a/tests/unit/test_upstream_stream_timeout.py +++ b/tests/unit/test_upstream_stream_timeout.py @@ -11,7 +11,7 @@ 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.cooldown import is_cooling_down, record_failure from routstr.upstream.stream_timeout import open_guarded_stream @@ -33,13 +33,6 @@ async def _stalls_after_first() -> AsyncIterator[bytes]: 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) From a1071f4a98d2813a15f46a88ef445e29c96e4d4e Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 29 Sep 2026 02:58:40 +0200 Subject: [PATCH 03/12] 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") From 9b07b4a7c898d73e4b3478368602fd1a7c0110f0 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 29 Sep 2026 20:33:08 +0200 Subject: [PATCH 04/12] fix: normalize streamed test chunks for mypy --- tests/unit/test_upstream_stream_timeout.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/tests/unit/test_upstream_stream_timeout.py b/tests/unit/test_upstream_stream_timeout.py index 7c42df41..867677b8 100644 --- a/tests/unit/test_upstream_stream_timeout.py +++ b/tests/unit/test_upstream_stream_timeout.py @@ -213,7 +213,12 @@ async def test_responses_idle_timeout_does_not_emit_completed( result = await provider.handle_streaming_responses_completion( response, key, 100, reservation_snapshot=MagicMock() ) - emitted = b"".join([chunk async for chunk in result.body_iterator]) + 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 From 40bee57be5c3cff5c2699f577e78521b1299bab2 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 29 Sep 2026 22:57:23 +0200 Subject: [PATCH 05/12] test: use production sqlite busy timeout in integration engine --- tests/integration/conftest.py | 5 +++++ 1 file changed, 5 insertions(+) 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 From e589b37bc56e9dca7d56082d83e35279a53c104a Mon Sep 17 00:00:00 2001 From: 9qeklajc <211699015+9qeklajc@users.noreply.github.com> Date: Wed, 30 Sep 2026 13:35:01 +0000 Subject: [PATCH 06/12] feat(upstream): add native DeepSeek provider, retire V4 pricing shim Add a first-class `deepseek` upstream provider that needs only an API key; the base URL is fixed to https://api.deepseek.com. DEEPSEEK_API_KEY seeds it on startup. Models come from DeepSeek's own /models and are priced from a peak-rate table in routstr/upstream/deepseek.py. A listed model the table misses is imported disabled instead of taking a litellm/OpenRouter price, via a new GenericUpstreamProvider.use_fallback_pricing switch (default True, so other providers are unchanged). The node bills one flat price per model, so the table holds the peak rate and never bills below DeepSeek's cost. Thinking-mode reasoning_content is forwarded unchanged: DeepSeek requires it on requests that carry tools and ignores it otherwise. Remove the temporary DeepSeek V4 pricing shim and its startup call. The pinned litellm 1.101.2 already ships every key it filled. --- docs/provider/configuration.md | 23 ++ routstr/core/main.py | 6 - routstr/upstream/__init__.py | 2 + routstr/upstream/deepseek.py | 105 +++++++++ routstr/upstream/deepseek_v4_pricing_shim.py | 73 ------ routstr/upstream/generic.py | 8 +- routstr/upstream/helpers.py | 1 + tests/integration/test_secret_bootstrap.py | 1 - tests/unit/test_cache_pricing.py | 9 - tests/unit/test_upstream_deepseek.py | 232 +++++++++++++++++++ 10 files changed, 369 insertions(+), 91 deletions(-) create mode 100644 routstr/upstream/deepseek.py delete mode 100644 routstr/upstream/deepseek_v4_pricing_shim.py create mode 100644 tests/unit/test_upstream_deepseek.py diff --git a/docs/provider/configuration.md b/docs/provider/configuration.md index 6a47854d..d2695fd9 100644 --- a/docs/provider/configuration.md +++ b/docs/provider/configuration.md @@ -48,6 +48,29 @@ Connect to your AI provider(s): | **Upstream URL** | API endpoint (e.g., `https://api.openai.com/v1`) | | **API Key** | Your provider's API key | +### DeepSeek + +Choose **DeepSeek** as the provider type and paste an API key from +[platform.deepseek.com](https://platform.deepseek.com/api_keys); the base URL +is fixed to `https://api.deepseek.com`. Setting `DEEPSEEK_API_KEY` seeds the +provider on startup instead. + +Models are listed from DeepSeek's own `/models` and priced from a rate table +in `routstr/upstream/deepseek.py`, not from litellm or OpenRouter: + +- **Peak rates only.** DeepSeek charges half price off-peak, but the node bills + one flat price per model, so it bills the peak rate. Clients overpay + off-peak; the node never bills below cost. Time-of-day pricing is planned. +- **Unknown models import disabled.** A model DeepSeek lists that the table + does not price shows up disabled in the Admin Dashboard. Enable it with a + manual price, or add it to the table. +- **Cache hits** bill at DeepSeek's cache-hit rate (about 2% of the input + rate). + +Thinking-mode `reasoning_content` is returned to clients unchanged in +responses, and forwarded unchanged when it appears in conversation history. +DeepSeek requires it on requests that carry `tools` and ignores it otherwise. + ### PPQ Auto Top-up PPQ providers can automatically purchase more credits when their USD balance diff --git a/routstr/core/main.py b/routstr/core/main.py index 584fa736..424025f0 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -34,7 +34,6 @@ from ..payment.price import update_prices_periodically from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically from ..refund import periodic_refund_reconcile from ..upstream.auto_topup import periodic_auto_topup -from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing from ..upstream.http_client import close_upstream_http_client from ..upstream.litellm_routing import configure_litellm from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout @@ -89,11 +88,6 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: # debug logging) before any upstream provider dispatches a request. configure_litellm() - # TEMPORARY: backfill DeepSeek V4 pricing missing from litellm's cost - # map (BerriAI/litellm#30430). Remove this call and - # deepseek_v4_pricing_shim.py once litellm ships these models. - register_deepseek_v4_pricing() - # Run database migrations on startup run_migrations() diff --git a/routstr/upstream/__init__.py b/routstr/upstream/__init__.py index 85d094e8..c57e0c09 100644 --- a/routstr/upstream/__init__.py +++ b/routstr/upstream/__init__.py @@ -1,6 +1,7 @@ from .anthropic import AnthropicUpstreamProvider from .azure import AzureUpstreamProvider from .base import BaseUpstreamProvider +from .deepseek import DeepSeekUpstreamProvider from .fireworks import FireworksUpstreamProvider from .gemini import GeminiUpstreamProvider from .generic import GenericUpstreamProvider @@ -19,6 +20,7 @@ from .xai import XAIUpstreamProvider upstream_provider_classes: list[type[BaseUpstreamProvider]] = [ AnthropicUpstreamProvider, AzureUpstreamProvider, + DeepSeekUpstreamProvider, FireworksUpstreamProvider, GeminiUpstreamProvider, GenericUpstreamProvider, diff --git a/routstr/upstream/deepseek.py b/routstr/upstream/deepseek.py new file mode 100644 index 00000000..01e5e334 --- /dev/null +++ b/routstr/upstream/deepseek.py @@ -0,0 +1,105 @@ +"""First-class upstream for the DeepSeek API. + +Pricing comes from ``_PEAK_RATES`` below, not from litellm or OpenRouter: +litellm's bundled ``deepseek-v4-flash`` entry is stale, the OpenRouter feed +carries resale prices below DeepSeek's own peak rate, and neither knows the +current ``deepseek-flash`` id. A model DeepSeek lists that the table does not +cover is imported disabled rather than priced from those sources. + +DeepSeek bills peak hours at twice the off-peak rate. The node has one flat +price per model, so the table holds the PEAK rates: a client may overpay +off-peak but the node never bills below its own cost. + +Rates: https://api-docs.deepseek.com/quick_start/pricing (checked 2026-09-30). +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +from .base import BaseUpstreamProvider +from .generic import GenericUpstreamProvider +from .pricing_resolver import ResolvedPricing + +if TYPE_CHECKING: + from ..core.db import UpstreamProviderRow + +_CONTEXT_LENGTH = 1_000_000 +_MAX_OUTPUT_TOKENS = 384_000 + +# USD per 1M tokens at DeepSeek's peak rate: (input cache miss, output, input +# cache hit). DeepSeek has no cache-write charge. +_FLASH = (0.30, 1.20, 0.006) +_PRO = (1.32, 3.96, 0.044) + +_PEAK_RATES: dict[str, tuple[float, float, float]] = { + "deepseek-flash": _FLASH, + # Retired ids DeepSeek still accepts, served and billed as deepseek-flash. + "deepseek-v4-flash": _FLASH, + "deepseek-v4-flash-vision-exp": _FLASH, + "deepseek-v4-pro": _PRO, +} + +# Pro is the only current model without vision support. +_TEXT_ONLY = {"deepseek-v4-pro"} + + +class DeepSeekUpstreamProvider(GenericUpstreamProvider): + """Upstream provider specifically configured for the DeepSeek API.""" + + provider_type = "deepseek" + default_base_url = "https://api.deepseek.com" + platform_url = "https://platform.deepseek.com/api_keys" + litellm_provider_prefix = "deepseek/" + use_fallback_pricing = False + + def __init__(self, api_key: str, provider_fee: float = 1.01): + super().__init__( + base_url=self.default_base_url, + api_key=api_key, + provider_fee=provider_fee, + upstream_name="DeepSeek", + ) + + @classmethod + def _build_from_row( + cls, provider_row: "UpstreamProviderRow" + ) -> "DeepSeekUpstreamProvider": + return cls(api_key=provider_row.api_key, provider_fee=provider_row.provider_fee) + + @classmethod + def get_provider_metadata(cls) -> dict[str, object]: + return { + "id": cls.provider_type, + "name": "DeepSeek", + "default_base_url": cls.default_base_url, + "fixed_base_url": True, + "platform_url": cls.platform_url, + } + + def _apply_provider_field(self, response_json: object) -> None: + # A first-party upstream: stamp "deepseek", not Generic's hostname. + BaseUpstreamProvider._apply_provider_field(self, response_json) + + def transform_model_name(self, model_id: str) -> str: + """Strip the 'deepseek/' prefix for DeepSeek API compatibility.""" + return model_id.removeprefix("deepseek/") + + def _native_pricing( + self, model_id: str, model_spec: dict + ) -> ResolvedPricing | None: + """Price ``model_id`` from the peak-rate table; ``None`` if absent.""" + rates = _PEAK_RATES.get(model_id) + if rates is None: + return None + input_usd, output_usd, cache_hit_usd = rates + input_modalities = ["text"] if model_id in _TEXT_ONLY else ["text", "image"] + return ResolvedPricing( + prompt=input_usd / 1_000_000, + completion=output_usd / 1_000_000, + context_length=_CONTEXT_LENGTH, + source="native", + max_completion_tokens=_MAX_OUTPUT_TOKENS, + input_cache_read=cache_hit_usd / 1_000_000, + input_modalities=input_modalities, + ) diff --git a/routstr/upstream/deepseek_v4_pricing_shim.py b/routstr/upstream/deepseek_v4_pricing_shim.py deleted file mode 100644 index ba0c392d..00000000 --- a/routstr/upstream/deepseek_v4_pricing_shim.py +++ /dev/null @@ -1,73 +0,0 @@ -"""TEMPORARY: local DeepSeek V4 pricing shim. - -litellm's bundled cost map does not yet ship ``deepseek-v4-flash`` / -``deepseek-v4-pro``. Without an entry, ``backfill_cache_pricing`` cannot find a -``cache_read_input_token_cost`` and cache reads fall back to the full input -rate — a large overcharge on cache hits (DeepSeek V4 hits are ~0.008-0.02x -input, i.e. cached tokens cost 50-120x less than regular input). - -This module injects the missing entries into ``litellm.model_cost`` at startup -so the existing backfill path resolves them. Rates mirror the canonical -``deepseek`` provider entries now in litellm's ``model_prices`` map -(``input_cost_per_token`` is the cache-*miss* rate; -``cache_read_input_token_cost`` is the cache-*hit* rate), sourced from -https://api-docs.deepseek.com/quick_start/pricing via -https://github.com/BerriAI/litellm/pull/26380 (issue -https://github.com/BerriAI/litellm/issues/30430). - -=== REMOVAL (once litellm ships these models) === -Delete this file and the single ``register_deepseek_v4_pricing()`` call in -``routstr/core/main.py``. Nothing else depends on it. Entries are only added -when absent, so a stale shim is harmless after upstream lands — but remove it. -""" - -import litellm - -from ..core import get_logger - -logger = get_logger(__name__) - -# USD per token. Mirrors the canonical ``deepseek`` provider entries in -# litellm's model_prices map (source: DeepSeek API pricing docs). Keep these in -# sync with ``litellm.model_cost["deepseek/deepseek-v4-*"]``. -_DEEPSEEK_V4_RATES: dict[str, dict[str, float]] = { - "deepseek-v4-flash": { - "input_cost_per_token": 1.4e-07, - "output_cost_per_token": 2.8e-07, - "cache_read_input_token_cost": 2.8e-09, - "cache_creation_input_token_cost": 0.0, - "input_cost_per_token_cache_hit": 2.8e-09, - }, - "deepseek-v4-pro": { - "input_cost_per_token": 4.35e-07, - "output_cost_per_token": 8.7e-07, - "cache_read_input_token_cost": 3.625e-09, - "cache_creation_input_token_cost": 0.0, - "input_cost_per_token_cache_hit": 3.625e-09, - }, -} - - -def register_deepseek_v4_pricing() -> None: - """Inject DeepSeek V4 pricing into ``litellm.model_cost`` if absent. - - Idempotent and non-destructive: a key already present in the cost map - (e.g. once litellm ships it) is left untouched. Registers both the bare - (``deepseek-v4-flash``) and prefixed (``deepseek/deepseek-v4-flash``) - spellings since ``backfill_cache_pricing`` tries both. - """ - added = [] - for bare, rates in _DEEPSEEK_V4_RATES.items(): - for key in (bare, f"deepseek/{bare}"): - if key in litellm.model_cost: - continue - entry: dict[str, object] = dict(rates) - entry["litellm_provider"] = "deepseek" - entry["mode"] = "chat" - litellm.model_cost[key] = entry - added.append(key) - if added: - logger.info( - "Registered temporary DeepSeek V4 pricing shim", - extra={"models": added}, - ) diff --git a/routstr/upstream/generic.py b/routstr/upstream/generic.py index c9edf109..1032e85d 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -28,7 +28,11 @@ class GenericUpstreamProvider(BaseUpstreamProvider): provider_type = "generic" default_base_url = "http://localhost:8888" - platform_url = None + platform_url: str | None = None + # Subclasses that own an authoritative price table set this False so a model + # the table misses imports disabled instead of taking a litellm/OpenRouter + # price that may undercut the upstream's own rate. + use_fallback_pricing = True def __init__( self, @@ -162,7 +166,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider): model_spec = model_data.get("model_spec", {}) resolved = self._native_pricing(model_id, model_spec) - if resolved is None: + if resolved is None and self.use_fallback_pricing: resolved = await resolver.resolve(model_id) if resolved is None: diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index dc544b3a..c94733bc 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -272,6 +272,7 @@ async def _seed_providers_from_settings( ("PERPLEXITY_API_KEY", "perplexity", None, None), ("FIREWORKS_API_KEY", "fireworks", None, None), ("XAI_API_KEY", "xai", None, None), + ("DEEPSEEK_API_KEY", "deepseek", None, None), ("TINFOIL_API_KEY", "tinfoil", None, None), ("TYPESAFE_API_KEY", "typesafe", None, None), ] diff --git a/tests/integration/test_secret_bootstrap.py b/tests/integration/test_secret_bootstrap.py index 12f0d9a1..6dcdb990 100644 --- a/tests/integration/test_secret_bootstrap.py +++ b/tests/integration/test_secret_bootstrap.py @@ -470,7 +470,6 @@ async def test_startup_runs_bootstrap_before_settings_initialize( return None monkeypatch.setattr(main, "configure_litellm", lambda: None) - monkeypatch.setattr(main, "register_deepseek_v4_pricing", lambda: None) monkeypatch.setattr(main, "run_migrations", lambda: None) monkeypatch.setattr(main, "init_db", noop_init_db) monkeypatch.setattr(main, "create_session", fake_create_session) diff --git a/tests/unit/test_cache_pricing.py b/tests/unit/test_cache_pricing.py index cc4a4dbc..cfb7e9a6 100644 --- a/tests/unit/test_cache_pricing.py +++ b/tests/unit/test_cache_pricing.py @@ -31,15 +31,6 @@ from routstr.payment.models import ( backfill_cache_pricing, ) from routstr.upstream import GenericUpstreamProvider -from routstr.upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing - - -@pytest.fixture(autouse=True) -def _deepseek_v4_pricing() -> None: - # litellm's bundled cost map lacks the DeepSeek V4 entries (they only - # appear when its remote map is reachable); production injects them at - # startup via this same shim. - register_deepseek_v4_pricing() def _make_model(model_id: str, pricing: Pricing) -> Model: diff --git a/tests/unit/test_upstream_deepseek.py b/tests/unit/test_upstream_deepseek.py new file mode 100644 index 00000000..f310f7ac --- /dev/null +++ b/tests/unit/test_upstream_deepseek.py @@ -0,0 +1,232 @@ +"""Unit tests for ``DeepSeekUpstreamProvider``. + +DeepSeek is priced from the provider's own peak-rate table, never from litellm +or OpenRouter: litellm's ``deepseek-v4-flash`` entry is stale and OpenRouter +resells below DeepSeek's peak rate, so either would bill under cost. These +tests pin the table prices (including the cache-hit rate), that a model the +table misses imports disabled without consulting the fallback chain, and that +``reasoning_content`` in history reaches DeepSeek untouched — thinking mode +with ``tools`` answers 400 when it is stripped. +""" + +from __future__ import annotations + +import json +from typing import Any +from unittest.mock import AsyncMock, Mock, patch + +import pytest + +from routstr.upstream import upstream_provider_classes +from routstr.upstream.deepseek import DeepSeekUpstreamProvider + + +class _FakeResponse: + def __init__(self, payload: dict[str, Any]) -> None: + self._payload = payload + + def raise_for_status(self) -> None: + return None + + def json(self) -> dict[str, Any]: + return self._payload + + +class _FakeAsyncClient: + def __init__(self, payload: dict[str, Any], calls: list[dict[str, Any]]) -> None: + self._payload = payload + self._calls = calls + + async def __aenter__(self) -> "_FakeAsyncClient": + return self + + async def __aexit__(self, *exc: object) -> bool: + return False + + async def get( + self, url: str, headers: dict[str, str] | None = None + ) -> _FakeResponse: + self._calls.append({"url": url, "headers": headers}) + return _FakeResponse(self._payload) + + +# Shape of DeepSeek's ``GET /models``: bare ids, no pricing. +CATALOG: dict[str, Any] = { + "object": "list", + "data": [ + {"id": "deepseek-flash", "object": "model", "owned_by": "deepseek"}, + {"id": "deepseek-v4-pro", "object": "model", "owned_by": "deepseek"}, + {"id": "deepseek-v4-flash", "object": "model", "owned_by": "deepseek"}, + {"id": "deepseek-chat", "object": "model", "owned_by": "deepseek"}, + ], +} + + +async def _fetch( + catalog: dict[str, Any] = CATALOG, +) -> tuple[dict[str, Any], list[dict[str, Any]], AsyncMock]: + calls: list[dict[str, Any]] = [] + fallback = AsyncMock(return_value=None) + provider = DeepSeekUpstreamProvider(api_key="sk-test") + with ( + patch( + "routstr.upstream.generic.httpx.AsyncClient", + lambda *args, **kwargs: _FakeAsyncClient(catalog, calls), + ), + patch("routstr.upstream.generic.FallbackPricingResolver.resolve", fallback), + ): + models = await provider.fetch_models() + return {m.id: m for m in models}, calls, fallback + + +def test_metadata_and_registration() -> None: + assert DeepSeekUpstreamProvider in upstream_provider_classes + assert DeepSeekUpstreamProvider.get_provider_metadata() == { + "id": "deepseek", + "name": "DeepSeek", + "default_base_url": "https://api.deepseek.com", + "fixed_base_url": True, + "platform_url": "https://platform.deepseek.com/api_keys", + } + + +def test_build_from_row_ignores_row_base_url() -> None: + row = Mock( + api_key="sk-row", provider_fee=1.05, base_url="https://elsewhere.example" + ) + provider = DeepSeekUpstreamProvider._build_from_row(row) + assert provider.api_key == "sk-row" + assert provider.provider_fee == 1.05 + assert provider.base_url == "https://api.deepseek.com" + + +def test_litellm_prefix_is_deepseek() -> None: + provider = DeepSeekUpstreamProvider(api_key="sk-test") + assert provider.get_litellm_provider_prefix() == "deepseek/" + + +@pytest.mark.parametrize( + "model_id,expected", + [ + ("deepseek/deepseek-v4-flash", "deepseek-v4-flash"), + ("deepseek-v4-flash", "deepseek-v4-flash"), + ("deepseek/deepseek-flash", "deepseek-flash"), + ], +) +def test_transform_model_name(model_id: str, expected: str) -> None: + provider = DeepSeekUpstreamProvider(api_key="sk-test") + assert provider.transform_model_name(model_id) == expected + + +def test_provider_field_names_deepseek_not_host() -> None: + provider = DeepSeekUpstreamProvider(api_key="sk-test") + payload: dict[str, Any] = {"id": "chatcmpl-1"} + provider._apply_provider_field(payload) + assert payload["provider"] == "deepseek" + + +@pytest.mark.asyncio +async def test_fetch_models_calls_deepseek_models_endpoint_with_key() -> None: + _, calls, _ = await _fetch() + assert calls == [ + { + "url": "https://api.deepseek.com/models", + "headers": {"Authorization": "Bearer sk-test"}, + } + ] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "model_id,prompt,completion,cache_read", + [ + ("deepseek-flash", 0.30, 1.20, 0.006), + # Retired alias DeepSeek serves and bills as deepseek-flash. + ("deepseek-v4-flash", 0.30, 1.20, 0.006), + ("deepseek-v4-pro", 1.32, 3.96, 0.044), + ], +) +async def test_table_models_priced_at_peak_rate( + model_id: str, prompt: float, completion: float, cache_read: float +) -> None: + models, _, _ = await _fetch() + model = models[model_id] + assert model.enabled is True + assert model.pricing.prompt == pytest.approx(prompt / 1_000_000) + assert model.pricing.completion == pytest.approx(completion / 1_000_000) + assert model.pricing.input_cache_read == pytest.approx(cache_read / 1_000_000) + assert model.context_length == 1_000_000 + + +@pytest.mark.asyncio +async def test_vision_follows_the_model() -> None: + models, _, _ = await _fetch() + assert "image" in models["deepseek-flash"].architecture.input_modalities + assert models["deepseek-v4-pro"].architecture.input_modalities == ["text"] + + +@pytest.mark.asyncio +async def test_unlisted_model_imports_disabled_without_fallback() -> None: + """litellm prices ``deepseek-chat``; the provider must not take that price.""" + models, _, fallback = await _fetch() + model = models["deepseek-chat"] + assert model.enabled is False + assert model.pricing.prompt == 0.0 + assert model.pricing.completion == 0.0 + fallback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_cache_rate_survives_fee_and_is_not_replaced_by_litellm() -> None: + """litellm's stale ``deepseek-v4-flash`` cache rate (1.4e-08 in the bundled + map) must not replace the table's; backfill only fills an absent rate. The + fee applies to the cache rate like every other component. + + The litellm entry is pinned here because the remote cost map already + carries the table's rate, which would let an overwrite go unnoticed.""" + models, _, _ = await _fetch() + provider = DeepSeekUpstreamProvider(api_key="sk-test", provider_fee=1.05) + stale = {"cache_read_input_token_cost": 1.4e-08} + with patch("routstr.payment.models.litellm_cost_entry", return_value=stale): + priced = provider._apply_provider_fee_to_model(models["deepseek-v4-flash"]) + assert priced.pricing.input_cache_read == pytest.approx(0.006e-6 * 1.05) + assert priced.pricing.prompt == pytest.approx(0.30e-6 * 1.05) + # A cache hit costs 2% of a miss, not the full input rate. + assert priced.pricing.input_cache_read / priced.pricing.prompt == pytest.approx( + 0.02 + ) + + +@pytest.mark.asyncio +async def test_reasoning_content_in_history_reaches_upstream() -> None: + models, _, _ = await _fetch() + provider = DeepSeekUpstreamProvider(api_key="sk-test") + messages = [ + {"role": "user", "content": "weather in Paris?"}, + { + "role": "assistant", + "content": "", + "reasoning_content": "Need the weather tool.", + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "get_weather", "arguments": "{}"}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": "18C"}, + ] + body = json.dumps( + { + "model": "deepseek/deepseek-flash", + "messages": messages, + "tools": [{"type": "function", "function": {"name": "get_weather"}}], + } + ).encode() + out = provider.prepare_request_body(body, models["deepseek-flash"]) + + assert out is not None + sent = json.loads(out) + assert sent["model"] == "deepseek-flash" + assert sent["messages"] == messages From 50d2c3139929d9241f8818678dd62c3663fb8505 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 30 Sep 2026 23:56:01 +0200 Subject: [PATCH 07/12] fix: upgrade dependencies to address Dependabot alerts --- pyproject.toml | 3 + ui/package.json | 6 +- ui/pnpm-lock.yaml | 136 ++++++++++++++++++++--------------------- ui/pnpm-workspace.yaml | 4 +- uv.lock | 22 ++++--- 5 files changed, 88 insertions(+), 83 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 7daf7489..cc0de477 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -113,6 +113,9 @@ override-dependencies = [ # Transitive deps whose dependents allow the patched version but don't require # it. Constraints raise the floor without bypassing any upstream pin. constraint-dependencies = [ + "anyio>=4.14.2", + "pyjwt>=2.15.0", + "urllib3>=2.8.0", "starlette>=1.3.1", "httpcore>=1.0.9", # 1.0.8 caps h11<0.15 # 1.76 is the first grpcio-tools release with CPython 3.14 wheels. diff --git a/ui/package.json b/ui/package.json index 1d34f131..209c4fad 100644 --- a/ui/package.json +++ b/ui/package.json @@ -39,7 +39,7 @@ "@radix-ui/react-toggle-group": "^1.1.11", "@radix-ui/react-tooltip": "^1.2.8", "@tanstack/react-query": "^5.90.21", - "axios": "^1.16.0", + "axios": "^1.20.0", "class-variance-authority": "^0.7.1", "clsx": "^2.1.1", "cmdk": "^1.1.1", @@ -48,7 +48,7 @@ "geist": "^1.7.0", "input-otp": "^1.4.2", "lucide-react": "^0.575.0", - "next": "16.3.4", + "next": "16.3.6", "next-themes": "^0.4.6", "qrcode": "^1.5.4", "radix-ui": "^1.4.3", @@ -74,7 +74,7 @@ "@types/react": "^19.2.14", "@types/react-dom": "^19.2.3", "eslint": "^9.7.0", - "eslint-config-next": "16.3.4", + "eslint-config-next": "16.3.6", "eslint-config-prettier": "^10.1.8", "eslint-plugin-prettier": "^5.5.5", "eslint-plugin-react": "^7.37.5", diff --git a/ui/pnpm-lock.yaml b/ui/pnpm-lock.yaml index 097f3490..836a06f6 100644 --- a/ui/pnpm-lock.yaml +++ b/ui/pnpm-lock.yaml @@ -7,8 +7,8 @@ settings: overrides: '@babel/core': 7.29.6 ajv@6: 6.14.0 - brace-expansion@1: 1.1.18 - brace-expansion@5: 5.0.9 + brace-expansion@1: 1.1.21 + brace-expansion@5: 5.0.12 flatted: 3.4.2 follow-redirects: 1.16.0 form-data: 4.0.6 @@ -103,8 +103,8 @@ importers: specifier: ^5.90.21 version: 5.90.21(react@19.2.4) axios: - specifier: ^1.16.0 - version: 1.18.1 + specifier: ^1.20.0 + version: 1.20.0 class-variance-authority: specifier: ^0.7.1 version: 0.7.1 @@ -122,7 +122,7 @@ importers: version: 8.6.0(react@19.2.4) geist: specifier: ^1.7.0 - version: 1.7.0(next@16.3.4(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)) + version: 1.7.0(next@16.3.6(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)) input-otp: specifier: ^1.4.2 version: 1.4.2(react-dom@19.2.4(react@19.2.4))(react@19.2.4) @@ -130,8 +130,8 @@ importers: specifier: ^0.575.0 version: 0.575.0(react@19.2.4) next: - specifier: 16.3.4 - version: 16.3.4(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4) + specifier: 16.3.6 + version: 16.3.6(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4) next-themes: specifier: ^0.4.6 version: 0.4.6(react-dom@19.2.4(react@19.2.4))(react@19.2.4) @@ -203,8 +203,8 @@ importers: specifier: ^9.7.0 version: 9.38.0(jiti@2.6.1) eslint-config-next: - specifier: 16.3.4 - version: 16.3.4(@typescript-eslint/parser@8.57.0(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3))(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3) + specifier: 16.3.6 + version: 16.3.6(@typescript-eslint/parser@8.57.0(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3))(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3) eslint-config-prettier: specifier: ^10.1.8 version: 10.1.8(eslint@9.38.0(jiti@2.6.1)) @@ -580,56 +580,56 @@ packages: '@napi-rs/wasm-runtime@0.2.12': resolution: {integrity: sha512-ZVWUcfwY4E/yPitQJl481FjFo3K22D6qF0DuFH6Y/nbnE11GY5uguDxZMGXPQ8WQ0128MXQD7TnfHyK4oWoIJQ==} - '@next/env@16.3.4': - resolution: {integrity: sha512-cjWZnUUa6jZq2kFaNe/ZyJdZonOZ/QoN0Zka2nz/FLOrfx14pQuM9c5RaSVkWMqgdt4ksgPAMWPyHSs/CyV48Q==} + '@next/env@16.3.6': + resolution: {integrity: sha512-x9Vblze1EbtltQYnNH38xCPWU3TVfBd1eXqA3+w9+BTpedkkdNpAaltXlGQ/nsc1+E0mVTNrtcbX3GoO09zeLQ==} - '@next/eslint-plugin-next@16.3.4': - resolution: {integrity: sha512-szW9y2Aumu4z88YXfTzcFsgUAg2k64uzbtcO5L9f1AKS4w/GUKJcbFllRflROVyNPgJtGOnvNxiyp3v6b+prIA==} + '@next/eslint-plugin-next@16.3.6': + resolution: {integrity: sha512-jowwDX+7DOlDIjJLgTMxudw+k37QnWu1JkZLkSi9MaJBfDYcfhAPMKBhXL0idYzFN/AGg//axnOR4cLkHX/Rng==} - '@next/swc-darwin-arm64@16.3.4': - resolution: {integrity: sha512-iBr3I5LZNk5/bgl5//iTgD2tcym14MX0Xo7fD//u9dYAEgGzza1y9oywluPtf74YnOswVdH1908aK9xVz7zQTw==} + '@next/swc-darwin-arm64@16.3.6': + resolution: {integrity: sha512-E/7GEqaUkt8mk/T8v9lAnrhzR06kdq1ZBkC12F8tAMkdIadwNp3H1KqHynDHrpcTlGCUdq/qu6vUL2aYVyYBdw==} engines: {node: '>= 10'} cpu: [arm64] os: [darwin] - '@next/swc-darwin-x64@16.3.4': - resolution: {integrity: sha512-2dpiSyl2Jw/NrBPaU2MAKGSa+2MR82pJIn4Sm5Rjr+gxAeuh0z158Su3Z2O8zn7UNNq+ej4bToed6RcRN/Lydg==} + '@next/swc-darwin-x64@16.3.6': + resolution: {integrity: sha512-yBE893/nDWTlaiBD1p+qgt7NUen4U5R6FXyH0s67Npq1S3E0cVSef1WIXC2xBRgQvwAvJq6DnS6Y6PrY0cy4Ew==} engines: {node: '>= 10'} cpu: [x64] os: [darwin] - '@next/swc-linux-arm64-gnu@16.3.4': - resolution: {integrity: sha512-+t+U8HZT+fApePCS5h89CSH3datz29MkzyfCn+6fpsZBG/oiEOhINcb9rtkv6sdpToLGFn2e6146NzaKCXkqrA==} + '@next/swc-linux-arm64-gnu@16.3.6': + resolution: {integrity: sha512-KJDpjBqBPYlvkivmyrp+Qys6k/7ksbqGQvRVc6ZEGfR+cjQxx+nUkJaWmNZJsmoOrqYNbaXByF8wa0lBwDhB3Q==} engines: {node: '>= 10'} cpu: [arm64] os: [linux] - '@next/swc-linux-arm64-musl@16.3.4': - resolution: {integrity: sha512-mx03GNs1ocQA5JQ4FxDMmIsNkdrZh8cuezKCrId28e5/gIPU/l7Kcy2+vmCCzdjnnmXJy+iOAu+7K0QppO6Urg==} + '@next/swc-linux-arm64-musl@16.3.6': + resolution: {integrity: sha512-mqNg2K+hvWskSRb/QM+Ix412DvBsuSF0XV+frTSw5vmoucNnIlynFwKYew8D01bfATErMOM7Bujrf0BA5DRKFA==} engines: {node: '>= 10'} cpu: [arm64] os: [linux] - '@next/swc-linux-x64-gnu@16.3.4': - resolution: {integrity: sha512-YIhGY6fSMfha52bnVxnzc9zaVBzJg+cqQTOD8tXIBSx4fuv0pVMxQTE0PaS59YhnMOiYiG09IMwxJAf/CFm/Dw==} + '@next/swc-linux-x64-gnu@16.3.6': + resolution: {integrity: sha512-nFncBNGAYouRHjRVaITs9beZRfhX4ssVwpnvPIAbkZVH6LtGoAVlH4bJ8Cnf9SOo9bsXgPFer/GdHtEE3JNOkw==} engines: {node: '>= 10'} cpu: [x64] os: [linux] - '@next/swc-linux-x64-musl@16.3.4': - resolution: {integrity: sha512-+eaaX6axpDb0yF1GCpiERe6njplvdC+nks/fKfcHu3XPGRrald8P3/X7yv7QLdjA51knnxwl9pxdIJsg+w1L+Q==} + '@next/swc-linux-x64-musl@16.3.6': + resolution: {integrity: sha512-5Mf3cHDGR/Iz0ng2Bj3zUR3p5QS9YK3Hn2QiAfavFmyF48zwThAjpFoiTKNIcOHLYS4zEk+gzyJ/9deQ2ZB8yQ==} engines: {node: '>= 10'} cpu: [x64] os: [linux] - '@next/swc-win32-arm64-msvc@16.3.4': - resolution: {integrity: sha512-0jcXW7Xs/uzICrmgV3MhDYDeRy++1CqnpDIerlPIqYO4bhzB4WNbX/aRnQclustsAyTkFKB0z6rbcjmNg5tR8A==} + '@next/swc-win32-arm64-msvc@16.3.6': + resolution: {integrity: sha512-0jkJy0C2kbrJWTk4YLa3xk80pVBpx8FCHJym7CnUfDAXe/FWv5qT7SQJbR0KuemyxaEDlEx5WT4VQJoTW+/9Qw==} engines: {node: '>= 10'} cpu: [arm64] os: [win32] - '@next/swc-win32-x64-msvc@16.3.4': - resolution: {integrity: sha512-vvBzwu1pYQCp92maZCFCIw/XgOTMR5tur9GjakwIo2cmwRTMKajRZZDS9+e4KsUZWKu1E007WUeAFXRRjZeuzw==} + '@next/swc-win32-x64-msvc@16.3.6': + resolution: {integrity: sha512-/YXjI1e5OXcZ7YpxRwgP/1jAV/SBKTzeVKqN2mk7mLpcICsyn3Gl5+dIfDTJp70M0ccMhyMMRso4v6mPDCGepg==} engines: {node: '>= 10'} cpu: [x64] os: [win32] @@ -1842,8 +1842,8 @@ packages: resolution: {integrity: sha512-BASOg+YwO2C+346x3LZOeoovTIoTrRqEsqMa6fmfAV0P+U9mFr9NsyOEpiYvFjbc64NMrSswhV50WdXzdb/Z5A==} engines: {node: '>=4'} - axios@1.18.1: - resolution: {integrity: sha512-3nTvFlvpn9Zu/RkHUqtc7/+al4UpRW5az71ap5zccp6e8RAYEzhMTecX8Dz1wWDYrPpUoB1HAQEGEAEvUr7S9g==} + axios@1.20.0: + resolution: {integrity: sha512-r8aOh8j9cGKpgQAqpzrUHnSIc6a59Y3Xf/cv8sy1DrHCkZHzQGEuoq1tARk6qSyDdtQGSDgpb9kFlruzPvrgwg==} axobject-query@4.1.0: resolution: {integrity: sha512-qIj0G9wZbMGNLjLmg1PT6v2mE9AH2zlnADJD/2tC6E00hgmhUOfEB6greHPAfLRSufHqROIUTkw6E+M3lH0PTQ==} @@ -1861,11 +1861,11 @@ packages: engines: {node: '>=6.0.0'} hasBin: true - brace-expansion@1.1.18: - resolution: {integrity: sha512-Edep/X9fGqVNmzKBVsDYIOtD+z1tuezV70LBjdCst9Tqu76lsnvRiZ6oTic1n+/BIwX6QDGAO94PN4N2SADvtw==} + brace-expansion@1.1.21: + resolution: {integrity: sha512-9zeA+KLZNNzglF2TPKRQEDyx6Yby7daAkuy8MiPzpXPsYDWi/DRM8jmwUDxokQjYqBpv5DgPiwD4h4ZZSy1Ujw==} - brace-expansion@5.0.9: - resolution: {integrity: sha512-ScQ4IuvIEF1TMlP7Zt+vjJ//9zlPb2SDcxWxM3bk8s6t6GGdJ7KO1dCcTidOPJKePW30LE/2cT7wCyPho9/Wxg==} + brace-expansion@5.0.12: + resolution: {integrity: sha512-YovQ3rzhaLMIrDjNDMkNS01tea93qhEhG5xy8f6+R0l+dw3Ki+5sCoIoI942iuLZTHWogWktgwVDhU09iNEimQ==} engines: {node: 20 || >=22} braces@3.0.3: @@ -2142,8 +2142,8 @@ packages: resolution: {integrity: sha512-TtpcNJ3XAzx3Gq8sWRzJaVajRs0uVxA2YAkdb1jm2YkPz4G6egUFAyA3n5vtEIZefPk5Wa4UXbKuS5fKkJWdgA==} engines: {node: '>=10'} - eslint-config-next@16.3.4: - resolution: {integrity: sha512-35/8RM10huEL9vlr8hUZMERMENHBrnyHN3ZZkF9efSgzGaqK34jIqry44A956//zriUhUAUW0XSkcolhrryqAA==} + eslint-config-next@16.3.6: + resolution: {integrity: sha512-1Upt3U7BDwU+ilpe2byZjAfts9oNq4d4fv/zXEvs8/4yS+cwOQW/WCxUNy8gCDquX67SzeehDvKblVC6ZBMocQ==} peerDependencies: eslint: '>=9.0.0' typescript: '>=3.3.1' @@ -2821,8 +2821,8 @@ packages: react: ^16.8 || ^17 || ^18 || ^19 || ^19.0.0-rc react-dom: ^16.8 || ^17 || ^18 || ^19 || ^19.0.0-rc - next@16.3.4: - resolution: {integrity: sha512-/Ztf6CeRH+ejEXUrYtqI4gkS66eFIHuSwqi60RgcpWKodxFZx2/dqVCMKBwILfAHXQ+F1b1vAudgj3mnxqtoIA==} + next@16.3.6: + resolution: {integrity: sha512-L+otWM/aQbYTx98aZhgEoMb4bZAXx1YVW4UMA/vuCyCoWG5HJyZUili8QAkqzrcC+5///tsz3s0M+SlyB5bLMw==} engines: {node: '>=20.9.0'} hasBin: true peerDependencies: @@ -3899,37 +3899,37 @@ snapshots: '@tybys/wasm-util': 0.10.1 optional: true - '@next/env@16.3.4': {} + '@next/env@16.3.6': {} - '@next/eslint-plugin-next@16.3.4(eslint@9.38.0(jiti@2.6.1))': + '@next/eslint-plugin-next@16.3.6(eslint@9.38.0(jiti@2.6.1))': dependencies: '@eslint-community/eslint-utils': 4.9.1(eslint@9.38.0(jiti@2.6.1)) fast-glob: 3.3.1 transitivePeerDependencies: - eslint - '@next/swc-darwin-arm64@16.3.4': + '@next/swc-darwin-arm64@16.3.6': optional: true - '@next/swc-darwin-x64@16.3.4': + '@next/swc-darwin-x64@16.3.6': optional: true - '@next/swc-linux-arm64-gnu@16.3.4': + '@next/swc-linux-arm64-gnu@16.3.6': optional: true - '@next/swc-linux-arm64-musl@16.3.4': + '@next/swc-linux-arm64-musl@16.3.6': optional: true - '@next/swc-linux-x64-gnu@16.3.4': + '@next/swc-linux-x64-gnu@16.3.6': optional: true - '@next/swc-linux-x64-musl@16.3.4': + '@next/swc-linux-x64-musl@16.3.6': optional: true - '@next/swc-win32-arm64-msvc@16.3.4': + '@next/swc-win32-arm64-msvc@16.3.6': optional: true - '@next/swc-win32-x64-msvc@16.3.4': + '@next/swc-win32-x64-msvc@16.3.6': optional: true '@nodelib/fs.scandir@2.1.5': @@ -5168,7 +5168,7 @@ snapshots: axe-core@4.11.1: {} - axios@1.18.1: + axios@1.20.0: dependencies: follow-redirects: 1.16.0 form-data: 4.0.6 @@ -5186,12 +5186,12 @@ snapshots: baseline-browser-mapping@2.11.21: {} - brace-expansion@1.1.18: + brace-expansion@1.1.21: dependencies: balanced-match: 1.0.2 concat-map: 0.0.1 - brace-expansion@5.0.9: + brace-expansion@5.0.12: dependencies: balanced-match: 4.0.4 @@ -5576,9 +5576,9 @@ snapshots: escape-string-regexp@4.0.0: {} - eslint-config-next@16.3.4(@typescript-eslint/parser@8.57.0(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3))(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3): + eslint-config-next@16.3.6(@typescript-eslint/parser@8.57.0(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3))(eslint@9.38.0(jiti@2.6.1))(typescript@5.9.3): dependencies: - '@next/eslint-plugin-next': 16.3.4(eslint@9.38.0(jiti@2.6.1)) + '@next/eslint-plugin-next': 16.3.6(eslint@9.38.0(jiti@2.6.1)) eslint: 9.38.0(jiti@2.6.1) eslint-import-resolver-node: 0.3.9 eslint-import-resolver-typescript: 3.10.1(eslint-plugin-import@2.32.0)(eslint@9.38.0(jiti@2.6.1)) @@ -5872,9 +5872,9 @@ snapshots: functions-have-names@1.2.3: {} - geist@1.7.0(next@16.3.4(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)): + geist@1.7.0(next@16.3.6(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4)): dependencies: - next: 16.3.4(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4) + next: 16.3.6(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4) generator-function@2.0.1: {} @@ -6263,11 +6263,11 @@ snapshots: minimatch@10.2.4: dependencies: - brace-expansion: 5.0.9 + brace-expansion: 5.0.12 minimatch@3.1.4: dependencies: - brace-expansion: 1.1.18 + brace-expansion: 1.1.21 minimist@1.2.8: {} @@ -6284,9 +6284,9 @@ snapshots: react: 19.2.4 react-dom: 19.2.4(react@19.2.4) - next@16.3.4(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4): + next@16.3.6(@babel/core@7.29.6)(@types/node@25.4.0)(react-dom@19.2.4(react@19.2.4))(react@19.2.4): dependencies: - '@next/env': 16.3.4 + '@next/env': 16.3.6 '@swc/helpers': 0.5.23 baseline-browser-mapping: 2.11.21 caniuse-lite: 1.0.30001810 @@ -6295,14 +6295,14 @@ snapshots: react-dom: 19.2.4(react@19.2.4) styled-jsx: 5.1.6(@babel/core@7.29.6)(react@19.2.4) optionalDependencies: - '@next/swc-darwin-arm64': 16.3.4 - '@next/swc-darwin-x64': 16.3.4 - '@next/swc-linux-arm64-gnu': 16.3.4 - '@next/swc-linux-arm64-musl': 16.3.4 - '@next/swc-linux-x64-gnu': 16.3.4 - '@next/swc-linux-x64-musl': 16.3.4 - '@next/swc-win32-arm64-msvc': 16.3.4 - '@next/swc-win32-x64-msvc': 16.3.4 + '@next/swc-darwin-arm64': 16.3.6 + '@next/swc-darwin-x64': 16.3.6 + '@next/swc-linux-arm64-gnu': 16.3.6 + '@next/swc-linux-arm64-musl': 16.3.6 + '@next/swc-linux-x64-gnu': 16.3.6 + '@next/swc-linux-x64-musl': 16.3.6 + '@next/swc-win32-arm64-msvc': 16.3.6 + '@next/swc-win32-x64-msvc': 16.3.6 sharp: 0.35.4(@types/node@25.4.0) transitivePeerDependencies: - '@babel/core' diff --git a/ui/pnpm-workspace.yaml b/ui/pnpm-workspace.yaml index 317f28e0..5eb49756 100644 --- a/ui/pnpm-workspace.yaml +++ b/ui/pnpm-workspace.yaml @@ -5,8 +5,8 @@ onlyBuiltDependencies: overrides: '@babel/core': 7.29.6 ajv@6: 6.14.0 - brace-expansion@1: 1.1.18 - brace-expansion@5: 5.0.9 + brace-expansion@1: 1.1.21 + brace-expansion@5: 5.0.12 flatted: 3.4.2 follow-redirects: 1.16.0 form-data: 4.0.6 diff --git a/uv.lock b/uv.lock index 86d2be5d..824deeb8 100644 --- a/uv.lock +++ b/uv.lock @@ -9,10 +9,13 @@ resolution-markers = [ [manifest] constraints = [ + { name = "anyio", specifier = ">=4.14.2" }, { name = "grpcio", specifier = ">=1.76.0,<2.0.0" }, { name = "grpcio-tools", specifier = ">=1.76.0,<2.0.0" }, { name = "httpcore", specifier = ">=1.0.9" }, + { name = "pyjwt", specifier = ">=2.15.0" }, { name = "starlette", specifier = ">=1.3.1" }, + { name = "urllib3", specifier = ">=2.8.0" }, ] overrides = [ { name = "cryptography", specifier = ">=49.0.0" }, @@ -211,16 +214,15 @@ wheels = [ [[package]] name = "anyio" -version = "4.9.0" +version = "4.14.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "idna" }, - { name = "sniffio" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/95/7d/4c1bd541d4dffa1b52bd83fb8527089e097a106fc90b467a7313b105f840/anyio-4.9.0.tar.gz", hash = "sha256:673c0c244e15788651a4ff38710fea9675823028a6f08a5eda409e0c9840a028", size = 190949, upload-time = "2025-03-17T00:02:54.77Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/a1/ee/48ca1a7c89ffec8b6a0c5d02b89c305671d5ffd8d3c94acf8b8c408575bb/anyio-4.9.0-py3-none-any.whl", hash = "sha256:9f76d541cad6e36af7beb62e978876f3b41e3e04f2c1fbf0884604c0a9c4d93c", size = 100916, upload-time = "2025-03-17T00:02:52.713Z" }, + { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" }, ] [[package]] @@ -2394,11 +2396,11 @@ wheels = [ [[package]] name = "pyjwt" -version = "2.13.0" +version = "2.15.1" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/3b/81/58d0ac84e1ef3a3843791d6954d94c0b33d526c75eeb1efbce9d0a4c4077/pyjwt-2.13.0.tar.gz", hash = "sha256:41571c89ca91598c79e8ef18a2d07367d4810fbbd6f637794879baf1b7703423", size = 107515, upload-time = "2026-05-21T19:54:36.618Z" } +sdist = { url = "https://files.pythonhosted.org/packages/43/ea/5194e52748b0da83d71e082d75496eaec6e58f419f5e184786ded517e6a9/pyjwt-2.15.1.tar.gz", hash = "sha256:4f259e80cdfb6b3fc18a7de51fd1ef9ec79652f25019bae68975ca2468a34df8", size = 121252, upload-time = "2026-09-28T18:40:42.598Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/a3/5e/ecf12fdb62546d64385c158514e9b2b671f7832108ef2ecd2020ce0af2d1/pyjwt-2.13.0-py3-none-any.whl", hash = "sha256:66adcc2aff09b3f1bbd95fc1e1577df8ac8723c978552fd43304c8a290ac5728", size = 31274, upload-time = "2026-05-21T19:54:35.362Z" }, + { url = "https://files.pythonhosted.org/packages/50/ca/44de4e75f8aadc457f0634be3b542815078ded46dca30efb960edeecad6e/pyjwt-2.15.1-py3-none-any.whl", hash = "sha256:42d59d631f7768a1028a64c7ff581a9bf7519804daf91fc5b6c56e30eec5e193", size = 33860, upload-time = "2026-09-28T18:40:41.429Z" }, ] [[package]] @@ -3229,11 +3231,11 @@ wheels = [ [[package]] name = "urllib3" -version = "2.7.0" +version = "2.8.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/53/0c/06f8b233b8fd13b9e5ee11424ef85419ba0d8ba0b3138bf360be2ff56953/urllib3-2.7.0.tar.gz", hash = "sha256:231e0ec3b63ceb14667c67be60f2f2c40a518cb38b03af60abc813da26505f4c", size = 433602, upload-time = "2026-05-07T16:13:18.596Z" } +sdist = { url = "https://files.pythonhosted.org/packages/e3/05/b17359e1cefb4f909b5e40b1b90a496d987258916dbbf88e842c729f510e/urllib3-2.8.0.tar.gz", hash = "sha256:63bf2ead4c879426ebf22ef2a781eeb4aa3b4ae798a0435506f8687fd5bb9b63", size = 458972, upload-time = "2026-09-15T19:29:36.253Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/7f/3e/5db95bcf282c52709639744ca2a8b149baccf648e39c8cc87553df9eae0c/urllib3-2.7.0-py3-none-any.whl", hash = "sha256:9fb4c81ebbb1ce9531cce37674bbc6f1360472bc18ca9a553ede278ef7276897", size = 131087, upload-time = "2026-05-07T16:13:17.151Z" }, + { url = "https://files.pythonhosted.org/packages/92/9d/c4e665119135114480843e7ab388fa94d8480650450e6f8e26b70d323a4c/urllib3-2.8.0-py3-none-any.whl", hash = "sha256:0cf3cae568d36aa9576b28dfb35f11328f1cb974ca7647d9475ebb86c75ac6e3", size = 135717, upload-time = "2026-09-15T19:29:34.577Z" }, ] [[package]] From 9e62ba302ec50b907a601b539c3f64b653899da2 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 1 Oct 2026 01:14:06 +0200 Subject: [PATCH 08/12] fix: shorten litellm errors and normalize buffered stream failures --- routstr/upstream/base.py | 51 +++++--- routstr/upstream/messages_dispatch.py | 122 ++++++++++++++------ tests/unit/test_messages_upstream_errors.py | 20 ++++ 3 files changed, 137 insertions(+), 56 deletions(-) create mode 100644 tests/unit/test_messages_upstream_errors.py diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 2f65ebd3..69bc1769 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -3024,25 +3024,40 @@ class BaseUpstreamProvider: input_cost = 0.0 output_cost = 0.0 - async for annotated in messages_dispatch.stream_annotated_events( - iterator, requested_model - ): - if annotated.model: - last_model_seen = annotated.model - # See _stream_litellm_messages for why this is max() not +=. - input_tokens = max(input_tokens, annotated.input_tokens) - output_tokens = max(output_tokens, annotated.output_tokens) - cache_read_input_tokens = max( - cache_read_input_tokens, annotated.cache_read_input_tokens + try: + annotated_events = messages_dispatch.stream_annotated_events( + iterator, requested_model ) - cache_creation_input_tokens = max( - cache_creation_input_tokens, - annotated.cache_creation_input_tokens, - ) - total_cost = max(total_cost, annotated.total_cost) - input_cost = max(input_cost, annotated.input_cost) - output_cost = max(output_cost, annotated.output_cost) - buffered.append(annotated) + async for annotated in annotated_events: + if annotated.model: + last_model_seen = annotated.model + # See _stream_litellm_messages for why this is max() not +=. + input_tokens = max(input_tokens, annotated.input_tokens) + output_tokens = max(output_tokens, annotated.output_tokens) + cache_read_input_tokens = max( + cache_read_input_tokens, annotated.cache_read_input_tokens + ) + cache_creation_input_tokens = max( + cache_creation_input_tokens, + annotated.cache_creation_input_tokens, + ) + total_cost = max(total_cost, annotated.total_cost) + input_cost = max(input_cost, annotated.input_cost) + output_cost = max(output_cost, annotated.output_cost) + buffered.append(annotated) + except Exception as exc: + # Buffering lets us return an HTTP error before sending headers. + if messages_dispatch.is_provider_exception(exc): + raise messages_dispatch.upstream_error_from_exception( + exc, + log_message="Upstream stream failed mid-flight", + log_extra={ + "model": last_model_seen or requested_model or "unknown", + "provider": self.provider_type or self.base_url, + "request_id": request_id, + }, + ) from exc + raise response_headers: dict[str, str] = { "Cache-Control": "no-cache", diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index d9df0277..85bfd312 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -485,6 +485,79 @@ def compute_refund(amount: int, unit: str, cost_msats: int) -> int: raise ValueError(f"Invalid unit: {unit}") +_MAX_UPSTREAM_MESSAGE_CHARS = 300 + + +def collapse_litellm_message(message: str) -> str: + """Keep the innermost provider message and cap its length.""" + tail = message.rsplit("Original exception:", 1)[-1].strip() + while True: + stripped = tail + for prefix in ("litellm.",): + if stripped.startswith(prefix): + stripped = stripped[len(prefix) :] + head, _, rest = stripped.partition(": ") + if rest and head.endswith(("Error", "Exception")): + stripped = rest.strip() + if stripped == tail: + break + tail = stripped + if len(tail) > _MAX_UPSTREAM_MESSAGE_CHARS: + tail = tail[: _MAX_UPSTREAM_MESSAGE_CHARS - 1].rstrip() + "…" + return tail + + +def is_provider_exception(exc: BaseException) -> bool: + """Distinguish SDK failures from bugs in our stream handling.""" + return type(exc).__module__.split(".", 1)[0] in {"litellm", "openai"} + + +def upstream_error_from_exception( + exc: Exception, + *, + log_message: str, + log_extra: dict[str, Any] | None = None, +) -> UpstreamError: + """Redact and classify provider failures, including mid-stream errors.""" + raw_message = getattr(exc, "message", None) or str(exc) or repr(exc) + # Redact provider account ids before the message reaches logs or the client. + exc_message = collapse_litellm_message(redact_org_ids(raw_message)) + exc_status = getattr(exc, "status_code", None) + exc_response = getattr(exc, "response", None) + response_text = None + if exc_response is not None: + try: + response_text = redact_org_ids( + getattr(exc_response, "text", str(exc_response)) + ) + except Exception: + response_text = "" + status_for_classify = exc_status if isinstance(exc_status, int) else 502 + rate_limit = classify_rate_limit( + status_for_classify, exc_message, getattr(exc, "headers", None) + ) + logger.error( + log_message, + extra={ + "error": exc_message, + "error_type": type(exc).__name__, + "status_code": exc_status, + "error_code": rate_limit.code if rate_limit else None, + "llm_provider": getattr(exc, "llm_provider", None), + "body": redact_org_ids(str(getattr(exc, "body", "") or "")) or None, + "response_text": response_text, + **(log_extra or {}), + }, + ) + return UpstreamError( + f"Upstream error via litellm: {exc_message}", + status_code=status_for_classify, + code=rate_limit.code if rate_limit else None, + details=rate_limit.as_details() if rate_limit else None, + from_upstream_response=True, + ) + + async def dispatch_anthropic_messages( *, request_body: bytes | None, @@ -606,44 +679,10 @@ async def dispatch_anthropic_messages( try: result = await litellm.anthropic.messages.acreate(**kwargs) except Exception as exc: - raw_message = getattr(exc, "message", None) or str(exc) or repr(exc) - # Redact provider account identifiers before the message reaches logs - # or the surfaced error. - exc_message = redact_org_ids(raw_message) - exc_status = getattr(exc, "status_code", None) - exc_response = getattr(exc, "response", None) - response_text = None - if exc_response is not None: - try: - response_text = redact_org_ids( - getattr(exc_response, "text", str(exc_response)) - ) - except Exception: - response_text = "" - status_for_classify = exc_status if isinstance(exc_status, int) else 502 - rate_limit = classify_rate_limit( - status_for_classify, exc_message, getattr(exc, "headers", None) - ) - logger.error( - "litellm dispatch failed", - extra={ - "error": exc_message, - "error_type": type(exc).__name__, - "status_code": exc_status, - "error_code": rate_limit.code if rate_limit else None, - "llm_provider": getattr(exc, "llm_provider", None), - "body": redact_org_ids(str(getattr(exc, "body", "") or "")) or None, - "response_text": response_text, - "model": litellm_model, - "api_base": base_url, - }, - ) - raise UpstreamError( - f"Upstream error via litellm: {exc_message}", - status_code=status_for_classify, - code=rate_limit.code if rate_limit else None, - details=rate_limit.as_details() if rate_limit else None, - from_upstream_response=True, + raise upstream_error_from_exception( + exc, + log_message="litellm dispatch failed", + log_extra={"model": litellm_model, "api_base": base_url}, ) from exc if transform_stream is not None and hasattr(result, "__aiter__"): @@ -661,6 +700,13 @@ async def dispatch_anthropic_messages( cast(AsyncIterator[Any], result) ) except Exception as exc: + if is_provider_exception(exc): + # Upstream failed part-way through, not an aggregation bug. + raise upstream_error_from_exception( + exc, + log_message="Upstream stream failed mid-flight", + log_extra={"model": litellm_model, "api_base": base_url}, + ) from exc logger.error( "Failed to aggregate streamed events into message", extra={ diff --git a/tests/unit/test_messages_upstream_errors.py b/tests/unit/test_messages_upstream_errors.py new file mode 100644 index 00000000..61b6febd --- /dev/null +++ b/tests/unit/test_messages_upstream_errors.py @@ -0,0 +1,20 @@ +import pytest + +from routstr.upstream.messages_dispatch import collapse_litellm_message + + +@pytest.mark.parametrize( + ("message", "expected"), + [ + ("You have no credits remaining.", "You have no credits remaining."), + ( + "litellm.MidStreamFallbackError: litellm.APIError: No credits. " + "Original exception: MidStreamFallbackError: No credits. " + "Original exception: APIError: litellm.APIError: No credits.", + "No credits.", + ), + ("x" * 301, "x" * 299 + "…"), + ], +) +def test_collapse_litellm_message(message: str, expected: str) -> None: + assert collapse_litellm_message(message) == expected From 5833482240b4ca76e94cc52678fa2a5b7be80620 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 1 Oct 2026 11:49:29 +0200 Subject: [PATCH 09/12] fix: disable upstream stream timeouts by default --- .env.example | 6 ++++++ routstr/core/settings.py | 9 +++++---- tests/unit/test_upstream_stream_timeout.py | 17 +++++++---------- 3 files changed, 18 insertions(+), 14 deletions(-) diff --git a/.env.example b/.env.example index 5688f6c3..6aaa1006 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 + # Logging # LOG_LEVEL=INFO # ENABLE_CONSOLE_LOGGING=true diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 3ddc9f6a..7e93a894 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -40,14 +40,15 @@ 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 + # 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. + # 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=60.0, ge=0, env="UPSTREAM_FIRST_TOKEN_TIMEOUT_SECONDS" + default=0.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" + 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. diff --git a/tests/unit/test_upstream_stream_timeout.py b/tests/unit/test_upstream_stream_timeout.py index 867677b8..7942337c 100644 --- a/tests/unit/test_upstream_stream_timeout.py +++ b/tests/unit/test_upstream_stream_timeout.py @@ -15,7 +15,7 @@ from routstr.core.error_scope import ( ERROR_SCOPE_UPSTREAM, ) from routstr.core.exceptions import UpstreamError -from routstr.core.settings import settings +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 @@ -149,15 +149,12 @@ async def test_zero_first_token_timeout_disables_the_guard( 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"] +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 From 05265474613905e612a50999a6950bf93f960f4f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 1 Oct 2026 11:49:57 +0200 Subject: [PATCH 10/12] test: cover mid-stream litellm failures on buffered messages paths --- routstr/upstream/messages_dispatch.py | 5 +- tests/unit/test_messages_upstream_errors.py | 148 ++++++++++++++++++-- 2 files changed, 141 insertions(+), 12 deletions(-) diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index 85bfd312..da7c6f9f 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -492,10 +492,7 @@ def collapse_litellm_message(message: str) -> str: """Keep the innermost provider message and cap its length.""" tail = message.rsplit("Original exception:", 1)[-1].strip() while True: - stripped = tail - for prefix in ("litellm.",): - if stripped.startswith(prefix): - stripped = stripped[len(prefix) :] + stripped = tail.removeprefix("litellm.") head, _, rest = stripped.partition(": ") if rest and head.endswith(("Error", "Exception")): stripped = rest.strip() diff --git a/tests/unit/test_messages_upstream_errors.py b/tests/unit/test_messages_upstream_errors.py index 61b6febd..62699f69 100644 --- a/tests/unit/test_messages_upstream_errors.py +++ b/tests/unit/test_messages_upstream_errors.py @@ -1,20 +1,152 @@ -import pytest +import os +from typing import Any, AsyncIterator +from unittest.mock import AsyncMock, patch -from routstr.upstream.messages_dispatch import collapse_litellm_message +import litellm +import pytest +from litellm.exceptions import MidStreamFallbackError + +os.environ.setdefault("UPSTREAM_BASE_URL", "http://test") +os.environ.setdefault("UPSTREAM_API_KEY", "test") + +from routstr.core.exceptions import UpstreamError # noqa: E402 +from routstr.payment.models import Architecture, Model, Pricing # noqa: E402 +from routstr.upstream.base import BaseUpstreamProvider # noqa: E402 +from routstr.upstream.messages_dispatch import ( # noqa: E402 + collapse_litellm_message, +) + +_MIDSTREAM_FAILURE = MidStreamFallbackError( + message="No credits.", + model="x", + llm_provider="openai", + original_exception=litellm.APIError( + status_code=500, message="No credits.", llm_provider="openai", model="x" + ), +) @pytest.mark.parametrize( ("message", "expected"), [ ("You have no credits remaining.", "You have no credits remaining."), - ( - "litellm.MidStreamFallbackError: litellm.APIError: No credits. " - "Original exception: MidStreamFallbackError: No credits. " - "Original exception: APIError: litellm.APIError: No credits.", - "No credits.", - ), + # upstream_error_from_exception reads `.message`, which omits the + # "Original exception:" chain that only `str()` appends. + (_MIDSTREAM_FAILURE.message, "No credits."), + (str(_MIDSTREAM_FAILURE), "No credits."), ("x" * 301, "x" * 299 + "…"), ], ) def test_collapse_litellm_message(message: str, expected: str) -> None: assert collapse_litellm_message(message) == expected + + +_RATE_LIMIT = litellm.RateLimitError( + message=( + "Rate limit reached for gpt-4o on tokens per min (TPM): Limit 30000, " + "Used 29000, Requested 2000. Please try again in 1.2s." + ), + llm_provider="openai", + model="gpt-4o", +) +_BAD_REQUEST = litellm.BadRequestError( + message="context length exceeded", model="gpt-4o", llm_provider="openai" +) + +_MID_STREAM_CASES = [ + pytest.param(_RATE_LIMIT, 429, "UPSTREAM_RATE_LIMIT", id="rate-limit"), + pytest.param(_BAD_REQUEST, 400, None, id="bad-request"), + pytest.param(_MIDSTREAM_FAILURE, 500, None, id="midstream-fallback"), +] + + +def _make_model() -> Model: + return Model( + id="gpt-4o", + name="gpt-4o", + created=0, + description="", + context_length=8192, + architecture=Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="x", + instruct_type=None, + ), + pricing=Pricing( + prompt=0.0, + completion=0.0, + request=0.0, + image=0.0, + web_search=0.0, + internal_reasoning=0.0, + max_cost=0.0, + ), + ) + + +def _failing_stream(exc: Exception) -> AsyncIterator[dict]: + async def gen() -> AsyncIterator[dict]: + yield { + "type": "message_start", + "message": {"id": "msg_1", "model": "gpt-4o", "usage": {}}, + } + raise exc + + return gen() + + +def _assert_upstream_error( + err: UpstreamError, status_code: int, code: str | None +) -> None: + assert err.status_code == status_code + assert err.code == code + assert err.from_upstream_response is True + assert "litellm." not in str(err) + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("exc", "status_code", "code"), _MID_STREAM_CASES) +async def test_non_streaming_aggregation_surfaces_mid_stream_failure( + exc: Exception, status_code: int, code: str | None +) -> None: + async def fake_acreate(**kwargs: Any) -> AsyncIterator[dict]: + return _failing_stream(exc) + + with ( + patch( + "litellm.anthropic.messages.acreate", + new=AsyncMock(side_effect=fake_acreate), + ), + pytest.raises(UpstreamError) as exc_info, + ): + await BaseUpstreamProvider( + base_url="http://test", api_key="k" + )._dispatch_anthropic_messages( + request_body=b'{"messages": [], "max_tokens": 8, "stream": false}', + model_obj=_make_model(), + ) + + _assert_upstream_error(exc_info.value, status_code, code) + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("exc", "status_code", "code"), _MID_STREAM_CASES) +async def test_x_cashu_buffered_stream_surfaces_mid_stream_failure( + exc: Exception, status_code: int, code: str | None +) -> None: + provider = BaseUpstreamProvider(base_url="http://test", api_key="k") + + with pytest.raises(UpstreamError) as exc_info: + await provider._stream_x_cashu_litellm_messages( + _failing_stream(exc), + amount=5_000, + unit="sat", + max_cost_for_model=10_000, + requested_model="gpt-4o", + mint=None, + request_id="req-test", + ) + + _assert_upstream_error(exc_info.value, status_code, code) From 6492261e49f7aa7b0b10d35acad30b6ebc5b2751 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 1 Oct 2026 11:49:59 +0200 Subject: [PATCH 11/12] fix: add backoff so litellm's deepseek /v1/messages stream works on prod installs --- docs/provider/configuration.md | 2 +- pyproject.toml | 1 + routstr/upstream/deepseek.py | 7 +-- tests/unit/test_upstream_deepseek.py | 65 ++++++++++++++++++++++++++++ uv.lock | 11 +++++ 5 files changed, 82 insertions(+), 4 deletions(-) diff --git a/docs/provider/configuration.md b/docs/provider/configuration.md index d2695fd9..20668e60 100644 --- a/docs/provider/configuration.md +++ b/docs/provider/configuration.md @@ -65,7 +65,7 @@ in `routstr/upstream/deepseek.py`, not from litellm or OpenRouter: does not price shows up disabled in the Admin Dashboard. Enable it with a manual price, or add it to the table. - **Cache hits** bill at DeepSeek's cache-hit rate (about 2% of the input - rate). + rate on flash, about 3% on pro). Thinking-mode `reasoning_content` is returned to clients unchanged in responses, and forwarded unchanged when it appears in conversation history. diff --git a/pyproject.toml b/pyproject.toml index 296fab3f..da1c339a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -22,6 +22,7 @@ dependencies = [ "pillow>=10", "openai>=1.98.0", "litellm>=1.101.2,<1.102", + "backoff>=2.2", # litellm's native Anthropic-messages streaming (e.g. deepseek/) imports litellm.proxy, which needs it "orjson>=3.10", ] diff --git a/routstr/upstream/deepseek.py b/routstr/upstream/deepseek.py index 01e5e334..a4eed7b3 100644 --- a/routstr/upstream/deepseek.py +++ b/routstr/upstream/deepseek.py @@ -1,9 +1,10 @@ """First-class upstream for the DeepSeek API. Pricing comes from ``_PEAK_RATES`` below, not from litellm or OpenRouter: -litellm's bundled ``deepseek-v4-flash`` entry is stale, the OpenRouter feed -carries resale prices below DeepSeek's own peak rate, and neither knows the -current ``deepseek-flash`` id. A model DeepSeek lists that the table does not +litellm's bundled ``deepseek-v4-flash`` entry is stale (input, output and cache +rates alike), the OpenRouter feed carries resale prices below DeepSeek's own +peak rate, and neither the bundled map nor OpenRouter knows the current +``deepseek-flash`` id. A model DeepSeek lists that the table does not cover is imported disabled rather than priced from those sources. DeepSeek bills peak hours at twice the off-peak rate. The node has one flat diff --git a/tests/unit/test_upstream_deepseek.py b/tests/unit/test_upstream_deepseek.py index f310f7ac..96ab14f9 100644 --- a/tests/unit/test_upstream_deepseek.py +++ b/tests/unit/test_upstream_deepseek.py @@ -12,9 +12,13 @@ with ``tools`` answers 400 when it is stripped. from __future__ import annotations import json +import threading +from collections.abc import Iterator +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import Any from unittest.mock import AsyncMock, Mock, patch +import litellm import pytest from routstr.upstream import upstream_provider_classes @@ -230,3 +234,64 @@ async def test_reasoning_content_in_history_reaches_upstream() -> None: sent = json.loads(out) assert sent["model"] == "deepseek-flash" assert sent["messages"] == messages + + +_ANTHROPIC_SSE = ( + b"event: message_start\n" + b'data: {"type":"message_start","message":{"id":"msg_1","type":"message",' + b'"role":"assistant","model":"deepseek-flash","content":[],' + b'"stop_reason":null,"usage":{"input_tokens":3,"output_tokens":0}}}\n\n' + b"event: message_stop\n" + b'data: {"type":"message_stop"}\n\n' +) + + +@pytest.fixture +def anthropic_stub() -> Iterator[tuple[str, list[tuple[str, dict[str, Any]]]]]: + """Loopback stand-in for DeepSeek's Anthropic-format endpoint.""" + seen: list[tuple[str, dict[str, Any]]] = [] + + class Handler(BaseHTTPRequestHandler): + def do_POST(self) -> None: + length = int(self.headers["Content-Length"]) + seen.append((self.path, json.loads(self.rfile.read(length)))) + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Content-Length", str(len(_ANTHROPIC_SSE))) + self.end_headers() + self.wfile.write(_ANTHROPIC_SSE) + + def log_message(self, *args: Any) -> None: + return None + + server = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + yield f"http://127.0.0.1:{server.server_address[1]}", seen + finally: + server.shutdown() + server.server_close() + + +@pytest.mark.asyncio +async def test_messages_stream_reaches_deepseek_anthropic_endpoint( + anthropic_stub: tuple[str, list[tuple[str, dict[str, Any]]]], +) -> None: + # litellm sends deepseek/ Messages calls to DeepSeek's /anthropic endpoint; + # its stream iterator imports litellm.proxy, which needs ``backoff``. + api_base, seen = anthropic_stub + stream = await litellm.anthropic.messages.acreate( + model=DeepSeekUpstreamProvider.litellm_provider_prefix + "deepseek-flash", + messages=[{"role": "user", "content": "hi"}], + max_tokens=8, + stream=True, + api_key="sk-test", + api_base=api_base, + ) + chunks = [chunk async for chunk in stream] # type: ignore[union-attr] + + assert b"message_stop" in b"".join(chunks) + assert len(seen) == 1 + assert seen[0][0] == "/anthropic/v1/messages" + assert seen[0][1]["model"] == "deepseek-flash" diff --git a/uv.lock b/uv.lock index 86d2be5d..ee38717c 100644 --- a/uv.lock +++ b/uv.lock @@ -282,6 +282,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/77/06/bb80f5f86020c4551da315d78b3ab75e8228f89f0162f2c3a819e407941a/attrs-25.3.0-py3-none-any.whl", hash = "sha256:427318ce031701fea540783410126f03899a97ffc6f61596ad581ac2e40e3bc3", size = 63815, upload-time = "2025-03-13T11:10:21.14Z" }, ] +[[package]] +name = "backoff" +version = "2.2.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/47/d7/5bbeb12c44d7c4f2fb5b56abce497eb5ed9f34d85701de869acedd602619/backoff-2.2.1.tar.gz", hash = "sha256:03f829f5bb1923180821643f8753b0502c3b682293992485b0eef2807afa5cba", size = 17001, upload-time = "2022-10-05T19:19:32.061Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/df/73/b6e24bd22e6720ca8ee9a85a0c4a2971af8497d8f3193fa05390cbd46e09/backoff-2.2.1-py3-none-any.whl", hash = "sha256:63579f9a0628e06278f7e47b7d7d5b6ce20dc65c5e96a6f3ca99a6adca0396e8", size = 15148, upload-time = "2022-10-05T19:19:30.546Z" }, +] + [[package]] name = "base58" version = "2.1.1" @@ -2711,6 +2720,7 @@ source = { editable = "." } dependencies = [ { name = "aiosqlite" }, { name = "alembic" }, + { name = "backoff" }, { name = "cashu" }, { name = "fastapi", extra = ["standard-no-fastapi-cloud-cli"] }, { name = "greenlet" }, @@ -2747,6 +2757,7 @@ dev = [ requires-dist = [ { name = "aiosqlite", specifier = ">=0.20" }, { name = "alembic", specifier = ">=1.13" }, + { name = "backoff", specifier = ">=2.2" }, { name = "cashu", specifier = ">=0.20" }, { name = "fastapi", extras = ["standard-no-fastapi-cloud-cli"], specifier = ">=0.141" }, { name = "greenlet", specifier = ">=3.2.1" }, From 60921cef2075a276b21c7c6a935a1954d1887be6 Mon Sep 17 00:00:00 2001 From: redshift <213178690+1ftredsh@users.noreply.github.com> Date: Thu, 1 Oct 2026 23:30:53 +0800 Subject: [PATCH 12/12] fix: allow /v1/messages/count_tokens through the proxy allowlist The exact-match endpoint allowlist in the proxy omitted `messages/count_tokens`, so every Claude Code / Anthropic SDK request was 404'd with `Path '/v1/messages/count_tokens' not found` before it ever reached the (fully supported) forwarding path. Add the endpoint to `_ALLOWED_ENDPOINTS` and a regression test that pins it as always reachable on POST. Regression history: - 933ba105 "add missing messages endpoint" (2026-04-01): added messages/count_tokens handling to the forwarding layer. - 164ed775 "support /message/count_tokens endpoint" (2026-05-09): added the local count_tokens handler. - 96661384 "update not found proxy" (2026-05-14): only GET was gated, so POST count_tokens passed implicitly. - 0217002e "Gate proxy forwarding behind a segment-anchored API path allowlist" (2026-08-23): POST gated, but the "v1/" prefix still carried count_tokens. - 5af04364 "Restrict proxy forwarding to an exact method/path allowlist" (2026-08-24): switched to an exact table, added "messages" but omitted "messages/count_tokens". This is where the endpoint got locked out. --- routstr/proxy.py | 4 ++++ tests/unit/test_proxy_path_allowlist.py | 21 +++++++++++++++++++++ 2 files changed, 25 insertions(+) diff --git a/routstr/proxy.py b/routstr/proxy.py index 681a3aeb..a2cf9492 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -277,6 +277,10 @@ _ALLOWED_ENDPOINTS: dict[str, frozenset[str]] = { "completions": frozenset({"POST"}), "responses": frozenset({"POST"}), "messages": frozenset({"POST"}), + # Anthropic token-counting subroute; the proxy's allowlist is exact, so the + # "messages" entry above does not carry it. Clients (Claude Code, the + # Anthropic SDKs) call it before every request. + "messages/count_tokens": frozenset({"POST"}), "embeddings": frozenset({"POST"}), # TypeSafe System One decision endpoint: POST {state, model, questions} # -> {answers, usage}. Non-streaming, JSON in/out; billed from the diff --git a/tests/unit/test_proxy_path_allowlist.py b/tests/unit/test_proxy_path_allowlist.py index 4cd1c67f..1019dbad 100644 --- a/tests/unit/test_proxy_path_allowlist.py +++ b/tests/unit/test_proxy_path_allowlist.py @@ -51,6 +51,8 @@ def test_ambiguous_paths_are_rejected(path: str) -> None: "v1/chat/completions", "chat/completions", "v1/responses", + "v1/messages", + "v1/messages/count_tokens", "v1/embeddings", "models", "v1/models/gpt-4", @@ -138,6 +140,7 @@ def test_known_prefix_does_not_carry_an_unknown_endpoint(path: str) -> None: ("completions", "POST"), ("v1/responses", "POST"), ("v1/messages", "POST"), + ("v1/messages/count_tokens", "POST"), ("v1/embeddings", "POST"), ("models", "GET"), ("attestation", "GET"), @@ -163,6 +166,24 @@ def test_method_must_match_the_endpoint(path: str, method: str) -> None: assert _forwarding_allowed(path, method) is False +@pytest.mark.parametrize( + "path", + [ + "messages/count_tokens", + "v1/messages/count_tokens", + "v1/messages/count_tokens/", + ], +) +def test_count_tokens_endpoint_stays_allowed(path: str) -> None: + # Regression guard: /v1/messages/count_tokens is supported end-to-end + # (local handler when the upstream lacks native Anthropic support, plain + # forward otherwise), but the exact-match allowlist once omitted it, so + # Claude Code and the Anthropic SDKs were 404'd on every request. It must + # always be reachable, on POST only. + assert _forwarding_allowed(path, "POST") is True + assert _forwarding_allowed(path, "GET") is False + + def test_operator_additions_are_parsed_per_endpoint() -> None: parsed = _parse_extra_allowed_endpoints("POST:v1/rerank, GET:batches ,post:audio/x") assert parsed == {