99 lines
2.9 KiB
Python
99 lines
2.9 KiB
Python
"""
|
|
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},
|
|
)
|