feat: log request stage timings and never suppress error responses

This commit is contained in:
9qeklajc
2026-09-29 00:56:10 +02:00
parent 9207658978
commit c0a9a0e4b9
4 changed files with 156 additions and 4 deletions
+32 -4
View File
@@ -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",
]
+3
View File
@@ -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(
+3
View File
@@ -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"]:
+118
View File
@@ -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