56 lines
1.8 KiB
Python
56 lines
1.8 KiB
Python
"""HTTP middleware registration."""
|
|
|
|
from uuid import uuid4
|
|
from time import perf_counter
|
|
|
|
from fastapi import FastAPI, Request
|
|
from starlette.middleware.base import BaseHTTPMiddleware
|
|
from starlette.responses import Response
|
|
from structlog.contextvars import bind_contextvars, clear_contextvars
|
|
|
|
from app.observability.metrics import get_metrics
|
|
|
|
|
|
class RequestCorrelationMiddleware(BaseHTTPMiddleware):
|
|
"""Assign and propagate request correlation identifiers."""
|
|
|
|
def __init__(self, app: FastAPI, *, request_id_header: str) -> None:
|
|
super().__init__(app)
|
|
self._request_id_header = request_id_header
|
|
|
|
async def dispatch(self, request: Request, call_next) -> Response:
|
|
request_id = str(uuid4())
|
|
request.state.request_id = request_id
|
|
bind_contextvars(request_id=request_id)
|
|
|
|
started_at = perf_counter()
|
|
status_code = 500
|
|
is_error = True
|
|
|
|
try:
|
|
response = await call_next(request)
|
|
status_code = response.status_code
|
|
is_error = status_code >= 500
|
|
response.headers[self._request_id_header] = request_id
|
|
return response
|
|
finally:
|
|
duration_ms = (perf_counter() - started_at) * 1000
|
|
get_metrics().record_http_request(
|
|
method=request.method,
|
|
route=request.url.path,
|
|
status_code=status_code,
|
|
duration_ms=duration_ms,
|
|
is_error=is_error,
|
|
)
|
|
clear_contextvars()
|
|
|
|
|
|
def install_middleware(app: FastAPI) -> None:
|
|
"""Register middleware components for the API."""
|
|
|
|
settings = app.state.settings
|
|
app.add_middleware(
|
|
RequestCorrelationMiddleware,
|
|
request_id_header=settings.controller.request_id_header,
|
|
)
|