mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: measure request duration across streamed bodies and keep prefix-skipped client errors suppressed
This commit is contained in:
+145
-40
@@ -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={
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user