| """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, |
| ) |
|
|
| |
| 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" |
|
|
| |
| 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 |
|
|
| |
| 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"}) |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|