lightning / helper /ratelimit.py
R.C.M.
Improved structure
99d08e9
Raw
History Blame
4.71 kB
import time
from typing import Optional, Dict
from fastapi import HTTPException, Request
from helper.misc import sanitize_client_id, get_usage_lock, get_usage_period_key, build_default_subject, bind_client_subject, resolve_bound_subject
from helper.subscriptions import fetch_subscription, usage_store, normalize_plan_key, TIER_CONFIG
import os
IDENTITY_CACHE_TTL_SECONDS = 60
identity_cache = {}
CLIENT_BIND_TTL_SECONDS = int(
os.getenv("CLIENT_BIND_TTL_SECONDS", str(8 * 24 * 60 * 60))
)
MAX_CLIENT_ID_LENGTH = 128
client_subject_bindings = {}
MAX_CHAT_PROMPT_CHARS = int(os.getenv("MAX_CHAT_PROMPT_CHARS", "120000"))
MAX_CHAT_PROMPT_BYTES = int(os.getenv("MAX_CHAT_PROMPT_BYTES", "500000"))
MAX_GROQ_PROMPT_CHARS = int(os.getenv("MAX_GROQ_PROMPT_CHARS", "90000"))
MAX_GROQ_PROMPT_BYTES = int(os.getenv("MAX_GROQ_PROMPT_BYTES", "350000"))
MAX_MEDIA_PROMPT_CHARS = int(os.getenv("MAX_MEDIA_PROMPT_CHARS", "4000"))
MAX_MEDIA_PROMPT_BYTES = int(os.getenv("MAX_MEDIA_PROMPT_BYTES", "16000"))
async def resolve_rate_limit_identity(
request: Request,
authorization: Optional[str],
client_id: Optional[str] = None,
) -> tuple[str, str]:
now = time.time()
normalized_client_id = sanitize_client_id(client_id)
default_subject = build_default_subject(request, normalized_client_id)
if not authorization or not authorization.startswith("Bearer "):
return "free", resolve_bound_subject(normalized_client_id, default_subject)
token = authorization.split(" ", 1)[1].strip()
if not token:
return "free", resolve_bound_subject(normalized_client_id, default_subject)
cached = identity_cache.get(token)
if cached and cached.get("expires_at", 0) > now:
plan_key = cached.get("plan_key", "free")
subject = cached.get("subject", default_subject)
bind_client_subject(normalized_client_id, subject, plan_key)
return plan_key, subject
try:
sub = await fetch_subscription(token)
except Exception:
return "free", resolve_bound_subject(normalized_client_id, default_subject)
if not isinstance(sub, dict) or sub.get("error"):
return "free", resolve_bound_subject(normalized_client_id, default_subject)
email = sub.get("email")
if isinstance(email, str) and email.strip():
subject = f"user:{email.strip().lower()}"
else:
subject = default_subject
plan_key = normalize_plan_key(sub.get("plan_key"))
identity_cache[token] = {
"plan_key": plan_key,
"subject": subject,
"expires_at": now + IDENTITY_CACHE_TTL_SECONDS,
}
bind_client_subject(normalized_client_id, subject, plan_key)
return plan_key, subject
async def enforce_rate_limit(
request: Request,
authorization: Optional[str],
metric: str,
client_id: Optional[str] = None,
) -> Dict[str, Optional[int | str]]:
if metric not in usage_store:
raise HTTPException(status_code=500, detail=f"Unknown limit metric: {metric}")
plan_key, subject = await resolve_rate_limit_identity(
request, authorization, client_id
)
plan = TIER_CONFIG.get(plan_key) or TIER_CONFIG["free"]
plan_limits = plan.get("limits", {})
limit = plan_limits.get(metric)
window_key = get_usage_period_key(metric)
lock = get_usage_lock(metric, subject)
async with lock:
bucket = usage_store[metric]
entry = bucket.get(subject)
if not entry or entry.get("window") != window_key:
entry = {"window": window_key, "count": 0}
bucket[subject] = entry
if limit is not None and entry["count"] >= int(limit):
raise HTTPException(
status_code=429,
detail=f"{metric} limit reached for {plan.get('name', 'current plan')}",
)
entry["count"] += 1
remaining = None if limit is None else max(0, int(limit) - entry["count"])
return {
"plan_key": plan_key,
"remaining": remaining,
"used": entry["count"],
"window": window_key,
}
async def check_audio_rate_limit(
request: Request,
authorization: Optional[str],
client_id: Optional[str] = None,
):
await enforce_rate_limit(request, authorization, "audioWeekly", client_id)
async def check_image_rate_limit(
request: Request,
authorization: Optional[str],
client_id: Optional[str] = None,
):
await enforce_rate_limit(request, authorization, "imagesDaily", client_id)
async def check_video_rate_limit(
request: Request,
authorization: Optional[str],
client_id: Optional[str] = None,
):
await enforce_rate_limit(request, authorization, "videosDaily", client_id)