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: