mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
perf: reduce request latency
This commit is contained in:
@@ -65,6 +65,17 @@ ROUTSTR_SECRET_KEY=
|
||||
# Network Configuration
|
||||
# CORS_ORIGINS=*
|
||||
# TOR_PROXY_URL=socks5://127.0.0.1:9050
|
||||
# PROXY_EXTRA_ALLOWED_PATHS=
|
||||
|
||||
# Upstream Connection Pools (one pool per upstream origin)
|
||||
# UPSTREAM_MAX_CONNECTIONS=200
|
||||
# UPSTREAM_MAX_KEEPALIVE_CONNECTIONS=50
|
||||
# UPSTREAM_KEEPALIVE_EXPIRY=60
|
||||
# UPSTREAM_POOL_TIMEOUT=5
|
||||
# UPSTREAM_CONNECT_TIMEOUT=30
|
||||
# UPSTREAM_READ_TIMEOUT=900
|
||||
# UPSTREAM_WRITE_TIMEOUT=30
|
||||
# UPSTREAM_CONNECT_RETRIES=1
|
||||
|
||||
# Logging
|
||||
# LOG_LEVEL=INFO
|
||||
|
||||
+90
-4
@@ -552,8 +552,8 @@ async def _validate_bearer_key_locked(
|
||||
|
||||
async def pay_for_request(
|
||||
key: ApiKey, cost_per_request: int, session: AsyncSession
|
||||
) -> int:
|
||||
"""Process payment for a request."""
|
||||
) -> ReservationSnapshot:
|
||||
"""Reserve funds and return the durable identity for this request."""
|
||||
# Ensure cost_per_request is at least the minimum allowed request cost
|
||||
cost_per_request = max(cost_per_request, settings.min_request_msat)
|
||||
|
||||
@@ -738,6 +738,53 @@ async def pay_for_request(
|
||||
extra={"reservation_id": reservation.release_id},
|
||||
)
|
||||
|
||||
try:
|
||||
# Identity checks only: this call just committed the reservation, so the
|
||||
# stale-reservation sweeper may legitimately have released it already.
|
||||
# Release is a terminal state that settlement handles; it is not a
|
||||
# mismatch between the record and the request.
|
||||
await _validate_reservation_snapshot(
|
||||
key, reservation, session, require_active=False
|
||||
)
|
||||
except BaseException:
|
||||
released = False
|
||||
try:
|
||||
released = await _transition_reservation_to_released(
|
||||
reservation,
|
||||
session,
|
||||
decrement_requests=True,
|
||||
idempotent_success=True,
|
||||
)
|
||||
except BaseException:
|
||||
try:
|
||||
await session.rollback()
|
||||
except BaseException:
|
||||
pass
|
||||
|
||||
if not released:
|
||||
try:
|
||||
async with create_session() as cleanup_session:
|
||||
released = await _transition_reservation_to_released(
|
||||
reservation,
|
||||
cleanup_session,
|
||||
decrement_requests=True,
|
||||
idempotent_success=True,
|
||||
)
|
||||
except BaseException:
|
||||
logger.exception(
|
||||
"Failed to release invalid billing reservation",
|
||||
extra={"reservation_id": reservation.release_id},
|
||||
)
|
||||
|
||||
if not released:
|
||||
logger.error(
|
||||
"Invalid billing reservation could not be released",
|
||||
extra={"reservation_id": reservation.release_id},
|
||||
)
|
||||
await _stop_reservation_heartbeat(reservation.release_id)
|
||||
_clear_current_reservation(reservation)
|
||||
raise
|
||||
|
||||
logger.info(
|
||||
"Payment processed successfully",
|
||||
extra={
|
||||
@@ -762,7 +809,7 @@ async def pay_for_request(
|
||||
},
|
||||
)
|
||||
|
||||
return cost_per_request
|
||||
return reservation
|
||||
|
||||
|
||||
async def revert_pay_for_request(
|
||||
@@ -1104,7 +1151,7 @@ async def _charge_reservation_rows(
|
||||
return True
|
||||
|
||||
|
||||
async def adjust_payment_for_tokens(
|
||||
async def _adjust_payment_for_tokens(
|
||||
key: ApiKey,
|
||||
response_data: dict,
|
||||
session: AsyncSession,
|
||||
@@ -1540,6 +1587,45 @@ async def adjust_payment_for_tokens(
|
||||
raise AssertionError("Unreachable: unhandled calculate_cost result")
|
||||
|
||||
|
||||
async def adjust_payment_for_tokens(
|
||||
key: ApiKey,
|
||||
response_data: dict,
|
||||
session: AsyncSession,
|
||||
deducted_max_cost: int,
|
||||
model_obj: "Model | None" = None,
|
||||
provider_fee: float | None = None,
|
||||
reservation_snapshot: ReservationSnapshot | None = None,
|
||||
) -> dict:
|
||||
"""Settle payment while exposing latency for every import path."""
|
||||
started = time.perf_counter()
|
||||
key_log_hash = key.hashed_key[:8] + "..."
|
||||
succeeded = False
|
||||
try:
|
||||
result = await _adjust_payment_for_tokens(
|
||||
key,
|
||||
response_data,
|
||||
session,
|
||||
deducted_max_cost,
|
||||
model_obj,
|
||||
provider_fee,
|
||||
reservation_snapshot,
|
||||
)
|
||||
succeeded = True
|
||||
return result
|
||||
finally:
|
||||
logger.info(
|
||||
"Payment settlement finished",
|
||||
extra={
|
||||
"key_hash": key_log_hash,
|
||||
"model": response_data.get("model", "unknown"),
|
||||
"settlement_duration_ms": round(
|
||||
(time.perf_counter() - started) * 1000, 2
|
||||
),
|
||||
"settlement_succeeded": succeeded,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def periodic_dead_key_prune() -> None:
|
||||
"""Periodically prune dead API keys. Interval <= 0 disables it.
|
||||
|
||||
|
||||
+207
-2
@@ -28,7 +28,12 @@ DO NOT modify or remove these messages without updating the usage tracking logic
|
||||
- The 'max_cost_for_model' field is extracted for refund calculation
|
||||
- Must include 'max_cost_for_model' in extra dict
|
||||
|
||||
6. Any ERROR level logs with "upstream" in the message
|
||||
6. "Payment settlement finished" (INFO) - routstr/auth.py and routstr/upstream/ehbp.py
|
||||
- Emitted once per settlement attempt, including EHBP settlements
|
||||
- Carries 'settlement_duration_ms' and 'settlement_succeeded'; the EHBP
|
||||
emitter adds 'settlement_type'
|
||||
|
||||
7. Any ERROR level logs with "upstream" in the message
|
||||
- Used to count upstream provider errors
|
||||
- Helps identify service reliability issues
|
||||
|
||||
@@ -37,11 +42,15 @@ If you need to modify these messages, ensure you also update the parsing logic i
|
||||
- routstr/core/log_manager.py
|
||||
"""
|
||||
|
||||
import copy
|
||||
import logging.config
|
||||
import logging.handlers
|
||||
import os
|
||||
import queue
|
||||
import re
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import tomllib
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
@@ -127,6 +136,202 @@ class DailyRotatingFileHandler(logging.handlers.TimedRotatingFileHandler):
|
||||
pass
|
||||
|
||||
|
||||
class QueuedDailyRotatingFileHandler(logging.Handler):
|
||||
"""Move rotating-file I/O off request threads.
|
||||
|
||||
When both locks are needed, acquire the logging module lock before the
|
||||
handler lock to match ``dictConfig``.
|
||||
"""
|
||||
|
||||
_queue: queue.Queue[logging.LogRecord]
|
||||
_target: DailyRotatingFileHandler
|
||||
_listener: logging.handlers.QueueListener
|
||||
_drain_timeout_seconds = 5.0
|
||||
_reopen_backoff_seconds = 5.0
|
||||
|
||||
def __init__(self, filename: str, **kwargs: Any) -> None:
|
||||
super().__init__()
|
||||
self._filename = filename
|
||||
self._kwargs = kwargs
|
||||
self._stopped = True
|
||||
self._next_open_attempt = 0.0
|
||||
self._open()
|
||||
|
||||
def _open(self) -> None:
|
||||
"""Attach a fresh rotating file handler and start draining it."""
|
||||
# A new queue per listener: QueueListener's stop sentinel is a shared
|
||||
# singleton, so two listeners on one queue would steal each other's.
|
||||
record_queue: queue.Queue[logging.LogRecord] = queue.Queue()
|
||||
target = DailyRotatingFileHandler(self._filename, **self._kwargs)
|
||||
target.setFormatter(self.formatter)
|
||||
listener = logging.handlers.QueueListener(record_queue, target)
|
||||
try:
|
||||
listener.start()
|
||||
except Exception:
|
||||
target.close()
|
||||
raise
|
||||
|
||||
self._queue = record_queue
|
||||
self._target = target
|
||||
self._listener = listener
|
||||
self._stopped = False
|
||||
self._closed = False
|
||||
with getattr(logging, "_lock"):
|
||||
handler_list = getattr(logging, "_handlerList")
|
||||
# This wrapper owns the target's shutdown and lock ordering.
|
||||
handler_list[:] = [
|
||||
reference for reference in handler_list if reference() is not target
|
||||
]
|
||||
if not any(reference() is self for reference in handler_list):
|
||||
getattr(logging, "_addHandlerRef")(self)
|
||||
|
||||
def _reopen_locked(self) -> bool:
|
||||
"""Reopen using the lock order required by ``dictConfig``."""
|
||||
with getattr(logging, "_lock"):
|
||||
self.acquire()
|
||||
try:
|
||||
if not self._stopped:
|
||||
return True
|
||||
if time.monotonic() < self._next_open_attempt:
|
||||
return False
|
||||
try:
|
||||
self._open()
|
||||
except Exception:
|
||||
self._next_open_attempt = (
|
||||
time.monotonic() + self._reopen_backoff_seconds
|
||||
)
|
||||
raise
|
||||
return True
|
||||
finally:
|
||||
self.release()
|
||||
|
||||
def setFormatter(self, fmt: logging.Formatter | None) -> None:
|
||||
super().setFormatter(fmt)
|
||||
self._target.setFormatter(fmt)
|
||||
|
||||
def handle(self, record: logging.LogRecord) -> bool:
|
||||
if not self.filter(record):
|
||||
return False
|
||||
|
||||
while True:
|
||||
self.acquire()
|
||||
try:
|
||||
if not self._stopped:
|
||||
self.emit(record)
|
||||
return True
|
||||
finally:
|
||||
self.release()
|
||||
|
||||
if sys.is_finalizing():
|
||||
# logging.shutdown() already ran; a new listener thread would
|
||||
# never drain, so write the record synchronously instead.
|
||||
self._emit_synchronously(record)
|
||||
return False
|
||||
|
||||
try:
|
||||
# Do not acquire the module lock while holding the handler lock.
|
||||
if not self._reopen_locked():
|
||||
return False
|
||||
except Exception:
|
||||
self.handleError(record)
|
||||
return False
|
||||
|
||||
def _emit_synchronously(self, record: logging.LogRecord) -> None:
|
||||
try:
|
||||
sys.stderr.write(self.format(record) + "\n")
|
||||
except Exception:
|
||||
self.handleError(record)
|
||||
|
||||
def emit(self, record: logging.LogRecord) -> None:
|
||||
try:
|
||||
if not self._stopped:
|
||||
self._queue.put_nowait(copy.copy(record))
|
||||
except Exception:
|
||||
# Handler.handle() does not catch exceptions raised by emit().
|
||||
self.handleError(record)
|
||||
|
||||
def flush(self) -> None:
|
||||
self.acquire()
|
||||
try:
|
||||
if self._stopped:
|
||||
return
|
||||
|
||||
deadline = time.monotonic() + self._drain_timeout_seconds
|
||||
with self._queue.all_tasks_done:
|
||||
while self._queue.unfinished_tasks:
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
break
|
||||
self._queue.all_tasks_done.wait(remaining)
|
||||
pending = self._queue.unfinished_tasks
|
||||
if pending:
|
||||
sys.stderr.write(
|
||||
f"Logging listener for {self._filename} still has {pending} "
|
||||
f"record(s) queued after {self._drain_timeout_seconds}s flush\n"
|
||||
)
|
||||
self._target.flush()
|
||||
finally:
|
||||
self.release()
|
||||
|
||||
def _stop_listener(self) -> bool:
|
||||
thread = self._listener._thread
|
||||
if thread is None:
|
||||
return True
|
||||
self._listener.enqueue_sentinel()
|
||||
thread.join(timeout=self._drain_timeout_seconds)
|
||||
if thread.is_alive():
|
||||
return False
|
||||
self._listener._thread = None
|
||||
return True
|
||||
|
||||
def _close_retired_listener(
|
||||
self,
|
||||
listener: logging.handlers.QueueListener,
|
||||
target: DailyRotatingFileHandler,
|
||||
) -> None:
|
||||
def finish() -> None:
|
||||
thread = listener._thread
|
||||
if thread is not None:
|
||||
thread.join()
|
||||
listener._thread = None
|
||||
try:
|
||||
target.flush()
|
||||
finally:
|
||||
target.close()
|
||||
|
||||
threading.Thread(target=finish, daemon=True).start()
|
||||
|
||||
def close(self) -> None:
|
||||
# Stop the listener under the handler lock, then close the target outside
|
||||
# it because FileHandler.close() also takes the logging module lock.
|
||||
self.acquire()
|
||||
try:
|
||||
target = None
|
||||
retired = None
|
||||
if not self._stopped:
|
||||
listener = self._listener
|
||||
current_target = self._target
|
||||
stopped = self._stop_listener()
|
||||
self._stopped = True
|
||||
if stopped:
|
||||
target = current_target
|
||||
else:
|
||||
retired = (listener, current_target)
|
||||
sys.stderr.write(
|
||||
f"Logging listener for {self._filename} did not stop "
|
||||
f"within {self._drain_timeout_seconds}s; reopening on next record\n"
|
||||
)
|
||||
finally:
|
||||
self.release()
|
||||
|
||||
if retired is not None:
|
||||
self._close_retired_listener(*retired)
|
||||
if target is not None:
|
||||
target.flush()
|
||||
target.close()
|
||||
super().close()
|
||||
|
||||
|
||||
def get_package_version() -> str:
|
||||
"""Read the package version from pyproject.toml."""
|
||||
try:
|
||||
@@ -369,7 +574,7 @@ def setup_logging() -> None:
|
||||
"handlers": {
|
||||
"console": console_handler,
|
||||
"file": {
|
||||
"()": DailyRotatingFileHandler,
|
||||
"()": QueuedDailyRotatingFileHandler,
|
||||
"level": log_level,
|
||||
"formatter": "json",
|
||||
"filename": "logs/app.log",
|
||||
|
||||
@@ -35,6 +35,7 @@ from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_perio
|
||||
from ..refund import periodic_refund_reconcile
|
||||
from ..upstream.auto_topup import periodic_auto_topup
|
||||
from ..upstream.deepseek_v4_pricing_shim import register_deepseek_v4_pricing
|
||||
from ..upstream.http_client import close_upstream_http_client
|
||||
from ..upstream.litellm_routing import configure_litellm
|
||||
from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout
|
||||
from .admin import admin_router
|
||||
@@ -260,6 +261,14 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
"Error stopping background tasks",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
finally:
|
||||
try:
|
||||
await close_upstream_http_client()
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error closing upstream HTTP connection pools",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
|
||||
class _ImmutableStaticFiles(StaticFiles):
|
||||
|
||||
@@ -177,6 +177,30 @@ class Settings(BaseSettings):
|
||||
default=30.0, gt=0, env="DATABASE_BUSY_TIMEOUT"
|
||||
)
|
||||
|
||||
# Per-origin upstream connection pools. These fields are env-only below.
|
||||
upstream_max_connections: int = Field(
|
||||
default=200, ge=1, env="UPSTREAM_MAX_CONNECTIONS"
|
||||
)
|
||||
upstream_max_keepalive_connections: int = Field(
|
||||
default=50, ge=0, env="UPSTREAM_MAX_KEEPALIVE_CONNECTIONS"
|
||||
)
|
||||
upstream_keepalive_expiry: float = Field(
|
||||
default=60.0, gt=0, env="UPSTREAM_KEEPALIVE_EXPIRY"
|
||||
)
|
||||
upstream_pool_timeout: float = Field(default=5.0, gt=0, env="UPSTREAM_POOL_TIMEOUT")
|
||||
upstream_read_timeout: float = Field(
|
||||
default=900.0, gt=0, env="UPSTREAM_READ_TIMEOUT"
|
||||
)
|
||||
upstream_connect_timeout: float = Field(
|
||||
default=30.0, gt=0, env="UPSTREAM_CONNECT_TIMEOUT"
|
||||
)
|
||||
upstream_write_timeout: float = Field(
|
||||
default=30.0, gt=0, env="UPSTREAM_WRITE_TIMEOUT"
|
||||
)
|
||||
upstream_connect_retries: int = Field(
|
||||
default=1, ge=0, env="UPSTREAM_CONNECT_RETRIES"
|
||||
)
|
||||
|
||||
# Logging
|
||||
log_level: str = Field(default="INFO", env="LOG_LEVEL")
|
||||
enable_console_logging: bool = Field(default=True, env="ENABLE_CONSOLE_LOGGING")
|
||||
@@ -238,6 +262,15 @@ ENV_ONLY_FIELDS = frozenset(
|
||||
"database_pool_pre_ping",
|
||||
"database_pool_hold_warn_seconds",
|
||||
"database_busy_timeout",
|
||||
# Reconfiguring a live pool would disrupt in-flight streams.
|
||||
"upstream_max_connections",
|
||||
"upstream_max_keepalive_connections",
|
||||
"upstream_keepalive_expiry",
|
||||
"upstream_pool_timeout",
|
||||
"upstream_read_timeout",
|
||||
"upstream_connect_timeout",
|
||||
"upstream_write_timeout",
|
||||
"upstream_connect_retries",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
+15
-10
@@ -10,7 +10,6 @@ from sqlmodel import select
|
||||
from .algorithm import create_model_mappings
|
||||
from .auth import (
|
||||
ReservationSnapshot,
|
||||
get_reservation_snapshot,
|
||||
pay_for_request,
|
||||
revert_pay_for_request,
|
||||
validate_bearer_key,
|
||||
@@ -487,7 +486,7 @@ async def _proxy(
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
response = await upstream.forward_get_request(request, path, headers)
|
||||
if (
|
||||
response.status_code in [424, 502, 429]
|
||||
response.status_code in [424, 502, 503, 429]
|
||||
and i < len(selected_upstreams) - 1
|
||||
):
|
||||
logger.warning(
|
||||
@@ -723,7 +722,10 @@ async def _proxy(
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
response = await upstream.forward_get_request(request, path, headers)
|
||||
|
||||
if response.status_code in [424, 502, 429] and i < len(candidates) - 1:
|
||||
if (
|
||||
response.status_code in [424, 502, 503, 429]
|
||||
and i < len(candidates) - 1
|
||||
):
|
||||
error_message = ""
|
||||
try:
|
||||
if hasattr(response, "body"):
|
||||
@@ -767,9 +769,8 @@ async def _proxy(
|
||||
|
||||
reservation_snapshot: ReservationSnapshot | None = None
|
||||
if is_ehbp or request_body_dict:
|
||||
await pay_for_request(key, max_cost_for_model, session)
|
||||
reservation_snapshot = await get_reservation_snapshot(key, session)
|
||||
# Snapshot validation performs SELECTs after pay_for_request commits.
|
||||
reservation_snapshot = await pay_for_request(key, max_cost_for_model, session)
|
||||
# pay_for_request refreshes the key after committing the reservation.
|
||||
# End that read transaction before waiting on upstream response headers.
|
||||
await _finish_read_transaction(session)
|
||||
|
||||
@@ -796,15 +797,17 @@ async def _proxy(
|
||||
key, session, max_cost_for_model, reservation_snapshot
|
||||
)
|
||||
try:
|
||||
await pay_for_request(key, candidate_max, session)
|
||||
reservation_snapshot = await pay_for_request(
|
||||
key, candidate_max, session
|
||||
)
|
||||
except HTTPException:
|
||||
if i == len(candidates) - 1:
|
||||
raise
|
||||
await pay_for_request(key, max_cost_for_model, session)
|
||||
reservation_snapshot = await get_reservation_snapshot(key, session)
|
||||
reservation_snapshot = await pay_for_request(
|
||||
key, max_cost_for_model, session
|
||||
)
|
||||
await _finish_read_transaction(session)
|
||||
continue
|
||||
reservation_snapshot = await get_reservation_snapshot(key, session)
|
||||
await _finish_read_transaction(session)
|
||||
max_cost_for_model = candidate_max
|
||||
|
||||
@@ -948,9 +951,11 @@ async def _proxy(
|
||||
|
||||
if response.status_code != 200:
|
||||
# 424 is an upstream failure re-reported by error_scope.
|
||||
# 502/503 are upstream errors, 429 rate limits.
|
||||
should_retry = response.status_code in [
|
||||
424,
|
||||
502,
|
||||
503,
|
||||
429,
|
||||
400,
|
||||
401,
|
||||
|
||||
+899
-613
File diff suppressed because it is too large
Load Diff
+55
-14
@@ -5,7 +5,7 @@ import math
|
||||
import time
|
||||
import traceback
|
||||
from dataclasses import dataclass, field
|
||||
from typing import AsyncIterator, Mapping
|
||||
from typing import AsyncIterator, Awaitable, Mapping
|
||||
from urllib.parse import urlsplit, urlunsplit
|
||||
|
||||
from fastapi import Request
|
||||
@@ -649,6 +649,37 @@ async def _release_failed_ehbp_charge(
|
||||
)
|
||||
|
||||
|
||||
async def _record_ehbp_settlement(
|
||||
operation: Awaitable[int],
|
||||
*,
|
||||
key: ApiKey,
|
||||
model_id: str,
|
||||
settlement_type: str,
|
||||
) -> int:
|
||||
"""Expose EHBP settlement latency alongside normal request settlement."""
|
||||
started = time.perf_counter()
|
||||
# A rollback can expire the ORM instance, so capture this before the operation.
|
||||
key_log_hash = key.hashed_key[:8] + "..."
|
||||
succeeded = False
|
||||
try:
|
||||
result = await operation
|
||||
succeeded = True
|
||||
return result
|
||||
finally:
|
||||
logger.info(
|
||||
"Payment settlement finished",
|
||||
extra={
|
||||
"key_hash": key_log_hash,
|
||||
"model": model_id,
|
||||
"settlement_type": settlement_type,
|
||||
"settlement_duration_ms": round(
|
||||
(time.perf_counter() - started) * 1000, 2
|
||||
),
|
||||
"settlement_succeeded": succeeded,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def finalize_ehbp_actual_cost_payment(
|
||||
key: ApiKey,
|
||||
session: AsyncSession,
|
||||
@@ -937,13 +968,18 @@ async def forward_ehbp_request(
|
||||
)
|
||||
billing_model = cost_info.pop("actual_model", None) or model_obj.id
|
||||
computed_msats = int(cost_info["total_msats"])
|
||||
charged_msats = await finalize_ehbp_actual_cost_payment(
|
||||
key,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
billing_model,
|
||||
cost_info,
|
||||
reservation_snapshot,
|
||||
charged_msats = await _record_ehbp_settlement(
|
||||
finalize_ehbp_actual_cost_payment(
|
||||
key,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
billing_model,
|
||||
cost_info,
|
||||
reservation_snapshot,
|
||||
),
|
||||
key=key,
|
||||
model_id=billing_model,
|
||||
settlement_type="ehbp_usage",
|
||||
)
|
||||
cost_data = {
|
||||
**cost_info,
|
||||
@@ -963,12 +999,17 @@ async def forward_ehbp_request(
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
charged_msats = await finalize_ehbp_max_cost_payment(
|
||||
key,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
model_obj.id,
|
||||
reservation_snapshot,
|
||||
charged_msats = await _record_ehbp_settlement(
|
||||
finalize_ehbp_max_cost_payment(
|
||||
key,
|
||||
session,
|
||||
max_cost_for_model,
|
||||
model_obj.id,
|
||||
reservation_snapshot,
|
||||
),
|
||||
key=key,
|
||||
model_id=model_obj.id,
|
||||
settlement_type="ehbp_unmeasured_release",
|
||||
)
|
||||
cost_data = {
|
||||
"total_msats": charged_msats,
|
||||
|
||||
@@ -44,6 +44,7 @@ Pipeline
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator, AsyncIterator
|
||||
@@ -55,6 +56,7 @@ from ..core import get_logger
|
||||
from ..core.error_scope import ERROR_SCOPE_NODE
|
||||
from ..core.exceptions import UpstreamError
|
||||
from ..payment.models import Model
|
||||
from .http_client import acquire_upstream_http_client
|
||||
from .messages_dispatch import (
|
||||
ANTHROPIC_ONLY_FIELDS,
|
||||
aggregate_anthropic_events_to_message,
|
||||
@@ -64,6 +66,46 @@ logger = get_logger(__name__)
|
||||
|
||||
DUMMY_THOUGHT_SIGNATURE = "skip_thought_signature_validator"
|
||||
|
||||
|
||||
class _ResponseOwnedIterator:
|
||||
"""Close the upstream response even if iteration never starts."""
|
||||
|
||||
def __init__(
|
||||
self, iterator: AsyncIterator[bytes], response: httpx.Response
|
||||
) -> None:
|
||||
self._iterator = iterator
|
||||
self._response = response
|
||||
self._cleanup_task: asyncio.Task[None] | None = None
|
||||
|
||||
def __aiter__(self) -> _ResponseOwnedIterator:
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> bytes:
|
||||
try:
|
||||
return await self._iterator.__anext__()
|
||||
except StopAsyncIteration:
|
||||
await self.aclose()
|
||||
raise
|
||||
except BaseException:
|
||||
try:
|
||||
await self.aclose()
|
||||
finally:
|
||||
raise
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
try:
|
||||
close = getattr(self._iterator, "aclose", None)
|
||||
if close is not None:
|
||||
await close()
|
||||
finally:
|
||||
await self._response.aclose()
|
||||
|
||||
async def aclose(self) -> None:
|
||||
if self._cleanup_task is None:
|
||||
self._cleanup_task = asyncio.create_task(self._cleanup())
|
||||
await asyncio.shield(self._cleanup_task)
|
||||
|
||||
|
||||
# Mapping: OpenAI finish_reason → Anthropic stop_reason
|
||||
_FINISH_TO_STOP = {
|
||||
"stop": "end_turn",
|
||||
@@ -299,17 +341,21 @@ async def _openai_chunks_to_anthropic_events(
|
||||
yield _sse_event("message_stop", {"type": "message_stop"})
|
||||
|
||||
|
||||
GEMINI_STREAM_READ_TIMEOUT_SECONDS = 120.0
|
||||
|
||||
|
||||
async def _post_and_stream(
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
payload: dict,
|
||||
log_extra: dict[str, Any] | None,
|
||||
) -> tuple[httpx.AsyncClient, httpx.Response]:
|
||||
"""POST to upstream chat-completions and return (client, response) for
|
||||
streaming. Caller is responsible for closing both."""
|
||||
) -> httpx.Response:
|
||||
"""POST to upstream chat-completions and return a streaming response."""
|
||||
url = f"{base_url.rstrip('/')}/chat/completions"
|
||||
client = httpx.AsyncClient(timeout=httpx.Timeout(120.0, read=120.0))
|
||||
try:
|
||||
client = acquire_upstream_http_client(url)
|
||||
# HTTPX replaces rather than merges per-request timeout settings.
|
||||
client_timeout = client.timeout
|
||||
request = client.build_request(
|
||||
"POST",
|
||||
url,
|
||||
@@ -319,10 +365,25 @@ async def _post_and_stream(
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "text/event-stream",
|
||||
},
|
||||
timeout=httpx.Timeout(
|
||||
connect=client_timeout.connect,
|
||||
read=GEMINI_STREAM_READ_TIMEOUT_SECONDS,
|
||||
write=client_timeout.write,
|
||||
pool=client_timeout.pool,
|
||||
),
|
||||
)
|
||||
response = await client.send(request, stream=True)
|
||||
except UpstreamError:
|
||||
raise
|
||||
except httpx.PoolTimeout as exc:
|
||||
logger.error(
|
||||
"Gemini messages dispatch pool exhausted",
|
||||
extra={"error": str(exc), "url": url, **(log_extra or {})},
|
||||
)
|
||||
raise UpstreamError(
|
||||
"Upstream connection pool is busy", status_code=503
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
await client.aclose()
|
||||
logger.error(
|
||||
"Gemini messages dispatch HTTP error",
|
||||
extra={"error": str(exc), "url": url, **(log_extra or {})},
|
||||
@@ -336,7 +397,6 @@ async def _post_and_stream(
|
||||
body_bytes = await response.aread()
|
||||
finally:
|
||||
await response.aclose()
|
||||
await client.aclose()
|
||||
body_text = body_bytes.decode("utf-8", errors="replace")
|
||||
logger.error(
|
||||
"Gemini messages dispatch upstream error",
|
||||
@@ -353,7 +413,7 @@ async def _post_and_stream(
|
||||
from_upstream_response=True,
|
||||
)
|
||||
|
||||
return client, response
|
||||
return response
|
||||
|
||||
|
||||
async def dispatch_gemini_messages(
|
||||
@@ -374,9 +434,7 @@ async def dispatch_gemini_messages(
|
||||
aggregates).
|
||||
"""
|
||||
if not request_body:
|
||||
raise UpstreamError(
|
||||
"Missing request body for /v1/messages", status_code=400
|
||||
)
|
||||
raise UpstreamError("Missing request body for /v1/messages", status_code=400)
|
||||
|
||||
try:
|
||||
body: dict = json.loads(request_body)
|
||||
@@ -444,9 +502,7 @@ async def dispatch_gemini_messages(
|
||||
},
|
||||
)
|
||||
|
||||
http_client, response = await _post_and_stream(
|
||||
base_url, api_key, openai_kwargs, log_extra
|
||||
)
|
||||
response = await _post_and_stream(base_url, api_key, openai_kwargs, log_extra)
|
||||
|
||||
async def line_iter() -> AsyncGenerator[str, None]:
|
||||
try:
|
||||
@@ -454,10 +510,9 @@ async def dispatch_gemini_messages(
|
||||
yield line
|
||||
finally:
|
||||
await response.aclose()
|
||||
await http_client.aclose()
|
||||
|
||||
anthropic_event_iter = _openai_chunks_to_anthropic_events(
|
||||
line_iter(), requested_model
|
||||
anthropic_event_iter = _ResponseOwnedIterator(
|
||||
_openai_chunks_to_anthropic_events(line_iter(), requested_model), response
|
||||
)
|
||||
|
||||
if not client_stream:
|
||||
@@ -478,6 +533,8 @@ async def dispatch_gemini_messages(
|
||||
f"Failed to aggregate upstream stream: {exc}",
|
||||
status_code=502,
|
||||
) from exc
|
||||
finally:
|
||||
await anthropic_event_iter.aclose()
|
||||
return client_stream, aggregated, requested_model
|
||||
|
||||
return client_stream, anthropic_event_iter, requested_model
|
||||
|
||||
@@ -0,0 +1,481 @@
|
||||
"""Per-origin HTTP client pools with event-loop-aware shutdown."""
|
||||
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import functools
|
||||
import ipaddress
|
||||
import ssl
|
||||
import threading
|
||||
import weakref
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, cast
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
|
||||
from ..core import get_logger
|
||||
from ..core.exceptions import UpstreamError
|
||||
from ..core.settings import settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
# Guards all module-level bookkeeping (_clients, _client_loop, _closing,
|
||||
# _pending_closes, _failed_closes, _close_completed). Multiple event loops can
|
||||
# live on different OS threads (tests and reload/shutdown paths exercise
|
||||
# this), so compound read-modify-write sequences on these dicts need a real
|
||||
# lock. Reentrant because _collect_completed_closes re-enters _schedule_close
|
||||
# when rehoming clients. Never held across an await.
|
||||
_state_lock = threading.RLock()
|
||||
|
||||
_clients: dict[str, httpx.AsyncClient] = {}
|
||||
_client_loop: asyncio.AbstractEventLoop | None = None
|
||||
_closing = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class _CloseSubmission:
|
||||
client: httpx.AsyncClient
|
||||
completion: concurrent.futures.Future[None]
|
||||
task: asyncio.Task[None] | None = None
|
||||
retired: bool = False
|
||||
settlement_lock: threading.Lock = field(default_factory=threading.Lock)
|
||||
settled_outcome: tuple[str, object | None] | None = None
|
||||
|
||||
|
||||
_pending_closes: dict[
|
||||
asyncio.AbstractEventLoop,
|
||||
dict[concurrent.futures.Future[None], _CloseSubmission],
|
||||
] = {}
|
||||
_failed_closes: dict[asyncio.AbstractEventLoop, set[httpx.AsyncClient]] = {}
|
||||
_close_completed: weakref.WeakKeyDictionary[httpx.AsyncClient, bool] = (
|
||||
weakref.WeakKeyDictionary()
|
||||
)
|
||||
|
||||
|
||||
class _StatelessCookies(httpx.Cookies):
|
||||
"""Prevent response cookies from leaking between callers sharing a pool."""
|
||||
|
||||
def extract_cookies(self, response: httpx.Response) -> None:
|
||||
return
|
||||
|
||||
|
||||
def upstream_origin_key(url: str) -> str:
|
||||
"""Return a canonical origin for an absolute HTTP(S) URL."""
|
||||
error = "Upstream URL must be an absolute HTTP(S) URL with a valid authority"
|
||||
if not isinstance(url, str):
|
||||
raise ValueError(error)
|
||||
try:
|
||||
parts = urlsplit(url)
|
||||
hostname = parts.hostname
|
||||
port = parts.port
|
||||
except ValueError as exc:
|
||||
raise ValueError(error) from exc
|
||||
|
||||
scheme = parts.scheme.lower()
|
||||
authority = parts.netloc.rsplit("@", 1)[-1]
|
||||
if (
|
||||
scheme not in {"http", "https"}
|
||||
or not hostname
|
||||
or "@" in parts.netloc
|
||||
or any(character.isspace() for character in hostname)
|
||||
or authority.endswith(":")
|
||||
):
|
||||
raise ValueError(error)
|
||||
|
||||
try:
|
||||
address = ipaddress.ip_address(hostname)
|
||||
except ValueError:
|
||||
# HTTPX URL serialization applies the same IDNA normalization used for
|
||||
# requests, so Unicode and punycode spellings share one pool key.
|
||||
try:
|
||||
normalized = httpx.URL(url).copy_with(
|
||||
username=None,
|
||||
password=None,
|
||||
path="/",
|
||||
query=None,
|
||||
fragment=None,
|
||||
)
|
||||
except httpx.InvalidURL as exc:
|
||||
raise ValueError(error) from exc
|
||||
return str(normalized).rstrip("/")
|
||||
|
||||
canonical_host = address.compressed
|
||||
if address.version == 6:
|
||||
canonical_host = f"[{canonical_host}]"
|
||||
default_port = 80 if scheme == "http" else 443
|
||||
port_suffix = f":{port}" if port is not None and port != default_port else ""
|
||||
return f"{scheme}://{canonical_host}{port_suffix}"
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _shared_ssl_context() -> ssl.SSLContext:
|
||||
# Loading the CA bundle costs tens of milliseconds; do it once per process
|
||||
# instead of once per origin pool.
|
||||
return httpx.create_ssl_context()
|
||||
|
||||
|
||||
def _build_client() -> httpx.AsyncClient:
|
||||
limits = httpx.Limits(
|
||||
max_connections=settings.upstream_max_connections,
|
||||
max_keepalive_connections=settings.upstream_max_keepalive_connections,
|
||||
keepalive_expiry=settings.upstream_keepalive_expiry,
|
||||
)
|
||||
client = httpx.AsyncClient(
|
||||
transport=httpx.AsyncHTTPTransport(
|
||||
verify=_shared_ssl_context(),
|
||||
limits=limits,
|
||||
retries=settings.upstream_connect_retries,
|
||||
),
|
||||
timeout=httpx.Timeout(
|
||||
connect=settings.upstream_connect_timeout,
|
||||
read=settings.upstream_read_timeout,
|
||||
write=settings.upstream_write_timeout,
|
||||
pool=settings.upstream_pool_timeout,
|
||||
),
|
||||
)
|
||||
# AsyncClient's public setter copies into a concrete Cookies jar, so replace
|
||||
# the backing jar directly to keep response cookies out of it.
|
||||
client._cookies = _StatelessCookies()
|
||||
return client
|
||||
|
||||
|
||||
def _close_is_pending(client: httpx.AsyncClient) -> bool:
|
||||
with _state_lock:
|
||||
return any(
|
||||
submission.client is client
|
||||
for closes in _pending_closes.values()
|
||||
for submission in closes.values()
|
||||
)
|
||||
|
||||
|
||||
def _forget_failed_client(client: httpx.AsyncClient) -> None:
|
||||
with _state_lock:
|
||||
for failed_loop, failed in list(_failed_closes.items()):
|
||||
failed.discard(client)
|
||||
if not failed:
|
||||
_failed_closes.pop(failed_loop, None)
|
||||
|
||||
|
||||
async def _close_client_resources(client: httpx.AsyncClient) -> None:
|
||||
if not client.is_closed:
|
||||
await client.aclose()
|
||||
return
|
||||
|
||||
# HTTPX marks the client closed before awaiting its transports. A retry after
|
||||
# cancellation or failure therefore has to resume at the transport boundary.
|
||||
raw_client = cast(Any, client)
|
||||
resources = [raw_client._transport]
|
||||
resources.extend(
|
||||
proxy for proxy in raw_client._mounts.values() if proxy is not None
|
||||
)
|
||||
seen: set[int] = set()
|
||||
for resource in resources:
|
||||
if id(resource) in seen:
|
||||
continue
|
||||
seen.add(id(resource))
|
||||
await resource.aclose()
|
||||
|
||||
|
||||
def _close_task_outcome(
|
||||
completed: asyncio.Task[None],
|
||||
) -> tuple[str, object | None]:
|
||||
if completed.cancelled():
|
||||
return ("cancelled", None)
|
||||
exception = completed.exception()
|
||||
if exception is not None:
|
||||
return ("exception", exception)
|
||||
return ("result", completed.result())
|
||||
|
||||
|
||||
def _matching_close_outcomes(
|
||||
first: tuple[str, object | None], second: tuple[str, object | None]
|
||||
) -> bool:
|
||||
if first[0] != second[0]:
|
||||
return False
|
||||
if first[0] == "cancelled":
|
||||
return True
|
||||
return first[1] is second[1]
|
||||
|
||||
|
||||
def _settle_close_submission(
|
||||
submission: _CloseSubmission, completed: asyncio.Task[None]
|
||||
) -> None:
|
||||
outcome = _close_task_outcome(completed)
|
||||
with submission.settlement_lock:
|
||||
if submission.settled_outcome is not None:
|
||||
if _matching_close_outcomes(submission.settled_outcome, outcome):
|
||||
return
|
||||
raise RuntimeError("Close submission settled with conflicting outcomes")
|
||||
if submission.completion.done():
|
||||
raise RuntimeError("Close submission completion changed before settlement")
|
||||
|
||||
if outcome[0] == "result":
|
||||
submission.completion.set_result(None)
|
||||
elif outcome[0] == "exception":
|
||||
submission.completion.set_exception(cast(BaseException, outcome[1]))
|
||||
else:
|
||||
submission.completion.set_exception(asyncio.CancelledError())
|
||||
submission.settled_outcome = outcome
|
||||
|
||||
|
||||
def _settle_submission_from_task(submission: _CloseSubmission) -> None:
|
||||
task = submission.task
|
||||
if task is not None and task.done():
|
||||
_settle_close_submission(submission, task)
|
||||
|
||||
|
||||
def _finish_close_submission(
|
||||
submission: _CloseSubmission, completed: asyncio.Task[None]
|
||||
) -> None:
|
||||
if not submission.retired:
|
||||
_settle_close_submission(submission, completed)
|
||||
|
||||
|
||||
def _submit_close(
|
||||
client: httpx.AsyncClient, loop: asyncio.AbstractEventLoop
|
||||
) -> _CloseSubmission:
|
||||
"""Submit a close without creating its coroutine until the loop runs it."""
|
||||
submission = _CloseSubmission(client, concurrent.futures.Future())
|
||||
|
||||
def start() -> None:
|
||||
if not submission.completion.set_running_or_notify_cancel():
|
||||
return
|
||||
submission.task = loop.create_task(_close_client_resources(client))
|
||||
|
||||
submission.task.add_done_callback(
|
||||
lambda completed: _finish_close_submission(submission, completed)
|
||||
)
|
||||
|
||||
loop.call_soon_threadsafe(start)
|
||||
return submission
|
||||
|
||||
|
||||
def _collect_completed_closes() -> None:
|
||||
current_loop = asyncio.get_running_loop()
|
||||
rehome: list[httpx.AsyncClient] = []
|
||||
with _state_lock:
|
||||
for loop, closes in list(_pending_closes.items()):
|
||||
for future, submission in list(closes.items()):
|
||||
client = submission.client
|
||||
_settle_submission_from_task(submission)
|
||||
if not future.done():
|
||||
if submission.task is None and not loop.is_running():
|
||||
submission.retired = True
|
||||
future.cancel()
|
||||
closes.pop(future)
|
||||
rehome.append(client)
|
||||
elif loop.is_closed():
|
||||
# A task on a closed loop cannot resume, so it cannot
|
||||
# race a retry at the owned transport boundary.
|
||||
submission.retired = True
|
||||
closes.pop(future)
|
||||
rehome.append(client)
|
||||
continue
|
||||
closes.pop(future)
|
||||
try:
|
||||
future.result()
|
||||
except concurrent.futures.CancelledError:
|
||||
rehome.append(client)
|
||||
except asyncio.CancelledError:
|
||||
_close_completed.pop(client, None)
|
||||
_failed_closes.setdefault(loop, set()).add(client)
|
||||
except Exception as exc:
|
||||
_close_completed.pop(client, None)
|
||||
_failed_closes.setdefault(loop, set()).add(client)
|
||||
logger.warning(
|
||||
"Failed to close upstream HTTP client",
|
||||
extra={"error": str(exc), "error_type": type(exc).__name__},
|
||||
)
|
||||
else:
|
||||
_close_completed[client] = True
|
||||
_forget_failed_client(client)
|
||||
if not closes:
|
||||
_pending_closes.pop(loop, None)
|
||||
|
||||
for client in rehome:
|
||||
_schedule_close(client, current_loop)
|
||||
|
||||
|
||||
def _resume_stopped_loop(
|
||||
loop: asyncio.AbstractEventLoop,
|
||||
tasks: list[asyncio.Task[None]],
|
||||
timeout: float,
|
||||
) -> bool:
|
||||
if loop.is_closed() or loop.is_running():
|
||||
return False
|
||||
|
||||
async def wait_for_tasks() -> None:
|
||||
await asyncio.wait(tasks, timeout=timeout)
|
||||
|
||||
waiter = wait_for_tasks()
|
||||
try:
|
||||
loop.run_until_complete(waiter)
|
||||
except RuntimeError:
|
||||
waiter.close()
|
||||
return False
|
||||
return all(task.done() for task in tasks)
|
||||
|
||||
|
||||
async def _drain_pending_closes(timeout: float = 5.0) -> None:
|
||||
deadline = asyncio.get_running_loop().time() + timeout
|
||||
while True:
|
||||
_collect_completed_closes()
|
||||
with _state_lock:
|
||||
if not _pending_closes:
|
||||
return
|
||||
if all(loop.is_closed() for loop in _pending_closes):
|
||||
return
|
||||
pending_snapshot = [
|
||||
(
|
||||
owner_loop,
|
||||
[
|
||||
submission.task
|
||||
for submission in closes.values()
|
||||
if submission.task is not None and not submission.task.done()
|
||||
],
|
||||
)
|
||||
for owner_loop, closes in _pending_closes.items()
|
||||
]
|
||||
|
||||
remaining = deadline - asyncio.get_running_loop().time()
|
||||
if remaining <= 0:
|
||||
logger.error(
|
||||
"Timed out draining upstream HTTP client closes; retaining them for retry"
|
||||
)
|
||||
return
|
||||
|
||||
resumed = False
|
||||
for owner_loop, tasks in pending_snapshot:
|
||||
if owner_loop.is_closed() or owner_loop.is_running():
|
||||
continue
|
||||
if not tasks:
|
||||
continue
|
||||
resumed = True
|
||||
await asyncio.to_thread(
|
||||
_resume_stopped_loop,
|
||||
owner_loop,
|
||||
tasks,
|
||||
remaining,
|
||||
)
|
||||
_collect_completed_closes()
|
||||
|
||||
if not resumed:
|
||||
await asyncio.sleep(min(0.01, remaining))
|
||||
|
||||
|
||||
def _schedule_close(
|
||||
client: httpx.AsyncClient, owner_loop: asyncio.AbstractEventLoop
|
||||
) -> None:
|
||||
"""Schedule closure on the owning loop, retaining unfinished work."""
|
||||
with _state_lock:
|
||||
if _close_completed.get(client, False):
|
||||
_forget_failed_client(client)
|
||||
return
|
||||
if _close_is_pending(client):
|
||||
return
|
||||
|
||||
current_loop = asyncio.get_running_loop()
|
||||
execution_loop = owner_loop
|
||||
if owner_loop.is_closed() or not owner_loop.is_running():
|
||||
execution_loop = current_loop
|
||||
logger.warning(
|
||||
"Closing upstream HTTP client outside its inactive event loop"
|
||||
)
|
||||
|
||||
try:
|
||||
submission = _submit_close(client, execution_loop)
|
||||
except RuntimeError:
|
||||
if execution_loop is current_loop:
|
||||
_failed_closes.setdefault(owner_loop, set()).add(client)
|
||||
return
|
||||
logger.warning("Upstream HTTP client event loop stopped during shutdown")
|
||||
submission = _submit_close(client, current_loop)
|
||||
execution_loop = current_loop
|
||||
|
||||
_forget_failed_client(client)
|
||||
_pending_closes.setdefault(execution_loop, {})[submission.completion] = (
|
||||
submission
|
||||
)
|
||||
|
||||
|
||||
def acquire_upstream_http_client(url: str) -> httpx.AsyncClient:
|
||||
"""Return the pooled client for ``url``, mapping failures to ``UpstreamError``.
|
||||
|
||||
Shutdown becomes a 503 so callers can fail over; a malformed provider URL
|
||||
becomes a 502 instead of an unhandled 500.
|
||||
"""
|
||||
try:
|
||||
return get_upstream_http_client(url)
|
||||
except RuntimeError as exc:
|
||||
raise UpstreamError(str(exc), status_code=503) from exc
|
||||
except ValueError as exc:
|
||||
raise UpstreamError(str(exc), status_code=502) from exc
|
||||
|
||||
|
||||
def get_upstream_http_client(url: str) -> httpx.AsyncClient:
|
||||
"""Return the shared client for an absolute upstream URL's origin."""
|
||||
global _client_loop
|
||||
loop = asyncio.get_running_loop()
|
||||
with _state_lock:
|
||||
if _closing:
|
||||
raise RuntimeError("Upstream HTTP client is shutting down")
|
||||
|
||||
_collect_completed_closes()
|
||||
if _client_loop is not loop:
|
||||
stale_clients = list(_clients.values())
|
||||
stale_loop = _client_loop
|
||||
_clients.clear()
|
||||
_client_loop = loop
|
||||
if stale_loop is not None:
|
||||
for stale_client in stale_clients:
|
||||
_schedule_close(stale_client, stale_loop)
|
||||
|
||||
key = upstream_origin_key(url)
|
||||
client = _clients.get(key)
|
||||
if client is not None and not client.is_closed:
|
||||
return client
|
||||
client = _build_client()
|
||||
_clients[key] = client
|
||||
logger.debug(
|
||||
"Opened upstream HTTP connection pool",
|
||||
extra={
|
||||
"origin": key,
|
||||
"max_connections": settings.upstream_max_connections,
|
||||
"max_keepalive_connections": settings.upstream_max_keepalive_connections,
|
||||
"pool_timeout": settings.upstream_pool_timeout,
|
||||
"read_timeout": settings.upstream_read_timeout,
|
||||
},
|
||||
)
|
||||
return client
|
||||
|
||||
|
||||
async def close_upstream_http_client() -> None:
|
||||
"""Close every pool, using its owner loop while that loop remains active."""
|
||||
global _client_loop, _closing
|
||||
|
||||
with _state_lock:
|
||||
_collect_completed_closes()
|
||||
clients = list(_clients.values())
|
||||
owner_loop = _client_loop
|
||||
failed_clients = [
|
||||
(failed_loop, client)
|
||||
for failed_loop, failed in _failed_closes.items()
|
||||
for client in failed
|
||||
]
|
||||
if not clients and not failed_clients and not _pending_closes:
|
||||
return
|
||||
|
||||
_closing = True
|
||||
_clients.clear()
|
||||
_client_loop = None
|
||||
if owner_loop is not None:
|
||||
for client in clients:
|
||||
_schedule_close(client, owner_loop)
|
||||
for failed_loop, client in failed_clients:
|
||||
_schedule_close(client, failed_loop)
|
||||
|
||||
try:
|
||||
await _drain_pending_closes()
|
||||
finally:
|
||||
with _state_lock:
|
||||
_closing = False
|
||||
@@ -55,9 +55,13 @@ async def test_reserve_increases_reserved_balance(
|
||||
cost = 100
|
||||
key = await _persist(integration_session, _make_key(balance=500))
|
||||
|
||||
await pay_for_request(key, cost, integration_session)
|
||||
reservation = await pay_for_request(key, cost, integration_session)
|
||||
await integration_session.refresh(key)
|
||||
|
||||
assert reservation.key_hash == key.hashed_key
|
||||
assert reservation.billing_key_hash == key.hashed_key
|
||||
assert reservation.reserved_msats == cost
|
||||
assert reservation.release_id
|
||||
assert key.reserved_balance == cost
|
||||
assert key.balance == 500 # balance column is NOT decremented on reserve
|
||||
assert key.total_balance == 500 - cost # available = balance - reserved
|
||||
@@ -75,11 +79,11 @@ async def test_revert_releases_reservation(
|
||||
cost = 150
|
||||
key = await _persist(integration_session, _make_key(balance=300))
|
||||
|
||||
await pay_for_request(key, cost, integration_session)
|
||||
reservation = await pay_for_request(key, cost, integration_session)
|
||||
await integration_session.refresh(key)
|
||||
assert key.reserved_balance == cost
|
||||
|
||||
await revert_pay_for_request(key, integration_session, cost)
|
||||
await revert_pay_for_request(key, integration_session, cost, reservation)
|
||||
await integration_session.refresh(key)
|
||||
|
||||
assert key.reserved_balance == 0
|
||||
|
||||
@@ -42,7 +42,6 @@ def log_dir(tmp_path: Path) -> Path:
|
||||
|
||||
@pytest.fixture
|
||||
def handler(log_dir: Path) -> Iterator[DailyRotatingFileHandler]:
|
||||
"""A file handler configured exactly like the production ``file`` handler."""
|
||||
handler = DailyRotatingFileHandler(
|
||||
str(log_dir / "app.log"),
|
||||
when="midnight",
|
||||
|
||||
@@ -68,11 +68,8 @@ async def _run_proxy(
|
||||
),
|
||||
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=1_000)),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
"get_reservation_snapshot",
|
||||
AsyncMock(return_value=reservation),
|
||||
proxy_module, "pay_for_request", AsyncMock(return_value=reservation)
|
||||
),
|
||||
patch.object(proxy_module, "revert_pay_for_request", AsyncMock()),
|
||||
):
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
from typing import Any
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
import routstr.auth as auth_module
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_payment_settlement_logs_its_duration(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
async def settle(*_args: Any, **_kwargs: Any) -> dict[str, int]:
|
||||
return {"total_cost": 1}
|
||||
|
||||
log_info = Mock()
|
||||
monkeypatch.setattr(auth_module, "_adjust_payment_for_tokens", settle)
|
||||
monkeypatch.setattr(auth_module.logger, "info", log_info)
|
||||
key = ApiKey(hashed_key="abcdefgh1234", balance=0)
|
||||
session = AsyncSession()
|
||||
|
||||
result = await auth_module.adjust_payment_for_tokens(key, {}, session, 10)
|
||||
await session.close()
|
||||
|
||||
assert result == {"total_cost": 1}
|
||||
log_info.assert_called_once()
|
||||
(message,) = log_info.call_args.args
|
||||
extra = log_info.call_args.kwargs["extra"]
|
||||
assert message == "Payment settlement finished"
|
||||
assert extra["settlement_duration_ms"] >= 0
|
||||
assert extra["settlement_succeeded"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_payment_settlement_logs_failure_without_swallowing_it(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
async def fail(*_args: Any, **_kwargs: Any) -> dict:
|
||||
raise RuntimeError("database locked")
|
||||
|
||||
log_info = Mock()
|
||||
monkeypatch.setattr(auth_module, "_adjust_payment_for_tokens", fail)
|
||||
monkeypatch.setattr(auth_module.logger, "info", log_info)
|
||||
key = ApiKey(hashed_key="abcdefgh1234", balance=0)
|
||||
session = AsyncSession()
|
||||
|
||||
with pytest.raises(RuntimeError, match="database locked"):
|
||||
await auth_module.adjust_payment_for_tokens(key, {}, session, 10)
|
||||
await session.close()
|
||||
|
||||
extra = log_info.call_args.kwargs["extra"]
|
||||
assert extra["settlement_duration_ms"] >= 0
|
||||
assert extra["settlement_succeeded"] is False
|
||||
@@ -0,0 +1,168 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncGenerator, AsyncIterator
|
||||
from typing import cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
|
||||
async def _chunks() -> AsyncIterator[bytes]:
|
||||
yield b"chunk"
|
||||
|
||||
|
||||
def _forwarding_case() -> tuple[
|
||||
BaseUpstreamProvider,
|
||||
MagicMock,
|
||||
MagicMock,
|
||||
MagicMock,
|
||||
MagicMock,
|
||||
MagicMock,
|
||||
MagicMock,
|
||||
]:
|
||||
provider = BaseUpstreamProvider("https://api.example.com", "test-key")
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.query_params = {}
|
||||
key = MagicMock()
|
||||
key.hashed_key = "key-hash"
|
||||
session = MagicMock()
|
||||
model = MagicMock()
|
||||
model.forwarded_model_id = None
|
||||
model.id = "model"
|
||||
|
||||
response = MagicMock(spec=httpx.Response)
|
||||
response.status_code = 200
|
||||
response.headers = {"content-type": "application/octet-stream"}
|
||||
response.aclose = AsyncMock()
|
||||
response.aiter_bytes = MagicMock(side_effect=_chunks)
|
||||
|
||||
client = MagicMock()
|
||||
client.build_request.return_value = MagicMock()
|
||||
client.send = AsyncMock(return_value=response)
|
||||
return provider, request, key, session, model, response, client
|
||||
|
||||
|
||||
async def _forward(
|
||||
method_name: str,
|
||||
*,
|
||||
reservation_snapshot: object | None,
|
||||
) -> tuple[StreamingResponse, MagicMock, BaseUpstreamProvider]:
|
||||
provider, request, key, session, model, response, client = _forwarding_case()
|
||||
prepare_method = (
|
||||
"prepare_request_body"
|
||||
if method_name == "forward_request"
|
||||
else "prepare_responses_request_body"
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.base.acquire_upstream_http_client", return_value=client
|
||||
),
|
||||
patch.object(provider, "normalize_request_path", return_value="audio/speech"),
|
||||
patch.object(
|
||||
provider,
|
||||
"build_request_url",
|
||||
return_value="https://api.example.com/audio/speech",
|
||||
),
|
||||
patch.object(provider, prepare_method, return_value=b"{}"),
|
||||
patch.object(provider, "prepare_params", return_value={}),
|
||||
):
|
||||
result = await getattr(provider, method_name)(
|
||||
request=request,
|
||||
path="audio/speech",
|
||||
headers={},
|
||||
request_body=b"{}",
|
||||
key=key,
|
||||
max_cost_for_model=1_000,
|
||||
session=session,
|
||||
model_obj=model,
|
||||
reservation_snapshot=reservation_snapshot,
|
||||
)
|
||||
|
||||
assert isinstance(result, StreamingResponse)
|
||||
return result, response, provider
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"method_name", ["forward_request", "forward_responses_request"]
|
||||
)
|
||||
async def test_cancellation_before_stream_handoff_closes_response_once(
|
||||
method_name: str,
|
||||
) -> None:
|
||||
provider, request, key, session, model, response, client = _forwarding_case()
|
||||
prepare_method = (
|
||||
"prepare_request_body"
|
||||
if method_name == "forward_request"
|
||||
else "prepare_responses_request_body"
|
||||
)
|
||||
lookup_started = asyncio.Event()
|
||||
|
||||
async def wait_for_reservation(*_: object) -> None:
|
||||
lookup_started.set()
|
||||
await asyncio.Future()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.upstream.base.acquire_upstream_http_client", return_value=client
|
||||
),
|
||||
patch.object(provider, "normalize_request_path", return_value="audio/speech"),
|
||||
patch.object(
|
||||
provider,
|
||||
"build_request_url",
|
||||
return_value="https://api.example.com/audio/speech",
|
||||
),
|
||||
patch.object(provider, prepare_method, return_value=b"{}"),
|
||||
patch.object(provider, "prepare_params", return_value={}),
|
||||
patch(
|
||||
"routstr.upstream.base.get_reservation_snapshot",
|
||||
side_effect=wait_for_reservation,
|
||||
),
|
||||
):
|
||||
task = asyncio.create_task(
|
||||
getattr(provider, method_name)(
|
||||
request=request,
|
||||
path="audio/speech",
|
||||
headers={},
|
||||
request_body=b"{}",
|
||||
key=key,
|
||||
max_cost_for_model=1_000,
|
||||
session=session,
|
||||
model_obj=model,
|
||||
)
|
||||
)
|
||||
await lookup_started.wait()
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
response.aclose.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"method_name", ["forward_request", "forward_responses_request"]
|
||||
)
|
||||
async def test_successful_stream_handoff_does_not_close_response_early(
|
||||
method_name: str,
|
||||
) -> None:
|
||||
result, response, provider = await _forward(
|
||||
method_name,
|
||||
reservation_snapshot=MagicMock(),
|
||||
)
|
||||
response.aclose.assert_not_awaited()
|
||||
|
||||
iterator = cast(AsyncGenerator[bytes, None], result.body_iterator)
|
||||
with patch.object(
|
||||
provider,
|
||||
"_finalize_generic_streaming_payment",
|
||||
new=AsyncMock(),
|
||||
):
|
||||
assert await anext(iterator) == b"chunk"
|
||||
await iterator.aclose()
|
||||
|
||||
response.aclose.assert_awaited_once_with()
|
||||
@@ -0,0 +1,291 @@
|
||||
import logging
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
import threading
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
import routstr.core.logging as routstr_logging
|
||||
from routstr.core.logging import QueuedDailyRotatingFileHandler
|
||||
|
||||
|
||||
def _log_text(tmp_path: Path) -> str:
|
||||
return "".join(path.read_text() for path in sorted(tmp_path.glob("app_*.log")))
|
||||
|
||||
|
||||
def _make_handler(
|
||||
tmp_path: Path, name: str
|
||||
) -> tuple[logging.Logger, QueuedDailyRotatingFileHandler]:
|
||||
handler = QueuedDailyRotatingFileHandler(
|
||||
str(tmp_path / "app.log"), when="midnight", backupCount=1
|
||||
)
|
||||
handler.setFormatter(logging.Formatter("%(message)s"))
|
||||
logger = logging.Logger(name)
|
||||
logger.addHandler(handler)
|
||||
return logger, handler
|
||||
|
||||
|
||||
def test_queued_file_handler_flushes_records_on_close(tmp_path: Path) -> None:
|
||||
logger, handler = _make_handler(tmp_path, "queued-file-test")
|
||||
try:
|
||||
logger.info("written from listener")
|
||||
handler.flush()
|
||||
|
||||
assert "written from listener" in _log_text(tmp_path)
|
||||
finally:
|
||||
handler.close()
|
||||
|
||||
|
||||
def test_queued_file_handler_loses_no_records_on_close(tmp_path: Path) -> None:
|
||||
logger, handler = _make_handler(tmp_path, "queued-file-drain-test")
|
||||
try:
|
||||
for index in range(400):
|
||||
logger.info("Payment processed successfully %d", index)
|
||||
finally:
|
||||
handler.close()
|
||||
|
||||
written = _log_text(tmp_path)
|
||||
assert written.count("Payment processed successfully") == 400
|
||||
|
||||
|
||||
def test_queued_file_handler_keeps_logging_after_close(tmp_path: Path) -> None:
|
||||
"""dictConfig closes live handlers; uvicorn runs one after app import."""
|
||||
logger, handler = _make_handler(tmp_path, "queued-file-reopen-test")
|
||||
logger.info("before close")
|
||||
handler.close()
|
||||
|
||||
logger.info("after close")
|
||||
handler.close()
|
||||
assert "after close" in _log_text(tmp_path)
|
||||
|
||||
handler_list = getattr(logging, "_handlerList")
|
||||
handler_list[:] = [
|
||||
reference for reference in handler_list if reference() is not handler
|
||||
]
|
||||
handler.close()
|
||||
logger.info("after reopen")
|
||||
assert any(reference() is handler for reference in handler_list)
|
||||
logging.shutdown(
|
||||
handlerList=[reference for reference in handler_list if reference() is handler]
|
||||
)
|
||||
assert "after reopen" in _log_text(tmp_path)
|
||||
|
||||
|
||||
def test_queued_file_handler_contains_reopen_failures(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
logger, handler = _make_handler(tmp_path, "queued-file-failure-test")
|
||||
handler.close()
|
||||
|
||||
attempts = 0
|
||||
|
||||
def fail_to_open(*args: object, **kwargs: object) -> None:
|
||||
nonlocal attempts
|
||||
attempts += 1
|
||||
raise OSError("disk unavailable")
|
||||
|
||||
errors: list[logging.LogRecord] = []
|
||||
monkeypatch.setattr(routstr_logging, "DailyRotatingFileHandler", fail_to_open)
|
||||
monkeypatch.setattr(type(handler), "handleError", lambda _self, r: errors.append(r))
|
||||
|
||||
for _ in range(50):
|
||||
logger.info("must not reach billing")
|
||||
|
||||
assert attempts == 1
|
||||
assert len(errors) == 1
|
||||
handler.close()
|
||||
|
||||
|
||||
def test_queued_file_handler_emit_does_not_raise_into_caller(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
logger, handler = _make_handler(tmp_path, "queued-file-emit-failure-test")
|
||||
handled: list[logging.LogRecord] = []
|
||||
|
||||
class BrokenQueue:
|
||||
def put_nowait(self, _record: logging.LogRecord) -> None:
|
||||
raise OSError("queue is gone")
|
||||
|
||||
monkeypatch.setattr(handler, "_queue", BrokenQueue())
|
||||
monkeypatch.setattr(type(handler), "handleError", lambda _s, r: handled.append(r))
|
||||
|
||||
logger.info("settlement line")
|
||||
|
||||
assert len(handled) == 1
|
||||
handler.close()
|
||||
|
||||
|
||||
def test_queued_file_handler_recovers_after_close_timeout(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
logger, handler = _make_handler(tmp_path, "queued-file-timeout-test")
|
||||
listener_blocked = threading.Event()
|
||||
allow_listener = threading.Event()
|
||||
old_target = handler._target
|
||||
original_handle = old_target.handle
|
||||
original_close = old_target.close
|
||||
target_closed = threading.Event()
|
||||
close_count = 0
|
||||
|
||||
def blocked_handle(record: logging.LogRecord) -> bool:
|
||||
listener_blocked.set()
|
||||
assert allow_listener.wait(timeout=10)
|
||||
return original_handle(record)
|
||||
|
||||
def track_close() -> None:
|
||||
nonlocal close_count
|
||||
close_count += 1
|
||||
original_close()
|
||||
target_closed.set()
|
||||
|
||||
monkeypatch.setattr(old_target, "handle", blocked_handle)
|
||||
monkeypatch.setattr(old_target, "close", track_close)
|
||||
handler._drain_timeout_seconds = 0.01
|
||||
logger.info("blocked record")
|
||||
assert listener_blocked.wait(timeout=10)
|
||||
|
||||
handler.close()
|
||||
logger.info("record after timeout")
|
||||
allow_listener.set()
|
||||
handler.close()
|
||||
|
||||
assert target_closed.wait(timeout=10)
|
||||
assert close_count == 1
|
||||
assert "record after timeout" in _log_text(tmp_path)
|
||||
|
||||
|
||||
def test_queued_file_handler_reopens_when_close_wins_emit_race(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
logger, handler = _make_handler(tmp_path, "queued-file-atomic-race-test")
|
||||
emitter_waiting = threading.Event()
|
||||
allow_emitter = threading.Event()
|
||||
original_acquire = handler.acquire
|
||||
emitter_thread: threading.Thread | None = None
|
||||
gated = True
|
||||
|
||||
def gated_acquire() -> None:
|
||||
nonlocal gated
|
||||
if gated and threading.current_thread() is emitter_thread:
|
||||
gated = False
|
||||
emitter_waiting.set()
|
||||
assert allow_emitter.wait(timeout=10)
|
||||
original_acquire()
|
||||
|
||||
monkeypatch.setattr(handler, "acquire", gated_acquire)
|
||||
emitter_thread = threading.Thread(target=logger.info, args=("racing record",))
|
||||
try:
|
||||
emitter_thread.start()
|
||||
assert emitter_waiting.wait(timeout=10)
|
||||
|
||||
handler.close()
|
||||
allow_emitter.set()
|
||||
emitter_thread.join(timeout=10)
|
||||
assert not emitter_thread.is_alive()
|
||||
|
||||
handler.close()
|
||||
assert "racing record" in _log_text(tmp_path)
|
||||
finally:
|
||||
allow_emitter.set()
|
||||
handler.close()
|
||||
|
||||
|
||||
def test_queued_file_handler_survives_close_racing_with_emit(tmp_path: Path) -> None:
|
||||
logger, handler = _make_handler(tmp_path, "queued-file-race-test")
|
||||
done = threading.Event()
|
||||
|
||||
def spam() -> None:
|
||||
while not done.is_set():
|
||||
logger.info("racing record")
|
||||
|
||||
def churn() -> None:
|
||||
for _ in range(50):
|
||||
handler.close()
|
||||
|
||||
emitter = threading.Thread(target=spam, daemon=True)
|
||||
closer = threading.Thread(target=churn, daemon=True)
|
||||
try:
|
||||
emitter.start()
|
||||
closer.start()
|
||||
|
||||
closer.join(timeout=10)
|
||||
done.set()
|
||||
emitter.join(timeout=10)
|
||||
|
||||
assert not closer.is_alive(), "close() deadlocked against a concurrent emit()"
|
||||
assert not emitter.is_alive(), "emit() deadlocked against a concurrent close()"
|
||||
|
||||
logger.info("final record")
|
||||
handler.flush()
|
||||
assert "final record" in _log_text(tmp_path)
|
||||
finally:
|
||||
done.set()
|
||||
handler.close()
|
||||
|
||||
|
||||
def test_queued_file_handler_does_not_deadlock_against_dictconfig(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
script = textwrap.dedent(
|
||||
"""
|
||||
import logging
|
||||
import logging.config
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
from routstr.core.logging import QueuedDailyRotatingFileHandler
|
||||
|
||||
log_dir = Path(sys.argv[1])
|
||||
handler = QueuedDailyRotatingFileHandler(
|
||||
str(log_dir / "app.log"), when="midnight", backupCount=1
|
||||
)
|
||||
handler.setFormatter(logging.Formatter("%(message)s"))
|
||||
logger = logging.Logger("queued-file-dictconfig-test")
|
||||
logger.addHandler(handler)
|
||||
emitted = threading.Event()
|
||||
|
||||
def spam():
|
||||
for _ in range(100):
|
||||
logger.info("racing record")
|
||||
emitted.set()
|
||||
handler.close()
|
||||
time.sleep(0.001)
|
||||
|
||||
def reconfigure():
|
||||
assert emitted.wait(timeout=10)
|
||||
for _ in range(10):
|
||||
logging.config.dictConfig(
|
||||
{
|
||||
"version": 1,
|
||||
"disable_existing_loggers": False,
|
||||
"handlers": {},
|
||||
"loggers": {},
|
||||
"root": {"level": "INFO"},
|
||||
}
|
||||
)
|
||||
time.sleep(0.001)
|
||||
|
||||
emitter = threading.Thread(target=spam)
|
||||
configurer = threading.Thread(target=reconfigure)
|
||||
emitter.start()
|
||||
configurer.start()
|
||||
emitter.join(timeout=20)
|
||||
configurer.join(timeout=20)
|
||||
assert not emitter.is_alive(), "logging deadlocked against dictConfig"
|
||||
assert not configurer.is_alive(), "dictConfig deadlocked against logging"
|
||||
handler.close()
|
||||
"""
|
||||
)
|
||||
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-c", script, str(tmp_path)],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert "racing record" in _log_text(tmp_path)
|
||||
@@ -1,5 +1,6 @@
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from pydantic.v1 import ValidationError
|
||||
@@ -7,7 +8,7 @@ from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlmodel import text
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.settings import Settings, SettingsService, settings
|
||||
from routstr.core.settings import ENV_ONLY_FIELDS, Settings, SettingsService, settings
|
||||
|
||||
NSEC_HEX = "1" * 64
|
||||
|
||||
@@ -72,6 +73,20 @@ def test_database_pool_defaults_provide_concurrency_headroom() -> None:
|
||||
assert s.database_pool_hold_warn_seconds == 10.0
|
||||
|
||||
|
||||
def test_env_only_settings_are_documented() -> None:
|
||||
env_example = Path(__file__).parents[2] / ".env.example"
|
||||
documented = {
|
||||
line.lstrip("# ").split("=", 1)[0]
|
||||
for line in env_example.read_text().splitlines()
|
||||
if "=" in line
|
||||
}
|
||||
aliases = {
|
||||
Settings.__fields__[field].field_info.extra["env"] for field in ENV_ONLY_FIELDS
|
||||
}
|
||||
|
||||
assert aliases <= documented
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field", "bad_value"),
|
||||
[
|
||||
|
||||
@@ -16,13 +16,15 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
import pytest
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
||||
from sqlalchemy.pool import StaticPool
|
||||
from sqlmodel import SQLModel
|
||||
from sqlmodel import SQLModel, select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
import routstr.auth as auth_module
|
||||
from routstr.auth import pay_for_request
|
||||
from routstr.balance import refund_wallet_endpoint
|
||||
from routstr.core.db import (
|
||||
ApiKey,
|
||||
ReservationRelease,
|
||||
release_stale_reservations,
|
||||
reset_all_reserved_balances,
|
||||
)
|
||||
@@ -55,10 +57,16 @@ async def session() -> "AsyncGenerator[AsyncSession, None]":
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pay_for_request_sets_reserved_at(session: AsyncSession) -> None:
|
||||
async def test_pay_for_request_sets_reserved_at(
|
||||
session: AsyncSession, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
key = ApiKey(hashed_key="paykey", balance=10_000)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
logger_info = MagicMock()
|
||||
payments_info = MagicMock()
|
||||
monkeypatch.setattr(auth_module.logger, "info", logger_info)
|
||||
monkeypatch.setattr(auth_module.payments_logger, "info", payments_info)
|
||||
|
||||
before = int(time.time())
|
||||
await pay_for_request(key, 1_000, session)
|
||||
@@ -67,6 +75,57 @@ async def test_pay_for_request_sets_reserved_at(session: AsyncSession) -> None:
|
||||
assert key.reserved_balance == 1_000
|
||||
assert key.reserved_at is not None
|
||||
assert key.reserved_at >= before
|
||||
success_logs = [
|
||||
call
|
||||
for call in logger_info.call_args_list
|
||||
if call.args == ("Payment processed successfully",)
|
||||
]
|
||||
assert len(success_logs) == 1
|
||||
payments_info.assert_called_once()
|
||||
assert payments_info.call_args.args == ("RESERVE",)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
async def test_pay_for_request_releases_reservation_when_validation_fails(
|
||||
session: AsyncSession, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
key = ApiKey(hashed_key="invalid-reservation", balance=10_000)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
|
||||
async def reject_reservation(*_args: object, **_kwargs: object) -> None:
|
||||
raise RuntimeError("reservation identity changed")
|
||||
|
||||
logger_info = MagicMock()
|
||||
payments_info = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
auth_module, "_validate_reservation_snapshot", reject_reservation
|
||||
)
|
||||
monkeypatch.setattr(auth_module.logger, "info", logger_info)
|
||||
monkeypatch.setattr(auth_module.payments_logger, "info", payments_info)
|
||||
|
||||
with pytest.raises(RuntimeError, match="identity changed"):
|
||||
await pay_for_request(key, 1_000, session)
|
||||
|
||||
assert not any(
|
||||
call.args == ("Payment processed successfully",)
|
||||
for call in logger_info.call_args_list
|
||||
)
|
||||
payments_info.assert_not_called()
|
||||
|
||||
await session.refresh(key)
|
||||
release = (
|
||||
await session.exec(
|
||||
select(ReservationRelease).where(
|
||||
ReservationRelease.key_hash == key.hashed_key
|
||||
)
|
||||
)
|
||||
).one()
|
||||
assert key.reserved_balance == 0
|
||||
assert key.total_requests == 0
|
||||
assert release.status == "released"
|
||||
assert release.id not in auth_module._reservation_heartbeats
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -355,10 +414,9 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
|
||||
),
|
||||
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=1_000)),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
"get_reservation_snapshot",
|
||||
"pay_for_request",
|
||||
AsyncMock(return_value=reservation_snapshot),
|
||||
),
|
||||
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
|
||||
|
||||
@@ -42,8 +42,6 @@ async def test_stream_with_id_injection() -> None:
|
||||
key.hashed_key = "test_hash"
|
||||
key.balance = 1000
|
||||
|
||||
background_tasks = MagicMock()
|
||||
|
||||
# We need to mock adjust_payment_for_tokens since it's called at the end
|
||||
with MagicMock():
|
||||
from routstr.upstream import base
|
||||
@@ -66,7 +64,6 @@ async def test_stream_with_id_injection() -> None:
|
||||
response=mock_response,
|
||||
key=key,
|
||||
max_cost_for_model=100,
|
||||
background_tasks=background_tasks,
|
||||
requested_model="test-model",
|
||||
reservation_snapshot=ReservationSnapshot(
|
||||
release_id="test-release",
|
||||
|
||||
@@ -6,13 +6,13 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import BackgroundTasks
|
||||
from sqlalchemy.exc import SQLAlchemyError
|
||||
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
|
||||
from sqlmodel import SQLModel
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
import routstr.auth as auth_module
|
||||
import routstr.upstream.gemini_messages as gemini_messages
|
||||
from routstr.auth import (
|
||||
ReservationSnapshot,
|
||||
adjust_payment_for_tokens,
|
||||
@@ -160,7 +160,7 @@ async def test_post_commit_failure_cannot_release_charged_reservation() -> None:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generic_background_settlement_uses_explicit_reservation() -> None:
|
||||
async def test_generic_stream_settlement_uses_explicit_reservation() -> None:
|
||||
engine = await _engine()
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test-key", provider_fee=1.0
|
||||
@@ -212,8 +212,285 @@ async def test_generic_background_settlement_uses_explicit_reservation() -> None
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
def _opaque_stream_response(*chunks: bytes) -> MagicMock:
|
||||
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
response = MagicMock(spec=httpx.Response)
|
||||
response.aiter_bytes = aiter_bytes
|
||||
response.aclose = AsyncMock()
|
||||
return response
|
||||
|
||||
|
||||
class _CountingAsyncByteStream(httpx.AsyncByteStream):
|
||||
def __init__(self, *chunks: bytes) -> None:
|
||||
self._chunks = chunks
|
||||
self.close_count = 0
|
||||
|
||||
async def __aiter__(self) -> AsyncGenerator[bytes, None]:
|
||||
for chunk in self._chunks:
|
||||
yield chunk
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.close_count += 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_release_is_terminal_and_suppresses_background_charge() -> None:
|
||||
async def test_generic_stream_completion_settles_and_closes_once() -> None:
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test-key"
|
||||
)
|
||||
finalize = AsyncMock()
|
||||
provider._finalize_generic_streaming_payment = finalize # type: ignore[method-assign]
|
||||
response = _opaque_stream_response(b"first", b"second")
|
||||
reservation = MagicMock(spec=ReservationSnapshot)
|
||||
|
||||
stream = provider._stream_generic_with_settlement(
|
||||
response,
|
||||
"key-hash",
|
||||
500,
|
||||
"audio/speech",
|
||||
None,
|
||||
provider.provider_fee,
|
||||
reservation,
|
||||
)
|
||||
assert [chunk async for chunk in stream] == [b"first", b"second"]
|
||||
await stream.aclose()
|
||||
|
||||
finalize.assert_awaited_once_with(
|
||||
"key-hash",
|
||||
500,
|
||||
"audio/speech",
|
||||
None,
|
||||
provider.provider_fee,
|
||||
reservation,
|
||||
)
|
||||
response.aclose.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generic_stream_abort_settles_and_closes_once() -> None:
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test-key"
|
||||
)
|
||||
finalize = AsyncMock()
|
||||
provider._finalize_generic_streaming_payment = finalize # type: ignore[method-assign]
|
||||
response = _opaque_stream_response(b"first", b"second")
|
||||
reservation = MagicMock(spec=ReservationSnapshot)
|
||||
|
||||
stream = provider._stream_generic_with_settlement(
|
||||
response,
|
||||
"key-hash",
|
||||
500,
|
||||
"audio/speech",
|
||||
None,
|
||||
provider.provider_fee,
|
||||
reservation,
|
||||
)
|
||||
assert await anext(stream) == b"first"
|
||||
await stream.aclose()
|
||||
await stream.aclose()
|
||||
|
||||
finalize.assert_awaited_once_with(
|
||||
"key-hash",
|
||||
500,
|
||||
"audio/speech",
|
||||
None,
|
||||
provider.provider_fee,
|
||||
reservation,
|
||||
)
|
||||
response.aclose.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_response_closes_iterator_when_downstream_send_is_cancelled() -> (
|
||||
None
|
||||
):
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test-key"
|
||||
)
|
||||
finalize = AsyncMock()
|
||||
provider._finalize_generic_streaming_payment = finalize # type: ignore[method-assign]
|
||||
upstream_response = _opaque_stream_response(b"first", b"second")
|
||||
reservation = MagicMock(spec=ReservationSnapshot)
|
||||
upstream_response.status_code = 201
|
||||
upstream_response.headers = {"x-upstream": "preserved"}
|
||||
response = provider._generic_streaming_response(
|
||||
upstream_response,
|
||||
"key-hash",
|
||||
500,
|
||||
"audio/speech",
|
||||
None,
|
||||
provider.provider_fee,
|
||||
reservation,
|
||||
)
|
||||
sent: list[dict[str, object]] = []
|
||||
|
||||
async def receive() -> dict[str, str]:
|
||||
return {"type": "http.disconnect"}
|
||||
|
||||
async def send(message: dict[str, object]) -> None:
|
||||
sent.append(message)
|
||||
if message["type"] == "http.response.body" and message.get("body"):
|
||||
raise asyncio.CancelledError
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"asgi": {"version": "3.0", "spec_version": "2.4"},
|
||||
"method": "GET",
|
||||
"path": "/v1/audio/speech",
|
||||
"raw_path": b"/v1/audio/speech",
|
||||
"query_string": b"",
|
||||
"headers": [],
|
||||
"client": ("127.0.0.1", 1),
|
||||
"server": ("testserver", 80),
|
||||
"scheme": "http",
|
||||
}
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await response(scope, receive, send) # type: ignore[arg-type]
|
||||
|
||||
assert sent[0]["type"] == "http.response.start"
|
||||
assert sent[0]["status"] == 201
|
||||
headers = cast(list[tuple[bytes, bytes]], sent[0]["headers"])
|
||||
assert (b"x-upstream", b"preserved") in headers
|
||||
finalize.assert_awaited_once_with(
|
||||
"key-hash",
|
||||
500,
|
||||
"audio/speech",
|
||||
None,
|
||||
provider.provider_fee,
|
||||
reservation,
|
||||
)
|
||||
upstream_response.aclose.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generic_stream_settles_when_response_start_fails() -> None:
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test-key"
|
||||
)
|
||||
finalize = AsyncMock()
|
||||
provider._finalize_generic_streaming_payment = finalize # type: ignore[method-assign]
|
||||
upstream_response = _opaque_stream_response(b"never-read")
|
||||
upstream_response.status_code = 201
|
||||
upstream_response.headers = {"x-upstream": "preserved"}
|
||||
reservation = MagicMock(spec=ReservationSnapshot)
|
||||
response = provider._generic_streaming_response(
|
||||
upstream_response,
|
||||
"key-hash",
|
||||
500,
|
||||
"audio/speech",
|
||||
None,
|
||||
provider.provider_fee,
|
||||
reservation,
|
||||
)
|
||||
|
||||
async def receive() -> dict[str, str]:
|
||||
return {"type": "http.disconnect"}
|
||||
|
||||
async def send(message: dict[str, object]) -> None:
|
||||
assert message["type"] == "http.response.start"
|
||||
raise RuntimeError("response start failed")
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"asgi": {"version": "3.0", "spec_version": "2.4"},
|
||||
"method": "GET",
|
||||
"path": "/v1/audio/speech",
|
||||
"raw_path": b"/v1/audio/speech",
|
||||
"query_string": b"",
|
||||
"headers": [],
|
||||
"client": ("127.0.0.1", 1),
|
||||
"server": ("testserver", 80),
|
||||
"scheme": "http",
|
||||
}
|
||||
|
||||
with pytest.raises(RuntimeError, match="response start failed"):
|
||||
await response(scope, receive, send) # type: ignore[arg-type]
|
||||
|
||||
finalize.assert_awaited_once_with(
|
||||
"key-hash",
|
||||
500,
|
||||
"audio/speech",
|
||||
None,
|
||||
provider.provider_fee,
|
||||
reservation,
|
||||
)
|
||||
upstream_response.aclose.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("api", ["chat", "responses", "messages"])
|
||||
async def test_parsed_stream_finalizes_when_response_start_fails(api: str) -> None:
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test-key"
|
||||
)
|
||||
upstream_response = _opaque_stream_response(b"never-read")
|
||||
upstream_response.status_code = 200
|
||||
upstream_response.headers = {"content-type": "text/event-stream"}
|
||||
key = MagicMock(spec=ApiKey)
|
||||
key.hashed_key = f"{api}-start-failure"
|
||||
key.balance = 10_000
|
||||
snapshot = ReservationSnapshot(
|
||||
release_id=f"{api}-start-failure-release",
|
||||
key_hash=key.hashed_key,
|
||||
billing_key_hash=key.hashed_key,
|
||||
reserved_msats=500,
|
||||
)
|
||||
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)
|
||||
adjust = AsyncMock(return_value={"input_tokens": 0, "output_tokens": 0})
|
||||
|
||||
with (
|
||||
patch("routstr.upstream.base.adjust_payment_for_tokens", adjust),
|
||||
patch("routstr.upstream.base.create_session", return_value=session_context),
|
||||
):
|
||||
if api == "chat":
|
||||
response = await provider.handle_streaming_chat_completion(
|
||||
upstream_response, key, 500, reservation_snapshot=snapshot
|
||||
)
|
||||
elif api == "responses":
|
||||
response = await provider.handle_streaming_responses_completion(
|
||||
upstream_response, key, 500, reservation_snapshot=snapshot
|
||||
)
|
||||
else:
|
||||
response = await provider.handle_streaming_messages_completion(
|
||||
upstream_response, key, 500, reservation_snapshot=snapshot
|
||||
)
|
||||
|
||||
async def receive() -> dict[str, str]:
|
||||
return {"type": "http.disconnect"}
|
||||
|
||||
async def send(message: dict[str, object]) -> None:
|
||||
assert message["type"] == "http.response.start"
|
||||
raise RuntimeError("response start failed")
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"asgi": {"version": "3.0", "spec_version": "2.4"},
|
||||
"method": "GET",
|
||||
"path": f"/v1/{api}",
|
||||
"raw_path": f"/v1/{api}".encode(),
|
||||
"query_string": b"",
|
||||
"headers": [],
|
||||
"client": ("127.0.0.1", 1),
|
||||
"server": ("testserver", 80),
|
||||
"scheme": "http",
|
||||
}
|
||||
with pytest.raises(RuntimeError, match="response start failed"):
|
||||
await response(scope, receive, send) # type: ignore[arg-type]
|
||||
|
||||
adjust.assert_awaited_once()
|
||||
upstream_response.aclose.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_release_is_terminal_before_error_propagates() -> None:
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test-key"
|
||||
)
|
||||
@@ -237,7 +514,6 @@ async def test_streaming_release_is_terminal_and_suppresses_background_charge()
|
||||
release = AsyncMock(return_value=True)
|
||||
reservation_snapshot = MagicMock()
|
||||
reservation_snapshot.reserved_msats = 500
|
||||
background_tasks = MagicMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
@@ -255,7 +531,6 @@ async def test_streaming_release_is_terminal_and_suppresses_background_charge()
|
||||
response=upstream_response,
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
background_tasks=background_tasks,
|
||||
)
|
||||
|
||||
with pytest.raises(SQLAlchemyError, match="database unavailable"):
|
||||
@@ -264,7 +539,6 @@ async def test_streaming_release_is_terminal_and_suppresses_background_charge()
|
||||
|
||||
session.rollback.assert_awaited_once()
|
||||
release.assert_awaited_once_with(reservation_snapshot, session, 500)
|
||||
background_tasks.add_task.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -357,8 +631,6 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once(
|
||||
)
|
||||
upstream_response.aiter_bytes = aiter_bytes
|
||||
upstream_response.aclose = AsyncMock()
|
||||
client = MagicMock()
|
||||
client.aclose = AsyncMock()
|
||||
key = MagicMock(spec=ApiKey)
|
||||
key.hashed_key = f"{api}-partial"
|
||||
key.balance = 10_000
|
||||
@@ -391,9 +663,7 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once(
|
||||
response=upstream_response,
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
background_tasks=BackgroundTasks(),
|
||||
reservation_snapshot=snapshot,
|
||||
client=client,
|
||||
)
|
||||
else:
|
||||
response = await provider.handle_streaming_responses_completion(
|
||||
@@ -401,7 +671,6 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once(
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
reservation_snapshot=snapshot,
|
||||
client=client,
|
||||
)
|
||||
emitted = bytearray()
|
||||
async for chunk in response.body_iterator:
|
||||
@@ -414,7 +683,6 @@ async def test_partial_remote_protocol_error_finalizes_and_closes_once(
|
||||
else:
|
||||
release.assert_not_awaited()
|
||||
upstream_response.aclose.assert_awaited_once()
|
||||
client.aclose.assert_awaited_once()
|
||||
assert b"[DONE]" not in emitted
|
||||
|
||||
|
||||
@@ -436,8 +704,6 @@ async def test_partial_stream_closes_when_billing_db_is_down(
|
||||
)
|
||||
upstream_response.aiter_bytes = aiter_bytes
|
||||
upstream_response.aclose = AsyncMock()
|
||||
client = MagicMock()
|
||||
client.aclose = AsyncMock()
|
||||
key = MagicMock(spec=ApiKey)
|
||||
key.hashed_key = f"{api}-database-down"
|
||||
key.balance = 10_000
|
||||
@@ -461,9 +727,7 @@ async def test_partial_stream_closes_when_billing_db_is_down(
|
||||
response=upstream_response,
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
background_tasks=BackgroundTasks(),
|
||||
reservation_snapshot=snapshot,
|
||||
client=client,
|
||||
)
|
||||
else:
|
||||
response = await provider.handle_streaming_responses_completion(
|
||||
@@ -471,13 +735,11 @@ async def test_partial_stream_closes_when_billing_db_is_down(
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
reservation_snapshot=snapshot,
|
||||
client=client,
|
||||
)
|
||||
async for _ in response.body_iterator:
|
||||
pass
|
||||
|
||||
upstream_response.aclose.assert_awaited_once()
|
||||
client.aclose.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -633,6 +895,100 @@ async def test_messages_streaming_releases_and_raises_on_billing_failure(
|
||||
release.assert_awaited_once_with(snapshot, session, 500)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gemini_messages_finalizes_when_response_start_fails() -> None:
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test-key"
|
||||
)
|
||||
key = MagicMock(spec=ApiKey)
|
||||
key.hashed_key = "gemini-start-failure"
|
||||
key.balance = 10_000
|
||||
snapshot = ReservationSnapshot(
|
||||
release_id="gemini-start-failure-release",
|
||||
key_hash=key.hashed_key,
|
||||
billing_key_hash=key.hashed_key,
|
||||
reserved_msats=500,
|
||||
)
|
||||
model = MagicMock(spec=Model)
|
||||
model.id = "gemini-test"
|
||||
model.forwarded_model_id = None
|
||||
upstream_stream = _CountingAsyncByteStream(
|
||||
b'data: {"choices":[{"delta":{"content":"unused"}}]}\n\n'
|
||||
)
|
||||
upstream_response = httpx.Response(
|
||||
200,
|
||||
request=httpx.Request("POST", "https://gemini.example/chat/completions"),
|
||||
stream=upstream_stream,
|
||||
)
|
||||
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)
|
||||
adjust = AsyncMock(return_value={"input_tokens": 0, "output_tokens": 0})
|
||||
post_and_stream = AsyncMock(return_value=upstream_response)
|
||||
|
||||
with (
|
||||
patch("routstr.upstream.base.adjust_payment_for_tokens", adjust),
|
||||
patch("routstr.upstream.base.create_session", return_value=session_context),
|
||||
patch.object(
|
||||
gemini_messages,
|
||||
"_translate_anthropic_to_openai",
|
||||
return_value={"messages": []},
|
||||
),
|
||||
patch.object(gemini_messages, "_post_and_stream", post_and_stream),
|
||||
):
|
||||
(
|
||||
client_stream,
|
||||
iterator,
|
||||
requested_model,
|
||||
) = await gemini_messages.dispatch_gemini_messages(
|
||||
request_body=json.dumps(
|
||||
{"model": model.id, "messages": [], "stream": True}
|
||||
).encode(),
|
||||
model_obj=model,
|
||||
base_url="https://gemini.example",
|
||||
api_key="test-key",
|
||||
transform_model_name=lambda name: name,
|
||||
)
|
||||
assert client_stream is True
|
||||
assert requested_model == model.id
|
||||
response = provider._stream_litellm_messages(
|
||||
iterator=iterator,
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
requested_model=requested_model,
|
||||
reservation_snapshot=snapshot,
|
||||
)
|
||||
|
||||
async def receive() -> dict[str, str]:
|
||||
return {"type": "http.disconnect"}
|
||||
|
||||
async def send(message: dict[str, object]) -> None:
|
||||
assert message["type"] == "http.response.start"
|
||||
raise RuntimeError("response start failed")
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"asgi": {"version": "3.0", "spec_version": "2.4"},
|
||||
"method": "GET",
|
||||
"path": "/v1/messages",
|
||||
"raw_path": b"/v1/messages",
|
||||
"query_string": b"",
|
||||
"headers": [],
|
||||
"client": ("127.0.0.1", 1),
|
||||
"server": ("testserver", 80),
|
||||
"scheme": "http",
|
||||
}
|
||||
with pytest.raises(RuntimeError, match="response start failed"):
|
||||
await response(scope, receive, send) # type: ignore[arg-type]
|
||||
|
||||
post_and_stream.assert_awaited_once()
|
||||
adjust.assert_awaited_once()
|
||||
assert upstream_response.is_closed
|
||||
assert upstream_stream.close_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cross_key_reservation_snapshot_is_rejected_without_mutation() -> None:
|
||||
engine = await _engine()
|
||||
@@ -717,7 +1073,6 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat()
|
||||
{"model": model.id, "messages": [{"role": "user", "content": "hi"}]}
|
||||
).encode()
|
||||
|
||||
background_tasks = BackgroundTasks()
|
||||
try:
|
||||
with (
|
||||
patch(
|
||||
@@ -742,7 +1097,6 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat()
|
||||
response=upstream_response,
|
||||
key=key,
|
||||
max_cost_for_model=500,
|
||||
background_tasks=background_tasks,
|
||||
model_obj=model,
|
||||
reservation_snapshot=snapshot,
|
||||
request_body=request_body,
|
||||
@@ -750,10 +1104,6 @@ async def test_client_disconnect_midstream_estimates_usage_and_stops_heartbeat()
|
||||
iterator = cast(AsyncGenerator[bytes, None], response.body_iterator)
|
||||
await iterator.__anext__() # first chunk reaches the client
|
||||
await iterator.aclose() # client aborts the socket here
|
||||
|
||||
# Starlette runs the response's background tasks after the abort.
|
||||
for task in background_tasks.tasks:
|
||||
await task()
|
||||
finally:
|
||||
await auth_module._stop_reservation_heartbeat(snapshot.release_id)
|
||||
|
||||
|
||||
@@ -42,7 +42,9 @@ def _make_response(chunks: list[bytes]) -> MagicMock:
|
||||
return mock_response
|
||||
|
||||
|
||||
async def _drive(chunks: list[bytes], requested_model: str | None = None) -> list[bytes]:
|
||||
async def _drive(
|
||||
chunks: list[bytes], requested_model: str | None = None
|
||||
) -> list[bytes]:
|
||||
"""Run the real streaming generator over ``chunks`` and collect output bytes."""
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test_key"
|
||||
@@ -66,7 +68,6 @@ async def _drive(chunks: list[bytes], requested_model: str | None = None) -> lis
|
||||
response=_make_response(chunks),
|
||||
key=key,
|
||||
max_cost_for_model=100,
|
||||
background_tasks=MagicMock(),
|
||||
requested_model=requested_model,
|
||||
reservation_snapshot=ReservationSnapshot(
|
||||
release_id="test-release",
|
||||
|
||||
@@ -1293,10 +1293,9 @@ async def test_bearer_key_config_422_releases_reservation_and_passes_through() -
|
||||
),
|
||||
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=1_000)),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
"get_reservation_snapshot",
|
||||
"pay_for_request",
|
||||
AsyncMock(return_value=reservation_snapshot),
|
||||
),
|
||||
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
|
||||
|
||||
@@ -14,15 +14,21 @@ These tests cover the two pure helpers that drive the dispatcher:
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import routstr.upstream.gemini_messages as gemini_messages
|
||||
from routstr.core.exceptions import UpstreamError
|
||||
from routstr.upstream.gemini_messages import (
|
||||
DUMMY_THOUGHT_SIGNATURE,
|
||||
_openai_chunks_to_anthropic_events,
|
||||
_ResponseOwnedIterator,
|
||||
inject_thought_signatures,
|
||||
)
|
||||
|
||||
@@ -81,9 +87,7 @@ def test_inject_thought_signatures_preserves_existing_signature() -> None:
|
||||
inject_thought_signatures(messages)
|
||||
|
||||
assert (
|
||||
messages[0]["tool_calls"][0]["extra_content"]["google"][
|
||||
"thought_signature"
|
||||
]
|
||||
messages[0]["tool_calls"][0]["extra_content"]["google"]["thought_signature"]
|
||||
== "real-signature"
|
||||
)
|
||||
|
||||
@@ -138,6 +142,46 @@ async def _lines(*chunks: dict | str) -> AsyncGenerator[str, None]:
|
||||
yield c
|
||||
|
||||
|
||||
class _TrackingStream(httpx.AsyncByteStream):
|
||||
def __init__(
|
||||
self,
|
||||
*chunks: bytes,
|
||||
error: Exception | None = None,
|
||||
started: asyncio.Event | None = None,
|
||||
) -> None:
|
||||
self._chunks = chunks
|
||||
self._error = error
|
||||
self._started = started
|
||||
self.close_count = 0
|
||||
|
||||
async def __aiter__(self) -> AsyncGenerator[bytes, None]:
|
||||
if self._started is not None:
|
||||
self._started.set()
|
||||
await asyncio.Event().wait()
|
||||
for chunk in self._chunks:
|
||||
yield chunk
|
||||
if self._error is not None:
|
||||
raise self._error
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.close_count += 1
|
||||
|
||||
|
||||
def _owned_events(
|
||||
response: httpx.Response,
|
||||
) -> _ResponseOwnedIterator:
|
||||
async def line_iter() -> AsyncGenerator[str, None]:
|
||||
try:
|
||||
async for line in response.aiter_lines():
|
||||
yield line
|
||||
finally:
|
||||
await response.aclose()
|
||||
|
||||
return _ResponseOwnedIterator(
|
||||
_openai_chunks_to_anthropic_events(line_iter(), "gemini-test"), response
|
||||
)
|
||||
|
||||
|
||||
def _parse_anthropic_sse(blocks: list[bytes]) -> list[dict]:
|
||||
"""Flatten a list of Anthropic SSE byte chunks into event dicts."""
|
||||
events: list[dict] = []
|
||||
@@ -150,6 +194,58 @@ def _parse_anthropic_sse(blocks: list[bytes]) -> list[dict]:
|
||||
return events
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_owner_closes_once_after_normal_completion() -> None:
|
||||
stream = _TrackingStream(
|
||||
b'data: {"model":"gemini-test","choices":[{"delta":{"content":"ok"},"finish_reason":"stop"}]}\n\n'
|
||||
)
|
||||
response = httpx.Response(
|
||||
200,
|
||||
request=httpx.Request("POST", "https://gemini.example/chat/completions"),
|
||||
stream=stream,
|
||||
)
|
||||
|
||||
assert [event async for event in _owned_events(response)]
|
||||
assert response.is_closed
|
||||
assert stream.close_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_owner_closes_once_after_body_failure() -> None:
|
||||
stream = _TrackingStream(error=RuntimeError("upstream body failed"))
|
||||
response = httpx.Response(
|
||||
200,
|
||||
request=httpx.Request("POST", "https://gemini.example/chat/completions"),
|
||||
stream=stream,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="upstream body failed"):
|
||||
await _owned_events(response).__anext__()
|
||||
|
||||
assert response.is_closed
|
||||
assert stream.close_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_owner_closes_once_after_cancellation() -> None:
|
||||
started = asyncio.Event()
|
||||
stream = _TrackingStream(started=started)
|
||||
response = httpx.Response(
|
||||
200,
|
||||
request=httpx.Request("POST", "https://gemini.example/chat/completions"),
|
||||
stream=stream,
|
||||
)
|
||||
task = asyncio.create_task(_owned_events(response).__anext__())
|
||||
await started.wait()
|
||||
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
assert response.is_closed
|
||||
assert stream.close_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_translator_emits_text_only_response() -> None:
|
||||
"""Plain text response: message_start → content_block_* (text) →
|
||||
@@ -187,9 +283,7 @@ async def test_translator_emits_text_only_response() -> None:
|
||||
]
|
||||
# Text deltas concatenate to "Hello, world".
|
||||
text_deltas = [
|
||||
e["delta"]["text"]
|
||||
for e in events
|
||||
if e["type"] == "content_block_delta"
|
||||
e["delta"]["text"] for e in events if e["type"] == "content_block_delta"
|
||||
]
|
||||
assert "".join(text_deltas) == "Hello, world"
|
||||
# Stop reason was mapped from openai's "stop".
|
||||
@@ -270,9 +364,7 @@ async def test_translator_emits_tool_use_block() -> None:
|
||||
# Argument deltas were forwarded as input_json_delta partials.
|
||||
deltas = [e for e in events if e["type"] == "content_block_delta"]
|
||||
assert all(d["delta"]["type"] == "input_json_delta" for d in deltas)
|
||||
assert "".join(d["delta"]["partial_json"] for d in deltas) == (
|
||||
'{"cmd": "ls"}'
|
||||
)
|
||||
assert "".join(d["delta"]["partial_json"] for d in deltas) == ('{"cmd": "ls"}')
|
||||
# tool_calls finish_reason → tool_use stop_reason.
|
||||
msg_delta = next(e for e in events if e["type"] == "message_delta")
|
||||
assert msg_delta["delta"]["stop_reason"] == "tool_use"
|
||||
@@ -306,8 +398,36 @@ async def test_translator_handles_done_sentinel_and_blank_lines() -> None:
|
||||
assert events[0]["type"] == "message_start"
|
||||
assert events[-1]["type"] == "message_stop"
|
||||
text = "".join(
|
||||
e["delta"]["text"]
|
||||
for e in events
|
||||
if e["type"] == "content_block_delta"
|
||||
e["delta"]["text"] for e in events if e["type"] == "content_block_delta"
|
||||
)
|
||||
assert text == "ok"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_and_stream_maps_pool_timeout_to_503() -> None:
|
||||
client = MagicMock()
|
||||
client.timeout = httpx.Timeout(10.0)
|
||||
client.build_request = MagicMock(return_value=MagicMock())
|
||||
client.send = AsyncMock(side_effect=httpx.PoolTimeout("pool busy"))
|
||||
with patch(
|
||||
"routstr.upstream.gemini_messages.acquire_upstream_http_client",
|
||||
return_value=client,
|
||||
):
|
||||
with pytest.raises(UpstreamError) as exc_info:
|
||||
await gemini_messages._post_and_stream(
|
||||
"https://gemini.example", "key", {"model": "m"}, None
|
||||
)
|
||||
assert exc_info.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_and_stream_surfaces_shutdown_as_503() -> None:
|
||||
with patch(
|
||||
"routstr.upstream.gemini_messages.acquire_upstream_http_client",
|
||||
side_effect=UpstreamError("shutting down", status_code=503),
|
||||
):
|
||||
with pytest.raises(UpstreamError) as exc_info:
|
||||
await gemini_messages._post_and_stream(
|
||||
"https://gemini.example", "key", {"model": "m"}, None
|
||||
)
|
||||
assert exc_info.value.status_code == 503
|
||||
|
||||
@@ -0,0 +1,735 @@
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import routstr.upstream.http_client as http_client_module
|
||||
from routstr.core.exceptions import UpstreamError
|
||||
from routstr.core.settings import settings
|
||||
from routstr.upstream.http_client import (
|
||||
acquire_upstream_http_client,
|
||||
close_upstream_http_client,
|
||||
get_upstream_http_client,
|
||||
upstream_origin_key,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_http_client_is_reused_until_shutdown() -> None:
|
||||
first = get_upstream_http_client("https://api.example.com/v1/chat")
|
||||
second = get_upstream_http_client("https://api.example.com/v1/models")
|
||||
|
||||
assert second is first
|
||||
assert not first.is_closed
|
||||
|
||||
await close_upstream_http_client()
|
||||
assert first.is_closed
|
||||
|
||||
replacement = get_upstream_http_client("https://api.example.com/v1/chat")
|
||||
try:
|
||||
assert replacement is not first
|
||||
assert not replacement.is_closed
|
||||
finally:
|
||||
await close_upstream_http_client()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_http_client_is_isolated_per_origin() -> None:
|
||||
try:
|
||||
first = get_upstream_http_client("https://one.example.com/v1/chat")
|
||||
second = get_upstream_http_client("https://two.example.com/v1/chat")
|
||||
other_port = get_upstream_http_client("https://one.example.com:8443/v1/chat")
|
||||
|
||||
assert first is not second
|
||||
assert first is not other_port
|
||||
finally:
|
||||
await close_upstream_http_client()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("url", "expected"),
|
||||
[
|
||||
("https://api.example.com/v1/chat?x=1", "https://api.example.com"),
|
||||
("HTTPS://API.EXAMPLE.COM:443/v1/chat", "https://api.example.com"),
|
||||
("http://API.EXAMPLE.COM:80/v1/chat", "http://api.example.com"),
|
||||
("http://api.example.com:8080/v1/chat", "http://api.example.com:8080"),
|
||||
("https://bücher.example/v1/chat", "https://xn--bcher-kva.example"),
|
||||
("https://xn--bcher-kva.example/v1/chat", "https://xn--bcher-kva.example"),
|
||||
("https://[2001:db8::1]/v1/chat", "https://[2001:db8::1]"),
|
||||
(
|
||||
"https://[2001:0DB8:0:0:0:0:0:1]:443/v1/chat",
|
||||
"https://[2001:db8::1]",
|
||||
),
|
||||
("https://[2001:db8::1]:8443/v1/chat", "https://[2001:db8::1]:8443"),
|
||||
],
|
||||
)
|
||||
def test_upstream_origin_key_returns_http_origin(url: str, expected: str) -> None:
|
||||
assert upstream_origin_key(url) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url",
|
||||
[
|
||||
"",
|
||||
"/v1/chat",
|
||||
"ftp://api.example.com",
|
||||
"https://:443",
|
||||
"https://user@",
|
||||
"https://example.com:",
|
||||
"https://example.com:not-a-port",
|
||||
"https://example.com:65536",
|
||||
"https://[2001:db8::1",
|
||||
"https://exa mple.com",
|
||||
"https://user@example.com",
|
||||
"https://user:secret@example.com",
|
||||
"https://:secret@example.com",
|
||||
"https://@example.com",
|
||||
"https://exa\u200bmple.com",
|
||||
None,
|
||||
],
|
||||
)
|
||||
def test_upstream_origin_key_rejects_invalid_urls(url: object) -> None:
|
||||
with pytest.raises(ValueError, match="absolute HTTP") as exc_info:
|
||||
upstream_origin_key(url) # type: ignore[arg-type]
|
||||
assert "secret" not in str(exc_info.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acquire_maps_invalid_provider_url_to_502() -> None:
|
||||
with pytest.raises(UpstreamError) as exc_info:
|
||||
acquire_upstream_http_client("ftp://api.example.com")
|
||||
assert exc_info.value.status_code == 502
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acquire_maps_shutdown_to_503() -> None:
|
||||
with patch.object(
|
||||
http_client_module,
|
||||
"get_upstream_http_client",
|
||||
side_effect=RuntimeError("Upstream HTTP client is shutting down"),
|
||||
):
|
||||
with pytest.raises(UpstreamError) as exc_info:
|
||||
acquire_upstream_http_client("https://api.example.com")
|
||||
assert exc_info.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("first_url", "second_url"),
|
||||
[
|
||||
("https://EXAMPLE.com:443/v1/chat", "https://example.com/v1/models"),
|
||||
(
|
||||
"https://bücher.example/v1/chat",
|
||||
"https://xn--bcher-kva.example/v1/models",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_equivalent_origins_share_one_client(
|
||||
first_url: str, second_url: str
|
||||
) -> None:
|
||||
try:
|
||||
first = get_upstream_http_client(first_url)
|
||||
second = get_upstream_http_client(second_url)
|
||||
assert second is first
|
||||
finally:
|
||||
await close_upstream_http_client()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_http_client_applies_configured_pool_bounds() -> None:
|
||||
with (
|
||||
patch.object(
|
||||
http_client_module.httpx,
|
||||
"Limits",
|
||||
wraps=httpx.Limits,
|
||||
) as build_limits,
|
||||
patch.object(
|
||||
http_client_module.httpx,
|
||||
"AsyncHTTPTransport",
|
||||
wraps=httpx.AsyncHTTPTransport,
|
||||
) as build_transport,
|
||||
):
|
||||
client = get_upstream_http_client("https://api.example.com")
|
||||
|
||||
try:
|
||||
assert client.timeout.pool == settings.upstream_pool_timeout
|
||||
assert client.timeout.read == settings.upstream_read_timeout
|
||||
assert client.timeout.connect == settings.upstream_connect_timeout
|
||||
assert client.timeout.write == settings.upstream_write_timeout
|
||||
build_limits.assert_called_once_with(
|
||||
max_connections=settings.upstream_max_connections,
|
||||
max_keepalive_connections=settings.upstream_max_keepalive_connections,
|
||||
keepalive_expiry=settings.upstream_keepalive_expiry,
|
||||
)
|
||||
build_transport.assert_called_once()
|
||||
assert (
|
||||
build_transport.call_args.kwargs["retries"]
|
||||
== settings.upstream_connect_retries
|
||||
)
|
||||
finally:
|
||||
await close_upstream_http_client()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_http_client_does_not_share_cookies() -> None:
|
||||
client = get_upstream_http_client("https://example.com")
|
||||
try:
|
||||
first = client.build_request("GET", "https://example.com/test")
|
||||
response = httpx.Response(
|
||||
200,
|
||||
headers={"set-cookie": "sticky=upstream; Path=/"},
|
||||
request=first,
|
||||
)
|
||||
client.cookies.extract_cookies(response)
|
||||
|
||||
later = client.build_request("GET", "https://example.com/test")
|
||||
explicit = client.build_request(
|
||||
"GET", "https://example.com/test", headers={"cookie": "user=provided"}
|
||||
)
|
||||
|
||||
assert "cookie" not in later.headers
|
||||
assert explicit.headers["cookie"] == "user=provided"
|
||||
finally:
|
||||
await close_upstream_http_client()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_closes_foreign_client_on_its_owner_loop(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
foreign_loop = asyncio.new_event_loop()
|
||||
loop_ready = threading.Event()
|
||||
close_finished = threading.Event()
|
||||
close_loops: list[asyncio.AbstractEventLoop] = []
|
||||
|
||||
def run_foreign_loop() -> None:
|
||||
asyncio.set_event_loop(foreign_loop)
|
||||
loop_ready.set()
|
||||
foreign_loop.run_forever()
|
||||
|
||||
thread = threading.Thread(target=run_foreign_loop)
|
||||
thread.start()
|
||||
assert loop_ready.wait(timeout=10)
|
||||
|
||||
async def make_client() -> httpx.AsyncClient:
|
||||
client = get_upstream_http_client("https://example.com")
|
||||
original_close = client.aclose
|
||||
|
||||
async def tracked_close() -> None:
|
||||
close_loops.append(asyncio.get_running_loop())
|
||||
await original_close()
|
||||
close_finished.set()
|
||||
|
||||
monkeypatch.setattr(client, "aclose", tracked_close)
|
||||
return client
|
||||
|
||||
client_future = asyncio.run_coroutine_threadsafe(make_client(), foreign_loop)
|
||||
client = await asyncio.to_thread(client_future.result, 10)
|
||||
try:
|
||||
await close_upstream_http_client()
|
||||
assert await asyncio.to_thread(close_finished.wait, 10)
|
||||
assert client.is_closed
|
||||
assert close_loops == [foreign_loop]
|
||||
|
||||
await close_upstream_http_client()
|
||||
assert not http_client_module._pending_closes
|
||||
finally:
|
||||
foreign_loop.call_soon_threadsafe(foreign_loop.stop)
|
||||
await asyncio.to_thread(thread.join, 10)
|
||||
assert not thread.is_alive()
|
||||
foreign_loop.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_rehomes_queued_close_when_owner_loop_stops(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
foreign_loop = asyncio.new_event_loop()
|
||||
loop_ready = threading.Event()
|
||||
blocker_started = threading.Event()
|
||||
allow_stop = threading.Event()
|
||||
|
||||
def run_foreign_loop() -> None:
|
||||
asyncio.set_event_loop(foreign_loop)
|
||||
loop_ready.set()
|
||||
foreign_loop.run_forever()
|
||||
|
||||
thread = threading.Thread(target=run_foreign_loop)
|
||||
thread.start()
|
||||
assert loop_ready.wait(timeout=10)
|
||||
|
||||
async def make_client() -> httpx.AsyncClient:
|
||||
return get_upstream_http_client("https://example.com")
|
||||
|
||||
client_future = asyncio.run_coroutine_threadsafe(make_client(), foreign_loop)
|
||||
client = await asyncio.to_thread(client_future.result, 10)
|
||||
original_close = client.aclose
|
||||
close_loops: list[asyncio.AbstractEventLoop] = []
|
||||
|
||||
async def tracked_close() -> None:
|
||||
close_loops.append(asyncio.get_running_loop())
|
||||
await original_close()
|
||||
|
||||
monkeypatch.setattr(client, "aclose", tracked_close)
|
||||
|
||||
def stop_before_next_iteration() -> None:
|
||||
blocker_started.set()
|
||||
assert allow_stop.wait(timeout=10)
|
||||
foreign_loop.stop()
|
||||
|
||||
foreign_loop.call_soon_threadsafe(stop_before_next_iteration)
|
||||
assert blocker_started.wait(timeout=10)
|
||||
|
||||
try:
|
||||
closing = asyncio.create_task(close_upstream_http_client())
|
||||
while not http_client_module._pending_closes:
|
||||
await asyncio.sleep(0)
|
||||
assert not client.is_closed
|
||||
|
||||
allow_stop.set()
|
||||
await asyncio.to_thread(thread.join, 10)
|
||||
assert not thread.is_alive()
|
||||
|
||||
await closing
|
||||
assert client.is_closed
|
||||
assert close_loops == [asyncio.get_running_loop()]
|
||||
assert not http_client_module._pending_closes
|
||||
finally:
|
||||
allow_stop.set()
|
||||
if thread.is_alive():
|
||||
foreign_loop.call_soon_threadsafe(foreign_loop.stop)
|
||||
await asyncio.to_thread(thread.join, 10)
|
||||
foreign_loop.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_finishes_transport_close_on_stopped_owner_loop(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
foreign_loop = asyncio.new_event_loop()
|
||||
loop_ready = threading.Event()
|
||||
transport_started = threading.Event()
|
||||
allow_transport_close = threading.Event()
|
||||
transport_finished = threading.Event()
|
||||
|
||||
class BlockingTransport(httpx.AsyncBaseTransport):
|
||||
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, request=request)
|
||||
|
||||
async def aclose(self) -> None:
|
||||
transport_started.set()
|
||||
while not allow_transport_close.is_set():
|
||||
await asyncio.sleep(0)
|
||||
transport_finished.set()
|
||||
|
||||
client = httpx.AsyncClient(transport=BlockingTransport())
|
||||
monkeypatch.setattr(http_client_module, "_build_client", lambda: client)
|
||||
|
||||
def run_foreign_loop() -> None:
|
||||
asyncio.set_event_loop(foreign_loop)
|
||||
loop_ready.set()
|
||||
foreign_loop.run_forever()
|
||||
|
||||
thread = threading.Thread(target=run_foreign_loop)
|
||||
thread.start()
|
||||
assert loop_ready.wait(timeout=10)
|
||||
|
||||
async def register_client() -> None:
|
||||
assert get_upstream_http_client("https://example.com") is client
|
||||
|
||||
registered = asyncio.run_coroutine_threadsafe(register_client(), foreign_loop)
|
||||
await asyncio.to_thread(registered.result, 10)
|
||||
|
||||
try:
|
||||
closing = asyncio.create_task(close_upstream_http_client())
|
||||
assert await asyncio.to_thread(transport_started.wait, 10)
|
||||
|
||||
foreign_loop.call_soon_threadsafe(foreign_loop.stop)
|
||||
await asyncio.to_thread(thread.join, 10)
|
||||
assert not thread.is_alive()
|
||||
|
||||
allow_transport_close.set()
|
||||
await closing
|
||||
|
||||
assert transport_finished.is_set()
|
||||
assert client.is_closed
|
||||
assert not http_client_module._pending_closes
|
||||
finally:
|
||||
allow_transport_close.set()
|
||||
if thread.is_alive():
|
||||
foreign_loop.call_soon_threadsafe(foreign_loop.stop)
|
||||
await asyncio.to_thread(thread.join, 10)
|
||||
if not foreign_loop.is_closed():
|
||||
foreign_loop.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_started_close_on_closed_owner_loop_retries_transport() -> None:
|
||||
class CountingTransport(httpx.AsyncBaseTransport):
|
||||
def __init__(self) -> None:
|
||||
self.close_count = 0
|
||||
|
||||
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, request=request)
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.close_count += 1
|
||||
|
||||
owner_loop = MagicMock(spec=asyncio.AbstractEventLoop)
|
||||
owner_loop.is_closed.return_value = True
|
||||
owner_loop.is_running.return_value = False
|
||||
transport = CountingTransport()
|
||||
client = httpx.AsyncClient(transport=transport)
|
||||
completion: concurrent.futures.Future[None] = concurrent.futures.Future()
|
||||
completion.set_running_or_notify_cancel()
|
||||
task = MagicMock(spec=asyncio.Task)
|
||||
task.done.return_value = False
|
||||
submission = http_client_module._CloseSubmission(
|
||||
client=client,
|
||||
completion=completion,
|
||||
task=task,
|
||||
)
|
||||
http_client_module._pending_closes[owner_loop] = {completion: submission}
|
||||
|
||||
http_client_module._collect_completed_closes()
|
||||
await http_client_module._drain_pending_closes()
|
||||
|
||||
assert submission.retired
|
||||
assert transport.close_count == 1
|
||||
assert client.is_closed
|
||||
assert not http_client_module._pending_closes
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_closes_client_after_owner_loop_stopped(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
created: list[httpx.AsyncClient] = []
|
||||
owner_loops: list[asyncio.AbstractEventLoop] = []
|
||||
|
||||
def create_on_stopped_loop() -> None:
|
||||
owner_loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(owner_loop)
|
||||
owner_loops.append(owner_loop)
|
||||
|
||||
async def make_client() -> None:
|
||||
created.append(get_upstream_http_client("https://example.com"))
|
||||
|
||||
owner_loop.run_until_complete(make_client())
|
||||
|
||||
thread = threading.Thread(target=create_on_stopped_loop)
|
||||
thread.start()
|
||||
await asyncio.to_thread(thread.join, 10)
|
||||
assert not thread.is_alive()
|
||||
|
||||
client = created[0]
|
||||
owner_loop = owner_loops[0]
|
||||
close_loops: list[asyncio.AbstractEventLoop] = []
|
||||
original_close = client.aclose
|
||||
|
||||
async def tracked_close() -> None:
|
||||
close_loops.append(asyncio.get_running_loop())
|
||||
await original_close()
|
||||
|
||||
monkeypatch.setattr(client, "aclose", tracked_close)
|
||||
try:
|
||||
await close_upstream_http_client()
|
||||
assert client.is_closed
|
||||
assert close_loops == [asyncio.get_running_loop()]
|
||||
assert not http_client_module._pending_closes
|
||||
finally:
|
||||
owner_loop.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_closes_client_after_owner_loop_closed() -> None:
|
||||
created: list[httpx.AsyncClient] = []
|
||||
|
||||
def create_and_close_loop() -> None:
|
||||
owner_loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(owner_loop)
|
||||
|
||||
async def make_client() -> None:
|
||||
created.append(get_upstream_http_client("https://example.com"))
|
||||
|
||||
owner_loop.run_until_complete(make_client())
|
||||
owner_loop.close()
|
||||
|
||||
thread = threading.Thread(target=create_and_close_loop)
|
||||
thread.start()
|
||||
await asyncio.to_thread(thread.join, 10)
|
||||
assert not thread.is_alive()
|
||||
|
||||
client = created[0]
|
||||
await close_upstream_http_client()
|
||||
assert client.is_closed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_retries_failed_client_close(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
client = get_upstream_http_client("https://example.com")
|
||||
original_close = client.aclose
|
||||
attempts = 0
|
||||
|
||||
async def flaky_close() -> None:
|
||||
nonlocal attempts
|
||||
attempts += 1
|
||||
if attempts == 1:
|
||||
raise RuntimeError("close failed")
|
||||
await original_close()
|
||||
|
||||
monkeypatch.setattr(client, "aclose", flaky_close)
|
||||
|
||||
await close_upstream_http_client()
|
||||
assert not client.is_closed
|
||||
assert any(
|
||||
client in failed for failed in http_client_module._failed_closes.values()
|
||||
)
|
||||
|
||||
await close_upstream_http_client()
|
||||
assert client.is_closed
|
||||
assert attempts == 2
|
||||
assert not http_client_module._failed_closes
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_prunes_externally_closed_failed_client(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
client = get_upstream_http_client("https://example.com")
|
||||
original_close = client.aclose
|
||||
|
||||
async def fail_close() -> None:
|
||||
raise RuntimeError("close failed")
|
||||
|
||||
monkeypatch.setattr(client, "aclose", fail_close)
|
||||
await close_upstream_http_client()
|
||||
assert http_client_module._failed_closes
|
||||
|
||||
await original_close()
|
||||
await close_upstream_http_client()
|
||||
|
||||
assert client.is_closed
|
||||
assert not http_client_module._failed_closes
|
||||
assert not http_client_module._pending_closes
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure", [RuntimeError("failed"), asyncio.CancelledError()])
|
||||
async def test_shutdown_retries_transport_after_httpx_marks_client_closed(
|
||||
failure: BaseException,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
class FailOnceTransport(httpx.AsyncBaseTransport):
|
||||
def __init__(self) -> None:
|
||||
self.attempts = 0
|
||||
self.completed = False
|
||||
|
||||
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, request=request)
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.attempts += 1
|
||||
if self.attempts == 1:
|
||||
raise failure
|
||||
self.completed = True
|
||||
|
||||
transport = FailOnceTransport()
|
||||
client = httpx.AsyncClient(transport=transport)
|
||||
monkeypatch.setattr(http_client_module, "_build_client", lambda: client)
|
||||
assert get_upstream_http_client("https://example.com") is client
|
||||
|
||||
await close_upstream_http_client()
|
||||
assert client.is_closed
|
||||
assert transport.attempts == 1
|
||||
assert not transport.completed
|
||||
assert any(
|
||||
client in failed for failed in http_client_module._failed_closes.values()
|
||||
)
|
||||
|
||||
await close_upstream_http_client()
|
||||
assert transport.attempts == 2
|
||||
assert transport.completed
|
||||
assert not http_client_module._failed_closes
|
||||
assert not http_client_module._pending_closes
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_collects_done_task_before_owner_loop_callback(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
foreign_loop = asyncio.new_event_loop()
|
||||
loop_ready = threading.Event()
|
||||
transport_finished = threading.Event()
|
||||
|
||||
class StopAfterCloseTransport(httpx.AsyncBaseTransport):
|
||||
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, request=request)
|
||||
|
||||
async def aclose(self) -> None:
|
||||
transport_finished.set()
|
||||
asyncio.get_running_loop().stop()
|
||||
|
||||
client = httpx.AsyncClient(transport=StopAfterCloseTransport())
|
||||
monkeypatch.setattr(http_client_module, "_build_client", lambda: client)
|
||||
|
||||
def run_foreign_loop() -> None:
|
||||
asyncio.set_event_loop(foreign_loop)
|
||||
loop_ready.set()
|
||||
foreign_loop.run_forever()
|
||||
|
||||
thread = threading.Thread(target=run_foreign_loop)
|
||||
thread.start()
|
||||
assert loop_ready.wait(timeout=10)
|
||||
|
||||
async def register_client() -> None:
|
||||
assert get_upstream_http_client("https://example.com") is client
|
||||
|
||||
registered_client = asyncio.run_coroutine_threadsafe(
|
||||
register_client(), foreign_loop
|
||||
)
|
||||
await asyncio.to_thread(registered_client.result, 10)
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(close_upstream_http_client(), timeout=1)
|
||||
assert transport_finished.is_set()
|
||||
assert client.is_closed
|
||||
assert not http_client_module._pending_closes
|
||||
assert not http_client_module._failed_closes
|
||||
finally:
|
||||
if thread.is_alive():
|
||||
foreign_loop.call_soon_threadsafe(foreign_loop.stop)
|
||||
await asyncio.to_thread(thread.join, 10)
|
||||
foreign_loop.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_close_submission_settlement_is_atomic_across_threads(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
task = asyncio.create_task(asyncio.sleep(0))
|
||||
await task
|
||||
|
||||
client = httpx.AsyncClient()
|
||||
completion: concurrent.futures.Future[None] = concurrent.futures.Future()
|
||||
completion.set_running_or_notify_cancel()
|
||||
submission = http_client_module._CloseSubmission(
|
||||
client=client,
|
||||
completion=completion,
|
||||
task=task,
|
||||
)
|
||||
loop = asyncio.get_running_loop()
|
||||
http_client_module._pending_closes[loop] = {completion: submission}
|
||||
|
||||
barrier = threading.Barrier(2)
|
||||
errors: list[BaseException] = []
|
||||
original_settle = http_client_module._settle_close_submission
|
||||
|
||||
def synchronized_settle(
|
||||
close_submission: http_client_module._CloseSubmission,
|
||||
completed: asyncio.Task[None],
|
||||
) -> None:
|
||||
barrier.wait(timeout=10)
|
||||
original_settle(close_submission, completed)
|
||||
|
||||
def run(action: Callable[[], None]) -> None:
|
||||
try:
|
||||
action()
|
||||
except BaseException as exc:
|
||||
errors.append(exc)
|
||||
|
||||
monkeypatch.setattr(
|
||||
http_client_module,
|
||||
"_settle_close_submission",
|
||||
synchronized_settle,
|
||||
)
|
||||
collector = threading.Thread(
|
||||
target=run,
|
||||
args=(lambda: http_client_module._settle_submission_from_task(submission),),
|
||||
)
|
||||
callback = threading.Thread(
|
||||
target=run,
|
||||
args=(lambda: http_client_module._finish_close_submission(submission, task),),
|
||||
)
|
||||
|
||||
try:
|
||||
collector.start()
|
||||
callback.start()
|
||||
collector.join(timeout=10)
|
||||
callback.join(timeout=10)
|
||||
|
||||
assert not collector.is_alive()
|
||||
assert not callback.is_alive()
|
||||
assert errors == []
|
||||
assert completion.result() is None
|
||||
|
||||
monkeypatch.setattr(
|
||||
http_client_module,
|
||||
"_settle_close_submission",
|
||||
original_settle,
|
||||
)
|
||||
http_client_module._collect_completed_closes()
|
||||
|
||||
assert http_client_module._close_completed.get(client) is True
|
||||
assert not http_client_module._pending_closes
|
||||
assert not http_client_module._failed_closes
|
||||
finally:
|
||||
http_client_module._pending_closes.pop(loop, None)
|
||||
http_client_module._close_completed.pop(client, None)
|
||||
await client.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_close_submission_rejects_conflicting_outcomes() -> None:
|
||||
succeeded = asyncio.create_task(asyncio.sleep(0))
|
||||
|
||||
async def fail() -> None:
|
||||
raise RuntimeError("different outcome")
|
||||
|
||||
failed = asyncio.create_task(fail())
|
||||
await succeeded
|
||||
with pytest.raises(RuntimeError, match="different outcome"):
|
||||
await failed
|
||||
|
||||
client = httpx.AsyncClient()
|
||||
completion: concurrent.futures.Future[None] = concurrent.futures.Future()
|
||||
completion.set_running_or_notify_cancel()
|
||||
submission = http_client_module._CloseSubmission(client, completion)
|
||||
|
||||
try:
|
||||
http_client_module._settle_close_submission(submission, succeeded)
|
||||
with pytest.raises(RuntimeError, match="conflicting outcomes"):
|
||||
http_client_module._settle_close_submission(submission, failed)
|
||||
finally:
|
||||
await client.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_http_client_cannot_reopen_during_shutdown(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
client = get_upstream_http_client("https://example.com")
|
||||
close_started = asyncio.Event()
|
||||
allow_close = asyncio.Event()
|
||||
original_close = client.aclose
|
||||
|
||||
async def delayed_close() -> None:
|
||||
close_started.set()
|
||||
await allow_close.wait()
|
||||
await original_close()
|
||||
|
||||
monkeypatch.setattr(client, "aclose", delayed_close)
|
||||
closing = asyncio.create_task(close_upstream_http_client())
|
||||
await close_started.wait()
|
||||
|
||||
with pytest.raises(RuntimeError, match="shutting down"):
|
||||
get_upstream_http_client("https://example.com")
|
||||
|
||||
allow_close.set()
|
||||
await closing
|
||||
@@ -394,10 +394,9 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None:
|
||||
),
|
||||
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=1_000)),
|
||||
patch.object(
|
||||
proxy_module,
|
||||
"get_reservation_snapshot",
|
||||
"pay_for_request",
|
||||
AsyncMock(return_value=reservation),
|
||||
),
|
||||
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
|
||||
|
||||
@@ -0,0 +1,329 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import Request
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
from starlette.types import Message, Send
|
||||
|
||||
from routstr.upstream.base import BaseUpstreamProvider, _OwnedUpstreamStream
|
||||
|
||||
|
||||
class _CountingStream(httpx.AsyncByteStream):
|
||||
def __init__(self, payload: bytes) -> None:
|
||||
self.payload = payload
|
||||
self.close_count = 0
|
||||
|
||||
async def __aiter__(self) -> AsyncIterator[bytes]:
|
||||
yield self.payload
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.close_count += 1
|
||||
|
||||
|
||||
class _CountingTransport(httpx.AsyncBaseTransport):
|
||||
def __init__(self, payload: bytes) -> None:
|
||||
self.stream = _CountingStream(payload)
|
||||
self.close_count = 0
|
||||
|
||||
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, request=request, stream=self.stream)
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.close_count += 1
|
||||
|
||||
|
||||
class _CountingClient(httpx.AsyncClient):
|
||||
def __init__(self, transport: _CountingTransport) -> None:
|
||||
super().__init__(transport=transport)
|
||||
self.close_count = 0
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.close_count += 1
|
||||
await super().aclose()
|
||||
|
||||
|
||||
def _request() -> Request:
|
||||
sent = False
|
||||
|
||||
async def receive() -> dict[str, object]:
|
||||
nonlocal sent
|
||||
if sent:
|
||||
return {"type": "http.disconnect"}
|
||||
sent = True
|
||||
return {"type": "http.request", "body": b"{}", "more_body": False}
|
||||
|
||||
return Request(
|
||||
{
|
||||
"type": "http",
|
||||
"asgi": {"version": "3.0", "spec_version": "2.4"},
|
||||
"method": "POST",
|
||||
"scheme": "http",
|
||||
"path": "/v1/audio/speech",
|
||||
"raw_path": b"/v1/audio/speech",
|
||||
"query_string": b"",
|
||||
"headers": [],
|
||||
"client": ("test", 1),
|
||||
"server": ("test", 80),
|
||||
},
|
||||
receive,
|
||||
)
|
||||
|
||||
|
||||
async def _forward(
|
||||
provider: BaseUpstreamProvider,
|
||||
method_name: str,
|
||||
) -> tuple[StreamingResponse, _CountingClient, _CountingTransport]:
|
||||
transport = _CountingTransport(b"live-stream")
|
||||
client = _CountingClient(transport)
|
||||
model = MagicMock()
|
||||
|
||||
with patch("routstr.upstream.base.httpx.AsyncClient", return_value=client):
|
||||
result = await getattr(provider, method_name)(
|
||||
request=_request(),
|
||||
path="v1/audio/speech",
|
||||
headers={},
|
||||
amount=10,
|
||||
unit="sat",
|
||||
max_cost_for_model=10_000,
|
||||
model_obj=model,
|
||||
)
|
||||
|
||||
assert isinstance(result, StreamingResponse)
|
||||
return result, client, transport
|
||||
|
||||
|
||||
async def _run_asgi_response(
|
||||
response: StreamingResponse,
|
||||
send: Send,
|
||||
) -> None:
|
||||
async def receive() -> dict[str, str]:
|
||||
return {"type": "http.disconnect"}
|
||||
|
||||
await response(
|
||||
{
|
||||
"type": "http",
|
||||
"asgi": {"version": "3.0", "spec_version": "2.4"},
|
||||
},
|
||||
receive,
|
||||
send,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"method_name",
|
||||
["forward_x_cashu_request", "forward_x_cashu_responses_request"],
|
||||
)
|
||||
async def test_x_cashu_opaque_stream_owns_client_until_normal_completion(
|
||||
method_name: str,
|
||||
) -> None:
|
||||
provider = BaseUpstreamProvider(base_url="http://upstream", api_key="test")
|
||||
response, client, transport = await _forward(provider, method_name)
|
||||
messages: list[dict[str, Any]] = []
|
||||
|
||||
assert client.close_count == 0
|
||||
assert transport.stream.close_count == 0
|
||||
|
||||
async def send(message: Message) -> None:
|
||||
messages.append(dict(message))
|
||||
|
||||
await _run_asgi_response(response, send)
|
||||
|
||||
assert (
|
||||
b"".join(
|
||||
message.get("body", b"")
|
||||
for message in messages
|
||||
if message["type"] == "http.response.body"
|
||||
)
|
||||
== b"live-stream"
|
||||
)
|
||||
assert transport.stream.close_count == 1
|
||||
assert client.close_count == 1
|
||||
assert transport.close_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"method_name",
|
||||
["forward_x_cashu_request", "forward_x_cashu_responses_request"],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"failure",
|
||||
[RuntimeError("downstream send failed"), asyncio.CancelledError()],
|
||||
)
|
||||
async def test_x_cashu_opaque_stream_closes_client_when_send_fails(
|
||||
method_name: str,
|
||||
failure: BaseException,
|
||||
) -> None:
|
||||
provider = BaseUpstreamProvider(base_url="http://upstream", api_key="test")
|
||||
response, client, transport = await _forward(provider, method_name)
|
||||
|
||||
async def send(message: Message) -> None:
|
||||
if message["type"] == "http.response.body" and message.get("body"):
|
||||
raise failure
|
||||
|
||||
with pytest.raises(type(failure)):
|
||||
await _run_asgi_response(response, send)
|
||||
|
||||
assert transport.stream.close_count == 1
|
||||
assert client.close_count == 1
|
||||
assert transport.close_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("method_name", "path", "payload"),
|
||||
[
|
||||
(
|
||||
"forward_x_cashu_request",
|
||||
"v1/chat/completions",
|
||||
b'data: {"model":"m","usage":{"prompt_tokens":1,"completion_tokens":1}}\n\ndata: [DONE]\n\n',
|
||||
),
|
||||
(
|
||||
"forward_x_cashu_responses_request",
|
||||
"v1/responses",
|
||||
b'data: {"type":"response.completed","response":{"model":"m","usage":{"input_tokens":1,"output_tokens":1}}}\n\ndata: [DONE]\n\n',
|
||||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"failure", [None, RuntimeError("send failed"), asyncio.CancelledError()]
|
||||
)
|
||||
async def test_x_cashu_real_processed_stream_releases_buffered_upstream_promptly(
|
||||
method_name: str,
|
||||
path: str,
|
||||
payload: bytes,
|
||||
failure: BaseException | None,
|
||||
) -> None:
|
||||
provider = BaseUpstreamProvider(base_url="http://upstream", api_key="test")
|
||||
transport = _CountingTransport(payload)
|
||||
client = _CountingClient(transport)
|
||||
|
||||
with (
|
||||
patch("routstr.upstream.base.httpx.AsyncClient", return_value=client),
|
||||
patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=None)),
|
||||
):
|
||||
result = await getattr(provider, method_name)(
|
||||
request=_request(),
|
||||
path=path,
|
||||
headers={},
|
||||
amount=10,
|
||||
unit="sat",
|
||||
max_cost_for_model=10_000,
|
||||
model_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert isinstance(result, StreamingResponse)
|
||||
assert transport.stream.close_count == 1
|
||||
assert client.close_count == 1
|
||||
assert transport.close_count == 1
|
||||
|
||||
messages: list[Message] = []
|
||||
|
||||
async def send(message: Message) -> None:
|
||||
messages.append(message)
|
||||
if failure is not None and message["type"] == "http.response.body":
|
||||
if message.get("body"):
|
||||
raise failure
|
||||
|
||||
if failure is None:
|
||||
await _run_asgi_response(result, send)
|
||||
assert any(message.get("body") for message in messages)
|
||||
else:
|
||||
with pytest.raises(type(failure)):
|
||||
await _run_asgi_response(result, send)
|
||||
|
||||
assert transport.stream.close_count == 1
|
||||
assert client.close_count == 1
|
||||
assert transport.close_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_owned_upstream_cleanup_survives_caller_cancellation() -> None:
|
||||
cleanup_started = asyncio.Event()
|
||||
allow_cleanup = asyncio.Event()
|
||||
cleanup_finished = asyncio.Event()
|
||||
client_close_count = 0
|
||||
|
||||
async def body() -> AsyncIterator[bytes]:
|
||||
yield b"body"
|
||||
|
||||
response = MagicMock(spec=httpx.Response)
|
||||
response.aclose = AsyncMock()
|
||||
client = MagicMock(spec=httpx.AsyncClient)
|
||||
|
||||
async def close_client() -> None:
|
||||
nonlocal client_close_count
|
||||
client_close_count += 1
|
||||
cleanup_started.set()
|
||||
await allow_cleanup.wait()
|
||||
cleanup_finished.set()
|
||||
|
||||
client.aclose = close_client
|
||||
owned = _OwnedUpstreamStream(body(), response, client)
|
||||
|
||||
first_close = asyncio.create_task(owned.aclose())
|
||||
await cleanup_started.wait()
|
||||
first_close.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await first_close
|
||||
|
||||
allow_cleanup.set()
|
||||
await asyncio.wait_for(cleanup_finished.wait(), timeout=1)
|
||||
await owned.aclose()
|
||||
|
||||
assert response.aclose.await_count == 1
|
||||
assert client_close_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("method_name", "path", "handler_name"),
|
||||
[
|
||||
(
|
||||
"forward_x_cashu_request",
|
||||
"v1/chat/completions",
|
||||
"handle_x_cashu_chat_completion",
|
||||
),
|
||||
(
|
||||
"forward_x_cashu_responses_request",
|
||||
"v1/responses",
|
||||
"handle_x_cashu_responses_completion",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_x_cashu_non_streaming_result_closes_upstream_promptly(
|
||||
method_name: str,
|
||||
path: str,
|
||||
handler_name: str,
|
||||
) -> None:
|
||||
provider = BaseUpstreamProvider(base_url="http://upstream", api_key="test")
|
||||
transport = _CountingTransport(b"{}")
|
||||
client = _CountingClient(transport)
|
||||
|
||||
with (
|
||||
patch("routstr.upstream.base.httpx.AsyncClient", return_value=client),
|
||||
patch.object(
|
||||
provider,
|
||||
handler_name,
|
||||
new=AsyncMock(return_value=Response(b"done")),
|
||||
),
|
||||
):
|
||||
result = await getattr(provider, method_name)(
|
||||
request=_request(),
|
||||
path=path,
|
||||
headers={},
|
||||
amount=10,
|
||||
unit="sat",
|
||||
max_cost_for_model=10_000,
|
||||
model_obj=MagicMock(),
|
||||
)
|
||||
|
||||
assert not isinstance(result, StreamingResponse)
|
||||
assert transport.stream.close_count == 1
|
||||
assert client.close_count == 1
|
||||
assert transport.close_count == 1
|
||||
Reference in New Issue
Block a user