test2 / middleware.py
simikkk's picture
Upload 8 files
dde961d verified
Raw
History Blame Contribute Delete
5.15 kB
"""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