mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
Merge remote-tracking branch 'origin/main'
This commit is contained in:
+162
-29
@@ -1,7 +1,7 @@
|
||||
import time
|
||||
import uuid
|
||||
from contextvars import ContextVar
|
||||
from typing import Callable
|
||||
from typing import AsyncIterator, Callable
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from fastapi import Request, Response
|
||||
@@ -9,6 +9,7 @@ from starlette.datastructures import Headers
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
|
||||
from .logging import get_logger
|
||||
from .settings import settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -86,11 +87,18 @@ _SKIP_LOG_EXACT: frozenset[str] = frozenset(
|
||||
)
|
||||
|
||||
|
||||
def _should_log(method: str, path: str) -> bool:
|
||||
def _should_log(method: str, path: str, status_code: int | None = None) -> bool:
|
||||
if method in _SKIP_LOG_METHODS:
|
||||
return False
|
||||
# Our own faults are never noise, whatever the path.
|
||||
if status_code is not None and status_code >= 500:
|
||||
return True
|
||||
if path in _SKIP_LOG_EXACT:
|
||||
return False
|
||||
# A 4xx storm on a UI-polled path is exactly what we need to see.
|
||||
return status_code is not None and status_code >= 400
|
||||
# Client errors on the skipped prefixes stay hidden: 404s under /_next/ are
|
||||
# driven by whoever scans the node, and the admin UI's timer-driven polling
|
||||
# turns one expired session into a 401 per poll.
|
||||
return not any(path.startswith(prefix) for prefix in _SKIP_LOG_PREFIXES)
|
||||
|
||||
|
||||
@@ -103,12 +111,116 @@ def _attribution(request: Request) -> dict[str, object]:
|
||||
}
|
||||
|
||||
|
||||
def mark(request: Request, name: str) -> None:
|
||||
"""Record that stage ``name`` finished, for the completion log's timings."""
|
||||
marks = getattr(request.state, "stage_marks", None)
|
||||
if marks is not None:
|
||||
marks[name] = time.monotonic()
|
||||
|
||||
|
||||
def _request_content_length(headers: Headers) -> int | None:
|
||||
"""Client-supplied length, dropped unless it is a plausible byte count."""
|
||||
raw = headers.get("content-length")
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
value = int(raw)
|
||||
except ValueError:
|
||||
return None
|
||||
return value if value >= 0 else None
|
||||
|
||||
|
||||
class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
"""Middleware to log proxy interactions and page navigation.
|
||||
|
||||
Skips logging for static assets and Next.js chunks to avoid noise.
|
||||
"""
|
||||
|
||||
def _log_completion(
|
||||
self,
|
||||
*,
|
||||
request: Request,
|
||||
request_id: str,
|
||||
path: str,
|
||||
status_code: int,
|
||||
duration: float,
|
||||
headers_duration: float | None,
|
||||
stage_start: float,
|
||||
stage_marks: dict[str, float],
|
||||
incoming_logged: bool,
|
||||
) -> None:
|
||||
if not _should_log(request.method, path, status_code):
|
||||
return
|
||||
|
||||
extra: dict[str, object] = {
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": path,
|
||||
"status_code": status_code,
|
||||
"duration_ms": round(duration * 1000, 2),
|
||||
"content_length": _request_content_length(request.headers),
|
||||
**_attribution(request),
|
||||
}
|
||||
if headers_duration is not None:
|
||||
extra["time_to_headers_ms"] = round(headers_duration * 1000, 2)
|
||||
if not incoming_logged:
|
||||
# Tells log consumers that join on request_id why the matching
|
||||
# "Incoming request" record is missing.
|
||||
extra["incoming_suppressed"] = True
|
||||
for name, marked_at in stage_marks.items():
|
||||
extra[f"{name}_ms"] = round((marked_at - stage_start) * 1000, 2)
|
||||
if status_code >= 400:
|
||||
error_detail = getattr(request.state, "error_detail", None)
|
||||
if isinstance(error_detail, dict):
|
||||
extra["error_type"] = error_detail.get("error_type")
|
||||
extra["error_code"] = error_detail.get("error_code")
|
||||
extra["error_message"] = error_detail.get("error_message")
|
||||
log = (
|
||||
logger.warning
|
||||
if duration > settings.slow_request_warn_seconds
|
||||
else logger.info
|
||||
)
|
||||
log("Request completed", extra=extra)
|
||||
|
||||
async def _timed_body(
|
||||
self,
|
||||
body_iterator: AsyncIterator[bytes],
|
||||
*,
|
||||
request: Request,
|
||||
request_id: str,
|
||||
client_app: str,
|
||||
path: str,
|
||||
status_code: int,
|
||||
stage_start: float,
|
||||
stage_marks: dict[str, float],
|
||||
headers_duration: float,
|
||||
incoming_logged: bool,
|
||||
) -> AsyncIterator[bytes]:
|
||||
try:
|
||||
async for chunk in body_iterator:
|
||||
yield chunk
|
||||
finally:
|
||||
duration = time.monotonic() - stage_start
|
||||
# dispatch() has already reset both context vars by now, and the
|
||||
# logging filters read request_id/client_app from them.
|
||||
request_token = request_id_context.set(request_id)
|
||||
app_token = client_app_context.set(client_app)
|
||||
try:
|
||||
self._log_completion(
|
||||
request=request,
|
||||
request_id=request_id,
|
||||
path=path,
|
||||
status_code=status_code,
|
||||
duration=duration,
|
||||
headers_duration=headers_duration,
|
||||
stage_start=stage_start,
|
||||
stage_marks=stage_marks,
|
||||
incoming_logged=incoming_logged,
|
||||
)
|
||||
finally:
|
||||
request_id_context.reset(request_token)
|
||||
client_app_context.reset(app_token)
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Callable) -> Response:
|
||||
# Generate request ID
|
||||
request_id = str(uuid.uuid4())
|
||||
@@ -117,15 +229,17 @@ class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
# Set request ID in context for logging
|
||||
token = request_id_context.set(request_id)
|
||||
|
||||
client_app_token = client_app_context.set(
|
||||
client_app_from_headers(request.headers)
|
||||
)
|
||||
client_app = client_app_from_headers(request.headers)
|
||||
client_app_token = client_app_context.set(client_app)
|
||||
|
||||
path = request.url.path
|
||||
should_log = _should_log(request.method, path)
|
||||
|
||||
# Start timing
|
||||
start_time = time.time()
|
||||
# Start timing. Monotonic throughout: a wall-clock step would otherwise
|
||||
# produce negative durations and bogus slow-request warnings.
|
||||
stage_start = time.monotonic()
|
||||
stage_marks: dict[str, float] = {}
|
||||
request.state.stage_marks = stage_marks
|
||||
|
||||
if should_log:
|
||||
logger.info(
|
||||
@@ -144,34 +258,52 @@ class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
try:
|
||||
response = await call_next(request)
|
||||
|
||||
if should_log:
|
||||
duration = time.time() - start_time
|
||||
extra: dict[str, object] = {
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": path,
|
||||
"status_code": response.status_code,
|
||||
"duration_ms": round(duration * 1000, 2),
|
||||
**_attribution(request),
|
||||
}
|
||||
if response.status_code >= 400:
|
||||
error_detail = getattr(request.state, "error_detail", None)
|
||||
if isinstance(error_detail, dict):
|
||||
extra["error_type"] = error_detail.get("error_type")
|
||||
extra["error_code"] = error_detail.get("error_code")
|
||||
extra["error_message"] = error_detail.get("error_message")
|
||||
logger.info(
|
||||
"Request completed",
|
||||
extra=extra,
|
||||
)
|
||||
headers_duration = time.monotonic() - stage_start
|
||||
|
||||
if hasattr(response, "headers"):
|
||||
response.headers["x-routstr-request-id"] = request_id
|
||||
# Headers are already on the wire before a streamed body ends,
|
||||
# so this can only ever be time-to-headers.
|
||||
response.headers["x-routstr-duration-ms"] = str(
|
||||
round(headers_duration * 1000, 2)
|
||||
)
|
||||
|
||||
body_iterator = getattr(response, "body_iterator", None)
|
||||
if body_iterator is None:
|
||||
self._log_completion(
|
||||
request=request,
|
||||
request_id=request_id,
|
||||
path=path,
|
||||
status_code=response.status_code,
|
||||
duration=headers_duration,
|
||||
headers_duration=None,
|
||||
stage_start=stage_start,
|
||||
stage_marks=stage_marks,
|
||||
incoming_logged=should_log,
|
||||
)
|
||||
return response
|
||||
|
||||
# A StreamingResponse is barely started here: most of the time a
|
||||
# slow completion spends in the node is spent relaying its body, so
|
||||
# the completion log has to wait for the iterator to drain.
|
||||
response.body_iterator = self._timed_body(
|
||||
body_iterator,
|
||||
request=request,
|
||||
request_id=request_id,
|
||||
client_app=client_app,
|
||||
path=path,
|
||||
status_code=response.status_code,
|
||||
stage_start=stage_start,
|
||||
stage_marks=stage_marks,
|
||||
headers_duration=headers_duration,
|
||||
incoming_logged=should_log,
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
except Exception as e:
|
||||
# Always log failures, even for skipped paths, so we don't lose errors.
|
||||
duration = time.time() - start_time
|
||||
duration = time.monotonic() - stage_start
|
||||
logger.error(
|
||||
"Request failed",
|
||||
extra={
|
||||
@@ -196,5 +328,6 @@ __all__ = [
|
||||
"LoggingMiddleware",
|
||||
"UNKNOWN_CLIENT_APP",
|
||||
"client_app_context",
|
||||
"mark",
|
||||
"request_id_context",
|
||||
]
|
||||
|
||||
@@ -125,6 +125,14 @@ class Settings(BaseSettings):
|
||||
# widens what the provider credential can be spent against, so wildcards
|
||||
# and prefixes are not supported.
|
||||
proxy_extra_allowed_paths: str = Field(default="", env="PROXY_EXTRA_ALLOWED_PATHS")
|
||||
# Bound the client request body: a slow or oversized upload otherwise blocks
|
||||
# the proxy before authentication and holds server resources for its duration.
|
||||
request_body_timeout_seconds: float = Field(
|
||||
default=30.0, gt=0, env="REQUEST_BODY_TIMEOUT_SECONDS"
|
||||
)
|
||||
max_request_body_bytes: int = Field(
|
||||
default=20 * 1024 * 1024, gt=0, env="MAX_REQUEST_BODY_BYTES"
|
||||
)
|
||||
tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL")
|
||||
providers_refresh_interval_seconds: int = Field(
|
||||
default=0, env="PROVIDERS_REFRESH_INTERVAL_SECONDS"
|
||||
@@ -189,6 +197,9 @@ class Settings(BaseSettings):
|
||||
# Logging
|
||||
log_level: str = Field(default="INFO", env="LOG_LEVEL")
|
||||
enable_console_logging: bool = Field(default=True, env="ENABLE_CONSOLE_LOGGING")
|
||||
slow_request_warn_seconds: float = Field(
|
||||
default=60.0, gt=0, env="SLOW_REQUEST_WARN_SECONDS"
|
||||
)
|
||||
|
||||
# Other
|
||||
chat_completions_api_version: str = Field(
|
||||
|
||||
+71
-16
@@ -3,7 +3,7 @@ import inspect
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
from sqlmodel import select
|
||||
|
||||
@@ -21,7 +21,6 @@ from .core.db import (
|
||||
ModelRow,
|
||||
UpstreamProviderRow,
|
||||
create_session,
|
||||
get_session,
|
||||
)
|
||||
from .core.error_scope import (
|
||||
ERROR_SCOPE_UPSTREAM,
|
||||
@@ -29,6 +28,7 @@ from .core.error_scope import (
|
||||
UPSTREAM_UNAVAILABLE,
|
||||
)
|
||||
from .core.exceptions import UpstreamError
|
||||
from .core.middleware import mark
|
||||
from .core.not_found import build_not_found_response
|
||||
from .core.settings import settings
|
||||
from .payment.helpers import (
|
||||
@@ -418,23 +418,78 @@ def _attribute_request(
|
||||
request.state.provider = upstream.provider_type
|
||||
|
||||
|
||||
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
|
||||
async def proxy(
|
||||
request: Request, path: str, session: AsyncSession = Depends(get_session)
|
||||
) -> Response | StreamingResponse:
|
||||
"""Run proxy setup in a short request session, never across response streaming."""
|
||||
class _BodyLimitExceeded(Exception):
|
||||
"""The client body is larger than ``max_request_body_bytes``."""
|
||||
|
||||
|
||||
async def _read_bounded_body(request: Request) -> bytes | Response:
|
||||
"""Read the request body under a size and time bound.
|
||||
|
||||
Returns the body, or the error response to send instead. Both bounds run
|
||||
before any authentication or DB work, so an oversized or slowly uploaded
|
||||
body cannot occupy the request for longer than the timeout.
|
||||
"""
|
||||
max_bytes = settings.max_request_body_bytes
|
||||
timeout = settings.request_body_timeout_seconds
|
||||
|
||||
async def read() -> bytes:
|
||||
declared = request.headers.get("content-length", "")
|
||||
if declared.isdigit() and int(declared) > max_bytes:
|
||||
raise _BodyLimitExceeded
|
||||
body = bytearray()
|
||||
async for chunk in request.stream():
|
||||
body += chunk
|
||||
# Chunked uploads declare no length, so the cap is enforced here.
|
||||
if len(body) > max_bytes:
|
||||
raise _BodyLimitExceeded
|
||||
return bytes(body)
|
||||
|
||||
try:
|
||||
return await _proxy(request, path, session)
|
||||
finally:
|
||||
# FastAPI yield dependencies normally close after the response body is
|
||||
# sent. Close explicitly so a long stream cannot retain DB resources.
|
||||
close_result = session.close()
|
||||
if inspect.isawaitable(close_result):
|
||||
await close_result
|
||||
body = await asyncio.wait_for(read(), timeout)
|
||||
except _BodyLimitExceeded:
|
||||
error_type, message, status = (
|
||||
"invalid_request",
|
||||
f"Request body exceeds the {max_bytes} byte limit",
|
||||
413,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
error_type, message, status = (
|
||||
"timeout",
|
||||
f"Request body not received within {timeout} seconds",
|
||||
408,
|
||||
)
|
||||
else:
|
||||
# Draining the stream leaves Starlette unable to serve a second read.
|
||||
# Cache the body so later readers (EHBP forwarding, upstream stream
|
||||
# passthrough) get it instead of "Stream consumed".
|
||||
request._body = body
|
||||
return body
|
||||
return create_error_response(error_type, message, status, request=request)
|
||||
|
||||
|
||||
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
|
||||
async def proxy(request: Request, path: str) -> Response | StreamingResponse:
|
||||
"""Run proxy setup in a short request session, never across response streaming."""
|
||||
# Read the body before opening a session: a slow uploader must not hold a
|
||||
# DB connection while its request trickles in.
|
||||
request_body = await _read_bounded_body(request)
|
||||
if isinstance(request_body, Response):
|
||||
return request_body
|
||||
mark(request, "body_read")
|
||||
|
||||
async with create_session() as session:
|
||||
try:
|
||||
return await _proxy(request, path, session, request_body)
|
||||
finally:
|
||||
# Close explicitly so a long stream cannot retain DB resources
|
||||
# while its response body is being sent.
|
||||
close_result = session.close()
|
||||
if inspect.isawaitable(close_result):
|
||||
await close_result
|
||||
|
||||
|
||||
async def _proxy(
|
||||
request: Request, path: str, session: AsyncSession
|
||||
request: Request, path: str, session: AsyncSession, request_body: bytes
|
||||
) -> Response | StreamingResponse:
|
||||
# Screen the path before any routing decision: reject ambiguous spellings,
|
||||
# then require a known API prefix so nothing unknown is forwarded with the
|
||||
@@ -449,7 +504,6 @@ async def _proxy(
|
||||
return build_not_found_response(request, path)
|
||||
|
||||
is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
|
||||
request_body = await request.body()
|
||||
|
||||
# EHBP (Encrypted HTTP Body Protocol) requests carry an Ehbp-Encapsulated-Key
|
||||
# header and a binary HPKE-sealed body. The proxy cannot parse the body to
|
||||
@@ -725,6 +779,7 @@ async def _proxy(
|
||||
key = await get_bearer_token_key(
|
||||
headers, path, session, auth, max_cost_for_model, model_id
|
||||
)
|
||||
mark(request, "auth")
|
||||
|
||||
else:
|
||||
if request.method not in ["GET"]:
|
||||
|
||||
@@ -33,7 +33,6 @@ async def test_authenticated_proxy_releases_db_connection_before_upstream_header
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = {"authorization": "Bearer test-key"}
|
||||
request.body = AsyncMock(return_value=json.dumps({"model": "test-model"}).encode())
|
||||
request.url.path = "/v1/chat/completions"
|
||||
request.state.request_id = "pool-hold-regression"
|
||||
|
||||
@@ -59,7 +58,10 @@ async def test_authenticated_proxy_releases_db_connection_before_upstream_header
|
||||
patch("routstr.proxy.get_bearer_token_key", AsyncMock(return_value=key)),
|
||||
):
|
||||
response = await proxy_module._proxy(
|
||||
request, "v1/chat/completions", integration_session
|
||||
request,
|
||||
"v1/chat/completions",
|
||||
integration_session,
|
||||
json.dumps({"model": "test-model"}).encode(),
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Helpers for driving ``routstr.proxy.proxy`` with mocked request and session."""
|
||||
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from routstr import proxy as proxy_module
|
||||
|
||||
|
||||
def mock_request_stream(request: MagicMock, body: bytes) -> None:
|
||||
"""Give a mocked request a readable body stream (the proxy reads the stream)."""
|
||||
|
||||
async def stream() -> AsyncIterator[bytes]:
|
||||
yield body
|
||||
|
||||
request.stream = stream
|
||||
|
||||
|
||||
def patch_proxy_session(session: Any) -> Any:
|
||||
"""Make the proxy route use ``session`` instead of opening its own."""
|
||||
|
||||
@asynccontextmanager
|
||||
async def factory() -> AsyncIterator[Any]:
|
||||
yield session
|
||||
|
||||
return patch.object(proxy_module, "create_session", factory)
|
||||
@@ -0,0 +1,140 @@
|
||||
"""Bounded request-body read: size cap, read timeout, and late DB session."""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi.responses import Response
|
||||
from starlette.requests import Request
|
||||
|
||||
from routstr import proxy as proxy_module
|
||||
from routstr.core.settings import settings
|
||||
|
||||
|
||||
def _make_request(headers: dict[str, str], chunks: list[bytes]) -> MagicMock:
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = headers
|
||||
request.state.request_id = "req-bounded-body"
|
||||
request.consumed = []
|
||||
|
||||
async def stream() -> AsyncIterator[bytes]:
|
||||
for chunk in chunks:
|
||||
request.consumed.append(chunk)
|
||||
yield chunk
|
||||
|
||||
request.stream = stream
|
||||
return request
|
||||
|
||||
|
||||
def _slow_request(delay: float) -> MagicMock:
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = {}
|
||||
request.state.request_id = "req-slow-body"
|
||||
|
||||
async def stream() -> AsyncIterator[bytes]:
|
||||
yield b"{"
|
||||
await asyncio.sleep(delay)
|
||||
yield b"}"
|
||||
|
||||
request.stream = stream
|
||||
return request
|
||||
|
||||
|
||||
async def _run(request: MagicMock) -> tuple[Any, MagicMock, AsyncMock]:
|
||||
"""Run the proxy route with the session factory and _proxy stubbed out."""
|
||||
session_factory = MagicMock()
|
||||
inner = AsyncMock(return_value=Response(status_code=200))
|
||||
with (
|
||||
patch.object(proxy_module, "create_session", session_factory),
|
||||
patch.object(proxy_module, "_proxy", inner),
|
||||
):
|
||||
response = await proxy_module.proxy(request, "v1/chat/completions")
|
||||
return response, session_factory, inner
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oversize_content_length_rejected_without_reading() -> None:
|
||||
request = _make_request({"content-length": "999999999"}, [b"x" * 16])
|
||||
|
||||
response, session_factory, inner = await _run(request)
|
||||
|
||||
assert response.status_code == 413
|
||||
assert request.consumed == []
|
||||
inner.assert_not_awaited()
|
||||
session_factory.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oversize_chunked_body_rejected_mid_stream() -> None:
|
||||
with patch.object(settings, "max_request_body_bytes", 8):
|
||||
request = _make_request({}, [b"1234", b"5678", b"9012", b"3456"])
|
||||
response, session_factory, inner = await _run(request)
|
||||
|
||||
assert response.status_code == 413
|
||||
# Reading stops as soon as the cap is exceeded; the last chunk is never read.
|
||||
assert request.consumed == [b"1234", b"5678", b"9012"]
|
||||
inner.assert_not_awaited()
|
||||
session_factory.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_slow_body_times_out() -> None:
|
||||
with patch.object(settings, "request_body_timeout_seconds", 0.05):
|
||||
request = _slow_request(delay=5)
|
||||
response, session_factory, inner = await _run(request)
|
||||
|
||||
assert response.status_code == 408
|
||||
inner.assert_not_awaited()
|
||||
session_factory.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_normal_request_reaches_proxy_with_body() -> None:
|
||||
body = b'{"model": "test-model"}'
|
||||
request = _make_request({"content-length": str(len(body))}, [body])
|
||||
|
||||
response, session_factory, inner = await _run(request)
|
||||
|
||||
assert response.status_code == 200
|
||||
session_factory.assert_called_once()
|
||||
inner.assert_awaited_once_with(request, "v1/chat/completions", ANY, body)
|
||||
|
||||
|
||||
def _starlette_request(body: bytes) -> Request:
|
||||
messages: list[dict[str, Any]] = [
|
||||
{"type": "http.request", "body": body, "more_body": False}
|
||||
]
|
||||
|
||||
async def receive() -> dict[str, Any]:
|
||||
return messages.pop(0) if messages else {"type": "http.disconnect"}
|
||||
|
||||
return Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"headers": [(b"content-length", str(len(body)).encode())],
|
||||
"path": "/v1/chat/completions",
|
||||
"query_string": b"",
|
||||
"state": {},
|
||||
},
|
||||
receive,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_body_stays_readable_after_bounded_read() -> None:
|
||||
"""EHBP forwarding and upstream passthrough re-read the same request."""
|
||||
body = b'{"model": "test-model"}'
|
||||
request = _starlette_request(body)
|
||||
|
||||
assert await proxy_module._read_bounded_body(request) == body
|
||||
|
||||
assert await request.body() == body
|
||||
streamed = bytearray()
|
||||
async for chunk in request.stream():
|
||||
streamed += chunk
|
||||
assert bytes(streamed) == body
|
||||
@@ -17,6 +17,8 @@ from routstr.core.error_scope import (
|
||||
)
|
||||
from routstr.upstream.model_paths import decode_model_path, encode_model_path
|
||||
|
||||
from .proxy_test_utils import mock_request_stream, patch_proxy_session
|
||||
|
||||
MODEL_ID = "test-model"
|
||||
|
||||
|
||||
@@ -38,7 +40,7 @@ def _make_request(headers: dict[str, str], body: bytes) -> MagicMock:
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = headers
|
||||
request.body = AsyncMock(return_value=body)
|
||||
mock_request_stream(request, body)
|
||||
request.state = MagicMock()
|
||||
request.state.request_id = "req-model-path"
|
||||
return request
|
||||
@@ -72,8 +74,9 @@ async def _run_proxy(
|
||||
proxy_module, "pay_for_request", AsyncMock(return_value=reservation)
|
||||
),
|
||||
patch.object(proxy_module, "revert_pay_for_request", AsyncMock()),
|
||||
patch_proxy_session(MagicMock()),
|
||||
):
|
||||
return await proxy_module.proxy(request, path, session=MagicMock())
|
||||
return await proxy_module.proxy(request, path)
|
||||
|
||||
|
||||
def test_decode_model_path_round_trips_encode() -> None:
|
||||
@@ -529,8 +532,9 @@ async def test_unsupported_endpoint_pins_fail_before_payment(
|
||||
patch.object(
|
||||
proxy_module, "get_candidates", return_value=[(MagicMock(), upstream)]
|
||||
),
|
||||
patch_proxy_session(MagicMock()),
|
||||
):
|
||||
response = await proxy_module.proxy(request, path, MagicMock())
|
||||
response = await proxy_module.proxy(request, path)
|
||||
assert response.status_code == 400
|
||||
assert json.loads(response.body)["error"]["type"] == "unsupported_request"
|
||||
payment.assert_not_called()
|
||||
|
||||
@@ -6,6 +6,8 @@ from fastapi.responses import StreamingResponse
|
||||
|
||||
from routstr import proxy as proxy_module
|
||||
|
||||
from .proxy_test_utils import mock_request_stream, patch_proxy_session
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_closes_request_session_before_returning_response() -> None:
|
||||
@@ -15,9 +17,11 @@ async def test_proxy_closes_request_session_before_returning_response() -> None:
|
||||
request.headers = {"accept": "application/json"}
|
||||
request.url.path = "/not-an-api-route"
|
||||
request.state.request_id = "test-request"
|
||||
mock_request_stream(request, b"")
|
||||
session = AsyncMock()
|
||||
|
||||
response = await proxy_module.proxy(request, "not-an-api-route", session=session)
|
||||
with patch_proxy_session(session):
|
||||
response = await proxy_module.proxy(request, "not-an-api-route")
|
||||
|
||||
assert response.status_code == 404
|
||||
session.close.assert_awaited_once()
|
||||
@@ -26,6 +30,8 @@ async def test_proxy_closes_request_session_before_returning_response() -> None:
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_session_is_closed_before_first_stream_chunk() -> None:
|
||||
request = MagicMock()
|
||||
request.headers = {}
|
||||
mock_request_stream(request, b"")
|
||||
session = AsyncMock()
|
||||
|
||||
async def stream() -> AsyncIterator[bytes]:
|
||||
@@ -33,10 +39,11 @@ async def test_proxy_session_is_closed_before_first_stream_chunk() -> None:
|
||||
yield b"chunk"
|
||||
|
||||
upstream_response = StreamingResponse(stream())
|
||||
with patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)):
|
||||
response = await proxy_module.proxy(
|
||||
request, "v1/chat/completions", session=session
|
||||
)
|
||||
with (
|
||||
patch("routstr.proxy._proxy", AsyncMock(return_value=upstream_response)),
|
||||
patch_proxy_session(session),
|
||||
):
|
||||
response = await proxy_module.proxy(request, "v1/chat/completions")
|
||||
|
||||
assert isinstance(response, StreamingResponse)
|
||||
chunks = [chunk async for chunk in response.body_iterator]
|
||||
|
||||
@@ -0,0 +1,191 @@
|
||||
"""Tests for stage timings, the duration header and skipped-path error logging."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI, HTTPException, Request
|
||||
from fastapi.responses import StreamingResponse
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from routstr.core.middleware import LoggingMiddleware, mark
|
||||
from routstr.core.settings import settings
|
||||
|
||||
|
||||
class _RecordingHandler(logging.Handler):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(level=logging.DEBUG)
|
||||
self.records: list[logging.LogRecord] = []
|
||||
|
||||
def emit(self, record: logging.LogRecord) -> None:
|
||||
self.records.append(record)
|
||||
|
||||
def completions(self) -> list[logging.LogRecord]:
|
||||
return [r for r in self.records if r.getMessage() == "Request completed"]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def records() -> Iterator[_RecordingHandler]:
|
||||
handler = _RecordingHandler()
|
||||
middleware_logger = logging.getLogger("routstr.core.middleware")
|
||||
middleware_logger.setLevel(logging.INFO)
|
||||
original_propagate = middleware_logger.propagate
|
||||
original_handlers = middleware_logger.handlers
|
||||
middleware_logger.propagate = False
|
||||
middleware_logger.handlers = [handler]
|
||||
try:
|
||||
yield handler
|
||||
finally:
|
||||
middleware_logger.handlers = original_handlers
|
||||
middleware_logger.propagate = original_propagate
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client() -> Iterator[TestClient]:
|
||||
app = FastAPI()
|
||||
app.add_middleware(LoggingMiddleware)
|
||||
|
||||
# /v1/wallet/info is in _SKIP_LOG_EXACT, so it exercises the suppression path.
|
||||
@app.get("/v1/wallet/info")
|
||||
async def wallet_info(fail: bool = False) -> dict[str, str]:
|
||||
if fail:
|
||||
raise HTTPException(status_code=400, detail="spent token")
|
||||
return {"status": "ok"}
|
||||
|
||||
@app.post("/v1/chat/completions")
|
||||
async def completions(request: Request) -> dict[str, str]:
|
||||
mark(request, "body_read")
|
||||
mark(request, "auth")
|
||||
return {"status": "ok"}
|
||||
|
||||
@app.post("/v1/chat/completions/stream")
|
||||
async def streamed(request: Request) -> StreamingResponse:
|
||||
mark(request, "body_read")
|
||||
|
||||
async def body() -> AsyncIterator[bytes]:
|
||||
yield b"data: one\n\n"
|
||||
await asyncio.sleep(0.05)
|
||||
yield b"data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(body(), media_type="text/event-stream")
|
||||
|
||||
@app.get("/admin/api/boom")
|
||||
async def boom() -> dict[str, str]:
|
||||
raise HTTPException(status_code=500, detail="boom")
|
||||
|
||||
with TestClient(app, raise_server_exceptions=False) as test_client:
|
||||
yield test_client
|
||||
|
||||
|
||||
def test_skipped_path_logs_4xx(client: TestClient, records: _RecordingHandler) -> None:
|
||||
assert client.get("/v1/wallet/info", params={"fail": True}).status_code == 400
|
||||
|
||||
completions = records.completions()
|
||||
assert len(completions) == 1
|
||||
record = completions[0]
|
||||
assert record.status_code == 400 # type: ignore[attr-defined]
|
||||
assert record.path == "/v1/wallet/info" # type: ignore[attr-defined]
|
||||
assert record.method == "GET" # type: ignore[attr-defined]
|
||||
assert record.duration_ms >= 0 # type: ignore[attr-defined]
|
||||
|
||||
|
||||
def test_skipped_path_does_not_log_2xx(
|
||||
client: TestClient, records: _RecordingHandler
|
||||
) -> None:
|
||||
assert client.get("/v1/wallet/info").status_code == 200
|
||||
|
||||
assert records.completions() == []
|
||||
|
||||
|
||||
def test_duration_header_present(client: TestClient) -> None:
|
||||
response = client.get("/v1/wallet/info")
|
||||
|
||||
assert float(response.headers["x-routstr-duration-ms"]) >= 0
|
||||
|
||||
|
||||
def test_stage_fields_on_completion_log(
|
||||
client: TestClient, records: _RecordingHandler
|
||||
) -> None:
|
||||
response = client.post("/v1/chat/completions", json={"model": "m"})
|
||||
assert response.status_code == 200
|
||||
|
||||
completions = records.completions()
|
||||
assert len(completions) == 1
|
||||
record = completions[0]
|
||||
assert record.body_read_ms >= 0 # type: ignore[attr-defined]
|
||||
assert record.auth_ms >= record.body_read_ms # type: ignore[attr-defined]
|
||||
assert record.content_length == int( # type: ignore[attr-defined]
|
||||
response.request.headers["content-length"]
|
||||
)
|
||||
|
||||
|
||||
def test_bogus_content_length_is_dropped(
|
||||
client: TestClient, records: _RecordingHandler
|
||||
) -> None:
|
||||
assert (
|
||||
client.get(
|
||||
"/v1/wallet/info",
|
||||
params={"fail": True},
|
||||
headers={"content-length": "not-a-number"},
|
||||
).status_code
|
||||
== 400
|
||||
)
|
||||
|
||||
assert records.completions()[0].content_length is None # type: ignore[attr-defined]
|
||||
|
||||
|
||||
def test_streamed_duration_covers_the_body(
|
||||
client: TestClient, records: _RecordingHandler
|
||||
) -> None:
|
||||
response = client.post("/v1/chat/completions/stream", json={"model": "m"})
|
||||
assert response.status_code == 200
|
||||
assert response.text.endswith("data: [DONE]\n\n")
|
||||
|
||||
completions = records.completions()
|
||||
assert len(completions) == 1
|
||||
record = completions[0]
|
||||
# The body sleeps 50ms, so a duration that stopped at the headers would be
|
||||
# well under it.
|
||||
assert record.duration_ms >= 50 # type: ignore[attr-defined]
|
||||
assert record.time_to_headers_ms < record.duration_ms # type: ignore[attr-defined]
|
||||
logged_request_id = record.request_id # type: ignore[attr-defined]
|
||||
assert logged_request_id == response.headers["x-routstr-request-id"]
|
||||
assert record.body_read_ms >= 0 # type: ignore[attr-defined]
|
||||
|
||||
|
||||
def test_slow_streamed_request_logs_warning(
|
||||
client: TestClient, records: _RecordingHandler, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "slow_request_warn_seconds", 0.02)
|
||||
|
||||
assert (
|
||||
client.post("/v1/chat/completions/stream", json={"model": "m"}).status_code
|
||||
== 200
|
||||
)
|
||||
|
||||
assert records.completions()[0].levelno == logging.WARNING
|
||||
|
||||
|
||||
def test_prefix_skipped_path_still_hides_client_errors(
|
||||
client: TestClient, records: _RecordingHandler
|
||||
) -> None:
|
||||
# /admin/api/* is polled on a timer, so an expired session must not turn
|
||||
# into one log line per poll; a 500 on the same prefix must still be logged.
|
||||
assert client.get("/admin/api/balances").status_code == 404
|
||||
assert records.completions() == []
|
||||
|
||||
assert client.get("/admin/api/boom").status_code == 500
|
||||
assert len(records.completions()) == 1
|
||||
|
||||
|
||||
def test_slow_request_logs_warning(
|
||||
client: TestClient, records: _RecordingHandler, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setattr(settings, "slow_request_warn_seconds", 0.0)
|
||||
|
||||
assert client.post("/v1/chat/completions", json={"model": "m"}).status_code == 200
|
||||
|
||||
completions = records.completions()
|
||||
assert len(completions) == 1
|
||||
assert completions[0].levelno == logging.WARNING
|
||||
@@ -29,6 +29,8 @@ from routstr.core.db import (
|
||||
reset_all_reserved_balances,
|
||||
)
|
||||
|
||||
from .proxy_test_utils import mock_request_stream, patch_proxy_session
|
||||
|
||||
|
||||
def _make_engine() -> AsyncEngine:
|
||||
return create_async_engine(
|
||||
@@ -387,7 +389,7 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = {"authorization": "Bearer sk-cancelkey"}
|
||||
request.body = AsyncMock(return_value=b'{"model": "test-model"}')
|
||||
mock_request_stream(request, b'{"model": "test-model"}')
|
||||
|
||||
upstream = MagicMock()
|
||||
upstream.provider_type = "test"
|
||||
@@ -420,8 +422,9 @@ async def test_proxy_reverts_reservation_on_client_disconnect() -> None:
|
||||
AsyncMock(return_value=reservation_snapshot),
|
||||
),
|
||||
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
|
||||
patch_proxy_session(session),
|
||||
):
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await proxy_module.proxy(request, "v1/chat/completions", session=session)
|
||||
await proxy_module.proxy(request, "v1/chat/completions")
|
||||
|
||||
revert_mock.assert_awaited_once_with(key, session, 1000, reservation_snapshot)
|
||||
|
||||
@@ -31,6 +31,8 @@ from routstr.upstream.tinfoil import (
|
||||
)
|
||||
from routstr.upstream.tinfoil_trailer import TrailerResponse
|
||||
|
||||
from .proxy_test_utils import patch_proxy_session
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# parse_tinfoil_usage_metrics
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -1303,10 +1305,9 @@ async def test_bearer_key_config_422_releases_reservation_and_passes_through() -
|
||||
"routstr.upstream.ehbp.forward_with_trailer",
|
||||
AsyncMock(return_value=upstream_resp),
|
||||
),
|
||||
patch_proxy_session(session),
|
||||
):
|
||||
response = await proxy_module.proxy(
|
||||
request, "v1/chat/completions", session=session
|
||||
)
|
||||
response = await proxy_module.proxy(request, "v1/chat/completions")
|
||||
|
||||
# The reservation was released despite the early passthrough return.
|
||||
revert_mock.assert_awaited_once_with(key, session, 1_000, reservation_snapshot)
|
||||
|
||||
@@ -28,6 +28,8 @@ from routstr.upstream.rate_limit import (
|
||||
classify_rate_limit,
|
||||
)
|
||||
|
||||
from .proxy_test_utils import mock_request_stream, patch_proxy_session
|
||||
|
||||
# The exact scenario from the issue, with a realistic (fake) org identifier.
|
||||
RAW_ORG_ID = "org-abc123XYZ456def"
|
||||
RATE_LIMIT_MESSAGE = (
|
||||
@@ -353,7 +355,7 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None:
|
||||
request = MagicMock()
|
||||
request.method = "POST"
|
||||
request.headers = {"authorization": "Bearer sk-rlkey"}
|
||||
request.body = AsyncMock(return_value=b'{"model": "test-model"}')
|
||||
mock_request_stream(request, b'{"model": "test-model"}')
|
||||
request.state = MagicMock()
|
||||
request.state.request_id = "req-rl"
|
||||
|
||||
@@ -400,10 +402,9 @@ async def test_proxy_loop_surfaces_rate_limit_and_reverts_once() -> None:
|
||||
AsyncMock(return_value=reservation),
|
||||
),
|
||||
patch.object(proxy_module, "revert_pay_for_request", revert_mock),
|
||||
patch_proxy_session(session),
|
||||
):
|
||||
response = await proxy_module.proxy(
|
||||
request, "v1/chat/completions", session=session
|
||||
)
|
||||
response = await proxy_module.proxy(request, "v1/chat/completions")
|
||||
|
||||
# Original 429 status and the stable code/details survive to the client.
|
||||
assert response.status_code == 429
|
||||
|
||||
Reference in New Issue
Block a user