fix: measure request duration across streamed bodies and keep prefix-skipped client errors suppressed

This commit is contained in:
9qeklajc
2026-09-29 01:04:08 +02:00
parent c0a9a0e4b9
commit fb763c3511
2 changed files with 225 additions and 47 deletions
+145 -40
View File
@@ -1,7 +1,7 @@
import time import time
import uuid import uuid
from contextvars import ContextVar from contextvars import ContextVar
from typing import Callable from typing import AsyncIterator, Callable
from urllib.parse import urlsplit from urllib.parse import urlsplit
from fastapi import Request, Response 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: def _should_log(method: str, path: str, status_code: int | None = None) -> bool:
if method in _SKIP_LOG_METHODS: if method in _SKIP_LOG_METHODS:
return False return False
# A 4xx/5xx storm on a suppressed path is exactly what we need to see, so # Our own faults are never noise, whatever the path.
# the path filters below only ever hide successful responses. if status_code is not None and status_code >= 500:
if status_code is not None and status_code >= 400:
return True return True
if path in _SKIP_LOG_EXACT: 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) 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() 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): class LoggingMiddleware(BaseHTTPMiddleware):
"""Middleware to log proxy interactions and page navigation. """Middleware to log proxy interactions and page navigation.
Skips logging for static assets and Next.js chunks to avoid noise. 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: async def dispatch(self, request: Request, call_next: Callable) -> Response:
# Generate request ID # Generate request ID
request_id = str(uuid.uuid4()) request_id = str(uuid.uuid4())
@@ -129,15 +229,14 @@ class LoggingMiddleware(BaseHTTPMiddleware):
# Set request ID in context for logging # Set request ID in context for logging
token = request_id_context.set(request_id) token = request_id_context.set(request_id)
client_app_token = client_app_context.set( client_app = client_app_from_headers(request.headers)
client_app_from_headers(request.headers) client_app_token = client_app_context.set(client_app)
)
path = request.url.path path = request.url.path
should_log = _should_log(request.method, path) should_log = _should_log(request.method, path)
# Start timing # Start timing. Monotonic throughout: a wall-clock step would otherwise
start_time = time.time() # produce negative durations and bogus slow-request warnings.
stage_start = time.monotonic() stage_start = time.monotonic()
stage_marks: dict[str, float] = {} stage_marks: dict[str, float] = {}
request.state.stage_marks = stage_marks request.state.stage_marks = stage_marks
@@ -159,46 +258,52 @@ class LoggingMiddleware(BaseHTTPMiddleware):
try: try:
response = await call_next(request) 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"): if hasattr(response, "headers"):
response.headers["x-routstr-request-id"] = request_id 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( 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 return response
except Exception as e: except Exception as e:
# Always log failures, even for skipped paths, so we don't lose errors. # 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( logger.error(
"Request failed", "Request failed",
extra={ extra={
+80 -7
View File
@@ -1,10 +1,12 @@
"""Tests for stage timings, the duration header and skipped-path error logging.""" """Tests for stage timings, the duration header and skipped-path error logging."""
import asyncio
import logging import logging
from collections.abc import Iterator from collections.abc import AsyncIterator, Iterator
import pytest import pytest
from fastapi import FastAPI, HTTPException, Request from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import StreamingResponse
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
from routstr.core.middleware import LoggingMiddleware, mark from routstr.core.middleware import LoggingMiddleware, mark
@@ -57,13 +59,26 @@ def client() -> Iterator[TestClient]:
mark(request, "auth") mark(request, "auth")
return {"status": "ok"} 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: with TestClient(app, raise_server_exceptions=False) as test_client:
yield test_client yield test_client
def test_skipped_path_logs_4xx( def test_skipped_path_logs_4xx(client: TestClient, records: _RecordingHandler) -> None:
client: TestClient, records: _RecordingHandler
) -> None:
assert client.get("/v1/wallet/info", params={"fail": True}).status_code == 400 assert client.get("/v1/wallet/info", params={"fail": True}).status_code == 400
completions = records.completions() completions = records.completions()
@@ -100,12 +115,70 @@ def test_stage_fields_on_completion_log(
record = completions[0] record = completions[0]
assert record.body_read_ms >= 0 # type: ignore[attr-defined] assert record.body_read_ms >= 0 # type: ignore[attr-defined]
assert record.auth_ms >= record.body_read_ms # type: ignore[attr-defined] assert record.auth_ms >= record.body_read_ms # type: ignore[attr-defined]
assert ( assert record.content_length == int( # type: ignore[attr-defined]
record.content_length # type: ignore[attr-defined] response.request.headers["content-length"]
== 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( def test_slow_request_logs_warning(
client: TestClient, records: _RecordingHandler, monkeypatch: pytest.MonkeyPatch client: TestClient, records: _RecordingHandler, monkeypatch: pytest.MonkeyPatch
) -> None: ) -> None: