Spaces:
Runtime error
Runtime error
Download middleware.py from sdudeja/agentic-extractor: direct link, hf CLI and curl.
- Browser
- Download file 9.6 kB
-
https://huggingface.co/spaces/sdudeja/agentic-extractor/resolve/main/middleware.py
- Command line
-
hf download hf://spaces/sdudeja/agentic-extractor/middleware.py
-
curl -L -o middleware.py https://huggingface.co/spaces/sdudeja/agentic-extractor/resolve/main/middleware.py
9.6 kB
| """ | |
| 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" | |