agentic-extractor / middleware.py
adudeja's picture
all changes from render
04c4194
Raw History Blame Contribute Delete
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"