Skip to content

Commit 6bca1a0

Browse files
authored
Update ratelimit.py
1 parent 7792a26 commit 6bca1a0

1 file changed

Lines changed: 141 additions & 15 deletions

File tree

‎app/infra/ratelimit.py‎

Lines changed: 141 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -1,32 +1,48 @@
1-
"""In-memory sliding-window rate limiter.
1+
"""Sliding-window rate limiting with a pluggable backend.
22
3-
Process-local: the counters live in this process's memory, so limits are
4-
per-instance. That's the right trade for a single-instance deployment;
5-
running multiple web instances would need a shared store (Redis) or
6-
sticky routing to enforce a global limit. Good enough to stop one
7-
client hammering the money-spending run endpoint or brute-forcing auth.
3+
Two backends behind one ``check(key, limit, window)`` interface:
4+
5+
* ``SlidingWindowLimiter`` -- process-local (in-memory). The default:
6+
correct for a single instance, and the right trade for dev. Limits
7+
are per-instance, so N replicas allow N x the intended rate.
8+
* ``RedisRateLimiter`` -- a shared sorted-set window in Redis, so the
9+
limit is *global* across every instance. Select it for a
10+
multi-instance deploy (``PAW_RATE_LIMIT__BACKEND=redis``).
11+
12+
Both throttle the money-spending run endpoint (per user) and login /
13+
register (per IP), raising 429 + ``Retry-After`` when a window is full.
814
"""
915

1016
from __future__ import annotations
1117

1218
import threading
1319
import time
20+
import uuid
1421
from collections import defaultdict, deque
22+
from typing import Any, Protocol, runtime_checkable
1523

24+
from .config import get_settings
1625

17-
class SlidingWindowLimiter:
26+
27+
@runtime_checkable
28+
class RateLimiter(Protocol):
1829
"""Allow at most ``limit`` events per ``window`` seconds per key."""
1930

31+
def check(self, key: str, limit: int, window: float) -> tuple[bool, float]:
32+
"""Record an attempt; return ``(allowed, retry_after_seconds)``."""
33+
34+
def reset(self) -> None:
35+
"""Forget all counters (used to isolate tests)."""
36+
37+
38+
class SlidingWindowLimiter:
39+
"""Process-local sliding window (in-memory)."""
40+
2041
def __init__(self) -> None:
2142
self._hits: dict[str, deque[float]] = defaultdict(deque)
2243
self._lock = threading.Lock()
2344

2445
def check(self, key: str, limit: int, window: float) -> tuple[bool, float]:
25-
"""Record an attempt for *key*. Returns ``(allowed, retry_after)``.
26-
27-
``retry_after`` is seconds until the oldest hit in the window
28-
expires (0 when allowed).
29-
"""
3046
now = time.monotonic()
3147
cutoff = now - window
3248
with self._lock:
@@ -39,10 +55,120 @@ def check(self, key: str, limit: int, window: float) -> tuple[bool, float]:
3955
return True, 0.0
4056

4157
def reset(self) -> None:
42-
"""Forget all counters (used to isolate tests)."""
4358
with self._lock:
4459
self._hits.clear()
4560

4661

47-
# Process-wide limiter shared by all rate-limit dependencies.
48-
limiter = SlidingWindowLimiter()
62+
class RedisRateLimiter:
63+
"""Shared sliding window backed by a Redis sorted set per key.
64+
65+
The trim + count + conditional-add is done in a single Lua ``eval``
66+
so it is atomic on the Redis server: concurrent checks across
67+
instances cannot all observe ``count < limit`` and then each add,
68+
which a client-side read-then-write pipeline would allow (and which
69+
would defeat the whole point of a *global* limiter). The client is
70+
injectable for tests.
71+
"""
72+
73+
# KEYS[1]=zset ARGV[1]=now ARGV[2]=cutoff ARGV[3]=limit
74+
# ARGV[4]=member ARGV[5]=ttl. Returns {allowed(0/1), retry_ms_basis}.
75+
_LUA = """
76+
redis.call('ZREMRANGEBYSCORE', KEYS[1], 0, tonumber(ARGV[2]))
77+
local count = redis.call('ZCARD', KEYS[1])
78+
if count >= tonumber(ARGV[3]) then
79+
local oldest = redis.call('ZRANGE', KEYS[1], 0, 0, 'WITHSCORES')
80+
return {0, oldest[2]}
81+
end
82+
redis.call('ZADD', KEYS[1], ARGV[1], ARGV[4])
83+
redis.call('EXPIRE', KEYS[1], tonumber(ARGV[5]))
84+
return {1, '0'}
85+
"""
86+
87+
def __init__(self, url: str | None = None, client: Any = None, prefix: str = "rl:") -> None:
88+
self._client = client
89+
self._url = url
90+
self._prefix = prefix
91+
92+
def _redis(self) -> Any:
93+
if self._client is None:
94+
import redis # lazy: only when the redis backend is selected
95+
96+
self._client = redis.Redis.from_url(self._url or "redis://localhost:6379/0")
97+
return self._client
98+
99+
def check(self, key: str, limit: int, window: float) -> tuple[bool, float]:
100+
r = self._redis()
101+
full_key = self._prefix + key
102+
now = time.time()
103+
cutoff = now - window
104+
# a globally-unique member: uuid4 avoids the id() reuse/collision
105+
# that could silently overwrite a live hit and under-count
106+
member = f"{now}:{uuid.uuid4().hex}"
107+
try:
108+
allowed, basis = r.eval(
109+
self._LUA, 1, full_key, now, cutoff, limit, member, int(window) + 1
110+
)
111+
except Exception:
112+
# a limiter outage must not take down the endpoint: fail open
113+
# (allow) rather than 500. Logged by the caller if needed.
114+
return True, 0.0
115+
if int(allowed) == 1:
116+
return True, 0.0
117+
oldest_score = float(basis) if basis else now
118+
return False, max(0.0, window - (now - oldest_score))
119+
120+
def reset(self) -> None:
121+
# scoped flush of our namespace; best-effort (tests / admin)
122+
r = self._redis()
123+
for k in r.scan_iter(match=self._prefix + "*"):
124+
r.delete(k)
125+
126+
127+
def _build_limiter() -> RateLimiter:
128+
backend = get_settings().rate_limit_backend
129+
if backend == "memory":
130+
return SlidingWindowLimiter()
131+
if backend == "redis":
132+
return RedisRateLimiter(url=get_settings().rate_limit_redis_url or None)
133+
raise ValueError(f"unknown rate-limit backend: {backend!r}")
134+
135+
136+
# Process-wide limiter shared by all rate-limit dependencies. Rebuilt
137+
# lazily so a config change (backend selection) is picked up in tests.
138+
_limiter: RateLimiter | None = None
139+
_limiter_lock = threading.Lock()
140+
141+
142+
def get_limiter() -> RateLimiter:
143+
global _limiter
144+
with _limiter_lock:
145+
if _limiter is None:
146+
_limiter = _build_limiter()
147+
return _limiter
148+
149+
150+
def reset_limiter_for_tests() -> None:
151+
"""Drop the cached limiter so the next call rebuilds from config."""
152+
global _limiter
153+
with _limiter_lock:
154+
_limiter = None
155+
156+
157+
class _LimiterProxy:
158+
"""Attribute proxy so existing ``from ratelimit import limiter`` call
159+
sites keep working while the real backend is selected lazily."""
160+
161+
def check(self, key: str, limit: int, window: float) -> tuple[bool, float]:
162+
return get_limiter().check(key, limit, window)
163+
164+
def reset(self) -> None:
165+
# reset is best-effort (test/admin cleanup): never raise, even if
166+
# the configured backend can't be built or reached right now.
167+
try:
168+
get_limiter().reset()
169+
except Exception:
170+
reset_limiter_for_tests()
171+
172+
173+
# Back-compat singleton used by the rate-limit dependencies.
174+
limiter = _LimiterProxy()

0 commit comments

Comments
 (0)