""" Middleware for AARS backend: - RequestIDMiddleware: generate/accept X-Request-ID - AccessLogMiddleware: log every HTTP request - error_handler: catch unhandled exceptions, log with stack trace """ import logging import re import time import traceback import uuid from fastapi import Request from fastapi.responses import JSONResponse from starlette.middleware.base import BaseHTTPMiddleware from starlette.types import ASGIApp from logging_config import ( get_access_logger, set_request_id, set_user_id, ) logger = logging.getLogger(__name__) REQUEST_ID_PATTERN = re.compile(r"^[A-Za-z0-9_-]{1,128}$") def _normalize_request_id(raw: str | None) -> str: if raw and REQUEST_ID_PATTERN.match(raw): return raw return str(uuid.uuid4()) class RequestIDMiddleware(BaseHTTPMiddleware): """Generate or accept X-Request-ID and echo in response headers.""" async def dispatch(self, request: Request, call_next): rid = _normalize_request_id(request.headers.get("X-Request-ID")) set_request_id(rid) request.state.request_id = rid response = await call_next(request) response.headers["X-Request-ID"] = rid return response class AccessLogMiddleware(BaseHTTPMiddleware): """Log every HTTP request with route, method, status, duration.""" async def dispatch(self, request: Request, call_next): start = time.perf_counter() rid = getattr(request.state, "request_id", "-") status = 500 # default for exception path try: response = await call_next(request) status = response.status_code return response finally: duration_ms = int((time.perf_counter() - start) * 1000) access = get_access_logger() access.info( "%s %s %d", request.method, request.url.path, status, extra={ "route": request.url.path, "method": request.method, "status": status, "duration_ms": duration_ms, }, ) async def unhandled_exception_handler(request: Request, exc: Exception): """Catch-all for unhandled exceptions. Log with stack + return 500 JSON.""" rid = getattr(request.state, "request_id", "-") logger.error( "unhandled exception: %s %s", request.method, request.url.path, exc_info=exc, extra={ "route": request.url.path, "method": request.method, "status": 500, "context": { "exception_type": type(exc).__name__, "query_params": dict(request.query_params), }, }, ) return JSONResponse( status_code=500, content={"detail": "Internal server error", "request_id": rid}, headers={"X-Request-ID": rid}, )