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

365 lines
10 KiB
Python

"""
ⒸAngelaMos | 2025
limiter.py
"""
from __future__ import annotations
import asyncio
import logging
import functools
from typing import (
Any,
ParamSpec,
TypeVar,
TYPE_CHECKING,
)
from collections.abc import Callable
from starlette.requests import Request
from fastapi_420.algorithms import (
create_algorithm,
)
from fastapi_420.config import (
RateLimiterSettings,
get_settings,
)
from fastapi_420.exceptions import (
EnhanceYourCalm,
StorageConnectionError,
StorageError,
)
from fastapi_420.fingerprinting import (
CompositeFingerprinter,
)
from fastapi_420.storage import (
MemoryStorage,
RedisStorage,
create_storage,
)
from fastapi_420.types import (
Layer,
RateLimitKey,
RateLimitResult,
RateLimitRule,
)
if TYPE_CHECKING:
from fastapi_420.algorithms.base import BaseAlgorithm
from fastapi_420.storage import Storage
logger = logging.getLogger("fastapi_420")
P = ParamSpec("P")
R = TypeVar("R")
class RateLimiter:
"""
Main rate limiter class for FastAPI applications.
Usage:
limiter = RateLimiter()
@app.get("/api/data")
@limiter.limit("100/minute", "1000/hour")
async def get_data(request: Request):
return {"data": "value"}
"""
def __init__(
self,
settings: RateLimiterSettings | None = None,
storage: Storage | None = None,
) -> None:
self._settings = settings or get_settings()
self._storage = storage
self._fallback_storage: MemoryStorage | None = None
self._algorithm: BaseAlgorithm | None = None
self._fingerprinter: CompositeFingerprinter | None = None
self._initialized = False
self._lock = asyncio.Lock()
async def init(self) -> None:
"""
Initialize storage, algorithm, and fingerprinter
"""
async with self._lock:
if self._initialized:
return
if self._storage is None:
self._storage = create_storage(self._settings.storage)
if self._settings.storage.FALLBACK_TO_MEMORY:
self._fallback_storage = MemoryStorage.from_settings(
self._settings.storage
)
await self._fallback_storage.start_cleanup_task()
if isinstance(self._storage, RedisStorage):
try:
await self._storage.connect()
except StorageConnectionError:
if self._settings.FAIL_OPEN and self._fallback_storage:
logger.warning(
"Redis unavailable, using memory fallback",
extra = {
"redis_url":
self._settings.storage.REDIS_URL
},
)
self._storage = self._fallback_storage
self._fallback_storage = None
else:
raise
if isinstance(self._storage,
MemoryStorage
) and self._storage != self._fallback_storage:
await self._storage.start_cleanup_task()
self._algorithm = create_algorithm(self._settings.ALGORITHM)
self._fingerprinter = CompositeFingerprinter.from_settings(
self._settings.fingerprint
)
self._initialized = True
logger.info(
"Rate limiter initialized",
extra = {
"algorithm":
self._settings.ALGORITHM.value,
"storage":
self._storage.storage_type.value,
"fingerprint_level":
self._settings.fingerprint.LEVEL.value,
},
)
async def close(self) -> None:
"""
Close storage connections
"""
if self._storage:
await self._storage.close()
if self._fallback_storage:
await self._fallback_storage.close()
self._initialized = False
def limit(
self,
*rules: str,
key_func: Callable[[Request],
str] | None = None,
) -> Callable[[Callable[P,
R]],
Callable[P,
R]]:
"""
Decorator to apply rate limits to an endpoint
Args:
rules: Rate limit strings like "100/minute", "1000/hour"
key_func: Optional custom function to generate rate limit key
"""
parsed_rules = [RateLimitRule.parse(rule) for rule in rules]
def decorator(func: Callable[P, R]) -> Callable[P, R]:
@functools.wraps(func)
async def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
request = self._extract_request(args, kwargs)
if request is None:
return await func(*args, **kwargs) # type: ignore[misc, no-any-return]
await self._check_rate_limits(
request = request,
rules = parsed_rules,
key_func = key_func,
)
return await func(*args, **kwargs) # type: ignore[misc, no-any-return]
return wrapper # type: ignore[return-value]
return decorator
async def check(
self,
request: Request,
*rules: str,
key_func: Callable[[Request],
str] | None = None,
raise_on_limit: bool = True,
) -> RateLimitResult:
"""
Manually check rate limit without decorator
Returns the result of the strictest (most restrictive) rule
"""
parsed_rules = [RateLimitRule.parse(rule) for rule in rules]
if not parsed_rules:
parsed_rules = self._settings.get_default_rules()
return await self._check_rate_limits(
request = request,
rules = parsed_rules,
key_func = key_func,
raise_on_limit = raise_on_limit,
)
async def _check_rate_limits(
self,
request: Request,
rules: list[RateLimitRule],
key_func: Callable[[Request],
str] | None = None,
raise_on_limit: bool = True,
) -> RateLimitResult:
"""
Check all rules and return/raise for the most restrictive failure
"""
if not self._initialized:
await self.init()
storage = await self._get_active_storage()
if storage is None:
if self._settings.FAIL_OPEN:
return RateLimitResult(
allowed = True,
limit = 0,
remaining = 0,
reset_after = 0,
)
raise StorageError(operation = "check", backend = None)
fingerprint = await self._fingerprinter.extract(request) # type: ignore[union-attr]
endpoint = self._get_endpoint(request)
if key_func:
identifier = key_func(request)
else:
identifier = fingerprint.to_composite_key(
self._settings.fingerprint.LEVEL
)
worst_result: RateLimitResult | None = None
for rule in rules:
key = RateLimitKey(
prefix = self._settings.KEY_PREFIX,
version = self._settings.KEY_VERSION,
layer = Layer.USER,
endpoint = endpoint,
identifier = identifier,
window = rule.window_seconds,
).build()
result = await self._algorithm.check( # type: ignore[union-attr]
storage = storage,
key = key,
rule = rule,
)
if not result.allowed: # noqa: SIM102
if worst_result is None or result.retry_after > (worst_result.retry_after or 0): # type: ignore[operator]
worst_result = result
if worst_result is not None:
if self._settings.LOG_VIOLATIONS:
logger.warning(
"Rate limit exceeded",
extra = {
"endpoint": endpoint,
"identifier": identifier[: 16],
"remaining": worst_result.remaining,
"reset_after": worst_result.reset_after,
},
)
if raise_on_limit:
raise EnhanceYourCalm(
result = worst_result,
message = self._settings.HTTP_420_MESSAGE,
detail = self._settings.HTTP_420_DETAIL,
)
return worst_result
best_result = result
return best_result
async def _get_active_storage(self) -> Storage | None:
"""
Get active storage, falling back to memory if primary fails
"""
if self._storage is None:
return self._fallback_storage
try:
is_healthy = await self._storage.health_check()
if is_healthy:
return self._storage
except Exception: # noqa: S110
pass
if self._fallback_storage:
logger.warning(
"Primary storage unavailable, using memory fallback",
extra = {
"primary_storage": self._storage.storage_type.value
},
)
return self._fallback_storage
return None
def _extract_request(
self,
args: tuple[Any,
...],
kwargs: dict[str,
Any],
) -> Request | None:
"""
Extract Request object from function arguments
"""
for arg in args:
if isinstance(arg, Request):
return arg
for value in kwargs.values():
if isinstance(value, Request):
return value
return None
def _get_endpoint(self, request: Request) -> str:
"""
Get endpoint identifier from request
"""
route = request.scope.get("route")
if route:
return f"{request.method}:{route.path}"
return f"{request.method}:{request.url.path}"
@property
def settings(self) -> RateLimiterSettings:
"""
Get current settings
"""
return self._settings
@property
def is_initialized(self) -> bool:
"""
Check if limiter is initialized
"""
return self._initialized