mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +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 asyncio
|
||||||
import inspect
|
import inspect
|
||||||
import json
|
import json
|
||||||
|
import re
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter, HTTPException, Request
|
from fastapi import APIRouter, HTTPException, Request
|
||||||
@@ -23,6 +24,7 @@ from .core.db import (
|
|||||||
create_session,
|
create_session,
|
||||||
)
|
)
|
||||||
from .core.error_scope import (
|
from .core.error_scope import (
|
||||||
|
ERROR_SCOPE_HEADER,
|
||||||
ERROR_SCOPE_UPSTREAM,
|
ERROR_SCOPE_UPSTREAM,
|
||||||
UPSTREAM_ERROR_STATUS,
|
UPSTREAM_ERROR_STATUS,
|
||||||
UPSTREAM_UNAVAILABLE,
|
UPSTREAM_UNAVAILABLE,
|
||||||
@@ -39,7 +41,12 @@ 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.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.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 (
|
||||||
@@ -116,8 +123,6 @@ def get_candidates(
|
|||||||
if candidates := _provider_map.get(model_id_lower):
|
if candidates := _provider_map.get(model_id_lower):
|
||||||
return candidates
|
return candidates
|
||||||
|
|
||||||
import re
|
|
||||||
|
|
||||||
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
|
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
|
||||||
if base_model_id != model_id_lower:
|
if base_model_id != model_id_lower:
|
||||||
if candidates := _provider_map.get(base_model_id):
|
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
|
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(
|
def _attribute_request(
|
||||||
request: Request, model_obj: Model, upstream: BaseUpstreamProvider
|
request: Request, model_obj: Model, upstream: BaseUpstreamProvider
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -705,7 +717,10 @@ async def _proxy(
|
|||||||
healthy = [
|
healthy = [
|
||||||
candidate
|
candidate
|
||||||
for candidate in candidates
|
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:
|
if healthy:
|
||||||
candidates = healthy
|
candidates = healthy
|
||||||
@@ -737,7 +752,7 @@ async def _proxy(
|
|||||||
model_id,
|
model_id,
|
||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
return await forward_ehbp_x_cashu_request(
|
response = await forward_ehbp_x_cashu_request(
|
||||||
request=request,
|
request=request,
|
||||||
x_cashu_token=x_cashu,
|
x_cashu_token=x_cashu,
|
||||||
path=path,
|
path=path,
|
||||||
@@ -746,7 +761,7 @@ async def _proxy(
|
|||||||
upstream=upstream,
|
upstream=upstream,
|
||||||
)
|
)
|
||||||
elif is_responses_api:
|
elif is_responses_api:
|
||||||
return await upstream.handle_x_cashu_responses(
|
response = await upstream.handle_x_cashu_responses(
|
||||||
request,
|
request,
|
||||||
x_cashu,
|
x_cashu,
|
||||||
path,
|
path,
|
||||||
@@ -755,7 +770,7 @@ async def _proxy(
|
|||||||
request_body=request_body,
|
request_body=request_body,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
return await upstream.handle_x_cashu(
|
response = await upstream.handle_x_cashu(
|
||||||
request,
|
request,
|
||||||
x_cashu,
|
x_cashu,
|
||||||
path,
|
path,
|
||||||
@@ -763,6 +778,12 @@ async def _proxy(
|
|||||||
model_obj,
|
model_obj,
|
||||||
request_body=request_body,
|
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:
|
except UpstreamError as e:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Upstream %s failed (x-cashu) for model=%s: %s",
|
"Upstream %s failed (x-cashu) for model=%s: %s",
|
||||||
@@ -775,6 +796,13 @@ async def _proxy(
|
|||||||
"status_code": e.status_code,
|
"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:
|
if i == len(candidates) - 1:
|
||||||
last_error = e
|
last_error = e
|
||||||
continue
|
continue
|
||||||
@@ -1044,8 +1072,11 @@ async def _proxy(
|
|||||||
break
|
break
|
||||||
|
|
||||||
if response.status_code != 200:
|
if response.status_code != 200:
|
||||||
if _counts_toward_cooldown(response.status_code):
|
if _upstream_response_failure(response):
|
||||||
record_failure(upstream.base_url, model_id)
|
record_failure(
|
||||||
|
provider_identity(upstream),
|
||||||
|
candidate_model_identity(model_obj, 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 [
|
||||||
@@ -1133,8 +1164,13 @@ async def _proxy(
|
|||||||
raise
|
raise
|
||||||
|
|
||||||
except UpstreamError as e:
|
except UpstreamError as e:
|
||||||
if _counts_toward_cooldown(e.status_code):
|
if e.scope == ERROR_SCOPE_UPSTREAM and _counts_toward_cooldown(
|
||||||
record_failure(upstream.base_url, model_id)
|
e.status_code
|
||||||
|
):
|
||||||
|
record_failure(
|
||||||
|
provider_identity(upstream),
|
||||||
|
candidate_model_identity(model_obj, 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,
|
||||||
|
|||||||
+56
-13
@@ -34,6 +34,7 @@ from ..core.error_scope import (
|
|||||||
ERROR_SCOPE_HEADER,
|
ERROR_SCOPE_HEADER,
|
||||||
ERROR_SCOPE_NODE,
|
ERROR_SCOPE_NODE,
|
||||||
ERROR_SCOPE_UPSTREAM,
|
ERROR_SCOPE_UPSTREAM,
|
||||||
|
UPSTREAM_ERROR_STATUS,
|
||||||
client_code_for_upstream_error,
|
client_code_for_upstream_error,
|
||||||
client_status_for_upstream_error,
|
client_status_for_upstream_error,
|
||||||
upstream_status_details,
|
upstream_status_details,
|
||||||
@@ -68,6 +69,7 @@ from .cache_breakpoints import (
|
|||||||
inject_anthropic_cache_breakpoints,
|
inject_anthropic_cache_breakpoints,
|
||||||
is_explicit_cache_model,
|
is_explicit_cache_model,
|
||||||
)
|
)
|
||||||
|
from .cooldown import model_identity, provider_identity, record_failure
|
||||||
from .count_tokens import MissingUsageEstimator, count_tokens_locally
|
from .count_tokens import MissingUsageEstimator, count_tokens_locally
|
||||||
from .http_client import acquire_upstream_http_client, build_x_cashu_client
|
from .http_client import acquire_upstream_http_client, build_x_cashu_client
|
||||||
from .litellm_routing import detect_litellm_prefix
|
from .litellm_routing import detect_litellm_prefix
|
||||||
@@ -85,7 +87,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
|
from .stream_timeout import GuardedStream, open_guarded_stream
|
||||||
|
|
||||||
if typing.TYPE_CHECKING:
|
if typing.TYPE_CHECKING:
|
||||||
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
|
from .ehbp import ConfidentialInferenceProfile, EHBPForwardingTarget
|
||||||
@@ -1127,6 +1129,17 @@ class BaseUpstreamProvider:
|
|||||||
)
|
)
|
||||||
return True
|
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(
|
async def handle_streaming_chat_completion(
|
||||||
self,
|
self,
|
||||||
response: httpx.Response,
|
response: httpx.Response,
|
||||||
@@ -1148,7 +1161,7 @@ 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)
|
guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
|
||||||
|
|
||||||
if reservation_snapshot is None:
|
if reservation_snapshot is None:
|
||||||
async with create_session() as snapshot_session:
|
async with create_session() as snapshot_session:
|
||||||
@@ -1436,7 +1449,9 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
|
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"
|
yield b"data: [DONE]\n\n"
|
||||||
|
|
||||||
except httpx.RemoteProtocolError as stream_error:
|
except httpx.RemoteProtocolError as stream_error:
|
||||||
@@ -1642,7 +1657,7 @@ 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)
|
guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
|
||||||
|
|
||||||
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
||||||
|
|
||||||
@@ -1844,7 +1859,9 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
if usage_chunk_data is None:
|
if usage_chunk_data is None:
|
||||||
usage_chunk_data = {
|
usage_chunk_data = {
|
||||||
"type": "response.completed",
|
"type": "response.failed"
|
||||||
|
if guarded_chunks.timed_out
|
||||||
|
else "response.completed",
|
||||||
"response": {
|
"response": {
|
||||||
"model": last_model_seen or "unknown",
|
"model": last_model_seen or "unknown",
|
||||||
"usage": {
|
"usage": {
|
||||||
@@ -1865,6 +1882,14 @@ class BaseUpstreamProvider:
|
|||||||
+ cost_data.get("output_tokens", 0),
|
+ 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:
|
try:
|
||||||
self.inject_cost_metadata(
|
self.inject_cost_metadata(
|
||||||
@@ -1880,7 +1905,12 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode()
|
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"
|
yield b"data: [DONE]\n\n"
|
||||||
|
|
||||||
except httpx.RemoteProtocolError as stream_error:
|
except httpx.RemoteProtocolError as stream_error:
|
||||||
@@ -2125,6 +2155,7 @@ class BaseUpstreamProvider:
|
|||||||
provider_fee: float | None,
|
provider_fee: float | None,
|
||||||
reservation_snapshot: ReservationSnapshot,
|
reservation_snapshot: ReservationSnapshot,
|
||||||
finalizer: PersistentStreamFinalizer | None = None,
|
finalizer: PersistentStreamFinalizer | None = None,
|
||||||
|
guarded_chunks: GuardedStream | None = None,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes, None]:
|
||||||
"""Relay an opaque stream and settle it even if the caller disconnects."""
|
"""Relay an opaque stream and settle it even if the caller disconnects."""
|
||||||
if finalizer is None:
|
if finalizer is None:
|
||||||
@@ -2142,14 +2173,22 @@ class BaseUpstreamProvider:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
# This generator is already the response body, so a first-chunk
|
if guarded_chunks is None:
|
||||||
# timeout here can only abort the stream, never fail over.
|
guarded_chunks = await self._guard_stream(
|
||||||
async for chunk in await open_guarded_stream(response, self.provider_type):
|
response, model_obj, sse=False
|
||||||
|
)
|
||||||
|
async for chunk in guarded_chunks:
|
||||||
yield chunk
|
yield chunk
|
||||||
|
if guarded_chunks.timed_out:
|
||||||
|
raise UpstreamError(
|
||||||
|
"Upstream stream stalled",
|
||||||
|
status_code=UPSTREAM_ERROR_STATUS,
|
||||||
|
code="UPSTREAM_TIMEOUT",
|
||||||
|
)
|
||||||
finally:
|
finally:
|
||||||
await finalizer.run()
|
await finalizer.run()
|
||||||
|
|
||||||
def _generic_streaming_response(
|
async def _generic_streaming_response(
|
||||||
self,
|
self,
|
||||||
response: httpx.Response,
|
response: httpx.Response,
|
||||||
key_hash: str,
|
key_hash: str,
|
||||||
@@ -2159,6 +2198,7 @@ class BaseUpstreamProvider:
|
|||||||
provider_fee: float | None,
|
provider_fee: float | None,
|
||||||
reservation_snapshot: ReservationSnapshot,
|
reservation_snapshot: ReservationSnapshot,
|
||||||
) -> ClosingStreamingResponse:
|
) -> ClosingStreamingResponse:
|
||||||
|
guarded_chunks = await self._guard_stream(response, model_obj, sse=False)
|
||||||
finalizer = PersistentStreamFinalizer(
|
finalizer = PersistentStreamFinalizer(
|
||||||
lambda: finalize_and_close_stream(
|
lambda: finalize_and_close_stream(
|
||||||
lambda: self._finalize_generic_streaming_payment(
|
lambda: self._finalize_generic_streaming_payment(
|
||||||
@@ -2181,6 +2221,7 @@ class BaseUpstreamProvider:
|
|||||||
provider_fee,
|
provider_fee,
|
||||||
reservation_snapshot,
|
reservation_snapshot,
|
||||||
finalizer,
|
finalizer,
|
||||||
|
guarded_chunks,
|
||||||
)
|
)
|
||||||
return ClosingStreamingResponse(
|
return ClosingStreamingResponse(
|
||||||
stream,
|
stream,
|
||||||
@@ -2199,7 +2240,7 @@ 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)
|
guarded_chunks = await self._guard_stream(response, model_obj, sse=True)
|
||||||
|
|
||||||
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
||||||
usage_finalized = False
|
usage_finalized = False
|
||||||
@@ -2469,6 +2510,8 @@ class BaseUpstreamProvider:
|
|||||||
maybe_cost_event = await finalize_without_usage()
|
maybe_cost_event = await finalize_without_usage()
|
||||||
if maybe_cost_event is not None:
|
if maybe_cost_event is not None:
|
||||||
yield maybe_cost_event
|
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:
|
except httpx.ReadError:
|
||||||
if not usage_finalized:
|
if not usage_finalized:
|
||||||
@@ -3421,7 +3464,7 @@ class BaseUpstreamProvider:
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
result = self._generic_streaming_response(
|
result = await self._generic_streaming_response(
|
||||||
response,
|
response,
|
||||||
key.hashed_key,
|
key.hashed_key,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
@@ -3706,7 +3749,7 @@ class BaseUpstreamProvider:
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
result = self._generic_streaming_response(
|
result = await self._generic_streaming_response(
|
||||||
response,
|
response,
|
||||||
key.hashed_key,
|
key.hashed_key,
|
||||||
max_cost_for_model,
|
max_cost_for_model,
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ cooldown that outlives a restart would hide a provider that has recovered.
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import time
|
import time
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
from ..core import get_logger
|
from ..core import get_logger
|
||||||
from ..core.settings import settings
|
from ..core.settings import settings
|
||||||
@@ -19,6 +20,24 @@ _failures: dict[tuple[str, str], list[float]] = {}
|
|||||||
_cooling_until: dict[tuple[str, str], 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:
|
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."""
|
"""Count a timeout or 5xx, opening a cooldown once too many land in a minute."""
|
||||||
if settings.upstream_cooldown_seconds <= 0:
|
if settings.upstream_cooldown_seconds <= 0:
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from collections.abc import AsyncIterator
|
from collections.abc import AsyncIterator, Callable
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
@@ -11,25 +11,93 @@ from ..core import get_logger
|
|||||||
from ..core.error_scope import UPSTREAM_ERROR_STATUS
|
from ..core.error_scope import UPSTREAM_ERROR_STATUS
|
||||||
from ..core.exceptions import UpstreamError
|
from ..core.exceptions import UpstreamError
|
||||||
from ..core.settings import settings
|
from ..core.settings import settings
|
||||||
|
from .sse_splitter import SSEEventSplitter
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
async def open_guarded_stream(
|
class GuardedStream(AsyncIterator[bytes]):
|
||||||
response: httpx.Response, provider_type: str
|
def __init__(
|
||||||
) -> AsyncIterator[bytes]:
|
self,
|
||||||
"""Await the upstream's first chunk, then hand back the whole stream.
|
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
|
def __aiter__(self) -> GuardedStream:
|
||||||
makes a slow-starting provider recoverable: the proxy's candidate loop only
|
return self
|
||||||
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
|
async def __anext__(self) -> bytes:
|
||||||
iterator simply ends and the caller's finalizer settles actual usage.
|
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__()
|
chunks = response.aiter_bytes().__aiter__()
|
||||||
|
guarded_chunks = _sse_events(chunks) if sse else chunks
|
||||||
timeout = settings.upstream_first_token_timeout_seconds
|
timeout = settings.upstream_first_token_timeout_seconds
|
||||||
try:
|
try:
|
||||||
first = await _next_chunk(chunks, timeout)
|
first = await _next_chunk(guarded_chunks, timeout)
|
||||||
except TimeoutError:
|
except TimeoutError:
|
||||||
await response.aclose()
|
await response.aclose()
|
||||||
raise UpstreamError(
|
raise UpstreamError(
|
||||||
@@ -37,7 +105,7 @@ async def open_guarded_stream(
|
|||||||
status_code=UPSTREAM_ERROR_STATUS,
|
status_code=UPSTREAM_ERROR_STATUS,
|
||||||
code="UPSTREAM_TIMEOUT",
|
code="UPSTREAM_TIMEOUT",
|
||||||
) from None
|
) 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:
|
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)
|
return await (asyncio.wait_for(step, timeout) if timeout > 0 else step)
|
||||||
except StopAsyncIteration:
|
except StopAsyncIteration:
|
||||||
return None
|
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)
|
reservation = MagicMock(spec=ReservationSnapshot)
|
||||||
upstream_response.status_code = 201
|
upstream_response.status_code = 201
|
||||||
upstream_response.headers = {"x-upstream": "preserved"}
|
upstream_response.headers = {"x-upstream": "preserved"}
|
||||||
response = provider._generic_streaming_response(
|
response = await provider._generic_streaming_response(
|
||||||
upstream_response,
|
upstream_response,
|
||||||
"key-hash",
|
"key-hash",
|
||||||
500,
|
500,
|
||||||
@@ -377,7 +377,7 @@ async def test_generic_stream_settles_when_response_start_fails() -> None:
|
|||||||
upstream_response.status_code = 201
|
upstream_response.status_code = 201
|
||||||
upstream_response.headers = {"x-upstream": "preserved"}
|
upstream_response.headers = {"x-upstream": "preserved"}
|
||||||
reservation = MagicMock(spec=ReservationSnapshot)
|
reservation = MagicMock(spec=ReservationSnapshot)
|
||||||
response = provider._generic_streaming_response(
|
response = await provider._generic_streaming_response(
|
||||||
upstream_response,
|
upstream_response,
|
||||||
"key-hash",
|
"key-hash",
|
||||||
500,
|
500,
|
||||||
|
|||||||
@@ -9,8 +9,14 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||||||
import httpx
|
import httpx
|
||||||
import pytest
|
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.exceptions import UpstreamError
|
||||||
from routstr.core.settings import settings
|
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.cooldown import is_cooling_down, record_failure
|
||||||
from routstr.upstream.stream_timeout import open_guarded_stream
|
from routstr.upstream.stream_timeout import open_guarded_stream
|
||||||
|
|
||||||
@@ -33,6 +39,12 @@ async def _stalls_after_first() -> AsyncIterator[bytes]:
|
|||||||
yield b"never delivered"
|
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
|
@pytest.fixture
|
||||||
def fast_timeouts(monkeypatch: pytest.MonkeyPatch) -> None:
|
def fast_timeouts(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
monkeypatch.setattr(settings, "upstream_first_token_timeout_seconds", 0.01)
|
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()
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_zero_first_token_timeout_disables_the_guard(
|
async def test_zero_first_token_timeout_disables_the_guard(
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
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"]
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_guarded_stream_passes_every_chunk_through() -> None:
|
async def test_guarded_stream_passes_every_chunk_through() -> None:
|
||||||
async def _chunks() -> AsyncIterator[bytes]:
|
async def _chunks() -> AsyncIterator[bytes]:
|
||||||
@@ -135,6 +276,7 @@ def _upstream(base_url: str, forward: AsyncMock) -> MagicMock:
|
|||||||
upstream = MagicMock()
|
upstream = MagicMock()
|
||||||
upstream.provider_type = "test"
|
upstream.provider_type = "test"
|
||||||
upstream.base_url = base_url
|
upstream.base_url = base_url
|
||||||
|
upstream.db_id = None
|
||||||
upstream.prepare_headers = MagicMock(side_effect=lambda h: h)
|
upstream.prepare_headers = MagicMock(side_effect=lambda h: h)
|
||||||
upstream.forward_request = forward
|
upstream.forward_request = forward
|
||||||
return upstream
|
return upstream
|
||||||
@@ -143,6 +285,7 @@ def _upstream(base_url: str, forward: AsyncMock) -> MagicMock:
|
|||||||
async def _run_proxy(
|
async def _run_proxy(
|
||||||
candidates: list[tuple[MagicMock, MagicMock]],
|
candidates: list[tuple[MagicMock, MagicMock]],
|
||||||
revert_mock: AsyncMock,
|
revert_mock: AsyncMock,
|
||||||
|
request: MagicMock | None = None,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
from routstr import proxy as proxy_module
|
from routstr import proxy as proxy_module
|
||||||
from routstr.auth import ReservationSnapshot
|
from routstr.auth import ReservationSnapshot
|
||||||
@@ -173,7 +316,7 @@ async def _run_proxy(
|
|||||||
),
|
),
|
||||||
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
|
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
|
||||||
):
|
):
|
||||||
request = _proxy_request()
|
request = request or _proxy_request()
|
||||||
return await proxy_module._proxy(
|
return await proxy_module._proxy(
|
||||||
request, "v1/chat/completions", MagicMock(), await request.body()
|
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))
|
healthy = _upstream("https://ok.example", AsyncMock(return_value=healthy_response))
|
||||||
candidates = [(MagicMock(), sick), (MagicMock(), healthy)]
|
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
|
assert await _run_proxy(candidates, AsyncMock()) is healthy_response
|
||||||
sick.forward_request.assert_not_awaited()
|
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_response.status_code = 200
|
||||||
only = _upstream("https://only.example", AsyncMock(return_value=only_response))
|
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
|
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