File size: 3,877 Bytes
aac350d
23d337e
 
aac350d
 
 
 
 
23d337e
aac350d
 
 
23d337e
aac350d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
23d337e
aac350d
23d337e
 
 
 
 
 
aac350d
 
 
 
 
23d337e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
"""
API middleware — CORS, request-id, rate limiting, request body size limit,
global exception handler.
"""

from __future__ import annotations

import time
import uuid
from collections import defaultdict

from fastapi import Request, Response
from fastapi.responses import JSONResponse
from starlette.middleware.base import BaseHTTPMiddleware


class RequestContextMiddleware(BaseHTTPMiddleware):
    """Adds a request-id to every request + measures duration."""

    async def dispatch(self, request: Request, call_next):
        request_id = request.headers.get("X-Request-ID") or uuid.uuid4().hex[:12]
        request.state.request_id = request_id
        t0 = time.perf_counter()
        response: Response = await call_next(request)
        elapsed_ms = (time.perf_counter() - t0) * 1000.0
        response.headers["X-Request-ID"] = request_id
        response.headers["X-Response-Time-ms"] = f"{elapsed_ms:.2f}"
        return response


class RateLimitMiddleware(BaseHTTPMiddleware):
    """Simple in-memory per-IP rate limiter.

    For production, replace with a Redis-backed limiter.
    """

    def __init__(self, app, requests_per_minute: int = 30):
        super().__init__(app)
        self._limit = requests_per_minute
        self._buckets: dict[str, list[float]] = defaultdict(list)

    async def dispatch(self, request: Request, call_next):
        if request.url.path.startswith("/health"):
            return await call_next(request)
        client = request.client.host if request.client else "unknown"
        now = time.time()
        window = 60.0
        recent = [t for t in self._buckets[client] if now - t < window]
        if len(recent) >= self._limit:
            return JSONResponse(
                status_code=429,
                content={
                    "success": False,
                    "error": "rate limit exceeded",
                    "error_type": "RateLimitError",
                    "request_id": getattr(request.state, "request_id", None),
                },
                media_type="application/json",
            )
        recent.append(now)
        self._buckets[client] = recent
        return await call_next(request)


class RequestSizeLimitMiddleware(BaseHTTPMiddleware):
    """Rejects request bodies larger than max_bytes."""

    def __init__(self, app, max_bytes: int = 25 * 1024 * 1024):
        super().__init__(app)
        self._max = max_bytes

    async def dispatch(self, request: Request, call_next):
        cl = request.headers.get("content-length")
        if cl and cl.isdigit() and int(cl) > self._max:
            return JSONResponse(
                status_code=413,
                content={
                    "success": False,
                    "error": f"Request body exceeds {self._max} bytes",
                    "error_type": "PayloadTooLarge",
                    "request_id": getattr(request.state, "request_id", None),
                },
                media_type="application/json",
            )
        return await call_next(request)


class GlobalExceptionMiddleware(BaseHTTPMiddleware):
    """Catches all uncaught exceptions and returns a structured error response.

    Prevents stack traces from leaking to clients.
    """

    async def dispatch(self, request: Request, call_next):
        try:
            return await call_next(request)
        except Exception as e:
            from loguru import logger
            logger.exception(f"Unhandled exception on {request.url.path}")
            return JSONResponse(
                status_code=500,
                content={
                    "success": False,
                    "error": str(e),
                    "error_type": type(e).__name__,
                    "request_id": getattr(request.state, "request_id", None),
                },
                media_type="application/json",
            )