chore: add token revocation, caching, concurrency, and WebDAV lock modules
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
from .usage_tracking import UsageTrackingMiddleware
|
||||
from .rate_limit import RateLimitMiddleware
|
||||
from .security import SecurityHeadersMiddleware
|
||||
|
||||
__all__ = ["UsageTrackingMiddleware", "RateLimitMiddleware", "SecurityHeadersMiddleware"]
|
||||
|
||||
@@ -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"
|
||||
@@ -0,0 +1,49 @@
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
|
||||
class SecurityHeadersMiddleware(BaseHTTPMiddleware):
|
||||
SECURITY_HEADERS = {
|
||||
"X-Content-Type-Options": "nosniff",
|
||||
"X-Frame-Options": "DENY",
|
||||
"X-XSS-Protection": "1; mode=block",
|
||||
"Referrer-Policy": "strict-origin-when-cross-origin",
|
||||
"Permissions-Policy": "geolocation=(), microphone=(), camera=()",
|
||||
}
|
||||
|
||||
HTTPS_HEADERS = {
|
||||
"Strict-Transport-Security": "max-age=31536000; includeSubDomains",
|
||||
}
|
||||
|
||||
CSP_POLICY = (
|
||||
"default-src 'self'; "
|
||||
"script-src 'self' 'unsafe-inline' 'unsafe-eval'; "
|
||||
"style-src 'self' 'unsafe-inline'; "
|
||||
"img-src 'self' data: blob:; "
|
||||
"font-src 'self'; "
|
||||
"connect-src 'self'; "
|
||||
"frame-ancestors 'none'; "
|
||||
"base-uri 'self'; "
|
||||
"form-action 'self'"
|
||||
)
|
||||
|
||||
def __init__(self, app, enable_hsts: bool = False, enable_csp: bool = True):
|
||||
super().__init__(app)
|
||||
self.enable_hsts = enable_hsts
|
||||
self.enable_csp = enable_csp
|
||||
|
||||
async def dispatch(self, request: Request, call_next):
|
||||
response: Response = await call_next(request)
|
||||
|
||||
for header, value in self.SECURITY_HEADERS.items():
|
||||
response.headers[header] = value
|
||||
|
||||
if self.enable_hsts:
|
||||
for header, value in self.HTTPS_HEADERS.items():
|
||||
response.headers[header] = value
|
||||
|
||||
if self.enable_csp and not request.url.path.startswith("/webdav"):
|
||||
response.headers["Content-Security-Policy"] = self.CSP_POLICY
|
||||
|
||||
return response
|
||||
@@ -1,13 +1,18 @@
|
||||
from fastapi import Request
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from ..billing.usage_tracker import UsageTracker
|
||||
|
||||
try:
|
||||
from ..billing.usage_tracker import UsageTracker
|
||||
BILLING_AVAILABLE = True
|
||||
except ImportError:
|
||||
BILLING_AVAILABLE = False
|
||||
|
||||
|
||||
class UsageTrackingMiddleware(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next):
|
||||
response = await call_next(request)
|
||||
|
||||
if hasattr(request.state, "user") and request.state.user:
|
||||
if BILLING_AVAILABLE and hasattr(request.state, "user") and request.state.user:
|
||||
user = request.state.user
|
||||
|
||||
if (
|
||||
|
||||
Reference in New Issue
Block a user