2026-08-31 20:04:12 -04:00
|
|
|
"""Fixed-window rate limiting for the endpoints that had none.
|
|
|
|
|
|
|
|
|
|
`/signin`, `/signup`, `/forgot-password`, `/verify-otp` and `/reset-password-otp`
|
|
|
|
|
were completely unprotected: unlimited password guesses, unlimited OTP guesses
|
|
|
|
|
against any address, and unlimited outbound mail triggered by anyone.
|
|
|
|
|
|
|
|
|
|
A fixed window is deliberately simple. It permits a burst of up to 2x the limit
|
|
|
|
|
across a window boundary, which is the well-known trade-off; for login throttling
|
|
|
|
|
that is entirely acceptable, and a sliding window is not worth the extra Redis
|
|
|
|
|
round-trips on a system scheduled for replacement.
|
|
|
|
|
|
|
|
|
|
Fails **open** when Redis is unavailable — the alternative is that a Redis outage
|
|
|
|
|
locks every user out of the product. That is a deliberate choice and it is why
|
|
|
|
|
this is a mitigation rather than a control: it raises the cost of an online
|
|
|
|
|
guessing attack, it does not make one impossible.
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
import logging
|
|
|
|
|
from typing import Callable, Optional
|
|
|
|
|
|
|
|
|
|
from fastapi import HTTPException, Request, status
|
|
|
|
|
|
|
|
|
|
from app.core.redis import sync_redis_client
|
|
|
|
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def get_client_ip(request: Request) -> str:
|
|
|
|
|
"""Best-effort client address.
|
|
|
|
|
|
|
|
|
|
X-Forwarded-For is only trustworthy behind a proxy that overwrites it. Take
|
|
|
|
|
the left-most entry, which is what a single trusted proxy sets, and fall back
|
|
|
|
|
to the socket address.
|
|
|
|
|
"""
|
|
|
|
|
forwarded = request.headers.get("x-forwarded-for")
|
|
|
|
|
if forwarded:
|
|
|
|
|
first = forwarded.split(",")[0].strip()
|
|
|
|
|
if first:
|
|
|
|
|
return first
|
|
|
|
|
real_ip = request.headers.get("x-real-ip")
|
|
|
|
|
if real_ip:
|
|
|
|
|
return real_ip.strip()
|
|
|
|
|
return request.client.host if request.client else "unknown"
|
2026-09-07 18:15:34 +05:30
|
|
|
|
2026-08-31 20:04:12 -04:00
|
|
|
|
|
|
|
|
|
|
|
|
|
def _consume(bucket: str, limit: int, window_seconds: int) -> Optional[int]:
|
|
|
|
|
"""Increment the bucket. Returns seconds-to-wait if over limit, else None."""
|
2026-09-07 18:15:34 +05:30
|
|
|
from app.config.settings import settings
|
|
|
|
|
if settings.APP_ENV in ["local", "localdev"]:
|
|
|
|
|
return None
|
|
|
|
|
|
2026-08-31 20:04:12 -04:00
|
|
|
client = getattr(sync_redis_client, "client", None)
|
|
|
|
|
if client is None:
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
key = f"ratelimit:{bucket}"
|
|
|
|
|
try:
|
|
|
|
|
pipe = client.pipeline()
|
|
|
|
|
pipe.incr(key)
|
|
|
|
|
pipe.ttl(key)
|
|
|
|
|
count, ttl = pipe.execute()
|
|
|
|
|
|
|
|
|
|
if ttl is None or ttl < 0:
|
|
|
|
|
client.expire(key, window_seconds)
|
|
|
|
|
ttl = window_seconds
|
|
|
|
|
|
|
|
|
|
if int(count) > limit:
|
|
|
|
|
return int(ttl) if int(ttl) > 0 else window_seconds
|
|
|
|
|
return None
|
|
|
|
|
except Exception as e:
|
|
|
|
|
logger.warning("Rate limit check failed for %s: %s", bucket, e)
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def rate_limit(
|
|
|
|
|
name: str,
|
|
|
|
|
limit: int,
|
|
|
|
|
window_seconds: int,
|
|
|
|
|
by_body_field: Optional[str] = None,
|
|
|
|
|
) -> Callable:
|
|
|
|
|
"""Dependency factory.
|
|
|
|
|
|
|
|
|
|
Limits by client address, and additionally by a body field (typically
|
|
|
|
|
`email`) when one is named — so an attacker spreading guesses across many
|
|
|
|
|
addresses is still capped per address, and one spreading across many source
|
|
|
|
|
addresses is still capped per target account.
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
async def dependency(request: Request) -> None:
|
|
|
|
|
buckets: list[str] = [f"{name}:ip:{get_client_ip(request)}"]
|
|
|
|
|
|
|
|
|
|
if by_body_field:
|
|
|
|
|
try:
|
|
|
|
|
body = await request.json()
|
|
|
|
|
value = body.get(by_body_field) if isinstance(body, dict) else None
|
|
|
|
|
if value:
|
|
|
|
|
buckets.append(f"{name}:{by_body_field}:{str(value).lower()}")
|
|
|
|
|
except Exception:
|
|
|
|
|
pass
|
|
|
|
|
|
|
|
|
|
for bucket in buckets:
|
|
|
|
|
retry_after = _consume(bucket, limit, window_seconds)
|
|
|
|
|
if retry_after is not None:
|
|
|
|
|
logger.warning("Rate limit exceeded: %s", bucket)
|
|
|
|
|
raise HTTPException(
|
|
|
|
|
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
|
|
|
|
detail="Too many requests. Please try again later.",
|
|
|
|
|
headers={"Retry-After": str(retry_after)},
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
return dependency
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
SIGNIN_LIMIT = rate_limit("signin", limit=10, window_seconds=900, by_body_field="email")
|
|
|
|
|
SIGNUP_LIMIT = rate_limit("signup", limit=5, window_seconds=3600)
|
|
|
|
|
FORGOT_PASSWORD_LIMIT = rate_limit(
|
|
|
|
|
"forgot_password", limit=5, window_seconds=900, by_body_field="email"
|
|
|
|
|
)
|
|
|
|
|
VERIFY_OTP_LIMIT = rate_limit(
|
|
|
|
|
"verify_otp", limit=10, window_seconds=900, by_body_field="email"
|
|
|
|
|
)
|
|
|
|
|
RESET_PASSWORD_OTP_LIMIT = rate_limit(
|
|
|
|
|
"reset_password_otp", limit=10, window_seconds=900, by_body_field="email"
|
|
|
|
|
)
|