From 57c6dec506a166973b35ac9caaa70accea405a5e Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 27 Sep 2026 02:20:19 +0200 Subject: [PATCH 1/5] fix: bound proxy request body reads by size and time --- routstr/core/settings.py | 8 ++ routstr/proxy.py | 78 ++++++++++--- .../test_proxy_session_lifecycle.py | 6 +- tests/unit/proxy_test_utils.py | 27 +++++ tests/unit/test_bounded_request_body.py | 103 ++++++++++++++++++ tests/unit/test_model_path_routing.py | 10 +- tests/unit/test_proxy_session_lifecycle.py | 17 ++- tests/unit/test_stale_reservations.py | 7 +- tests/unit/test_tinfoil_integration.py | 7 +- tests/unit/test_upstream_rate_limit.py | 9 +- 10 files changed, 237 insertions(+), 35 deletions(-) create mode 100644 tests/unit/proxy_test_utils.py create mode 100644 tests/unit/test_bounded_request_body.py diff --git a/routstr/core/settings.py b/routstr/core/settings.py index da503a50..39448a58 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -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" diff --git a/routstr/proxy.py b/routstr/proxy.py index ae36d1d4..d9e7bc8d 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -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, @@ -418,23 +417,71 @@ 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 + return 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, + ) + 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 + + 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 +496,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 diff --git a/tests/integration/test_proxy_session_lifecycle.py b/tests/integration/test_proxy_session_lifecycle.py index 9a70f294..4233b653 100644 --- a/tests/integration/test_proxy_session_lifecycle.py +++ b/tests/integration/test_proxy_session_lifecycle.py @@ -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 diff --git a/tests/unit/proxy_test_utils.py b/tests/unit/proxy_test_utils.py new file mode 100644 index 00000000..a4fe4d3b --- /dev/null +++ b/tests/unit/proxy_test_utils.py @@ -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) diff --git a/tests/unit/test_bounded_request_body.py b/tests/unit/test_bounded_request_body.py new file mode 100644 index 00000000..e6c1e363 --- /dev/null +++ b/tests/unit/test_bounded_request_body.py @@ -0,0 +1,103 @@ +"""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 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) diff --git a/tests/unit/test_model_path_routing.py b/tests/unit/test_model_path_routing.py index 5b23720f..e2c3e52f 100644 --- a/tests/unit/test_model_path_routing.py +++ b/tests/unit/test_model_path_routing.py @@ -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() diff --git a/tests/unit/test_proxy_session_lifecycle.py b/tests/unit/test_proxy_session_lifecycle.py index 5d0416d5..7cc01b03 100644 --- a/tests/unit/test_proxy_session_lifecycle.py +++ b/tests/unit/test_proxy_session_lifecycle.py @@ -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] diff --git a/tests/unit/test_stale_reservations.py b/tests/unit/test_stale_reservations.py index 2b5d8e31..558fdda5 100644 --- a/tests/unit/test_stale_reservations.py +++ b/tests/unit/test_stale_reservations.py @@ -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) diff --git a/tests/unit/test_tinfoil_integration.py b/tests/unit/test_tinfoil_integration.py index ec5e8c11..14715f58 100644 --- a/tests/unit/test_tinfoil_integration.py +++ b/tests/unit/test_tinfoil_integration.py @@ -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) diff --git a/tests/unit/test_upstream_rate_limit.py b/tests/unit/test_upstream_rate_limit.py index c111b74a..0b73199c 100644 --- a/tests/unit/test_upstream_rate_limit.py +++ b/tests/unit/test_upstream_rate_limit.py @@ -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 From f3101a015808a099872ff22d996f5515733bb871 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 28 Sep 2026 09:21:58 +0200 Subject: [PATCH 2/5] fix streaming --- routstr/proxy.py | 8 +++++- tests/unit/test_bounded_request_body.py | 37 +++++++++++++++++++++++++ 2 files changed, 44 insertions(+), 1 deletion(-) diff --git a/routstr/proxy.py b/routstr/proxy.py index d9e7bc8d..ce50f528 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -444,7 +444,7 @@ async def _read_bounded_body(request: Request) -> bytes | Response: return bytes(body) try: - return await asyncio.wait_for(read(), timeout) + body = await asyncio.wait_for(read(), timeout) except _BodyLimitExceeded: error_type, message, status = ( "invalid_request", @@ -457,6 +457,12 @@ async def _read_bounded_body(request: Request) -> bytes | Response: 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) diff --git a/tests/unit/test_bounded_request_body.py b/tests/unit/test_bounded_request_body.py index e6c1e363..8b6cd06b 100644 --- a/tests/unit/test_bounded_request_body.py +++ b/tests/unit/test_bounded_request_body.py @@ -7,6 +7,7 @@ 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 @@ -101,3 +102,39 @@ async def test_normal_request_reaches_proxy_with_body() -> None: 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 From c0a9a0e4b9e501d2944566e1b62907180681c46f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 27 Sep 2026 03:10:10 +0200 Subject: [PATCH 3/5] feat: log request stage timings and never suppress error responses --- routstr/core/middleware.py | 36 +++++++- routstr/core/settings.py | 3 + routstr/proxy.py | 3 + tests/unit/test_request_stage_timing.py | 118 ++++++++++++++++++++++++ 4 files changed, 156 insertions(+), 4 deletions(-) create mode 100644 tests/unit/test_request_stage_timing.py diff --git a/routstr/core/middleware.py b/routstr/core/middleware.py index 65708d54..58351239 100644 --- a/routstr/core/middleware.py +++ b/routstr/core/middleware.py @@ -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,9 +87,13 @@ _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 + # A 4xx/5xx storm on a suppressed path is exactly what we need to see, so + # the path filters below only ever hide successful responses. + if status_code is not None and status_code >= 400: + return True if path in _SKIP_LOG_EXACT: return False return not any(path.startswith(prefix) for prefix in _SKIP_LOG_PREFIXES) @@ -103,6 +108,13 @@ 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() + + class LoggingMiddleware(BaseHTTPMiddleware): """Middleware to log proxy interactions and page navigation. @@ -126,6 +138,9 @@ class LoggingMiddleware(BaseHTTPMiddleware): # Start timing start_time = time.time() + stage_start = time.monotonic() + stage_marks: dict[str, float] = {} + request.state.stage_marks = stage_marks if should_log: logger.info( @@ -144,28 +159,40 @@ class LoggingMiddleware(BaseHTTPMiddleware): try: response = await call_next(request) - if should_log: - duration = time.time() - start_time + duration = time.time() - start_time + + if _should_log(request.method, path, response.status_code): extra: dict[str, object] = { "request_id": request_id, "method": request.method, "path": path, "status_code": response.status_code, "duration_ms": round(duration * 1000, 2), + "content_length": request.headers.get("content-length"), **_attribution(request), } + for name, marked_at in stage_marks.items(): + extra[f"{name}_ms"] = round((marked_at - stage_start) * 1000, 2) 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( + log = ( + logger.warning + if duration > settings.slow_request_warn_seconds + else logger.info + ) + log( "Request completed", extra=extra, ) if hasattr(response, "headers"): response.headers["x-routstr-request-id"] = request_id + response.headers["x-routstr-duration-ms"] = str( + round(duration * 1000, 2) + ) return response @@ -196,5 +223,6 @@ __all__ = [ "LoggingMiddleware", "UNKNOWN_CLIENT_APP", "client_app_context", + "mark", "request_id_context", ] diff --git a/routstr/core/settings.py b/routstr/core/settings.py index 39448a58..5dd2d57d 100644 --- a/routstr/core/settings.py +++ b/routstr/core/settings.py @@ -197,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( diff --git a/routstr/proxy.py b/routstr/proxy.py index ce50f528..38e2dfdc 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -28,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 ( @@ -474,6 +475,7 @@ async def proxy(request: Request, path: str) -> Response | StreamingResponse: 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: @@ -777,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"]: diff --git a/tests/unit/test_request_stage_timing.py b/tests/unit/test_request_stage_timing.py new file mode 100644 index 00000000..98c4a782 --- /dev/null +++ b/tests/unit/test_request_stage_timing.py @@ -0,0 +1,118 @@ +"""Tests for stage timings, the duration header and skipped-path error logging.""" + +import logging +from collections.abc import Iterator + +import pytest +from fastapi import FastAPI, HTTPException, Request +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"} + + 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 # type: ignore[attr-defined] + == response.request.headers["content-length"] + ) + + +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 From 1e60cbea54c1e970b1bffa909be8b2e9d1eb54d1 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 29 Sep 2026 00:45:25 +0200 Subject: [PATCH 4/5] fix: measure request duration across streamed bodies and keep prefix-skipped client errors suppressed --- routstr/core/middleware.py | 185 +++++++++++++++++++----- tests/unit/test_request_stage_timing.py | 86 ++++++++++- 2 files changed, 224 insertions(+), 47 deletions(-) diff --git a/routstr/core/middleware.py b/routstr/core/middleware.py index 58351239..e158b590 100644 --- a/routstr/core/middleware.py +++ b/routstr/core/middleware.py @@ -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 @@ -90,12 +90,15 @@ _SKIP_LOG_EXACT: frozenset[str] = frozenset( def _should_log(method: str, path: str, status_code: int | None = None) -> bool: if method in _SKIP_LOG_METHODS: return False - # A 4xx/5xx storm on a suppressed path is exactly what we need to see, so - # the path filters below only ever hide successful responses. - if status_code is not None and status_code >= 400: + # 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) @@ -115,12 +118,109 @@ def mark(request: Request, name: str) -> 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()) @@ -129,15 +229,14 @@ 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 @@ -159,46 +258,52 @@ class LoggingMiddleware(BaseHTTPMiddleware): try: response = await call_next(request) - duration = time.time() - start_time + headers_duration = time.monotonic() - stage_start - if _should_log(request.method, path, response.status_code): - extra: dict[str, object] = { - "request_id": request_id, - "method": request.method, - "path": path, - "status_code": response.status_code, - "duration_ms": round(duration * 1000, 2), - "content_length": request.headers.get("content-length"), - **_attribution(request), - } - for name, marked_at in stage_marks.items(): - extra[f"{name}_ms"] = round((marked_at - stage_start) * 1000, 2) - 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") - log = ( - logger.warning - if duration > settings.slow_request_warn_seconds - else logger.info - ) - log( - "Request completed", - extra=extra, - ) 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(duration * 1000, 2) + 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={ diff --git a/tests/unit/test_request_stage_timing.py b/tests/unit/test_request_stage_timing.py index 98c4a782..578388aa 100644 --- a/tests/unit/test_request_stage_timing.py +++ b/tests/unit/test_request_stage_timing.py @@ -1,10 +1,12 @@ """Tests for stage timings, the duration header and skipped-path error logging.""" +import asyncio import logging -from collections.abc import Iterator +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 @@ -57,13 +59,26 @@ def client() -> Iterator[TestClient]: 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: +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() @@ -100,12 +115,69 @@ def test_stage_fields_on_completion_log( 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 # type: ignore[attr-defined] - == response.request.headers["content-length"] + 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] + assert record.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: From fb763c35119d7c44c7c007e73b6d9ca4d6867053 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 29 Sep 2026 00:45:25 +0200 Subject: [PATCH 5/5] fix: measure request duration across streamed bodies and keep prefix-skipped client errors suppressed --- routstr/core/middleware.py | 185 +++++++++++++++++++----- tests/unit/test_request_stage_timing.py | 87 ++++++++++- 2 files changed, 225 insertions(+), 47 deletions(-) diff --git a/routstr/core/middleware.py b/routstr/core/middleware.py index 58351239..e158b590 100644 --- a/routstr/core/middleware.py +++ b/routstr/core/middleware.py @@ -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 @@ -90,12 +90,15 @@ _SKIP_LOG_EXACT: frozenset[str] = frozenset( def _should_log(method: str, path: str, status_code: int | None = None) -> bool: if method in _SKIP_LOG_METHODS: return False - # A 4xx/5xx storm on a suppressed path is exactly what we need to see, so - # the path filters below only ever hide successful responses. - if status_code is not None and status_code >= 400: + # 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) @@ -115,12 +118,109 @@ def mark(request: Request, name: str) -> 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()) @@ -129,15 +229,14 @@ 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 @@ -159,46 +258,52 @@ class LoggingMiddleware(BaseHTTPMiddleware): try: response = await call_next(request) - duration = time.time() - start_time + headers_duration = time.monotonic() - stage_start - if _should_log(request.method, path, response.status_code): - extra: dict[str, object] = { - "request_id": request_id, - "method": request.method, - "path": path, - "status_code": response.status_code, - "duration_ms": round(duration * 1000, 2), - "content_length": request.headers.get("content-length"), - **_attribution(request), - } - for name, marked_at in stage_marks.items(): - extra[f"{name}_ms"] = round((marked_at - stage_start) * 1000, 2) - 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") - log = ( - logger.warning - if duration > settings.slow_request_warn_seconds - else logger.info - ) - log( - "Request completed", - extra=extra, - ) 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(duration * 1000, 2) + 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={ diff --git a/tests/unit/test_request_stage_timing.py b/tests/unit/test_request_stage_timing.py index 98c4a782..38845c3a 100644 --- a/tests/unit/test_request_stage_timing.py +++ b/tests/unit/test_request_stage_timing.py @@ -1,10 +1,12 @@ """Tests for stage timings, the duration header and skipped-path error logging.""" +import asyncio import logging -from collections.abc import Iterator +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 @@ -57,13 +59,26 @@ def client() -> Iterator[TestClient]: 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: +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() @@ -100,12 +115,70 @@ def test_stage_fields_on_completion_log( 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 # type: ignore[attr-defined] - == response.request.headers["content-length"] + 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: