mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: guard meaningful stream events and scope cooldown failures
This commit is contained in:
+47
-11
@@ -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
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user