mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-06 04:38:22 +00:00
Merge pull request #776 from Routstr/feat/upstream-first-token-timeout
feat: upstream first-token timeout, stream idle timeout and provider cooldown
This commit is contained in:
@@ -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
|
||||
|
||||
# Request and reservation lifetime limits (seconds)
|
||||
# STALE_RESERVATION_TIMEOUT_SECONDS=300
|
||||
# MAX_REQUEST_LIFETIME_SECONDS=1800
|
||||
|
||||
@@ -40,6 +40,22 @@ class Settings(BaseSettings):
|
||||
upstream_5xx_retry_attempts: int = Field(
|
||||
default=1, ge=0, env="UPSTREAM_5XX_RETRY_ATTEMPTS"
|
||||
)
|
||||
# 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. Reasoning models can stay
|
||||
# silent for minutes, so set these above the longest expected think time.
|
||||
upstream_first_token_timeout_seconds: float = Field(
|
||||
default=0.0, ge=0, env="UPSTREAM_FIRST_TOKEN_TIMEOUT_SECONDS"
|
||||
)
|
||||
upstream_stream_idle_timeout_seconds: float = Field(
|
||||
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.
|
||||
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")
|
||||
|
||||
+62
-5
@@ -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,
|
||||
@@ -40,6 +42,12 @@ from .payment.helpers import (
|
||||
)
|
||||
from .payment.models import Model
|
||||
from .upstream import BaseUpstreamProvider
|
||||
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 +124,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):
|
||||
@@ -405,6 +411,18 @@ _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
|
||||
|
||||
|
||||
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:
|
||||
@@ -695,6 +713,20 @@ 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(
|
||||
provider_identity(candidate[1]),
|
||||
candidate_model_identity(candidate[0], 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.
|
||||
@@ -722,7 +754,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,
|
||||
@@ -731,7 +763,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,
|
||||
@@ -740,7 +772,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,
|
||||
@@ -748,6 +780,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",
|
||||
@@ -760,6 +798,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
|
||||
@@ -1030,6 +1075,11 @@ async def _proxy(
|
||||
break
|
||||
|
||||
if response.status_code != 200:
|
||||
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 [
|
||||
@@ -1117,6 +1167,13 @@ async def _proxy(
|
||||
raise
|
||||
|
||||
except UpstreamError as e:
|
||||
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,
|
||||
|
||||
+62
-10
@@ -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,6 +87,7 @@ from .stream_ownership import (
|
||||
close_upstream_exchange,
|
||||
finalize_and_close_stream,
|
||||
)
|
||||
from .stream_timeout import GuardedStream, open_guarded_stream
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
|
||||
@@ -1150,6 +1153,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,
|
||||
@@ -1171,6 +1185,8 @@ class BaseUpstreamProvider:
|
||||
Returns:
|
||||
StreamingResponse with cost data injected at the end
|
||||
"""
|
||||
guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
|
||||
|
||||
if reservation_snapshot is None:
|
||||
async with create_session() as snapshot_session:
|
||||
snapshot_key = await snapshot_session.get(key.__class__, key.hashed_key)
|
||||
@@ -1374,7 +1390,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
|
||||
@@ -1460,7 +1476,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:
|
||||
@@ -1666,6 +1684,8 @@ class BaseUpstreamProvider:
|
||||
Returns:
|
||||
StreamingResponse with cost data injected at the end
|
||||
"""
|
||||
guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
|
||||
|
||||
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
||||
|
||||
logger.debug(
|
||||
@@ -1818,7 +1838,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
|
||||
@@ -1867,7 +1887,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",
|
||||
"provider": provider_seen,
|
||||
"response": {
|
||||
"model": last_model_seen or "unknown",
|
||||
@@ -1889,6 +1911,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(
|
||||
@@ -1904,7 +1934,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:
|
||||
@@ -2149,6 +2184,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:
|
||||
@@ -2166,12 +2202,22 @@ class BaseUpstreamProvider:
|
||||
)
|
||||
)
|
||||
try:
|
||||
async for chunk in response.aiter_bytes():
|
||||
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,
|
||||
@@ -2181,6 +2227,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(
|
||||
@@ -2203,6 +2250,7 @@ class BaseUpstreamProvider:
|
||||
provider_fee,
|
||||
reservation_snapshot,
|
||||
finalizer,
|
||||
guarded_chunks,
|
||||
)
|
||||
return ClosingStreamingResponse(
|
||||
stream,
|
||||
@@ -2221,6 +2269,8 @@ class BaseUpstreamProvider:
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
request_body: bytes | None = None,
|
||||
) -> StreamingResponse:
|
||||
guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
|
||||
|
||||
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
||||
usage_finalized = False
|
||||
last_model_seen: str | None = None
|
||||
@@ -2314,7 +2364,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")
|
||||
@@ -2493,6 +2543,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:
|
||||
@@ -3451,7 +3503,7 @@ class BaseUpstreamProvider:
|
||||
},
|
||||
)
|
||||
|
||||
result = self._generic_streaming_response(
|
||||
result = await self._generic_streaming_response(
|
||||
response,
|
||||
key.hashed_key,
|
||||
max_cost_for_model,
|
||||
@@ -3736,7 +3788,7 @@ class BaseUpstreamProvider:
|
||||
},
|
||||
)
|
||||
|
||||
result = self._generic_streaming_response(
|
||||
result = await self._generic_streaming_response(
|
||||
response,
|
||||
key.hashed_key,
|
||||
max_cost_for_model,
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
"""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 typing import Any
|
||||
|
||||
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 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:
|
||||
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,117 @@
|
||||
"""Timeout guards for upstream streaming responses."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator, Callable
|
||||
|
||||
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
|
||||
from .sse_splitter import SSEEventSplitter
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
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(guarded_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 GuardedStream(first, guarded_chunks, provider_type, on_idle_timeout)
|
||||
|
||||
|
||||
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
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -0,0 +1,506 @@
|
||||
"""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.error_scope import (
|
||||
ERROR_SCOPE_HEADER,
|
||||
ERROR_SCOPE_NODE,
|
||||
ERROR_SCOPE_UPSTREAM,
|
||||
)
|
||||
from routstr.core.exceptions import UpstreamError
|
||||
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
|
||||
|
||||
|
||||
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"
|
||||
|
||||
|
||||
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)
|
||||
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_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,
|
||||
) -> 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"]
|
||||
|
||||
|
||||
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
|
||||
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.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
|
||||
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]:
|
||||
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.db_id = None
|
||||
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,
|
||||
request: MagicMock | None = None,
|
||||
) -> 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),
|
||||
):
|
||||
request = request or _proxy_request()
|
||||
return await proxy_module._proxy(
|
||||
request, "v1/chat/completions", MagicMock(), await request.body()
|
||||
)
|
||||
|
||||
|
||||
@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("test|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("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")
|
||||
Reference in New Issue
Block a user