Cybersecurity-Projects/PROJECTS/advanced/api-rate-limiter/src/fastapi_420/middleware.py

226 lines
6.6 KiB
Python

"""
ⒸAngelaMos | 2025
middleware.py
ASGI middleware for automatic rate limiting across all routes
Two middleware classes for different throttling strategies.
RateLimitMiddleware applies hard limits and returns HTTP 420
when exceeded, with support for path inclusion/exclusion lists
and path-specific limit overrides. SlowDownMiddleware takes a
softer approach, adding progressive delays to responses as
clients approach their limit instead of blocking them outright.
Key exports:
RateLimitMiddleware - hard blocking with HTTP 420 responses
SlowDownMiddleware - gradual throttling via response delays
Connects to:
exceptions.py - uses HTTP_420_ENHANCE_YOUR_CALM, EnhanceYourCalm
limiter.py - imports RateLimiter
"""
from __future__ import annotations
import re
import logging
import asyncio
from typing import (
TYPE_CHECKING,
)
from collections.abc import Callable
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request
from starlette.responses import (
JSONResponse,
Response,
)
from fastapi_420.exceptions import (
HTTP_420_ENHANCE_YOUR_CALM,
EnhanceYourCalm,
)
from fastapi_420.limiter import RateLimiter
if TYPE_CHECKING:
from starlette.types import ASGIApp
logger = logging.getLogger("fastapi_420")
class RateLimitMiddleware(BaseHTTPMiddleware):
"""
ASGI middleware for automatic rate limiting on all routes
Usage:
from fastapi import FastAPI
from fastapi_420.middleware import RateLimitMiddleware
from fastapi_420.limiter import RateLimiter
app = FastAPI()
limiter = RateLimiter()
app.add_middleware(
RateLimitMiddleware,
limiter=limiter,
default_limit="100/minute",
)
"""
def __init__(
self,
app: ASGIApp,
limiter: RateLimiter,
default_limit: str = "100/minute",
exclude_paths: list[str] | None = None,
exclude_patterns: list[str] | None = None,
include_paths: list[str] | None = None,
path_limits: dict[str,
str] | None = None,
key_func: Callable[[Request],
str] | None = None,
) -> None:
super().__init__(app)
self.limiter = limiter
self.default_limit = default_limit
self.exclude_paths = set(exclude_paths or [])
self.exclude_patterns = [
re.compile(p) for p in (exclude_patterns or [])
]
self.include_paths = set(include_paths) if include_paths else None
self.path_limits = path_limits or {}
self.key_func = key_func
self.exclude_paths.add("/health")
self.exclude_paths.add("/healthz")
self.exclude_paths.add("/ready")
self.exclude_paths.add("/metrics")
async def dispatch(self, request: Request, call_next: Callable) -> Response: # type: ignore[type-arg]
"""
Process request and apply rate limiting
"""
if not await self._should_limit(request):
return await call_next(request) # type: ignore[no-any-return]
limit = self._get_limit_for_path(request.url.path)
result = await self.limiter.check(
request,
limit,
key_func = self.key_func,
raise_on_limit = False,
)
if not result.allowed:
exc = EnhanceYourCalm(
result = result,
message = self.limiter.settings.HTTP_420_MESSAGE,
detail = self.limiter.settings.HTTP_420_DETAIL,
)
return self._create_420_response(exc)
response = await call_next(request)
if self.limiter.settings.INCLUDE_HEADERS:
for header_name, header_value in result.headers.items():
response.headers[header_name] = header_value
return response # type: ignore[no-any-return]
async def _should_limit(self, request: Request) -> bool:
"""
Determine if request should be rate limited
"""
path = request.url.path
if path in self.exclude_paths:
return False
for pattern in self.exclude_patterns:
if pattern.match(path):
return False
if self.include_paths is not None: # noqa: SIM102
if path not in self.include_paths:
for include_path in self.include_paths:
if path.startswith(include_path):
break
else:
return False
return True
def _get_limit_for_path(self, path: str) -> str:
"""
Get rate limit for specific path
"""
if path in self.path_limits:
return self.path_limits[path]
for pattern_path, limit in self.path_limits.items():
if path.startswith(pattern_path):
return limit
return self.default_limit
def _create_420_response(self, exc: EnhanceYourCalm) -> JSONResponse:
"""
Create HTTP 420 response
"""
headers = {}
if exc.result:
headers.update(exc.result.headers)
return JSONResponse(
status_code = HTTP_420_ENHANCE_YOUR_CALM,
content = exc.detail,
headers = headers,
)
class SlowDownMiddleware(BaseHTTPMiddleware):
"""
Alternative middleware that adds delays instead of blocking
Useful for gradual throttling rather than hard limits
"""
def __init__(
self,
app: ASGIApp,
limiter: RateLimiter,
threshold_limit: str = "50/minute",
max_delay_seconds: float = 5.0,
delay_increment: float = 0.5,
) -> None:
super().__init__(app)
self.limiter = limiter
self.threshold_limit = threshold_limit
self.max_delay_seconds = max_delay_seconds
self.delay_increment = delay_increment
async def dispatch(self, request: Request, call_next: Callable) -> Response: # type: ignore[type-arg]
"""
Process request with potential delays
"""
result = await self.limiter.check(
request,
self.threshold_limit,
raise_on_limit = False,
)
if result.remaining <= 0:
delay = min(
self.max_delay_seconds,
result.retry_after or self.delay_increment,
)
await asyncio.sleep(delay)
response = await call_next(request)
if self.limiter.settings.INCLUDE_HEADERS:
for header_name, header_value in result.headers.items():
response.headers[header_name] = header_value
return response # type: ignore[no-any-return]