Progress.
This commit is contained in:
@@ -0,0 +1,190 @@
|
||||
import asyncio
|
||||
import time
|
||||
from typing import Dict, Tuple, Optional
|
||||
from dataclasses import dataclass, field
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response, JSONResponse
|
||||
import logging
|
||||
from mywebdav.settings import settings
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RateLimitBucket:
|
||||
count: int = 0
|
||||
window_start: float = field(default_factory=time.time)
|
||||
|
||||
|
||||
class RateLimiter:
|
||||
LIMITS = {
|
||||
"login": (5, 60),
|
||||
"register": (3, 60),
|
||||
"api": (100, 60),
|
||||
"upload": (20, 60),
|
||||
"download": (50, 60),
|
||||
"webdav": (200, 60),
|
||||
"default": (100, 60),
|
||||
}
|
||||
|
||||
def __init__(self, db_manager=None):
|
||||
self.db_manager = db_manager
|
||||
self.buckets: Dict[str, RateLimitBucket] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
self._cleanup_task: Optional[asyncio.Task] = None
|
||||
self._running = False
|
||||
|
||||
async def start(self, db_manager=None):
|
||||
if db_manager:
|
||||
self.db_manager = db_manager
|
||||
self._running = True
|
||||
self._cleanup_task = asyncio.create_task(self._background_cleanup())
|
||||
logger.info("RateLimiter started")
|
||||
|
||||
async def stop(self):
|
||||
self._running = False
|
||||
if self._cleanup_task:
|
||||
self._cleanup_task.cancel()
|
||||
try:
|
||||
await self._cleanup_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
logger.info("RateLimiter stopped")
|
||||
|
||||
def _get_bucket_key(self, key: str, limit_type: str) -> str:
|
||||
return f"{limit_type}:{key}"
|
||||
|
||||
async def check_rate_limit(self, key: str, limit_type: str = "default") -> Tuple[bool, int, int]:
|
||||
max_requests, window_seconds = self.LIMITS.get(limit_type, self.LIMITS["default"])
|
||||
bucket_key = self._get_bucket_key(key, limit_type)
|
||||
now = time.time()
|
||||
|
||||
async with self._lock:
|
||||
if bucket_key not in self.buckets:
|
||||
self.buckets[bucket_key] = RateLimitBucket(count=1, window_start=now)
|
||||
return True, max_requests - 1, window_seconds
|
||||
|
||||
bucket = self.buckets[bucket_key]
|
||||
|
||||
if now - bucket.window_start >= window_seconds:
|
||||
bucket.count = 1
|
||||
bucket.window_start = now
|
||||
return True, max_requests - 1, window_seconds
|
||||
|
||||
if bucket.count >= max_requests:
|
||||
retry_after = int(window_seconds - (now - bucket.window_start))
|
||||
return False, 0, retry_after
|
||||
|
||||
bucket.count += 1
|
||||
remaining = max_requests - bucket.count
|
||||
return True, remaining, window_seconds
|
||||
|
||||
async def reset_rate_limit(self, key: str, limit_type: str = "default"):
|
||||
bucket_key = self._get_bucket_key(key, limit_type)
|
||||
async with self._lock:
|
||||
if bucket_key in self.buckets:
|
||||
del self.buckets[bucket_key]
|
||||
|
||||
async def _background_cleanup(self):
|
||||
while self._running:
|
||||
try:
|
||||
await asyncio.sleep(60)
|
||||
await self._cleanup_expired()
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Error in rate limit cleanup: {e}")
|
||||
|
||||
async def _cleanup_expired(self):
|
||||
now = time.time()
|
||||
async with self._lock:
|
||||
expired = []
|
||||
for bucket_key, bucket in self.buckets.items():
|
||||
limit_type = bucket_key.split(":")[0]
|
||||
_, window_seconds = self.LIMITS.get(limit_type, self.LIMITS["default"])
|
||||
if now - bucket.window_start > window_seconds * 2:
|
||||
expired.append(bucket_key)
|
||||
for key in expired:
|
||||
del self.buckets[key]
|
||||
if expired:
|
||||
logger.debug(f"Cleaned up {len(expired)} expired rate limit buckets")
|
||||
|
||||
|
||||
_rate_limiter: Optional[RateLimiter] = None
|
||||
|
||||
|
||||
async def init_rate_limiter(db_manager=None) -> RateLimiter:
|
||||
global _rate_limiter
|
||||
_rate_limiter = RateLimiter()
|
||||
await _rate_limiter.start(db_manager)
|
||||
return _rate_limiter
|
||||
|
||||
|
||||
async def shutdown_rate_limiter():
|
||||
global _rate_limiter
|
||||
if _rate_limiter:
|
||||
await _rate_limiter.stop()
|
||||
_rate_limiter = None
|
||||
|
||||
|
||||
def get_rate_limiter() -> RateLimiter:
|
||||
if not _rate_limiter:
|
||||
raise RuntimeError("Rate limiter not initialized")
|
||||
return _rate_limiter
|
||||
|
||||
|
||||
class RateLimitMiddleware(BaseHTTPMiddleware):
|
||||
ROUTE_LIMITS = {
|
||||
"/api/auth/login": "login",
|
||||
"/api/auth/register": "register",
|
||||
"/files/upload": "upload",
|
||||
"/files/download": "download",
|
||||
"/webdav": "webdav",
|
||||
}
|
||||
|
||||
async def dispatch(self, request: Request, call_next):
|
||||
if not settings.RATE_LIMIT_ENABLED:
|
||||
return await call_next(request)
|
||||
|
||||
|
||||
if not _rate_limiter:
|
||||
return await call_next(request)
|
||||
|
||||
client_ip = self._get_client_ip(request)
|
||||
limit_type = self._get_limit_type(request.url.path)
|
||||
|
||||
allowed, remaining, retry_after = await _rate_limiter.check_rate_limit(
|
||||
client_ip, limit_type
|
||||
)
|
||||
|
||||
if not allowed:
|
||||
return JSONResponse(
|
||||
status_code=429,
|
||||
content={
|
||||
"detail": "Too many requests",
|
||||
"retry_after": retry_after
|
||||
},
|
||||
headers={
|
||||
"Retry-After": str(retry_after),
|
||||
"X-RateLimit-Remaining": "0",
|
||||
}
|
||||
)
|
||||
|
||||
response = await call_next(request)
|
||||
|
||||
response.headers["X-RateLimit-Remaining"] = str(remaining)
|
||||
response.headers["X-RateLimit-Reset"] = str(retry_after)
|
||||
|
||||
return response
|
||||
|
||||
def _get_client_ip(self, request: Request) -> str:
|
||||
forwarded = request.headers.get("X-Forwarded-For")
|
||||
if forwarded:
|
||||
return forwarded.split(",")[0].strip()
|
||||
return request.client.host if request.client else "unknown"
|
||||
|
||||
def _get_limit_type(self, path: str) -> str:
|
||||
for route_prefix, limit_type in self.ROUTE_LIMITS.items():
|
||||
if path.startswith(route_prefix):
|
||||
return limit_type
|
||||
return "api"
|
||||
Reference in New Issue
Block a user