mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
feat: add upstream first-token timeout, stream idle timeout and provider cooldown
This commit is contained in:
@@ -40,6 +40,21 @@ class Settings(BaseSettings):
|
|||||||
upstream_5xx_retry_attempts: int = Field(
|
upstream_5xx_retry_attempts: int = Field(
|
||||||
default=1, ge=0, env="UPSTREAM_5XX_RETRY_ATTEMPTS"
|
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
|
# Node info
|
||||||
name: str = Field(default="ARoutstrNode", env="NAME")
|
name: str = Field(default="ARoutstrNode", env="NAME")
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ from .payment.helpers import (
|
|||||||
)
|
)
|
||||||
from .payment.models import Model
|
from .payment.models import Model
|
||||||
from .upstream import BaseUpstreamProvider
|
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.ehbp import forward_ehbp_request, forward_ehbp_x_cashu_request
|
||||||
from .upstream.helpers import init_upstreams
|
from .upstream.helpers import init_upstreams
|
||||||
from .upstream.model_paths import (
|
from .upstream.model_paths import (
|
||||||
@@ -405,6 +406,11 @@ _RETRYABLE_UPSTREAM_5XX = frozenset({502, 503, 504})
|
|||||||
_UPSTREAM_5XX_RETRY_BACKOFF_SECONDS = 0.5
|
_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)
|
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
|
||||||
async def proxy(
|
async def proxy(
|
||||||
request: Request, path: str, session: AsyncSession = Depends(get_session)
|
request: Request, path: str, session: AsyncSession = Depends(get_session)
|
||||||
@@ -621,6 +627,17 @@ async def _proxy(
|
|||||||
request=request,
|
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
|
# Reserve/max-cost checks use the best-ranked candidate; the failover loop
|
||||||
# below rebinds (model_obj, upstream) per candidate so forwarding and
|
# below rebinds (model_obj, upstream) per candidate so forwarding and
|
||||||
# settlement always use the model of the provider actually being tried.
|
# settlement always use the model of the provider actually being tried.
|
||||||
@@ -950,6 +967,8 @@ async def _proxy(
|
|||||||
break
|
break
|
||||||
|
|
||||||
if response.status_code != 200:
|
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.
|
# 424 is an upstream failure re-reported by error_scope.
|
||||||
# 502/503 are upstream errors, 429 rate limits.
|
# 502/503 are upstream errors, 429 rate limits.
|
||||||
should_retry = response.status_code in [
|
should_retry = response.status_code in [
|
||||||
@@ -1037,6 +1056,8 @@ async def _proxy(
|
|||||||
raise
|
raise
|
||||||
|
|
||||||
except UpstreamError as e:
|
except UpstreamError as e:
|
||||||
|
if _counts_toward_cooldown(e.status_code):
|
||||||
|
record_failure(upstream.base_url, model_id)
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Upstream %s failed for model=%s: %s",
|
"Upstream %s failed for model=%s: %s",
|
||||||
upstream.provider_type,
|
upstream.provider_type,
|
||||||
|
|||||||
@@ -85,6 +85,7 @@ from .stream_ownership import (
|
|||||||
close_upstream_exchange,
|
close_upstream_exchange,
|
||||||
finalize_and_close_stream,
|
finalize_and_close_stream,
|
||||||
)
|
)
|
||||||
|
from .stream_timeout import open_guarded_stream
|
||||||
|
|
||||||
if typing.TYPE_CHECKING:
|
if typing.TYPE_CHECKING:
|
||||||
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
|
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
|
||||||
@@ -1147,6 +1148,8 @@ class BaseUpstreamProvider:
|
|||||||
Returns:
|
Returns:
|
||||||
StreamingResponse with cost data injected at the end
|
StreamingResponse with cost data injected at the end
|
||||||
"""
|
"""
|
||||||
|
guarded_chunks = await open_guarded_stream(response, self.provider_type)
|
||||||
|
|
||||||
if reservation_snapshot is None:
|
if reservation_snapshot is None:
|
||||||
async with create_session() as snapshot_session:
|
async with create_session() as snapshot_session:
|
||||||
snapshot_key = await snapshot_session.get(key.__class__, key.hashed_key)
|
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
|
# multiple events can arrive together; buffering makes parsing
|
||||||
# boundary-independent for every provider.
|
# boundary-independent for every provider.
|
||||||
splitter = SSEEventSplitter()
|
splitter = SSEEventSplitter()
|
||||||
async for chunk in response.aiter_bytes():
|
async for chunk in guarded_chunks:
|
||||||
for raw_event in splitter.feed(chunk):
|
for raw_event in splitter.feed(chunk):
|
||||||
for out in _process_event(raw_event):
|
for out in _process_event(raw_event):
|
||||||
yield out
|
yield out
|
||||||
@@ -1639,6 +1642,8 @@ class BaseUpstreamProvider:
|
|||||||
Returns:
|
Returns:
|
||||||
StreamingResponse with cost data injected at the end
|
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)
|
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
@@ -1790,7 +1795,7 @@ class BaseUpstreamProvider:
|
|||||||
# Buffer across network chunks; dispatch only on the SSE event
|
# Buffer across network chunks; dispatch only on the SSE event
|
||||||
# delimiter so parsing is independent of byte boundaries.
|
# delimiter so parsing is independent of byte boundaries.
|
||||||
splitter = SSEEventSplitter()
|
splitter = SSEEventSplitter()
|
||||||
async for chunk in response.aiter_bytes():
|
async for chunk in guarded_chunks:
|
||||||
for raw_event in splitter.feed(chunk):
|
for raw_event in splitter.feed(chunk):
|
||||||
for out in _process_event(raw_event):
|
for out in _process_event(raw_event):
|
||||||
yield out
|
yield out
|
||||||
@@ -2137,7 +2142,9 @@ class BaseUpstreamProvider:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
try:
|
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
|
yield chunk
|
||||||
finally:
|
finally:
|
||||||
await finalizer.run()
|
await finalizer.run()
|
||||||
@@ -2192,6 +2199,8 @@ class BaseUpstreamProvider:
|
|||||||
reservation_snapshot: ReservationSnapshot | None = None,
|
reservation_snapshot: ReservationSnapshot | None = None,
|
||||||
request_body: bytes | None = None,
|
request_body: bytes | None = None,
|
||||||
) -> StreamingResponse:
|
) -> StreamingResponse:
|
||||||
|
guarded_chunks = await open_guarded_stream(response, self.provider_type)
|
||||||
|
|
||||||
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
||||||
usage_finalized = False
|
usage_finalized = False
|
||||||
last_model_seen: str | None = None
|
last_model_seen: str | None = None
|
||||||
@@ -2284,7 +2293,7 @@ class BaseUpstreamProvider:
|
|||||||
total_cost = max(total_cost, _coerce_usd(usage_or_root.get(field)))
|
total_cost = max(total_cost, _coerce_usd(usage_or_root.get(field)))
|
||||||
|
|
||||||
try:
|
try:
|
||||||
async for chunk in response.aiter_bytes():
|
async for chunk in guarded_chunks:
|
||||||
stored_chunks.append(chunk)
|
stored_chunks.append(chunk)
|
||||||
try:
|
try:
|
||||||
decoded_chunk = chunk.decode("utf-8", errors="ignore")
|
decoded_chunk = chunk.decode("utf-8", errors="ignore")
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user