"""OmniParse AI — Rate limiting + security middleware.""" import time, re, os from collections import defaultdict from fastapi import Request, HTTPException from fastapi.responses import Response from starlette.middleware.base import BaseHTTPMiddleware from config import ( RATE_GLOBAL, RATE_AUTH, RATE_UPLOAD, RATE_API, MAX_FILE_SIZE_MB, MAX_FILES_PER_REQ, STRIPE_WEBHOOK_SECRET, ) # ── Token bucket ──────────────────────────────────────────────────────────── class TokenBucket: def __init__(self, rate: float, cap: int): self.tokens = float(cap); self.cap = cap; self.rate = rate self.last = time.monotonic() def _refill(self): now = time.monotonic(); e = now - self.last self.tokens = min(self.cap, self.tokens + e * self.rate); self.last = now def consume(self, n=1) -> bool: self._refill() if self.tokens >= n: self.tokens -= n; return True return False class RateLimiter: def __init__(self): self._buckets: dict = defaultdict(dict) def _get(self, key: str, limit: str) -> TokenBucket: if limit not in self._buckets[key]: cnt, period = limit.split("/"); cnt = int(cnt) secs = {"second":1,"seconds":1,"minute":60,"minutes":60,"hour":3600,"hours":3600}.get(period,60) self._buckets[key][limit] = TokenBucket(cnt/secs, cnt) return self._buckets[key][limit] def limited(self, key: str, limit: str) -> bool: return not self._get(key, limit).consume() limiter = RateLimiter() def client_ip(req: Request) -> str: fwd = req.headers.get("X-Forwarded-For") if fwd: return fwd.split(",")[0].strip() rip = req.headers.get("X-Real-IP") if rip: return rip.strip() return req.client.host if req.client else "unknown" # ── Security headers middleware ───────────────────────────────────────────── class SecurityHeadersMiddleware(BaseHTTPMiddleware): async def dispatch(self, req: Request, call_next): resp = await call_next(req) resp.headers["X-Content-Type-Options"] = "nosniff" resp.headers["X-Frame-Options"] = "DENY" resp.headers["X-XSS-Protection"] = "1; mode=block" resp.headers["Referrer-Policy"] = "strict-origin-when-cross-origin" resp.headers["Permissions-Policy"] = "camera=(), microphone=(), geolocation=()" if req.url.scheme == "https": resp.headers["Strict-Transport-Security"] = "max-age=31536000; includeSubDomains" return resp # ── Rate limit helpers ────────────────────────────────────────────────────── def global_limit(req: Request): ip = client_ip(req) if limiter.limited(ip, RATE_GLOBAL): raise HTTPException(429, detail="Too many requests. Slow down.", headers={"Retry-After":"60"}) def auth_limit(req: Request): ip = client_ip(req) if limiter.limited(f"auth:{ip}", RATE_AUTH): raise HTTPException(429, detail="Too many login attempts. Wait and retry.", headers={"Retry-After":"60"}) def upload_limit(req: Request): ip = client_ip(req) if limiter.limited(f"upload:{ip}", RATE_UPLOAD): raise HTTPException(429, detail="Upload rate exceeded. Wait before uploading more.", headers={"Retry-After":"30"}) def api_limit(req: Request): ip = client_ip(req) if limiter.limited(f"api:{ip}", RATE_API): raise HTTPException(429, detail="API rate limit exceeded.", headers={"Retry-After":"30"}) # ── File validation ───────────────────────────────────────────────────────── def validate_file_size(size: int): if size > MAX_FILE_SIZE_MB * 1024 * 1024: raise HTTPException(413, detail=f"File too large. Maximum {MAX_FILE_SIZE_MB} MB.") def validate_file_count(n: int): if n > MAX_FILES_PER_REQ: raise HTTPException(400, detail=f"Maximum {MAX_FILES_PER_REQ} files per upload.") def validate_filename(fn: str) -> str: clean = os.path.basename(fn) clean = re.sub(r'[^a-zA-Z0-9._\- ()]', '_', clean) if not clean or clean.startswith('.'): clean = "uploaded_file" if len(clean) > 255: name, ext = os.path.splitext(clean); clean = name[:250] + ext return clean # ── Stripe webhook signature verification ────────────────────────────────── async def verify_stripe_webhook(req: Request, raw_body: bytes) -> bool: """Verify Stripe webhook signature. Returns True if valid.""" if not STRIPE_WEBHOOK_SECRET: return False try: import stripe sig = req.headers.get("Stripe-Signature", "") stripe.Webhook.construct_event(raw_body, sig, STRIPE_WEBHOOK_SECRET) return True except Exception: return False