File size: 5,146 Bytes
f44391a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
"""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