""" DocuLens - Security middleware. Provides: - API key authentication (header or query param) - In-memory rate limiting (sliding window per key) - Input validation helpers (file size, file type, filename sanitization) """ import os import re import time import logging import hashlib import secrets from collections import defaultdict from typing import Optional from fastapi import Request, HTTPException, Depends, Security from fastapi.security import APIKeyHeader logger = logging.getLogger(__name__) # --------------------------------------------------------------------------- # Configuration # --------------------------------------------------------------------------- # API key authentication mode: # "optional" - requests without a key are allowed (dev/demo mode) # "required" - every /api/v1/* request must carry a valid key AUTH_MODE = os.environ.get("AUTH_MODE", "optional") # Comma-separated list of valid API keys. In production, set this env var # or store keys in Supabase (see _load_keys_from_db). _ENV_API_KEYS = os.environ.get("API_KEYS", "") # Rate-limit defaults (requests per window per key) RATE_LIMIT_REQUESTS = int(os.environ.get("RATE_LIMIT_REQUESTS", "60")) RATE_LIMIT_WINDOW = int(os.environ.get("RATE_LIMIT_WINDOW", "60")) # seconds # File upload constraints MAX_FILE_SIZE_MB = int(os.environ.get("MAX_FILE_SIZE_MB", "20")) MAX_FILE_SIZE_BYTES = MAX_FILE_SIZE_MB * 1024 * 1024 ALLOWED_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".tiff", ".tif", ".webp", ".pdf"} ALLOWED_CONTENT_TYPES = { "image/jpeg", "image/png", "image/gif", "image/bmp", "image/tiff", "image/webp", "application/pdf", } # --------------------------------------------------------------------------- # API Key Store # --------------------------------------------------------------------------- def _parse_env_keys() -> set[str]: """Parse API_KEYS env var into a set of key hashes.""" if not _ENV_API_KEYS: return set() keys = set() for k in _ENV_API_KEYS.split(","): k = k.strip() if k: keys.add(_hash_key(k)) return keys def _hash_key(raw: str) -> str: """SHA-256 hash of a raw API key for safe comparison.""" return hashlib.sha256(raw.encode()).hexdigest() # In-memory key store (hashed). Populated on first use. _valid_key_hashes: Optional[set[str]] = None def _get_valid_keys() -> set[str]: """Return the set of valid key hashes, loading lazily.""" global _valid_key_hashes if _valid_key_hashes is None: _valid_key_hashes = _parse_env_keys() if _valid_key_hashes: logger.info(f"Loaded {len(_valid_key_hashes)} API key(s) from environment") else: logger.info("No API keys configured — auth depends on AUTH_MODE setting") return _valid_key_hashes def generate_api_key(prefix: str = "dmai") -> str: """Generate a new API key. Useful for bootstrapping.""" token = secrets.token_urlsafe(32) return f"{prefix}_{token}" # --------------------------------------------------------------------------- # API Key Authentication Dependency # --------------------------------------------------------------------------- _api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False) async def verify_api_key( request: Request, api_key: Optional[str] = Security(_api_key_header), ) -> Optional[str]: """FastAPI dependency that validates the API key. Behaviour depends on AUTH_MODE: - "required": rejects requests without a valid key (401) - "optional": allows keyless requests but still validates if a key is present (returns None for anonymous access) Returns the raw API key on success (useful for per-key rate limiting). """ # Also accept ?api_key= query param as fallback if not api_key: api_key = request.query_params.get("api_key") valid_keys = _get_valid_keys() if api_key: hashed = _hash_key(api_key) if valid_keys and hashed not in valid_keys: logger.warning("Rejected invalid API key") raise HTTPException(status_code=401, detail="Invalid API key") return api_key # No key provided if AUTH_MODE == "required" and valid_keys: raise HTTPException( status_code=401, detail="API key required. Pass it in the X-API-Key header.", ) return None # anonymous access allowed # --------------------------------------------------------------------------- # Rate Limiting (in-memory sliding window) # --------------------------------------------------------------------------- class RateLimiter: """Simple in-memory sliding-window rate limiter. Not distributed — suitable for single-instance deployments. For multi-replica production, swap to Redis-backed limiter. """ def __init__(self, max_requests: int, window_seconds: int): self.max_requests = max_requests self.window = window_seconds self._hits: dict[str, list[float]] = defaultdict(list) def check(self, key: str) -> tuple[bool, dict]: """Check if the key is within rate limits. Returns (allowed, headers) where headers contains standard rate-limit response headers. """ now = time.time() cutoff = now - self.window # Prune old entries self._hits[key] = [t for t in self._hits[key] if t > cutoff] remaining = max(0, self.max_requests - len(self._hits[key])) headers = { "X-RateLimit-Limit": str(self.max_requests), "X-RateLimit-Remaining": str(remaining), "X-RateLimit-Reset": str(int(cutoff + self.window)), } if len(self._hits[key]) >= self.max_requests: return False, headers self._hits[key].append(now) remaining = max(0, self.max_requests - len(self._hits[key])) headers["X-RateLimit-Remaining"] = str(remaining) return True, headers def cleanup(self): """Remove stale entries (call periodically to avoid memory leaks).""" now = time.time() cutoff = now - self.window stale = [k for k, v in self._hits.items() if not v or v[-1] < cutoff] for k in stale: del self._hits[k] # Singleton limiter _rate_limiter = RateLimiter(RATE_LIMIT_REQUESTS, RATE_LIMIT_WINDOW) async def check_rate_limit( request: Request, api_key: Optional[str] = Depends(verify_api_key), ): """FastAPI dependency that enforces rate limits. Uses the API key as the rate-limit bucket. Anonymous requests are bucketed by IP address. """ bucket = api_key or _get_client_ip(request) allowed, headers = _rate_limiter.check(bucket) # Attach headers to response (via request state) request.state.rate_limit_headers = headers if not allowed: raise HTTPException( status_code=429, detail="Rate limit exceeded. Try again shortly.", headers=headers, ) return api_key def _get_client_ip(request: Request) -> str: """Extract client IP, respecting X-Forwarded-For behind proxies.""" forwarded = request.headers.get("x-forwarded-for") if forwarded: return forwarded.split(",")[0].strip() return request.client.host if request.client else "unknown" # --------------------------------------------------------------------------- # Input Validation Helpers # --------------------------------------------------------------------------- def validate_file_upload(filename: Optional[str], file_size: int, content_type: Optional[str] = None): """Validate an uploaded file's name, size, and type. Raises HTTPException on validation failure. """ # Size check if file_size > MAX_FILE_SIZE_BYTES: raise HTTPException( status_code=413, detail=f"File too large. Maximum size is {MAX_FILE_SIZE_MB}MB.", ) # Extension check if filename: ext = os.path.splitext(filename)[1].lower() if ext and ext not in ALLOWED_EXTENSIONS: raise HTTPException( status_code=400, detail=f"Unsupported file type '{ext}'. Allowed: {', '.join(sorted(ALLOWED_EXTENSIONS))}", ) # Content-type check (advisory — files can lie about content type) if content_type and content_type not in ALLOWED_CONTENT_TYPES: # Only warn, don't block — content_type is unreliable logger.warning(f"Unexpected content type: {content_type} for file {filename}") def sanitize_filename(filename: Optional[str]) -> str: """Sanitize a filename to prevent path traversal and injection. Strips directory components, removes dangerous characters, and truncates to a safe length. """ if not filename: return "unnamed_upload" # Strip any directory components (path traversal defense) name = os.path.basename(filename) # Remove null bytes and control characters name = re.sub(r"[\x00-\x1f\x7f]", "", name) # Replace dangerous characters but keep dots, hyphens, underscores name = re.sub(r"[^\w.\-]", "_", name) # Collapse multiple underscores/dots name = re.sub(r"_{2,}", "_", name) name = re.sub(r"\.{2,}", ".", name) # Strip leading dots (hidden files) and leading/trailing underscores name = name.lstrip(".").strip("_") # Truncate to reasonable length if len(name) > 200: stem, ext = os.path.splitext(name) name = stem[:200 - len(ext)] + ext return name or "unnamed_upload"