Merge remote-tracking branch 'origin/main'

This commit is contained in:
thefux
2026-09-29 06:00:32 +00:00
12 changed files with 639 additions and 64 deletions
+162 -29
View File
@@ -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",
]
+11
View File
@@ -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
View File
@@ -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
+27
View File
@@ -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)
+140
View File
@@ -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
+7 -3
View File
@@ -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()
+12 -5
View File
@@ -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]
+191
View File
@@ -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
+5 -2
View File
@@ -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)
+4 -3
View File
@@ -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)
+5 -4
View File
@@ -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