from __future__ import annotations import time import uuid from typing import Any from fastapi import Request, Response from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint from app.core.logging_config import get_logger logger = get_logger(__name__) class AuditMiddleware(BaseHTTPMiddleware): """Middleware to log all API requests for audit purposes.""" EXCLUDED_PATHS = {"/api/v1/health", "/metrics", "/docs", "/openapi.json", "/redoc"} async def dispatch(self, request: Request, call_next: RequestResponseEndpoint) -> Response: if request.url.path in self.EXCLUDED_PATHS: return await call_next(request) request_id = str(uuid.uuid4()) start_time = time.monotonic() # Extract client info client_ip = request.client.host if request.client else "unknown" user_agent = request.headers.get("user-agent", "unknown") # Add request ID to request state request.state.request_id = request_id logger.info( "request_started", request_id=request_id, method=request.method, path=request.url.path, client_ip=client_ip, user_agent=user_agent[:200], ) try: response = await call_next(request) duration_ms = (time.monotonic() - start_time) * 1000 logger.info( "request_completed", request_id=request_id, method=request.method, path=request.url.path, status_code=response.status_code, duration_ms=round(duration_ms, 2), client_ip=client_ip, ) response.headers["X-Request-ID"] = request_id response.headers["X-Process-Time"] = f"{duration_ms:.2f}ms" return response except Exception as exc: duration_ms = (time.monotonic() - start_time) * 1000 logger.exception( "request_failed", request_id=request_id, method=request.method, path=request.url.path, duration_ms=round(duration_ms, 2), error=str(exc), ) raise