fix: guard meaningful stream events and scope cooldown failures

This commit is contained in:
9qeklajc
2026-09-29 02:58:40 +02:00
parent cd2fe38ae6
commit a1071f4a98
6 changed files with 455 additions and 61 deletions
+47 -11
View File
@@ -1,6 +1,7 @@
import asyncio
import 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,
+56 -13
View File
@@ -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,
+19
View File
@@ -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:
+80 -32
View File
@@ -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
@@ -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,
+251 -3
View File
@@ -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")