From 10b105bc7bffc07fd7555a9dc92fd6851c3af31d Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 26 Sep 2026 02:34:10 +0200 Subject: [PATCH] perf: reduce request latency --- .env.example | 11 + routstr/auth.py | 94 +- routstr/core/logging.py | 209 ++- routstr/core/main.py | 9 + routstr/core/settings.py | 33 + routstr/proxy.py | 25 +- routstr/upstream/base.py | 1512 ++++++++++------- routstr/upstream/ehbp.py | 69 +- routstr/upstream/gemini_messages.py | 89 +- routstr/upstream/http_client.py | 481 ++++++ .../integration/test_reservation_lifecycle.py | 10 +- tests/unit/test_log_secret_redaction.py | 1 - tests/unit/test_model_path_routing.py | 5 +- tests/unit/test_payment_settlement_timing.py | 55 + .../unit/test_pre_handoff_stream_ownership.py | 168 ++ tests/unit/test_queued_logging.py | 291 ++++ tests/unit/test_settings.py | 17 +- tests/unit/test_stale_reservations.py | 66 +- tests/unit/test_stream_id_injection.py | 3 - .../test_streaming_billing_finalization.py | 398 ++++- tests/unit/test_streaming_sse_providers.py | 5 +- tests/unit/test_tinfoil_integration.py | 3 +- tests/unit/test_upstream_gemini.py | 144 +- tests/unit/test_upstream_http_client.py | 735 ++++++++ tests/unit/test_upstream_rate_limit.py | 3 +- tests/unit/test_x_cashu_stream_ownership.py | 329 ++++ 26 files changed, 4048 insertions(+), 717 deletions(-) create mode 100644 routstr/upstream/http_client.py create mode 100644 tests/unit/test_payment_settlement_timing.py create mode 100644 tests/unit/test_pre_handoff_stream_ownership.py create mode 100644 tests/unit/test_queued_logging.py create mode 100644 tests/unit/test_upstream_http_client.py create mode 100644 tests/unit/test_x_cashu_stream_ownership.py diff --git a/.env.example b/.env.example index d540fce2..ab850738 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/routstr/auth.py b/routstr/auth.py index bf29e86e..9bfc7b3d 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -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. diff --git a/routstr/core/logging.py b/routstr/core/logging.py index 0fba407b..04bbb7f0 100644 --- a/routstr/core/logging.py +++ b/routstr/core/logging.py @@ -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", diff --git a/routstr/core/main.py b/routstr/core/main.py index bf1f4eea..584fa736 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -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): diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 77d623e7..77a22963 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -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", } ) diff --git a/routstr/proxy.py b/routstr/proxy.py index 1b9a3947..26bf20e7 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -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, diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 5679aae3..122a8ddc 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -11,9 +11,10 @@ from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, from typing import Any, Mapping, Self, cast import httpx -from fastapi import BackgroundTasks, HTTPException, Request +from fastapi import HTTPException, Request from fastapi.responses import Response, StreamingResponse from pydantic.v1 import BaseModel +from starlette.types import Receive, Scope, Send from ..auth import ( ReservationSnapshot, @@ -41,6 +42,7 @@ from ..core.error_scope import ( ) from ..core.exceptions import UpstreamError from ..core.redaction import redact_org_ids +from ..core.settings import settings from ..payment.cost_calculation import ( CostData, CostDataError, @@ -70,6 +72,7 @@ from .cache_breakpoints import ( is_explicit_cache_model, ) from .count_tokens import MissingUsageEstimator, count_tokens_locally +from .http_client import acquire_upstream_http_client from .litellm_routing import detect_litellm_prefix from .model_paths import public_provider_url from .rate_limit import UPSTREAM_RATE_LIMIT, classify_rate_limit @@ -92,19 +95,200 @@ async def _aclose_if_needed(resource: object | None) -> None: await result +async def _shielded_aclose(resource: object | None) -> None: + await asyncio.shield(_aclose_if_needed(resource)) + + +class _ResponseHandoff: + """Close a response unless ownership is transferred to a stream.""" + + def __init__(self) -> None: + self._response: object | None = None + + def acquire(self, response: object) -> None: + self._response = response + + def handoff(self) -> None: + self._response = None + + async def close(self, *, suppress_errors: bool = False) -> None: + response = self._response + self._response = None + if response is None: + return + try: + await _shielded_aclose(response) + except BaseException: + if not suppress_errors: + raise + logger.exception("Failed to close upstream response before handoff") + + async def _finalize_and_close_stream( finalize: Callable[[], Awaitable[None]] | None, response: object | None, - client: httpx.AsyncClient | None, ) -> None: + """Settle billing, then return the response connection to its pool.""" try: if finalize is not None: await finalize() finally: + await _aclose_if_needed(response) + + +class _PersistentStreamFinalizer: + """Run one stream finalizer to completion across cancellation boundaries.""" + + def __init__(self, finalize: Callable[[], Awaitable[None]]) -> None: + self._finalize = finalize + self._task: asyncio.Future[None] | None = None + self._lock = asyncio.Lock() + + async def run(self) -> None: + async with self._lock: + if self._task is None: + self._task = asyncio.ensure_future(self._finalize()) + task = self._task + await asyncio.shield(task) + + +class _FinalizingAsyncIterator: + """Tie iterator shutdown to a finalizer created before streaming starts.""" + + def __init__( + self, + iterator: AsyncIterator[bytes], + finalizer: _PersistentStreamFinalizer, + ) -> None: + self._iterator = iterator + self._finalizer = finalizer + + def __aiter__(self) -> Self: + return self + + async def __anext__(self) -> bytes: try: - await _aclose_if_needed(response) + return await self._iterator.__anext__() + except BaseException: + await self._finalizer.run() + raise + + async def aclose(self) -> None: + try: + await _aclose_if_needed(self._iterator) finally: - await _aclose_if_needed(client) + await self._finalizer.run() + + +class _ClosingStreamingResponse(StreamingResponse): + """Close the body iterator even when downstream ASGI sends fail.""" + + def __init__( + self, + content: AsyncIterator[bytes], + *, + finalizer: _PersistentStreamFinalizer | None = None, + **kwargs: Any, + ) -> None: + if finalizer is not None: + content = _FinalizingAsyncIterator(content, finalizer) + super().__init__(content, **kwargs) + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + try: + await super().__call__(scope, receive, send) + finally: + await asyncio.shield(_aclose_if_needed(self.body_iterator)) + + +class _OwnedUpstreamStream: + """Keep a one-shot HTTP client alive for the lifetime of its response.""" + + def __init__( + self, + iterator: AsyncIterator[bytes], + response: httpx.Response, + client: httpx.AsyncClient, + ) -> None: + self._iterator = iterator + self._response = response + self._client = client + self._cleanup_complete = False + self._cleanup_task: asyncio.Task[None] | None = None + self._close_lock = asyncio.Lock() + + def __aiter__(self) -> Self: + return self + + async def __anext__(self) -> bytes: + try: + return await self._iterator.__anext__() + except StopAsyncIteration: + await self.aclose() + raise + + async def _cleanup(self) -> None: + try: + await _aclose_if_needed(self._iterator) + finally: + try: + await self._response.aclose() + finally: + await self._client.aclose() + self._cleanup_complete = True + + async def aclose(self) -> None: + async with self._close_lock: + if self._cleanup_complete: + return + if self._cleanup_task is None or self._cleanup_task.done(): + self._cleanup_task = asyncio.create_task(self._cleanup()) + cleanup_task = self._cleanup_task + await asyncio.shield(cleanup_task) + + +def _attach_upstream_stream_owner( + result: StreamingResponse, + response: httpx.Response, + client: httpx.AsyncClient, +) -> StreamingResponse: + result.body_iterator = _OwnedUpstreamStream( + cast(AsyncIterator[bytes], result.body_iterator), response, client + ) + return result + + +async def _close_upstream_exchange( + response: httpx.Response | None, client: httpx.AsyncClient +) -> None: + try: + if response is not None: + await response.aclose() + finally: + await client.aclose() + + +def _build_x_cashu_client() -> httpx.AsyncClient: + """Build a per-request client for x-cashu forwarding. + + This path intentionally bypasses the shared per-origin pools from + ``http_client.py``: the response and client are handed off to + ``_OwnedUpstreamStream``/``_close_upstream_exchange``, which close the + client once the exchange finishes. Closing a pooled client would tear + down the shared pool for every caller, so ownership stays per-request + here at the cost of a fresh connection per call. + """ + return httpx.AsyncClient( + transport=httpx.AsyncHTTPTransport( + 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, + ), + ) CostMetadata = CostData | MaxCostData | dict[str, Any] @@ -1146,11 +1330,9 @@ class BaseUpstreamProvider: response: httpx.Response, key: ApiKey, max_cost_for_model: int, - background_tasks: BackgroundTasks, requested_model: str | None = None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, - client: httpx.AsyncClient | None = None, request_body: bytes | None = None, legacy_completion: bool = False, ) -> StreamingResponse: @@ -1184,51 +1366,60 @@ class BaseUpstreamProvider: }, ) + usage_finalized = False + last_model_seen: str | None = None + + async def finalize_db_only() -> None: + nonlocal usage_finalized + if usage_finalized: + return + try: + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + return + try: + await adjust_payment_for_tokens( + fresh_key, + usage_estimator.response_data(last_model_seen), + new_session, + max_cost_for_model, + model_obj, + self.provider_fee, + reservation_snapshot, + ) + usage_finalized = True + except Exception: + logger.exception( + "Fallback stream billing finalization failed; releasing reservation", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, new_session, reservation_snapshot + ) + ) + except Exception: + logger.exception( + "Fallback stream billing recovery could not access the database", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + + stream_finalizer = _PersistentStreamFinalizer( + lambda: _finalize_and_close_stream( + None if usage_finalized else finalize_db_only, + response, + ) + ) + async def stream_with_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: - usage_finalized: bool = False - last_model_seen: str | None = None + nonlocal usage_finalized, last_model_seen usage_chunk_data: dict | None = None done_seen: bool = False stream_id: str | None = None - async def finalize_db_only() -> None: - nonlocal usage_finalized - if usage_finalized: - return - try: - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - return - try: - await adjust_payment_for_tokens( - fresh_key, - usage_estimator.response_data(last_model_seen), - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - except Exception: - logger.exception( - "Fallback stream billing finalization failed; releasing reservation", - extra={"key_hash": key.hashed_key[:8] + "..."}, - ) - usage_finalized = ( - await self._release_failed_streaming_reservation( - fresh_key, new_session, reservation_snapshot - ) - ) - except Exception: - logger.exception( - "Fallback stream billing recovery could not access the database", - extra={"key_hash": key.hashed_key[:8] + "..."}, - ) - def _process_event( raw_event: bytes, final: bool = False ) -> Iterator[bytes]: @@ -1476,23 +1667,16 @@ class BaseUpstreamProvider: ) raise finally: - # Shielded so a client disconnect cannot cancel billing - # finalization or leak the upstream connection. - await asyncio.shield( - _finalize_and_close_stream( - None if usage_finalized else finalize_db_only, - response, - client, - ) - ) + await stream_finalizer.run() # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) response_headers.pop("content-encoding", None) response_headers.pop("content-length", None) - return StreamingResponse( + return _ClosingStreamingResponse( stream_with_cost(max_cost_for_model), + finalizer=stream_finalizer, status_code=response.status_code, headers=response_headers, ) @@ -1655,7 +1839,6 @@ class BaseUpstreamProvider: requested_model: str | None = None, model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, - client: httpx.AsyncClient | None = None, request_body: bytes | None = None, ) -> StreamingResponse: """Handle streaming Responses API responses with token usage tracking and cost adjustment. @@ -1679,51 +1862,60 @@ class BaseUpstreamProvider: }, ) + usage_finalized = False + last_model_seen: str | None = None + + async def finalize_db_only() -> None: + nonlocal usage_finalized + if usage_finalized: + return + try: + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + return + try: + await adjust_payment_for_tokens( + fresh_key, + usage_estimator.response_data(last_model_seen), + new_session, + max_cost_for_model, + model_obj, + self.provider_fee, + reservation_snapshot, + ) + usage_finalized = True + except Exception: + logger.exception( + "Fallback Responses billing finalization failed; releasing reservation", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + usage_finalized = ( + await self._release_failed_streaming_reservation( + fresh_key, new_session, reservation_snapshot + ) + ) + except Exception: + logger.exception( + "Fallback Responses billing recovery could not access the database", + extra={"key_hash": key.hashed_key[:8] + "..."}, + ) + + stream_finalizer = _PersistentStreamFinalizer( + lambda: _finalize_and_close_stream( + None if usage_finalized else finalize_db_only, + response, + ) + ) + async def stream_with_responses_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: - usage_finalized: bool = False - last_model_seen: str | None = None + nonlocal usage_finalized, last_model_seen reasoning_tokens: int = 0 usage_chunk_data: dict | None = None done_seen: bool = False - async def finalize_db_only() -> None: - nonlocal usage_finalized - if usage_finalized: - return - try: - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - return - try: - await adjust_payment_for_tokens( - fresh_key, - usage_estimator.response_data(last_model_seen), - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - except Exception: - logger.exception( - "Fallback Responses billing finalization failed; releasing reservation", - extra={"key_hash": key.hashed_key[:8] + "..."}, - ) - usage_finalized = ( - await self._release_failed_streaming_reservation( - fresh_key, new_session, reservation_snapshot - ) - ) - except Exception: - logger.exception( - "Fallback Responses billing recovery could not access the database", - extra={"key_hash": key.hashed_key[:8] + "..."}, - ) - def _process_event( raw_event: bytes, final: bool = False ) -> Iterator[bytes]: @@ -1927,23 +2119,16 @@ class BaseUpstreamProvider: ) raise finally: - # Shielded so a client disconnect cannot cancel billing - # finalization or leak the upstream connection. - await asyncio.shield( - _finalize_and_close_stream( - None if usage_finalized else finalize_db_only, - response, - client, - ) - ) + await stream_finalizer.run() # Remove inaccurate encoding headers from upstream response response_headers = dict(response.headers) response_headers.pop("content-encoding", None) response_headers.pop("content-length", None) - return StreamingResponse( + return _ClosingStreamingResponse( stream_with_responses_cost(max_cost_for_model), + finalizer=stream_finalizer, status_code=response.status_code, headers=response_headers, ) @@ -2107,12 +2292,12 @@ class BaseUpstreamProvider: provider_fee: float | None, reservation_snapshot: ReservationSnapshot, ) -> None: - """Background task to finalize payment for generic streaming requests.""" + """Finalize payment for a generic streaming request.""" async with create_session() as session: key = await session.get(ApiKey, key_hash) if not key: logger.warning( - "Key not found during background payment finalization", + "Key not found during generic streaming payment finalization", extra={"key_hash": key_hash[:8] + "..."}, ) return @@ -2131,7 +2316,7 @@ class BaseUpstreamProvider: reservation_snapshot=reservation_snapshot, ) logger.debug( - "Finalized generic streaming payment in background", + "Finalized generic streaming payment", extra={ "path": path, "key_hash": key_hash[:8] + "...", @@ -2139,7 +2324,7 @@ class BaseUpstreamProvider: ) except Exception as e: logger.error( - "Error finalizing generic streaming payment in background", + "Error finalizing generic streaming payment", extra={ "error": str(e), "key_hash": key_hash[:8] + "...", @@ -2147,6 +2332,78 @@ class BaseUpstreamProvider: }, ) + async def _stream_generic_with_settlement( + self, + response: httpx.Response, + key_hash: str, + max_cost: int, + path: str, + model_obj: Model | None, + provider_fee: float | None, + reservation_snapshot: ReservationSnapshot, + finalizer: _PersistentStreamFinalizer | None = None, + ) -> AsyncGenerator[bytes, None]: + """Relay an opaque stream and settle it even if the caller disconnects.""" + if finalizer is None: + finalizer = _PersistentStreamFinalizer( + lambda: _finalize_and_close_stream( + lambda: self._finalize_generic_streaming_payment( + key_hash, + max_cost, + path, + model_obj, + provider_fee, + reservation_snapshot, + ), + response, + ) + ) + try: + async for chunk in response.aiter_bytes(): + yield chunk + finally: + await finalizer.run() + + def _generic_streaming_response( + self, + response: httpx.Response, + key_hash: str, + max_cost: int, + path: str, + model_obj: Model | None, + provider_fee: float | None, + reservation_snapshot: ReservationSnapshot, + ) -> _ClosingStreamingResponse: + finalizer = _PersistentStreamFinalizer( + lambda: _finalize_and_close_stream( + lambda: self._finalize_generic_streaming_payment( + key_hash, + max_cost, + path, + model_obj, + provider_fee, + reservation_snapshot, + ), + response, + ) + ) + stream = self._stream_generic_with_settlement( + response, + key_hash, + max_cost, + path, + model_obj, + provider_fee, + reservation_snapshot, + finalizer, + ) + return _ClosingStreamingResponse( + stream, + finalizer=finalizer, + status_code=response.status_code, + headers=dict(response.headers), + ) + async def handle_streaming_messages_completion( self, response: httpx.Response, @@ -2158,13 +2415,59 @@ class BaseUpstreamProvider: request_body: bytes | None = None, ) -> StreamingResponse: usage_estimator = MissingUsageEstimator(request_body, model_obj) + usage_finalized = False + last_model_seen: str | None = None + + async def finalize_without_usage() -> bytes | None: + nonlocal usage_finalized + if usage_finalized: + return None + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + usage_finalized = True + return None + try: + cost_data = await adjust_payment_for_tokens( + fresh_key, + usage_estimator.response_data(last_model_seen), + new_session, + max_cost_for_model, + model_obj, + self.provider_fee, + reservation_snapshot, + ) + usage_finalized = True + return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() + except BaseException as e: + logger.critical( + "Error during Messages API usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + raise + + async def finalize_db_only() -> None: + if not usage_finalized: + await finalize_without_usage() + + stream_finalizer = _PersistentStreamFinalizer( + lambda: _finalize_and_close_stream(finalize_db_only, response) + ) async def stream_with_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: + nonlocal usage_finalized, last_model_seen stored_chunks: list[bytes] = [] - usage_finalized: bool = False - last_model_seen: str | None = None input_tokens: int = 0 output_tokens: int = 0 cache_read_input_tokens: int = 0 @@ -2202,45 +2505,6 @@ class BaseUpstreamProvider: for field in ("total_cost", "cost"): total_cost = max(total_cost, _coerce_usd(usage_or_root.get(field))) - async def finalize_without_usage() -> bytes | None: - nonlocal usage_finalized - if usage_finalized: - return None - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - usage_finalized = True - return None - try: - cost_data = await adjust_payment_for_tokens( - fresh_key, - usage_estimator.response_data(last_model_seen), - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - return f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() - except BaseException as e: - logger.critical( - "Error during Messages API usage finalization — CRITICAL", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "error": str(e), - }, - exc_info=True, - ) - usage_finalized = ( - await self._release_failed_streaming_reservation( - fresh_key, - new_session, - reservation_snapshot, - ) - ) - raise - try: async for chunk in response.aiter_bytes(): stored_chunks.append(chunk) @@ -2428,15 +2692,15 @@ class BaseUpstreamProvider: await finalize_without_usage() raise finally: - if not usage_finalized: - await finalize_without_usage() + await stream_finalizer.run() response_headers = dict(response.headers) response_headers.pop("content-encoding", None) response_headers.pop("content-length", None) - return StreamingResponse( + return _ClosingStreamingResponse( stream_with_cost(max_cost_for_model), + finalizer=stream_finalizer, status_code=response.status_code, headers=response_headers, ) @@ -2725,10 +2989,71 @@ class BaseUpstreamProvider: with cost reconciliation appended at end of stream.""" usage_estimator = MissingUsageEstimator(request_body, model_obj) + usage_finalized = False + last_model_seen: str | None = None + + async def finalize_without_usage() -> bytes | None: + nonlocal usage_finalized + if usage_finalized: + return None + logger.warning( + "Finalizing /v1/messages stream with locally estimated " + "usage because the upstream omitted `usage` from SSE. " + "Check that the upstream emits a final usage chunk; the " + "reservation ceiling will not be used as the charge.", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "model": last_model_seen or "unknown", + "provider": self.provider_type or self.base_url, + "max_cost_msats": max_cost_for_model, + }, + ) + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + usage_finalized = True + return None + try: + cost_data = await adjust_payment_for_tokens( + fresh_key, + usage_estimator.response_data(last_model_seen), + new_session, + max_cost_for_model, + model_obj, + self.provider_fee, + reservation_snapshot, + ) + usage_finalized = True + return ( + f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n" + ).encode() + except BaseException as e: + logger.critical( + "Error during LiteLLM Messages usage finalization — CRITICAL", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + exc_info=True, + ) + usage_finalized = await self._release_failed_streaming_reservation( + fresh_key, + new_session, + reservation_snapshot, + ) + raise + + async def finalize_stream() -> None: + try: + if not usage_finalized: + await finalize_without_usage() + finally: + await _aclose_if_needed(iterator) + + stream_finalizer = _PersistentStreamFinalizer(finalize_stream) async def stream_with_cost() -> AsyncGenerator[bytes, None]: - usage_finalized = False - last_model_seen: str | None = None + nonlocal usage_finalized, last_model_seen input_tokens = 0 output_tokens = 0 cache_read_input_tokens = 0 @@ -2737,59 +3062,6 @@ class BaseUpstreamProvider: input_cost = 0.0 output_cost = 0.0 - async def finalize_without_usage() -> bytes | None: - nonlocal usage_finalized - if usage_finalized: - return None - logger.warning( - "Finalizing /v1/messages stream with locally estimated " - "usage because the upstream omitted `usage` from SSE. " - "Check that the upstream emits a final usage chunk; the " - "reservation ceiling will not be used as the charge.", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "model": last_model_seen or "unknown", - "provider": self.provider_type or self.base_url, - "max_cost_msats": max_cost_for_model, - }, - ) - async with create_session() as new_session: - fresh_key = await new_session.get(key.__class__, key.hashed_key) - if not fresh_key: - usage_finalized = True - return None - try: - cost_data = await adjust_payment_for_tokens( - fresh_key, - usage_estimator.response_data(last_model_seen), - new_session, - max_cost_for_model, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - usage_finalized = True - return ( - f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n" - ).encode() - except BaseException as e: - logger.critical( - "Error during LiteLLM Messages usage finalization — CRITICAL", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "error": str(e), - }, - exc_info=True, - ) - usage_finalized = ( - await self._release_failed_streaming_reservation( - fresh_key, - new_session, - reservation_snapshot, - ) - ) - raise - try: async for annotated in messages_dispatch.stream_annotated_events( iterator, requested_model @@ -2892,11 +3164,11 @@ class BaseUpstreamProvider: await finalize_without_usage() raise finally: - if not usage_finalized: - await finalize_without_usage() + await stream_finalizer.run() - return StreamingResponse( + return _ClosingStreamingResponse( stream_with_cost(), + finalizer=stream_finalizer, media_type="text/event-stream", headers={"Cache-Control": "no-cache", "Connection": "keep-alive"}, ) @@ -3059,7 +3331,7 @@ class BaseUpstreamProvider: for annotated in buffered: yield annotated.sse_bytes - return StreamingResponse( + return _ClosingStreamingResponse( replay(), media_type="text/event-stream", headers=response_headers, @@ -3138,12 +3410,11 @@ class BaseUpstreamProvider: }, ) - client = httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) + response: httpx.Response | None = None + response_handoff = _ResponseHandoff() try: + client = acquire_upstream_http_client(url) if transformed_body is not None: response = await client.send( client.build_request( @@ -3166,6 +3437,7 @@ class BaseUpstreamProvider: ), stream=True, ) + response_handoff.acquire(response) if response.status_code != 200: if response.status_code >= 500: @@ -3200,8 +3472,7 @@ class BaseUpstreamProvider: "body_preview": body_preview, }, ) - await response.aclose() - await client.aclose() + await response_handoff.close() raise UpstreamError( f"Upstream {self.provider_type} returned {response.status_code} " f"for model {original_model_id or 'unknown'}: " @@ -3217,8 +3488,7 @@ class BaseUpstreamProvider: request, path, response, model_id=original_model_id ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() return mapped_error if ( @@ -3251,10 +3521,7 @@ class BaseUpstreamProvider: reservation_snapshot=reservation_snapshot, request_body=request_body, ) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - result.background = background_tasks + response_handoff.handoff() return result if response.status_code == 200: @@ -3271,8 +3538,7 @@ class BaseUpstreamProvider: request_body=request_body, ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() if path.endswith("messages/count_tokens"): if response.status_code == 200: @@ -3289,8 +3555,7 @@ class BaseUpstreamProvider: request_body=request_body, ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() if completion_path is not None: client_wants_streaming = False @@ -3327,19 +3592,18 @@ class BaseUpstreamProvider: ) if is_streaming and response.status_code == 200: - background_tasks = BackgroundTasks() - return await self.handle_streaming_chat_completion( + result = await self.handle_streaming_chat_completion( response, key, max_cost_for_model, - background_tasks, requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, - client=client, request_body=request_body, legacy_completion=completion_path == "completions", ) + response_handoff.handoff() + return result # Handle both non-streaming chat completions and embeddings if response.status_code == 200: @@ -3356,25 +3620,11 @@ class BaseUpstreamProvider: legacy_completion=completion_path == "completions", ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() if reservation_snapshot is None: reservation_snapshot = await get_reservation_snapshot(key, session) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - background_tasks.add_task( - self._finalize_generic_streaming_payment, - key.hashed_key, - max_cost_for_model, - path, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - logger.debug( "Streaming non-chat response", extra={ @@ -3384,18 +3634,24 @@ class BaseUpstreamProvider: }, ) - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, + result = self._generic_streaming_response( + response, + key.hashed_key, + max_cost_for_model, + path, + model_obj, + self.provider_fee, + reservation_snapshot, ) + response_handoff.handoff() + return result except UpstreamError: + await response_handoff.close() raise except httpx.RequestError as exc: - await client.aclose() + await response_handoff.close() error_type = type(exc).__name__ error_details = str(exc) @@ -3413,19 +3669,26 @@ class BaseUpstreamProvider: ) # Don't revert here — proxy.py owns payment revert to avoid double-revert - if isinstance(exc, httpx.ConnectError): + if isinstance(exc, httpx.PoolTimeout): + error_message = "Upstream connection pool is busy" + status_code = 503 + elif isinstance(exc, httpx.ConnectError): error_message = "Unable to connect to upstream service" + status_code = 502 elif isinstance(exc, httpx.TimeoutException): error_message = "Upstream service request timed out" + status_code = 502 elif isinstance(exc, httpx.NetworkError): error_message = "Network error while connecting to upstream service" + status_code = 502 else: error_message = f"Error connecting to upstream service: {error_type}" + status_code = 502 - raise UpstreamError(error_message, status_code=502) + raise UpstreamError(error_message, status_code=status_code) except Exception as exc: - await client.aclose() + await response_handoff.close() tb = traceback.format_exc() logger.error( @@ -3449,6 +3712,10 @@ class BaseUpstreamProvider: scope=ERROR_SCOPE_NODE, ) + except BaseException: + await response_handoff.close(suppress_errors=True) + raise + supports_ehbp: bool = False def get_confidential_inference_profile( @@ -3519,12 +3786,11 @@ class BaseUpstreamProvider: }, ) - client = httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) + response: httpx.Response | None = None + response_handoff = _ResponseHandoff() try: + client = acquire_upstream_http_client(url) if transformed_body is not None: response = await client.send( client.build_request( @@ -3547,6 +3813,7 @@ class BaseUpstreamProvider: ), stream=True, ) + response_handoff.acquire(response) if response.status_code != 200: if response.status_code >= 500: @@ -3580,8 +3847,7 @@ class BaseUpstreamProvider: "body_preview": body_preview, }, ) - await response.aclose() - await client.aclose() + await response_handoff.close() raise UpstreamError( f"Upstream {self.provider_type} returned {response.status_code} " f"for model {original_model_id or 'unknown'}: " @@ -3597,8 +3863,7 @@ class BaseUpstreamProvider: request, path, response, model_id=original_model_id ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() return mapped_error if path.startswith("responses"): @@ -3615,16 +3880,17 @@ class BaseUpstreamProvider: ) if is_streaming and response.status_code == 200: - return await self.handle_streaming_responses_completion( + result = await self.handle_streaming_responses_completion( response, key, max_cost_for_model, requested_model=original_model_id, model_obj=model_obj, reservation_snapshot=reservation_snapshot, - client=client, request_body=transformed_body, ) + response_handoff.handoff() + return result if response.status_code == 200: try: @@ -3639,25 +3905,11 @@ class BaseUpstreamProvider: request_body=transformed_body, ) finally: - await response.aclose() - await client.aclose() + await response_handoff.close() if reservation_snapshot is None: reservation_snapshot = await get_reservation_snapshot(key, session) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - background_tasks.add_task( - self._finalize_generic_streaming_payment, - key.hashed_key, - max_cost_for_model, - path, - model_obj, - self.provider_fee, - reservation_snapshot, - ) - logger.debug( "Streaming non-Responses API response", extra={ @@ -3667,18 +3919,24 @@ class BaseUpstreamProvider: }, ) - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, + result = self._generic_streaming_response( + response, + key.hashed_key, + max_cost_for_model, + path, + model_obj, + self.provider_fee, + reservation_snapshot, ) + response_handoff.handoff() + return result except UpstreamError: + await response_handoff.close() raise except httpx.RequestError as exc: - await client.aclose() + await response_handoff.close() error_type = type(exc).__name__ error_details = str(exc) @@ -3696,19 +3954,26 @@ class BaseUpstreamProvider: ) # Don't revert here — proxy.py owns payment revert to avoid double-revert - if isinstance(exc, httpx.ConnectError): + if isinstance(exc, httpx.PoolTimeout): + error_message = "Upstream connection pool is busy" + status_code = 503 + elif isinstance(exc, httpx.ConnectError): error_message = "Unable to connect to upstream service" + status_code = 502 elif isinstance(exc, httpx.TimeoutException): error_message = "Upstream service request timed out" + status_code = 502 elif isinstance(exc, httpx.NetworkError): error_message = "Network error while connecting to upstream service" + status_code = 502 else: error_message = f"Error connecting to upstream service: {error_type}" + status_code = 502 - raise UpstreamError(error_message, status_code=502) + raise UpstreamError(error_message, status_code=status_code) except Exception as exc: - await client.aclose() + await response_handoff.close() tb = traceback.format_exc() logger.error( @@ -3732,6 +3997,10 @@ class BaseUpstreamProvider: scope=ERROR_SCOPE_NODE, ) + except BaseException: + await response_handoff.close(suppress_errors=True) + raise + async def forward_get_request( self, request: Request, @@ -3761,66 +4030,91 @@ class BaseUpstreamProvider: }, ) - async with httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) as client: - try: - response = await client.send( - client.build_request( - request.method, - url, - headers=headers, - content=request.stream(), - params=self.prepare_params(path, request.query_params), - ), + response: httpx.Response | None = None + try: + client = acquire_upstream_http_client(url) + response = await client.send( + client.build_request( + request.method, + url, + headers=headers, + content=request.stream(), + params=self.prepare_params(path, request.query_params), + ), + ) + + logger.debug( + "GET request forwarded", + extra={ + "path": path, + "status_code": response.status_code, + "provider": self.provider_type, + }, + ) + if response.status_code != 200: + return await self.forward_upstream_error_response( + request, path, response ) - logger.debug( - "GET request forwarded", - extra={ - "path": path, - "status_code": response.status_code, - "provider": self.provider_type, - }, - ) - if response.status_code != 200: - try: - mapped = await self.forward_upstream_error_response( - request, path, response - ) - finally: - await response.aclose() - return mapped - - response_headers = dict(response.headers) - response_headers.pop("content-encoding", None) - response_headers.pop("content-length", None) - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=response_headers, - ) - except Exception as exc: - tb = traceback.format_exc() - logger.error( - "Error forwarding GET request", - extra={ - "error": str(exc), - "error_type": type(exc).__name__, - "method": request.method, - "url": url, - "path": path, - "query_params": dict(request.query_params), - "traceback": tb, - }, - ) - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, - ) + response_headers = dict(response.headers) + response_headers.pop("content-encoding", None) + response_headers.pop("content-length", None) + return Response( + content=response.content, + status_code=response.status_code, + headers=response_headers, + ) + except UpstreamError: + raise + except httpx.PoolTimeout: + logger.warning( + "Upstream connection pool exhausted on GET", + extra={"path": path, "url": url, "provider": self.provider_type}, + ) + return create_error_response( + "service_unavailable", + "Upstream connection pool is busy", + 503, + request=request, + ) + except httpx.RequestError as exc: + logger.warning( + "Upstream request error on GET", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "path": path, + "url": url, + "provider": self.provider_type, + }, + ) + return create_error_response( + "upstream_error", + "Unable to reach upstream service", + 502, + request=request, + ) + except Exception as exc: + logger.error( + "Error forwarding GET request", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "method": request.method, + "url": url, + "path": path, + "query_params": dict(request.query_params), + "traceback": traceback.format_exc(), + }, + ) + return create_error_response( + "internal_error", + "An unexpected server error occurred", + 500, + request=request, + ) + finally: + await _aclose_if_needed(response) async def get_x_cashu_cost( self, @@ -4146,7 +4440,7 @@ class BaseUpstreamProvider: for line in lines: yield (line + "\n").encode("utf-8") - return StreamingResponse( + return _ClosingStreamingResponse( generate(), status_code=response.status_code, headers=response_headers, @@ -4402,7 +4696,7 @@ class BaseUpstreamProvider: "unit": unit, }, ) - return StreamingResponse( + return _ClosingStreamingResponse( response.aiter_bytes(), status_code=response.status_code, headers=dict(response.headers), @@ -4490,154 +4784,150 @@ class BaseUpstreamProvider: }, ) - async with httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) as client: - try: - response = await client.send( - client.build_request( - request.method, - url, - headers=headers, - content=transformed_body if transformed_body else request_body, - params=self.prepare_params(path, request.query_params), - ), - stream=True, - ) + client = _build_x_cashu_client() + response: httpx.Response | None = None + try: + response = await client.send( + client.build_request( + request.method, + url, + headers=headers, + content=transformed_body if transformed_body else request_body, + params=self.prepare_params(path, request.query_params), + ), + stream=True, + ) - if response.status_code != 200: - logger.error( - "Received upstream response", - extra={ - "reason_phrase": response.reason_phrase, - "status_code": response.status_code, - "path": path, - "response_headers": dict(response.headers), - }, - ) - else: - logger.debug( - "Received upstream response", - extra={ - "status_code": response.status_code, - "path": path, - "response_headers": dict(response.headers), - }, - ) - - if response.status_code != 200: - logger.warning( - "Upstream request failed, processing refund", - extra={ - "status_code": response.status_code, - "path": path, - "amount": amount, - "unit": unit, - }, - ) - - refund_token = await self.send_refund( - amount, - unit, - mint, - request_id=getattr(request.state, "request_id", None), - ) - - logger.info( - "Refund processed for failed upstream request", - extra={ - "status_code": response.status_code, - "refund_amount": amount, - "unit": unit, - "refund_token_preview": refund_token[:20] + "..." - if len(refund_token) > 20 - else refund_token, - }, - ) - - error_response = Response( - content=json.dumps( - { - "error": { - "message": "Error forwarding request to upstream", - "type": "upstream_error", - # Pass the status as the code so a provider - # 4xx keeps the legacy numeric ``code``. - "code": client_code_for_upstream_error( - response.status_code, response.status_code - ), - "upstream_status": response.status_code, - "refund_token": refund_token, - } - } - ), - status_code=client_status_for_upstream_error( - response.status_code - ), - media_type="application/json", - ) - error_response.headers["X-Cashu"] = refund_token - error_response.headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM - return error_response - - if _x_cashu_path_has_settlement_handler(path): - logger.debug( - "Processing completion/embeddings/messages response", - extra={"path": path, "amount": amount, "unit": unit}, - ) - - result = await self.handle_x_cashu_chat_completion( - response, - amount, - unit, - max_cost_for_model, - mint, - request_id=getattr(request.state, "request_id", None), - model_obj=model_obj, - request_body=request_body, - ) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - result.background = background_tasks - return result - - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - - logger.debug( - "Streaming non-chat response", - extra={"path": path, "status_code": response.status_code}, - ) - - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, - ) - except Exception as exc: - tb = traceback.format_exc() + if response.status_code != 200: logger.error( - "Unexpected error in upstream forwarding", + "Received upstream response", extra={ - "error": str(exc), - "error_type": type(exc).__name__, - "method": request.method, - "url": url, + "reason_phrase": response.reason_phrase, + "status_code": response.status_code, "path": path, - "query_params": dict(request.query_params), - "traceback": tb, + "response_headers": dict(response.headers), }, ) - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, + else: + logger.debug( + "Received upstream response", + extra={ + "status_code": response.status_code, + "path": path, + "response_headers": dict(response.headers), + }, ) + if response.status_code != 200: + logger.warning( + "Upstream request failed, processing refund", + extra={ + "status_code": response.status_code, + "path": path, + "amount": amount, + "unit": unit, + }, + ) + + refund_token = await self.send_refund( + amount, + unit, + mint, + request_id=getattr(request.state, "request_id", None), + ) + + logger.info( + "Refund processed for failed upstream request", + extra={ + "status_code": response.status_code, + "refund_amount": amount, + "unit": unit, + "refund_token_preview": refund_token[:20] + "..." + if len(refund_token) > 20 + else refund_token, + }, + ) + + error_response = Response( + content=json.dumps( + { + "error": { + "message": "Error forwarding request to upstream", + "type": "upstream_error", + # Pass the status as the code so a provider + # 4xx keeps the legacy numeric ``code``. + "code": client_code_for_upstream_error( + response.status_code, response.status_code + ), + "upstream_status": response.status_code, + "refund_token": refund_token, + } + } + ), + status_code=client_status_for_upstream_error(response.status_code), + media_type="application/json", + ) + error_response.headers["X-Cashu"] = refund_token + error_response.headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM + await _close_upstream_exchange(response, client) + return error_response + + if _x_cashu_path_has_settlement_handler(path): + logger.debug( + "Processing completion/embeddings/messages response", + extra={"path": path, "amount": amount, "unit": unit}, + ) + + result = await self.handle_x_cashu_chat_completion( + response, + amount, + unit, + max_cost_for_model, + mint, + request_id=getattr(request.state, "request_id", None), + model_obj=model_obj, + request_body=request_body, + ) + if isinstance(result, StreamingResponse) and not response.is_closed: + return _attach_upstream_stream_owner(result, response, client) + await _close_upstream_exchange(response, client) + return result + + logger.debug( + "Streaming non-chat response", + extra={"path": path, "status_code": response.status_code}, + ) + + return _ClosingStreamingResponse( + _OwnedUpstreamStream(response.aiter_bytes(), response, client), + status_code=response.status_code, + headers=dict(response.headers), + ) + except asyncio.CancelledError: + await _close_upstream_exchange(response, client) + raise + except Exception as exc: + await _close_upstream_exchange(response, client) + tb = traceback.format_exc() + logger.error( + "Unexpected error in upstream forwarding", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "method": request.method, + "url": url, + "path": path, + "query_params": dict(request.query_params), + "traceback": tb, + }, + ) + return create_error_response( + "internal_error", + "An unexpected server error occurred", + 500, + request=request, + ) + async def handle_x_cashu_responses( self, request: Request, @@ -4810,142 +5100,138 @@ class BaseUpstreamProvider: }, ) - async with httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(retries=1), - timeout=None, - ) as client: - try: - response = await client.send( - client.build_request( - request.method, - url, - headers=headers, - content=transformed_body if transformed_body else request_body, - params=self.prepare_params(path, request.query_params), - ), - stream=True, - ) + client = _build_x_cashu_client() + response: httpx.Response | None = None + try: + response = await client.send( + client.build_request( + request.method, + url, + headers=headers, + content=transformed_body if transformed_body else request_body, + params=self.prepare_params(path, request.query_params), + ), + stream=True, + ) - logger.debug( - "Received upstream Responses API response", + logger.debug( + "Received upstream Responses API response", + extra={ + "status_code": response.status_code, + "path": path, + "response_headers": dict(response.headers), + }, + ) + + if response.status_code != 200: + logger.warning( + "Upstream Responses API request failed, processing refund", extra={ "status_code": response.status_code, "path": path, - "response_headers": dict(response.headers), + "amount": amount, + "unit": unit, }, ) - if response.status_code != 200: - logger.warning( - "Upstream Responses API request failed, processing refund", - extra={ - "status_code": response.status_code, - "path": path, - "amount": amount, - "unit": unit, - }, - ) - - refund_token = await self.send_refund( - amount, - unit, - mint, - request_id=getattr(request.state, "request_id", None), - ) - - logger.info( - "Refund processed for failed upstream Responses API request", - extra={ - "status_code": response.status_code, - "refund_amount": amount, - "unit": unit, - "refund_token_preview": refund_token[:20] + "..." - if len(refund_token) > 20 - else refund_token, - }, - ) - - error_response = Response( - content=json.dumps( - { - "error": { - "message": "Error forwarding Responses API request to upstream", - "type": "upstream_error", - # Pass the status as the code so a provider - # 4xx keeps the legacy numeric ``code``. - "code": client_code_for_upstream_error( - response.status_code, response.status_code - ), - "upstream_status": response.status_code, - "refund_token": refund_token, - } - } - ), - status_code=client_status_for_upstream_error( - response.status_code - ), - media_type="application/json", - ) - error_response.headers["X-Cashu"] = refund_token - error_response.headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM - return error_response - - if path.startswith("responses"): - logger.debug( - "Processing Responses API response", - extra={"path": path, "amount": amount, "unit": unit}, - ) - - result = await self.handle_x_cashu_responses_completion( - response, - amount, - unit, - max_cost_for_model, - mint, - request_id=getattr(request.state, "request_id", None), - model_obj=model_obj, - request_body=request_body, - ) - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - result.background = background_tasks - return result - - background_tasks = BackgroundTasks() - background_tasks.add_task(response.aclose) - background_tasks.add_task(client.aclose) - - logger.debug( - "Streaming non-responses response", - extra={"path": path, "status_code": response.status_code}, + refund_token = await self.send_refund( + amount, + unit, + mint, + request_id=getattr(request.state, "request_id", None), ) - return StreamingResponse( - response.aiter_bytes(), - status_code=response.status_code, - headers=dict(response.headers), - background=background_tasks, - ) - except Exception as exc: - tb = traceback.format_exc() - logger.error( - "Unexpected error in upstream Responses API forwarding", + logger.info( + "Refund processed for failed upstream Responses API request", extra={ - "error": str(exc), - "error_type": type(exc).__name__, - "method": request.method, - "url": url, - "path": path, - "query_params": dict(request.query_params), - "traceback": tb, + "status_code": response.status_code, + "refund_amount": amount, + "unit": unit, + "refund_token_preview": refund_token[:20] + "..." + if len(refund_token) > 20 + else refund_token, }, ) - return create_error_response( - "internal_error", - "An unexpected server error occurred", - 500, - request=request, + + error_response = Response( + content=json.dumps( + { + "error": { + "message": "Error forwarding Responses API request to upstream", + "type": "upstream_error", + # Pass the status as the code so a provider + # 4xx keeps the legacy numeric ``code``. + "code": client_code_for_upstream_error( + response.status_code, response.status_code + ), + "upstream_status": response.status_code, + "refund_token": refund_token, + } + } + ), + status_code=client_status_for_upstream_error(response.status_code), + media_type="application/json", ) + error_response.headers["X-Cashu"] = refund_token + error_response.headers[ERROR_SCOPE_HEADER] = ERROR_SCOPE_UPSTREAM + await _close_upstream_exchange(response, client) + return error_response + + if path.startswith("responses"): + logger.debug( + "Processing Responses API response", + extra={"path": path, "amount": amount, "unit": unit}, + ) + + result = await self.handle_x_cashu_responses_completion( + response, + amount, + unit, + max_cost_for_model, + mint, + request_id=getattr(request.state, "request_id", None), + model_obj=model_obj, + request_body=request_body, + ) + if isinstance(result, StreamingResponse) and not response.is_closed: + return _attach_upstream_stream_owner(result, response, client) + await _close_upstream_exchange(response, client) + return result + + logger.debug( + "Streaming non-responses response", + extra={"path": path, "status_code": response.status_code}, + ) + + return _ClosingStreamingResponse( + _OwnedUpstreamStream(response.aiter_bytes(), response, client), + status_code=response.status_code, + headers=dict(response.headers), + ) + except asyncio.CancelledError: + await _close_upstream_exchange(response, client) + raise + except Exception as exc: + await _close_upstream_exchange(response, client) + tb = traceback.format_exc() + logger.error( + "Unexpected error in upstream Responses API forwarding", + extra={ + "error": str(exc), + "error_type": type(exc).__name__, + "method": request.method, + "url": url, + "path": path, + "query_params": dict(request.query_params), + "traceback": tb, + }, + ) + return create_error_response( + "internal_error", + "An unexpected server error occurred", + 500, + request=request, + ) async def handle_x_cashu_responses_completion( self, @@ -5029,7 +5315,7 @@ class BaseUpstreamProvider: "unit": unit, }, ) - return StreamingResponse( + return _ClosingStreamingResponse( response.aiter_bytes(), status_code=response.status_code, headers=dict(response.headers), @@ -5224,7 +5510,7 @@ class BaseUpstreamProvider: for fields, data in events: yield _render_sse_event(fields, data).encode("utf-8") - return StreamingResponse( + return _ClosingStreamingResponse( generate(), status_code=response.status_code, headers=response_headers, diff --git a/routstr/upstream/ehbp.py b/routstr/upstream/ehbp.py index 673416f6..c1ee8cfc 100644 --- a/routstr/upstream/ehbp.py +++ b/routstr/upstream/ehbp.py @@ -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, diff --git a/routstr/upstream/gemini_messages.py b/routstr/upstream/gemini_messages.py index 8bc87db7..f2c454e0 100644 --- a/routstr/upstream/gemini_messages.py +++ b/routstr/upstream/gemini_messages.py @@ -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 diff --git a/routstr/upstream/http_client.py b/routstr/upstream/http_client.py new file mode 100644 index 00000000..0a09c8ca --- /dev/null +++ b/routstr/upstream/http_client.py @@ -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 diff --git a/tests/integration/test_reservation_lifecycle.py b/tests/integration/test_reservation_lifecycle.py index 1f60ce9d..d3fb9f9f 100644 --- a/tests/integration/test_reservation_lifecycle.py +++ b/tests/integration/test_reservation_lifecycle.py @@ -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 diff --git a/tests/unit/test_log_secret_redaction.py b/tests/unit/test_log_secret_redaction.py index 0f660dda..3b51644d 100644 --- a/tests/unit/test_log_secret_redaction.py +++ b/tests/unit/test_log_secret_redaction.py @@ -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", diff --git a/tests/unit/test_model_path_routing.py b/tests/unit/test_model_path_routing.py index 4faec94b..5b23720f 100644 --- a/tests/unit/test_model_path_routing.py +++ b/tests/unit/test_model_path_routing.py @@ -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()), ): diff --git a/tests/unit/test_payment_settlement_timing.py b/tests/unit/test_payment_settlement_timing.py new file mode 100644 index 00000000..4893fc7d --- /dev/null +++ b/tests/unit/test_payment_settlement_timing.py @@ -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 diff --git a/tests/unit/test_pre_handoff_stream_ownership.py b/tests/unit/test_pre_handoff_stream_ownership.py new file mode 100644 index 00000000..894798e0 --- /dev/null +++ b/tests/unit/test_pre_handoff_stream_ownership.py @@ -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() diff --git a/tests/unit/test_queued_logging.py b/tests/unit/test_queued_logging.py new file mode 100644 index 00000000..83019cef --- /dev/null +++ b/tests/unit/test_queued_logging.py @@ -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) diff --git a/tests/unit/test_settings.py b/tests/unit/test_settings.py index fb3c18e0..61f657ac 100644 --- a/tests/unit/test_settings.py +++ b/tests/unit/test_settings.py @@ -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"), [ diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 498b57f5..2b5d8e31 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -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), diff --git a/tests/unit/test_stream_id_injection.py b/tests/unit/test_stream_id_injection.py index 2d682bc5..29912a4f 100644 --- a/tests/unit/test_stream_id_injection.py +++ b/tests/unit/test_stream_id_injection.py @@ -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", diff --git a/tests/unit/test_streaming_billing_finalization.py b/tests/unit/test_streaming_billing_finalization.py index d60b598a..10c5a3e8 100644 --- a/tests/unit/test_streaming_billing_finalization.py +++ b/tests/unit/test_streaming_billing_finalization.py @@ -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) diff --git a/tests/unit/test_streaming_sse_providers.py b/tests/unit/test_streaming_sse_providers.py index ffb7e266..afbf22ba 100644 --- a/tests/unit/test_streaming_sse_providers.py +++ b/tests/unit/test_streaming_sse_providers.py @@ -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", diff --git a/tests/unit/test_tinfoil_integration.py b/tests/unit/test_tinfoil_integration.py index 36e76765..ec5e8c11 100644 --- a/tests/unit/test_tinfoil_integration.py +++ b/tests/unit/test_tinfoil_integration.py @@ -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), diff --git a/tests/unit/test_upstream_gemini.py b/tests/unit/test_upstream_gemini.py index 13683887..b9cf0b73 100644 --- a/tests/unit/test_upstream_gemini.py +++ b/tests/unit/test_upstream_gemini.py @@ -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 diff --git a/tests/unit/test_upstream_http_client.py b/tests/unit/test_upstream_http_client.py new file mode 100644 index 00000000..82b46057 --- /dev/null +++ b/tests/unit/test_upstream_http_client.py @@ -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 diff --git a/tests/unit/test_upstream_rate_limit.py b/tests/unit/test_upstream_rate_limit.py index 95b29d83..c111b74a 100644 --- a/tests/unit/test_upstream_rate_limit.py +++ b/tests/unit/test_upstream_rate_limit.py @@ -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), diff --git a/tests/unit/test_x_cashu_stream_ownership.py b/tests/unit/test_x_cashu_stream_ownership.py new file mode 100644 index 00000000..f3b7d0a6 --- /dev/null +++ b/tests/unit/test_x_cashu_stream_ownership.py @@ -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