complianceos / api /middleware.py
bhriguverma's picture
Upload api/middleware.py with huggingface_hub
2a001cf verified
Raw
History Blame Contribute Delete
5.32 kB
"""
api/middleware.py — Rate limiting middleware and tenant context helpers.
Rate limiter: sliding-window, per-API-key, 100 req/60s default.
Tenant context: extracted from User object (set by auth dependency).
"""
from __future__ import annotations
import logging
import time
from collections import defaultdict, deque
from typing import TYPE_CHECKING
from fastapi import Request, Response
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.responses import JSONResponse
if TYPE_CHECKING:
from api.auth import User
log = logging.getLogger(__name__)
# ─── Rate limiter ─────────────────────────────────────────────────────────────
class RateLimiter:
"""
Sliding-window per-key rate limiter (in-memory).
Thread-safe for CPython GIL; for multi-process deployment use Redis instead.
"""
def __init__(self, max_requests: int = 100, window_seconds: int = 60) -> None:
self._requests: dict[str, deque] = defaultdict(deque)
self.max_requests = max_requests
self.window_seconds = window_seconds
def is_allowed(self, key: str) -> bool:
"""
Check if key is within rate limit.
Returns True and records the request if allowed.
Returns False (and does NOT record) if the limit is exceeded.
"""
now = time.time()
window = self._requests[key]
# Evict expired timestamps
cutoff = now - self.window_seconds
while window and window[0] < cutoff:
window.popleft()
if len(window) >= self.max_requests:
return False
window.append(now)
return True
def remaining(self, key: str) -> int:
"""Return how many requests are left in the current window."""
now = time.time()
window = self._requests[key]
cutoff = now - self.window_seconds
# Count non-expired
active = sum(1 for t in window if t >= cutoff)
return max(0, self.max_requests - active)
def reset_in(self, key: str) -> float:
"""Return seconds until the oldest request expires (i.e. window resets)."""
window = self._requests.get(key)
if not window:
return 0.0
oldest = window[0]
return max(0.0, (oldest + self.window_seconds) - time.time())
def reset_key(self, key: str) -> None:
"""Clear the rate limit state for a key (admin action)."""
self._requests.pop(key, None)
# Singleton used by the middleware and can be imported by tests
rate_limiter = RateLimiter(max_requests=100, window_seconds=60)
# ─── FastAPI middleware ────────────────────────────────────────────────────────
class RateLimitMiddleware(BaseHTTPMiddleware):
"""
Starlette/FastAPI middleware that enforces per-API-key rate limits.
Health and docs endpoints are exempt.
"""
EXEMPT_PATHS = {"/health", "/docs", "/redoc", "/openapi.json", "/favicon.ico"}
def __init__(self, app, limiter: RateLimiter | None = None) -> None:
super().__init__(app)
self._limiter = limiter or rate_limiter
async def dispatch(self, request: Request, call_next) -> Response:
# Exempt health/docs
if request.url.path in self.EXEMPT_PATHS:
return await call_next(request)
api_key = request.headers.get("X-API-Key", "anonymous")
allowed = self._limiter.is_allowed(api_key)
if not allowed:
reset_secs = self._limiter.reset_in(api_key)
log.warning("Rate limit exceeded for key %r (resets in %.1fs)", api_key[:8], reset_secs)
return JSONResponse(
status_code=429,
content={
"detail": "Rate limit exceeded. Max 100 requests per 60 seconds.",
"retry_after_seconds": round(reset_secs, 1),
},
headers={
"Retry-After": str(int(reset_secs) + 1),
"X-RateLimit-Limit": str(self._limiter.max_requests),
"X-RateLimit-Remaining": "0",
"X-RateLimit-Reset": str(int(time.time() + reset_secs)),
},
)
response = await call_next(request)
# Inject rate-limit headers on allowed responses
remaining = self._limiter.remaining(api_key)
response.headers["X-RateLimit-Limit"] = str(self._limiter.max_requests)
response.headers["X-RateLimit-Remaining"] = str(remaining)
response.headers["X-RateLimit-Window"] = str(self._limiter.window_seconds)
return response
# ─── Tenant context helpers ───────────────────────────────────────────────────
def get_tenant_id(user: "User") -> str:
"""Extract tenant_id from authenticated User object."""
return user.tenant_id
def tenant_scoped_key(user: "User", suffix: str) -> str:
"""Build a namespaced cache / rate-limit key like 'tenant-default:queries'."""
return f"{user.tenant_id}:{suffix}"