mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
feat: log request stage timings and never suppress error responses
This commit is contained in:
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"]:
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user