91 lines
2.8 KiB
Python
91 lines
2.8 KiB
Python
"""
|
|
Middleware de rate limiting — par client et par plan.
|
|
|
|
Utilise slowapi (basé sur limits) avec identification par client_id
|
|
plutôt que par IP, pour compter les requêtes par client API.
|
|
|
|
Limites par plan (par heure) :
|
|
- free : 20 uploads, 50 AI
|
|
- standard : 100 uploads, 200 AI
|
|
- premium : 500 uploads, 1000 AI
|
|
|
|
Supporte le stockage Redis pour la persistence des compteurs.
|
|
"""
|
|
import logging
|
|
from contextvars import ContextVar
|
|
|
|
from slowapi import Limiter
|
|
from slowapi.util import get_remote_address
|
|
from starlette.requests import Request
|
|
|
|
from app.config import settings
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# ContextVar pour transmettre le plan du client au rate limiter dynamique
|
|
_current_client_plan: ContextVar[str] = ContextVar("current_client_plan", default="free")
|
|
|
|
|
|
def _get_client_id_from_request(request: Request) -> str:
|
|
"""
|
|
Extrait le client_id depuis la state de la requête.
|
|
Encode le plan pour isoler les buckets par plan.
|
|
Fallback vers l'IP si le client n'est pas encore authentifié.
|
|
"""
|
|
client_id = getattr(request.state, "client_id", None)
|
|
plan = getattr(request.state, "client_plan", "free")
|
|
if client_id:
|
|
return f"{client_id}:{plan}"
|
|
return get_remote_address(request)
|
|
|
|
|
|
# Stockage Redis pour les compteurs (si configuré), sinon mémoire
|
|
if settings.RATE_LIMIT_STORAGE_URL:
|
|
limiter = Limiter(
|
|
key_func=_get_client_id_from_request,
|
|
storage_uri=settings.RATE_LIMIT_STORAGE_URL,
|
|
)
|
|
else:
|
|
limiter = Limiter(key_func=_get_client_id_from_request)
|
|
|
|
|
|
def get_upload_rate_limit(plan: str) -> str:
|
|
"""Retourne la limite de rate pour les uploads selon le plan."""
|
|
limits = {
|
|
"free": f"{settings.RATE_LIMIT_FREE_UPLOAD}/hour",
|
|
"standard": f"{settings.RATE_LIMIT_STANDARD_UPLOAD}/hour",
|
|
"premium": f"{settings.RATE_LIMIT_PREMIUM_UPLOAD}/hour",
|
|
}
|
|
return limits.get(plan, limits["free"])
|
|
|
|
|
|
def get_ai_rate_limit(plan: str) -> str:
|
|
"""Retourne la limite de rate pour les endpoints AI selon le plan."""
|
|
limits = {
|
|
"free": f"{settings.RATE_LIMIT_FREE_AI}/hour",
|
|
"standard": f"{settings.RATE_LIMIT_STANDARD_AI}/hour",
|
|
"premium": f"{settings.RATE_LIMIT_PREMIUM_AI}/hour",
|
|
}
|
|
return limits.get(plan, limits["free"])
|
|
|
|
|
|
def dynamic_upload_limit() -> str:
|
|
"""Callable pour slowapi — lit le plan depuis le ContextVar."""
|
|
plan = _current_client_plan.get()
|
|
return get_upload_rate_limit(plan)
|
|
|
|
|
|
def dynamic_ai_limit() -> str:
|
|
"""Callable pour slowapi — lit le plan depuis le ContextVar."""
|
|
plan = _current_client_plan.get()
|
|
return get_ai_rate_limit(plan)
|
|
|
|
|
|
# Aliases pour compatibilité avec les tests existants
|
|
def upload_rate_limit_key(request) -> str:
|
|
return _get_client_id_from_request(request)
|
|
|
|
|
|
def ai_rate_limit_key(request) -> str:
|
|
return _get_client_id_from_request(request)
|