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
1016from __future__ import annotations
1117
1218import threading
1319import time
20+ import uuid
1421from 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