From c0a9a0e4b9e501d2944566e1b62907180681c46f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 27 Sep 2026 03:10:10 +0200 Subject: [PATCH] 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