"""FlowDeck — Security middleware: CSP headers + rate limiting.""" from __future__ import annotations import time from collections import defaultdict from starlette.middleware.base import BaseHTTPMiddleware from starlette.requests import Request from starlette.responses import JSONResponse # ── Constants ──────────────────────────────────────────────── # Allowed extensions for file uploads ALLOWED_EXTENSIONS: frozenset[str] = frozenset({ # Images ".png", ".jpg", ".jpeg", ".gif", ".webp", ".svg", ".bmp", ".ico", # Documents ".pdf", ".md", ".markdown", ".txt", ".log", ".csv", # Code ".py", ".js", ".jsx", ".ts", ".tsx", ".html", ".htm", ".xml", ".css", ".json", ".yaml", ".yml", ".toml", ".sql", ".sh", ".bash", ".zsh", ".ps1", ".rs", ".go", ".java", ".rb", ".php", ".c", ".cpp", ".h", ".swift", ".kt", ".scala", ".r", # Archives ".zip", ".tar", ".gz", ".rar", ".7z", # Misc ".env", ".cfg", ".conf", ".ini", ".dockerfile", ".makefile", }) MAX_UPLOAD_SIZE = 10 * 1024 * 1024 # 10 MB def validate_upload(filename: str, size: int) -> str | None: """Validate upload filename and size. Returns error message or None.""" if size > MAX_UPLOAD_SIZE: return f"File '{filename}' exceeds maximum size of 10 MB" ext = _ext(filename) if ext and ext not in ALLOWED_EXTENSIONS: return f"File extension '{ext}' is not allowed" return None def _ext(filename: str) -> str: """Extract lowercase extension from filename.""" if "." in filename: return "." + filename.rsplit(".", 1)[-1].lower() return "" # ── Content-Security-Policy Middleware ──────────────────────── class ContentSecurityPolicyMiddleware(BaseHTTPMiddleware): """Sets Content-Security-Policy headers on all HTML responses. A permissive-yet-safe policy for a Notion-style app that needs: - inline scripts (Alpine.js, HTMX) - inline styles - font loading - images from various sources - media (audio/video) - websocket connections for HMR/SSE """ CSP_HEADER = "Content-Security-Policy" CSP_VALUE = ( "default-src 'self'; " "script-src 'self' 'unsafe-inline' 'unsafe-eval'; " "style-src 'self' 'unsafe-inline' https://fonts.googleapis.com; " "img-src 'self' data: blob: https:; " "font-src 'self' data: https://fonts.gstatic.com; " "connect-src 'self' https: wss: ws:; " "media-src 'self' blob:; " "frame-src 'self'; " "object-src 'none'; " "base-uri 'self'; " "form-action 'self'; " ) async def dispatch(self, request: Request, call_next): response = await call_next(request) # Only set CSP on HTML responses content_type = response.headers.get("content-type", "") if "text/html" in content_type: response.headers[self.CSP_HEADER] = self.CSP_VALUE return response # ── Rate Limiting Middleware ────────────────────────────────── class RateLimitMiddleware(BaseHTTPMiddleware): """Simple in-memory sliding-window rate limiter per IP. Default: 100 requests per minute per IP for API routes. Non-API routes (HTML pages, static files) are not rate-limited. """ # Paths that should be rate-limited RATE_LIMITED_PREFIXES: tuple[str, ...] = ( "/api/", "/board/api/", "/auth/", ) # Paths exempt from rate limiting even under an API prefix EXEMPT_PATHS: frozenset[str] = frozenset({ "/api/health", "/api/frontend-error", "/api/frontend-errors", }) def __init__(self, app, max_requests: int = 100, window_seconds: int = 60): super().__init__(app) self.max_requests = max_requests self.window_seconds = window_seconds self._store: dict[str, tuple[float, int]] = defaultdict(lambda: (0.0, 0)) async def dispatch(self, request: Request, call_next): path = request.url.path # Respect the global rate-limit toggle (disabled in tests/local). from app.config import settings if not settings.rate_limit_enabled: return await call_next(request) # Only rate-limit API routes if not any(path.startswith(p) for p in self.RATE_LIMITED_PREFIXES): return await call_next(request) # Exempt health check and error capture if path in self.EXEMPT_PATHS: return await call_next(request) ip = request.client.host if request.client else "unknown" now = time.time() window_start, count = self._store[ip] if now - window_start > self.window_seconds: self._store[ip] = (now, 1) return await call_next(request) if count >= self.max_requests: return JSONResponse( {"error": "Rate limit exceeded", "detail": f"Max {self.max_requests} req/min per IP"}, status_code=429, ) self._store[ip] = (window_start, count + 1) return await call_next(request)