p-redis-limiter 0.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,156 @@
1
+ Metadata-Version: 2.4
2
+ Name: p-redis-limiter
3
+ Version: 0.1.0
4
+ Summary: A fast, atomic, multi-tier Token Bucket rate limiter for Python and FastAPI using Redis.
5
+ Author: Mehran
6
+ License: MIT
7
+ Keywords: redis,rate-limiter,token-bucket,fastapi,throttling
8
+ Classifier: Development Status :: 4 - Beta
9
+ Classifier: Intended Audience :: Developers
10
+ Classifier: Programming Language :: Python :: 3
11
+ Classifier: Programming Language :: Python :: 3.9
12
+ Classifier: Programming Language :: Python :: 3.10
13
+ Classifier: Programming Language :: Python :: 3.11
14
+ Classifier: Programming Language :: Python :: 3.12
15
+ Classifier: Programming Language :: Python :: 3.13
16
+ Classifier: Programming Language :: Python :: 3.14
17
+ Classifier: License :: OSI Approved :: MIT License
18
+ Classifier: Operating System :: OS Independent
19
+ Classifier: Framework :: FastAPI
20
+ Requires-Python: >=3.9
21
+ Description-Content-Type: text/markdown
22
+ Requires-Dist: redis>=4.2.0
23
+ Provides-Extra: dev
24
+ Requires-Dist: pytest>=7.0.0; extra == "dev"
25
+ Requires-Dist: pytest-asyncio>=0.20.0; extra == "dev"
26
+ Requires-Dist: httpx>=0.23.0; extra == "dev"
27
+ Requires-Dist: fastapi>=0.70.0; extra == "dev"
28
+
29
+ # p-redis-limiter
30
+
31
+ Redis token bucket rate limiter for Python and FastAPI.
32
+
33
+ ## Installation
34
+
35
+ ```bash
36
+ pip install git+https://github.com/mehranpng/p-redis-limiter.git
37
+ ```
38
+
39
+ ## Usage
40
+
41
+ ### 1. Global Middleware (All Endpoints)
42
+
43
+ Apply a rate limit across all endpoints:
44
+
45
+ ```python
46
+ from fastapi import FastAPI
47
+ from redis import Redis
48
+ from p_redis_limiter import RateLimitMiddleware, Rate
49
+
50
+ app = FastAPI()
51
+ r = Redis(host="localhost", port=6379, decode_responses=True)
52
+
53
+ # 40 requests per 60 seconds
54
+ app.add_middleware(
55
+ RateLimitMiddleware,
56
+ redis=r,
57
+ rates=Rate(40, 60),
58
+ )
59
+
60
+ @app.get("/")
61
+ def index():
62
+ return {"ok": True}
63
+ ```
64
+
65
+ ### 2. Specific Endpoint Only
66
+
67
+ To rate limit a specific route instead of the whole application, use `RateLimiter` with `Depends`:
68
+
69
+ ```python
70
+ from fastapi import FastAPI, Depends
71
+ from redis import Redis
72
+ from p_redis_limiter import RateLimiter, Rate
73
+
74
+ app = FastAPI()
75
+ r = Redis(host="localhost", port=6379, decode_responses=True)
76
+
77
+ # 5 requests per 60 seconds for login only
78
+ login_limiter = RateLimiter(r, Rate(5, 60), prefix="rate:login")
79
+
80
+ @app.get("/")
81
+ def home():
82
+ return {"message": "unlimited"}
83
+
84
+ @app.post("/login", dependencies=[Depends(login_limiter.as_dependency())])
85
+ def login():
86
+ return {"message": "login successful"}
87
+ ```
88
+
89
+ ### 3. Multiple Rate Limits
90
+
91
+ You can define multiple rules (e.g. 5 req/sec burst limit and 100 req/min). If one rule fails, tokens are not deducted from the others:
92
+
93
+ ```python
94
+ rates = [
95
+ Rate(5, 1), # 5 requests per 1 second
96
+ Rate(100, 60), # 100 requests per 60 seconds
97
+ ]
98
+
99
+ # Or string shorthand:
100
+ rates = ["5/s", "100/m"]
101
+ ```
102
+
103
+ ### 4. Cache TTL
104
+
105
+ By default, Redis keys expire automatically after `max(window * 2, 60)` seconds. You can specify a custom TTL:
106
+
107
+ ```python
108
+ # Per rate rule:
109
+ Rate(40, 60, ttl=120)
110
+
111
+ # Disable TTL (persist in Redis forever):
112
+ Rate(40, 60, ttl=-1)
113
+
114
+ # Or globally on middleware / limiter:
115
+ RateLimitMiddleware(redis=r, rates=Rate(40, 60), ttl=120)
116
+ ```
117
+
118
+ ### 5. Manual Usage (Specify IP Yourself)
119
+
120
+ If you prefer to get the IP yourself, simply pass your IP variable directly:
121
+
122
+ ```python
123
+ from p_redis_limiter import RateLimiter, Rate
124
+
125
+ limiter = RateLimiter(r, Rate(5, 60))
126
+
127
+ @app.post("/login")
128
+ def login(request: Request):
129
+ user_ip = get_my_ip(request) # your own IP variable
130
+
131
+ # Simple boolean check:
132
+ if not limiter.is_allowed(user_ip):
133
+ return JSONResponse({"detail": "Too many requests"}, status_code=429)
134
+
135
+ return {"ok": True}
136
+ ```
137
+
138
+ Or get full details (`remaining`, `retry_after`):
139
+
140
+ ```python
141
+ result = limiter.check(user_ip)
142
+ if not result.allowed:
143
+ print(f"Blocked! Retry after {result.retry_after}s")
144
+ ```
145
+
146
+ ## Options
147
+
148
+ | Parameter | Type | Default | Description |
149
+ |---|---|---|---|
150
+ | `redis` | `Redis` | Required | Redis client instance (`redis.Redis` or `redis.asyncio.Redis`) |
151
+ | `rates` | `Rate` / `list` / `str` | Required | Rate rules, e.g. `Rate(40, 60)` or `["5/s", "100/m"]` |
152
+ | `ttl` | `int` | Auto | Custom Redis key TTL in seconds |
153
+ | `prefix` | `str` | `"rate"` | Prefix for Redis keys |
154
+ | `identifier` | `Callable` / `str` | Client IP | Custom function or header name (auto-detects Cloudflare, Nginx, ALB, direct IP) |
155
+ | `exclude_paths` | `list[str]` | `None` | List of paths to exclude from rate limiting |
156
+ | `error_detail` | `str` | `"Too many requests"` | Response detail message on 429 |
@@ -0,0 +1,128 @@
1
+ # p-redis-limiter
2
+
3
+ Redis token bucket rate limiter for Python and FastAPI.
4
+
5
+ ## Installation
6
+
7
+ ```bash
8
+ pip install git+https://github.com/mehranpng/p-redis-limiter.git
9
+ ```
10
+
11
+ ## Usage
12
+
13
+ ### 1. Global Middleware (All Endpoints)
14
+
15
+ Apply a rate limit across all endpoints:
16
+
17
+ ```python
18
+ from fastapi import FastAPI
19
+ from redis import Redis
20
+ from p_redis_limiter import RateLimitMiddleware, Rate
21
+
22
+ app = FastAPI()
23
+ r = Redis(host="localhost", port=6379, decode_responses=True)
24
+
25
+ # 40 requests per 60 seconds
26
+ app.add_middleware(
27
+ RateLimitMiddleware,
28
+ redis=r,
29
+ rates=Rate(40, 60),
30
+ )
31
+
32
+ @app.get("/")
33
+ def index():
34
+ return {"ok": True}
35
+ ```
36
+
37
+ ### 2. Specific Endpoint Only
38
+
39
+ To rate limit a specific route instead of the whole application, use `RateLimiter` with `Depends`:
40
+
41
+ ```python
42
+ from fastapi import FastAPI, Depends
43
+ from redis import Redis
44
+ from p_redis_limiter import RateLimiter, Rate
45
+
46
+ app = FastAPI()
47
+ r = Redis(host="localhost", port=6379, decode_responses=True)
48
+
49
+ # 5 requests per 60 seconds for login only
50
+ login_limiter = RateLimiter(r, Rate(5, 60), prefix="rate:login")
51
+
52
+ @app.get("/")
53
+ def home():
54
+ return {"message": "unlimited"}
55
+
56
+ @app.post("/login", dependencies=[Depends(login_limiter.as_dependency())])
57
+ def login():
58
+ return {"message": "login successful"}
59
+ ```
60
+
61
+ ### 3. Multiple Rate Limits
62
+
63
+ You can define multiple rules (e.g. 5 req/sec burst limit and 100 req/min). If one rule fails, tokens are not deducted from the others:
64
+
65
+ ```python
66
+ rates = [
67
+ Rate(5, 1), # 5 requests per 1 second
68
+ Rate(100, 60), # 100 requests per 60 seconds
69
+ ]
70
+
71
+ # Or string shorthand:
72
+ rates = ["5/s", "100/m"]
73
+ ```
74
+
75
+ ### 4. Cache TTL
76
+
77
+ By default, Redis keys expire automatically after `max(window * 2, 60)` seconds. You can specify a custom TTL:
78
+
79
+ ```python
80
+ # Per rate rule:
81
+ Rate(40, 60, ttl=120)
82
+
83
+ # Disable TTL (persist in Redis forever):
84
+ Rate(40, 60, ttl=-1)
85
+
86
+ # Or globally on middleware / limiter:
87
+ RateLimitMiddleware(redis=r, rates=Rate(40, 60), ttl=120)
88
+ ```
89
+
90
+ ### 5. Manual Usage (Specify IP Yourself)
91
+
92
+ If you prefer to get the IP yourself, simply pass your IP variable directly:
93
+
94
+ ```python
95
+ from p_redis_limiter import RateLimiter, Rate
96
+
97
+ limiter = RateLimiter(r, Rate(5, 60))
98
+
99
+ @app.post("/login")
100
+ def login(request: Request):
101
+ user_ip = get_my_ip(request) # your own IP variable
102
+
103
+ # Simple boolean check:
104
+ if not limiter.is_allowed(user_ip):
105
+ return JSONResponse({"detail": "Too many requests"}, status_code=429)
106
+
107
+ return {"ok": True}
108
+ ```
109
+
110
+ Or get full details (`remaining`, `retry_after`):
111
+
112
+ ```python
113
+ result = limiter.check(user_ip)
114
+ if not result.allowed:
115
+ print(f"Blocked! Retry after {result.retry_after}s")
116
+ ```
117
+
118
+ ## Options
119
+
120
+ | Parameter | Type | Default | Description |
121
+ |---|---|---|---|
122
+ | `redis` | `Redis` | Required | Redis client instance (`redis.Redis` or `redis.asyncio.Redis`) |
123
+ | `rates` | `Rate` / `list` / `str` | Required | Rate rules, e.g. `Rate(40, 60)` or `["5/s", "100/m"]` |
124
+ | `ttl` | `int` | Auto | Custom Redis key TTL in seconds |
125
+ | `prefix` | `str` | `"rate"` | Prefix for Redis keys |
126
+ | `identifier` | `Callable` / `str` | Client IP | Custom function or header name (auto-detects Cloudflare, Nginx, ALB, direct IP) |
127
+ | `exclude_paths` | `list[str]` | `None` | List of paths to exclude from rate limiting |
128
+ | `error_detail` | `str` | `"Too many requests"` | Response detail message on 429 |
@@ -0,0 +1,16 @@
1
+ """p-redis-limiter: Atomic Token Bucket Rate Limiter with Redis for Python & FastAPI."""
2
+
3
+ from p_redis_limiter.limiter import RateLimitResult, RateLimiter
4
+ from p_redis_limiter.middleware import RateLimitMiddleware, default_client_ip
5
+ from p_redis_limiter.rate import Rate
6
+
7
+ __version__ = "0.1.0"
8
+
9
+ __all__ = [
10
+ "Rate",
11
+ "RateLimiter",
12
+ "RateLimitResult",
13
+ "RateLimitMiddleware",
14
+ "default_client_ip",
15
+ "__version__",
16
+ ]
@@ -0,0 +1,264 @@
1
+ from __future__ import annotations
2
+
3
+ import asyncio
4
+ import inspect
5
+ import math
6
+ import time
7
+ from dataclasses import dataclass
8
+ from typing import Any, Callable, Iterable, List, Optional, Union
9
+
10
+ from p_redis_limiter.lua import TOKEN_BUCKET_LUA
11
+ from p_redis_limiter.rate import Rate
12
+
13
+ try:
14
+ from starlette.requests import Request
15
+ from starlette.exceptions import HTTPException
16
+ except ImportError:
17
+ Request = Any
18
+ HTTPException = Exception
19
+
20
+
21
+ @dataclass(frozen=True)
22
+ class RateLimitResult:
23
+ """Result of a rate limit check.
24
+
25
+ Attributes:
26
+ allowed: True if request is allowed, False if rate limited.
27
+ remaining: Remaining tokens in the bucket (minimum across all tiers).
28
+ retry_after: Seconds to wait before retrying (0.0 if allowed).
29
+ reset_in: Seconds until all buckets are fully refilled to capacity.
30
+ """
31
+
32
+ allowed: bool
33
+ remaining: int
34
+ retry_after: float
35
+ reset_in: float
36
+
37
+
38
+ def _default_key_builder(prefix: str, identifier: str, rate: Rate, is_multi: bool) -> str:
39
+ if not is_multi:
40
+ return f"{prefix}:{identifier}"
41
+ return f"{prefix}:{identifier}:{rate.window_tag}"
42
+
43
+
44
+ def _is_async_redis_client(client: Any) -> bool:
45
+ try:
46
+ from redis.asyncio import Redis as AsyncRedis, RedisCluster as AsyncRedisCluster
47
+ if isinstance(client, (AsyncRedis, AsyncRedisCluster)):
48
+ return True
49
+ except ImportError:
50
+ pass
51
+
52
+ module_name = getattr(client.__class__, "__module__", "")
53
+ if "asyncio" in module_name:
54
+ return True
55
+
56
+ for attr in ("execute_command", "eval", "evalsha"):
57
+ func = getattr(client, attr, None)
58
+ if func and inspect.iscoroutinefunction(func):
59
+ return True
60
+
61
+ return False
62
+
63
+
64
+ class RateLimiter:
65
+ """Atomic Token Bucket Rate Limiter powered by Redis.
66
+
67
+ Supports single or multiple rate limits, automatic/custom TTL, and both
68
+ synchronous and asynchronous Redis clients.
69
+
70
+ Usage examples:
71
+ limiter = RateLimiter(r, 40, 60)
72
+ limiter = RateLimiter(r, Rate(40, 60))
73
+ limiter = RateLimiter(r, "40/m")
74
+
75
+ limiter = RateLimiter(r, rates=[Rate(5, 1), Rate(100, 60)])
76
+ limiter = RateLimiter(r, rates=["5/s", "100/m"])
77
+ """
78
+
79
+ def __init__(
80
+ self,
81
+ redis: Any,
82
+ *args: Any,
83
+ rates: Optional[Union[Rate, str, tuple[int, Union[int, float]], Iterable[Any]]] = None,
84
+ requests: Optional[int] = None,
85
+ window: Optional[Union[int, float]] = None,
86
+ prefix: str = "rate",
87
+ ttl: Optional[int] = None,
88
+ key_builder: Optional[Callable[[str, str, Rate, bool], str]] = None,
89
+ ) -> None:
90
+ self.redis = redis
91
+ self.prefix = prefix
92
+ self.ttl = ttl
93
+ self._key_builder = key_builder or _default_key_builder
94
+ self._is_async = _is_async_redis_client(redis)
95
+
96
+ parsed_rates: List[Rate] = []
97
+
98
+ if len(args) == 2 and isinstance(args[0], int) and isinstance(args[1], (int, float)):
99
+ parsed_rates.append(Rate(requests=args[0], window=float(args[1]), ttl=ttl))
100
+ elif len(args) == 1:
101
+ rates = args[0]
102
+
103
+ if rates is not None:
104
+ if isinstance(rates, (Rate, str, tuple)):
105
+ parsed_rates.append(Rate.of(rates, default_ttl=ttl))
106
+ elif isinstance(rates, Iterable):
107
+ for r in rates:
108
+ parsed_rates.append(Rate.of(r, default_ttl=ttl))
109
+ elif requests is not None and window is not None:
110
+ parsed_rates.append(Rate(requests=requests, window=float(window), ttl=ttl))
111
+
112
+ if not parsed_rates:
113
+ raise ValueError(
114
+ "At least one rate must be provided. "
115
+ "E.g.: RateLimiter(r, 40, 60) or RateLimiter(r, rates=[Rate(5, 1), Rate(100, 60)])"
116
+ )
117
+
118
+ self.rates: tuple[Rate, ...] = tuple(parsed_rates)
119
+ self.is_multi: bool = len(self.rates) > 1
120
+
121
+ if hasattr(self.redis, "register_script"):
122
+ self._script = self.redis.register_script(TOKEN_BUCKET_LUA)
123
+ else:
124
+ self._script = None
125
+
126
+ @property
127
+ def primary_rate(self) -> Rate:
128
+ """The primary (first) rate configured."""
129
+ return self.rates[0]
130
+
131
+ def get_key(self, identifier: str, rate: Rate) -> str:
132
+ """Build the Redis key for an identifier and rate tier."""
133
+ return self._key_builder(self.prefix, identifier, rate, self.is_multi)
134
+
135
+ def _prepare_call(self, identifier: str, cost: int) -> tuple[list[str], list[Any]]:
136
+ if cost <= 0:
137
+ raise ValueError("cost must be a positive integer >= 1")
138
+
139
+ now = time.time()
140
+ keys = [self.get_key(identifier, r) for r in self.rates]
141
+
142
+ args: list[Any] = [now, cost]
143
+ for r in self.rates:
144
+ effective_ttl_s = r.effective_ttl(self.ttl)
145
+ ttl_ms = int(effective_ttl_s * 1000)
146
+ args.extend([r.requests, r.window, ttl_ms])
147
+
148
+ return keys, args
149
+
150
+ def _parse_result(self, raw: Any) -> RateLimitResult:
151
+ allowed = bool(raw[0])
152
+ remaining = int(raw[1])
153
+ retry_after = round(float(raw[2]), 4)
154
+ reset_in = round(float(raw[3]), 4)
155
+
156
+ return RateLimitResult(
157
+ allowed=allowed,
158
+ remaining=remaining,
159
+ retry_after=retry_after,
160
+ reset_in=reset_in,
161
+ )
162
+
163
+ def check(self, identifier: str, cost: int = 1) -> RateLimitResult:
164
+ """Check and consume rate limit synchronously.
165
+
166
+ Raises RuntimeError if initialized with an async Redis client.
167
+ """
168
+ if self._is_async:
169
+ raise RuntimeError(
170
+ "Cannot use synchronous check() with an async Redis client. "
171
+ "Use 'await limiter.check_async(...)' instead."
172
+ )
173
+
174
+ keys, args = self._prepare_call(identifier, cost)
175
+
176
+ if self._script is not None:
177
+ raw = self._script(keys=keys, args=args)
178
+ else:
179
+ raw = self.redis.eval(TOKEN_BUCKET_LUA, len(keys), *keys, *args)
180
+
181
+ return self._parse_result(raw)
182
+
183
+ async def check_async(self, identifier: str, cost: int = 1) -> RateLimitResult:
184
+ """Check and consume rate limit asynchronously.
185
+
186
+ Works with both async Redis clients (native await) and sync Redis clients
187
+ (dispatched to thread pool).
188
+ """
189
+ keys, args = self._prepare_call(identifier, cost)
190
+
191
+ if self._is_async:
192
+ if self._script is not None:
193
+ raw = await self._script(keys=keys, args=args)
194
+ else:
195
+ raw = await self.redis.eval(TOKEN_BUCKET_LUA, len(keys), *keys, *args)
196
+ else:
197
+ if self._script is not None:
198
+ raw = await asyncio.to_thread(self._script, keys=keys, args=args)
199
+ else:
200
+ raw = await asyncio.to_thread(
201
+ self.redis.eval, TOKEN_BUCKET_LUA, len(keys), *keys, *args
202
+ )
203
+
204
+ return self._parse_result(raw)
205
+
206
+ def is_allowed(self, identifier: str, cost: int = 1) -> bool:
207
+ """Convenience method returning True if request is allowed, False otherwise."""
208
+ return self.check(identifier, cost).allowed
209
+
210
+ async def is_allowed_async(self, identifier: str, cost: int = 1) -> bool:
211
+ """Async convenience method returning True if request is allowed, False otherwise."""
212
+ res = await self.check_async(identifier, cost)
213
+ return res.allowed
214
+
215
+ def check_request(
216
+ self, request: Any, ip: Optional[str] = None, cost: int = 1
217
+ ) -> RateLimitResult:
218
+ from p_redis_limiter.middleware import default_client_ip
219
+ ident = ip if ip is not None else default_client_ip(request)
220
+ return self.check(ident, cost=cost)
221
+
222
+ async def check_request_async(
223
+ self, request: Any, ip: Optional[str] = None, cost: int = 1
224
+ ) -> RateLimitResult:
225
+ from p_redis_limiter.middleware import default_client_ip
226
+ ident = ip if ip is not None else default_client_ip(request)
227
+ return await self.check_async(ident, cost=cost)
228
+
229
+ def as_dependency(
230
+ self,
231
+ cost: int = 1,
232
+ identifier: Optional[Union[str, Callable[[Any], str]]] = None,
233
+ identifier_func: Optional[Callable[[Any], str]] = None,
234
+ error_detail: str = "Too many requests",
235
+ ):
236
+ from p_redis_limiter.middleware import default_client_ip
237
+
238
+ target_ident = identifier or identifier_func
239
+ if callable(target_ident):
240
+ ident_fn = target_ident
241
+ elif isinstance(target_ident, str):
242
+ ident_fn = lambda req: req.headers.get(target_ident) or default_client_ip(req)
243
+ else:
244
+ ident_fn = default_client_ip
245
+
246
+ async def _rate_limit_dependency(request: Request):
247
+ ident = ident_fn(request)
248
+ result = await self.check_async(ident, cost=cost)
249
+
250
+ if not result.allowed:
251
+ headers = {
252
+ "Retry-After": str(max(1, int(math.ceil(result.retry_after)))),
253
+ "X-RateLimit-Limit": str(self.primary_rate.requests),
254
+ "X-RateLimit-Remaining": "0",
255
+ "X-RateLimit-Reset": str(max(1, int(math.ceil(result.reset_in)))),
256
+ }
257
+ raise HTTPException(
258
+ status_code=429,
259
+ detail=error_detail,
260
+ headers=headers,
261
+ )
262
+ return result
263
+
264
+ return _rate_limit_dependency
@@ -0,0 +1,99 @@
1
+ """Atomic Multi-tier Token Bucket Lua Script for Redis."""
2
+
3
+ TOKEN_BUCKET_LUA = """
4
+ local now = tonumber(ARGV[1])
5
+ local cost = tonumber(ARGV[2])
6
+ local num_keys = #KEYS
7
+
8
+ local allowed = 1
9
+ local min_remaining = 999999999
10
+ local max_retry_after = 0
11
+ local max_reset_in = 0
12
+
13
+ local buckets_tokens = {}
14
+ local buckets_last_refill = {}
15
+ local capacities = {}
16
+ local windows = {}
17
+ local ttls = {}
18
+
19
+ for i = 1, num_keys do
20
+ local offset = 2 + (i - 1) * 3
21
+ local capacity = tonumber(ARGV[offset + 1])
22
+ local window = tonumber(ARGV[offset + 2])
23
+ local ttl_ms = tonumber(ARGV[offset + 3])
24
+
25
+ capacities[i] = capacity
26
+ windows[i] = window
27
+ ttls[i] = ttl_ms
28
+
29
+ local key = KEYS[i]
30
+ local bucket = redis.call('HMGET', key, 'tokens', 'last_refill')
31
+ local tokens = tonumber(bucket[1])
32
+ local last_refill = tonumber(bucket[2])
33
+
34
+ if tokens == nil or last_refill == nil then
35
+ tokens = capacity
36
+ last_refill = now
37
+ else
38
+ local time_passed = now - last_refill
39
+ if time_passed > 0 then
40
+ local refill_rate = capacity / window
41
+ tokens = math.min(capacity, tokens + (time_passed * refill_rate))
42
+ last_refill = now
43
+ end
44
+ end
45
+
46
+ buckets_tokens[i] = tokens
47
+ buckets_last_refill[i] = last_refill
48
+
49
+ if tokens < cost then
50
+ allowed = 0
51
+ local needed = cost - tokens
52
+ local refill_rate = capacity / window
53
+ local retry_after = needed / refill_rate
54
+ if retry_after > max_retry_after then
55
+ max_retry_after = retry_after
56
+ end
57
+ end
58
+ end
59
+
60
+ for i = 1, num_keys do
61
+ local key = KEYS[i]
62
+ local tokens = buckets_tokens[i]
63
+ local last_refill = buckets_last_refill[i]
64
+ local capacity = capacities[i]
65
+ local window = windows[i]
66
+ local ttl_ms = ttls[i]
67
+
68
+ if allowed == 1 then
69
+ tokens = tokens - cost
70
+ local rem = math.floor(tokens)
71
+ if rem < min_remaining then
72
+ min_remaining = rem
73
+ end
74
+ else
75
+ min_remaining = 0
76
+ end
77
+
78
+ local missing = capacity - tokens
79
+ if missing > 0 then
80
+ local reset_in = missing / (capacity / window)
81
+ if reset_in > max_reset_in then
82
+ max_reset_in = reset_in
83
+ end
84
+ end
85
+
86
+ redis.call('HSET', key, 'tokens', tostring(tokens), 'last_refill', tostring(last_refill))
87
+ if ttl_ms > 0 then
88
+ redis.call('PEXPIRE', key, ttl_ms)
89
+ else
90
+ redis.call('PERSIST', key)
91
+ end
92
+ end
93
+
94
+ if min_remaining == 999999999 then
95
+ min_remaining = 0
96
+ end
97
+
98
+ return {allowed, min_remaining, tostring(max_retry_after), tostring(max_reset_in)}
99
+ """
@@ -0,0 +1,148 @@
1
+ from __future__ import annotations
2
+
3
+ import inspect
4
+ import math
5
+ from typing import Any, Awaitable, Callable, Iterable, Optional, Set, Union
6
+
7
+ try:
8
+ from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
9
+ from starlette.requests import Request
10
+ from starlette.responses import JSONResponse, Response
11
+ from starlette.types import ASGIApp
12
+ STARLETTE_AVAILABLE = True
13
+ except ImportError:
14
+ STARLETTE_AVAILABLE = False
15
+ BaseHTTPMiddleware = object
16
+ RequestResponseEndpoint = Any
17
+ Request = Any
18
+ Response = Any
19
+ JSONResponse = Any
20
+ ASGIApp = Any
21
+
22
+ from p_redis_limiter.limiter import RateLimitResult, RateLimiter
23
+ from p_redis_limiter.rate import Rate
24
+
25
+
26
+ def default_client_ip(request: Request) -> str:
27
+ cf_ip = request.headers.get("CF-Connecting-IP")
28
+ if cf_ip:
29
+ return cf_ip.strip()
30
+
31
+ true_client_ip = request.headers.get("True-Client-IP")
32
+ if true_client_ip:
33
+ return true_client_ip.strip()
34
+
35
+ real_ip = request.headers.get("X-Real-IP")
36
+ if real_ip:
37
+ return real_ip.strip()
38
+
39
+ forwarded = request.headers.get("X-Forwarded-For")
40
+ if forwarded:
41
+ return forwarded.split(",")[0].strip()
42
+
43
+ if request.client and request.client.host:
44
+ return request.client.host
45
+
46
+ return "127.0.0.1"
47
+
48
+
49
+ class RateLimitMiddleware(BaseHTTPMiddleware):
50
+ """FastAPI & Starlette Middleware for Token Bucket Rate Limiting.
51
+
52
+ Usage:
53
+ app.add_middleware(
54
+ RateLimitMiddleware,
55
+ redis=r,
56
+ rates=[Rate(5, 1), Rate(40, 60)],
57
+ )
58
+ """
59
+
60
+ def __init__(
61
+ self,
62
+ app: ASGIApp,
63
+ redis: Any,
64
+ rates: Optional[Union[Rate, str, tuple[int, Union[int, float]], Iterable[Any]]] = None,
65
+ *,
66
+ requests: Optional[int] = None,
67
+ window: Optional[Union[int, float]] = None,
68
+ prefix: str = "rate",
69
+ ttl: Optional[int] = None,
70
+ identifier: Optional[Union[str, Callable[[Request], str]]] = None,
71
+ error_detail: str = "Too many requests",
72
+ status_code: int = 429,
73
+ set_headers: bool = True,
74
+ exclude_paths: Optional[Iterable[str]] = None,
75
+ on_blocked: Optional[
76
+ Callable[[Request, RateLimitResult], Union[Response, Awaitable[Response]]]
77
+ ] = None,
78
+ ) -> None:
79
+ if not STARLETTE_AVAILABLE:
80
+ raise ImportError(
81
+ "Starlette/FastAPI is required to use RateLimitMiddleware. "
82
+ "Install it with: pip install fastapi"
83
+ )
84
+ super().__init__(app)
85
+
86
+ if isinstance(redis, RateLimiter):
87
+ self.limiter = redis
88
+ else:
89
+ self.limiter = RateLimiter(
90
+ redis=redis,
91
+ rates=rates,
92
+ requests=requests,
93
+ window=window,
94
+ prefix=prefix,
95
+ ttl=ttl,
96
+ )
97
+
98
+ if callable(identifier):
99
+ self.identifier = identifier
100
+ elif isinstance(identifier, str):
101
+ self.identifier = lambda req: req.headers.get(identifier) or default_client_ip(req)
102
+ else:
103
+ self.identifier = default_client_ip
104
+ self.error_detail = error_detail
105
+ self.status_code = status_code
106
+ self.set_headers = set_headers
107
+ self.exclude_paths: Set[str] = set(exclude_paths) if exclude_paths else set()
108
+ self.on_blocked = on_blocked
109
+
110
+ async def dispatch(
111
+ self, request: Request, call_next: RequestResponseEndpoint
112
+ ) -> Response:
113
+ if self.exclude_paths and request.url.path in self.exclude_paths:
114
+ return await call_next(request)
115
+
116
+ ident = self.identifier(request)
117
+ result = await self.limiter.check_async(ident)
118
+
119
+ if not result.allowed:
120
+ if self.on_blocked is not None:
121
+ custom_resp = self.on_blocked(request, result)
122
+ if inspect.isawaitable(custom_resp):
123
+ return await custom_resp
124
+ return custom_resp
125
+
126
+ headers = {}
127
+ if self.set_headers:
128
+ headers = {
129
+ "Retry-After": str(max(1, int(math.ceil(result.retry_after)))),
130
+ "X-RateLimit-Limit": str(self.limiter.primary_rate.requests),
131
+ "X-RateLimit-Remaining": "0",
132
+ "X-RateLimit-Reset": str(max(1, int(math.ceil(result.reset_in)))),
133
+ }
134
+
135
+ return JSONResponse(
136
+ {"detail": self.error_detail},
137
+ status_code=self.status_code,
138
+ headers=headers if headers else None,
139
+ )
140
+
141
+ response = await call_next(request)
142
+
143
+ if self.set_headers:
144
+ response.headers["X-RateLimit-Limit"] = str(self.limiter.primary_rate.requests)
145
+ response.headers["X-RateLimit-Remaining"] = str(result.remaining)
146
+ response.headers["X-RateLimit-Reset"] = str(max(1, int(math.ceil(result.reset_in))))
147
+
148
+ return response
File without changes
@@ -0,0 +1,121 @@
1
+ from __future__ import annotations
2
+
3
+ import math
4
+ import re
5
+ from dataclasses import dataclass
6
+ from typing import Any, Optional, Union
7
+
8
+ _TIME_UNITS = {
9
+ "s": 1.0,
10
+ "sec": 1.0,
11
+ "second": 1.0,
12
+ "seconds": 1.0,
13
+ "m": 60.0,
14
+ "min": 60.0,
15
+ "minute": 60.0,
16
+ "minutes": 60.0,
17
+ "h": 3600.0,
18
+ "hr": 3600.0,
19
+ "hour": 3600.0,
20
+ "hours": 3600.0,
21
+ "d": 86400.0,
22
+ "day": 86400.0,
23
+ "days": 86400.0,
24
+ }
25
+
26
+ _RATE_REGEX = re.compile(
27
+ r"^\s*(\d+)\s*\/\s*(?:(\d+(?:\.\d+)?)\s*)?([a-zA-Z]+)?\s*$"
28
+ )
29
+
30
+
31
+ @dataclass(frozen=True)
32
+ class Rate:
33
+ """Represents a rate limit rule: allowed `requests` within `window` seconds.
34
+
35
+ Args:
36
+ requests: Number of requests allowed within the window.
37
+ window: Window duration in seconds (int or float).
38
+ ttl: Optional custom TTL in seconds for Redis cache.
39
+ If not set, defaults to max(ceil(window * 2), 60).
40
+ """
41
+
42
+ requests: int
43
+ window: float
44
+ ttl: Optional[int] = None
45
+
46
+ def __post_init__(self) -> None:
47
+ if not isinstance(self.requests, int) or self.requests <= 0:
48
+ raise ValueError(f"requests must be a positive integer > 0, got: {self.requests}")
49
+ if self.window <= 0:
50
+ raise ValueError(f"window must be a positive number > 0, got: {self.window}")
51
+ if self.ttl is not None and self.ttl < -1:
52
+ raise ValueError(f"ttl must be a positive integer, 0, or -1 (to disable TTL), got: {self.ttl}")
53
+
54
+ def effective_ttl(self, default_ttl: Optional[int] = None) -> int:
55
+ target = self.ttl if self.ttl is not None else default_ttl
56
+ if target is not None:
57
+ return 0 if target in (0, -1) else int(target)
58
+ return max(int(math.ceil(self.window * 2)), 60)
59
+
60
+ @property
61
+ def window_tag(self) -> str:
62
+ """Compact string representation of the window duration for key suffix."""
63
+ if float(self.window).is_integer():
64
+ return f"{int(self.window)}s"
65
+ return f"{self.window}s"
66
+
67
+ @classmethod
68
+ def parse(cls, rate_str: str, ttl: Optional[int] = None) -> Rate:
69
+ """Parse human-friendly rate string such as '5/s', '40/minute', '100/10s'.
70
+
71
+ Examples:
72
+ Rate.parse("5/s") -> Rate(requests=5, window=1.0)
73
+ Rate.parse("40/minute") -> Rate(requests=40, window=60.0)
74
+ Rate.parse("10/2m") -> Rate(requests=10, window=120.0)
75
+ Rate.parse("100/30s", ttl=120) -> Rate(requests=100, window=30.0, ttl=120)
76
+ """
77
+ match = _RATE_REGEX.match(rate_str.strip())
78
+ if not match:
79
+ raise ValueError(
80
+ f"Invalid rate limit string: '{rate_str}'. "
81
+ "Expected format like '5/s', '40/minute', '100/30s', etc."
82
+ )
83
+
84
+ req_str, mult_str, unit_str = match.groups()
85
+ requests = int(req_str)
86
+
87
+ multiplier = float(mult_str) if mult_str else 1.0
88
+
89
+ if unit_str:
90
+ unit_lower = unit_str.lower()
91
+ if unit_lower not in _TIME_UNITS:
92
+ valid = ", ".join(_TIME_UNITS.keys())
93
+ raise ValueError(
94
+ f"Unknown time unit '{unit_str}' in rate '{rate_str}'. Valid units: {valid}"
95
+ )
96
+ base_seconds = _TIME_UNITS[unit_lower]
97
+ else:
98
+ base_seconds = 1.0
99
+
100
+ window = multiplier * base_seconds
101
+ return cls(requests=requests, window=window, ttl=ttl)
102
+
103
+ @classmethod
104
+ def of(cls, value: Union[Rate, str, tuple[int, Union[int, float]], list[Any]], default_ttl: Optional[int] = None) -> Rate:
105
+ """Convert various input types into a Rate instance."""
106
+ if isinstance(value, Rate):
107
+ if value.ttl is None and default_ttl is not None:
108
+ return cls(requests=value.requests, window=value.window, ttl=default_ttl)
109
+ return value
110
+ if isinstance(value, str):
111
+ return cls.parse(value, ttl=default_ttl)
112
+ if isinstance(value, (tuple, list)) and len(value) >= 2:
113
+ custom_ttl = value[2] if len(value) > 2 else default_ttl
114
+ return cls(requests=int(value[0]), window=float(value[1]), ttl=custom_ttl)
115
+ raise TypeError(
116
+ f"Cannot convert {type(value).__name__} to Rate. "
117
+ "Pass a Rate instance, a string (e.g. '40/m'), or a (requests, window) tuple."
118
+ )
119
+
120
+ def __str__(self) -> str:
121
+ return f"{self.requests}/{self.window_tag}"
@@ -0,0 +1,156 @@
1
+ Metadata-Version: 2.4
2
+ Name: p-redis-limiter
3
+ Version: 0.1.0
4
+ Summary: A fast, atomic, multi-tier Token Bucket rate limiter for Python and FastAPI using Redis.
5
+ Author: Mehran
6
+ License: MIT
7
+ Keywords: redis,rate-limiter,token-bucket,fastapi,throttling
8
+ Classifier: Development Status :: 4 - Beta
9
+ Classifier: Intended Audience :: Developers
10
+ Classifier: Programming Language :: Python :: 3
11
+ Classifier: Programming Language :: Python :: 3.9
12
+ Classifier: Programming Language :: Python :: 3.10
13
+ Classifier: Programming Language :: Python :: 3.11
14
+ Classifier: Programming Language :: Python :: 3.12
15
+ Classifier: Programming Language :: Python :: 3.13
16
+ Classifier: Programming Language :: Python :: 3.14
17
+ Classifier: License :: OSI Approved :: MIT License
18
+ Classifier: Operating System :: OS Independent
19
+ Classifier: Framework :: FastAPI
20
+ Requires-Python: >=3.9
21
+ Description-Content-Type: text/markdown
22
+ Requires-Dist: redis>=4.2.0
23
+ Provides-Extra: dev
24
+ Requires-Dist: pytest>=7.0.0; extra == "dev"
25
+ Requires-Dist: pytest-asyncio>=0.20.0; extra == "dev"
26
+ Requires-Dist: httpx>=0.23.0; extra == "dev"
27
+ Requires-Dist: fastapi>=0.70.0; extra == "dev"
28
+
29
+ # p-redis-limiter
30
+
31
+ Redis token bucket rate limiter for Python and FastAPI.
32
+
33
+ ## Installation
34
+
35
+ ```bash
36
+ pip install git+https://github.com/mehranpng/p-redis-limiter.git
37
+ ```
38
+
39
+ ## Usage
40
+
41
+ ### 1. Global Middleware (All Endpoints)
42
+
43
+ Apply a rate limit across all endpoints:
44
+
45
+ ```python
46
+ from fastapi import FastAPI
47
+ from redis import Redis
48
+ from p_redis_limiter import RateLimitMiddleware, Rate
49
+
50
+ app = FastAPI()
51
+ r = Redis(host="localhost", port=6379, decode_responses=True)
52
+
53
+ # 40 requests per 60 seconds
54
+ app.add_middleware(
55
+ RateLimitMiddleware,
56
+ redis=r,
57
+ rates=Rate(40, 60),
58
+ )
59
+
60
+ @app.get("/")
61
+ def index():
62
+ return {"ok": True}
63
+ ```
64
+
65
+ ### 2. Specific Endpoint Only
66
+
67
+ To rate limit a specific route instead of the whole application, use `RateLimiter` with `Depends`:
68
+
69
+ ```python
70
+ from fastapi import FastAPI, Depends
71
+ from redis import Redis
72
+ from p_redis_limiter import RateLimiter, Rate
73
+
74
+ app = FastAPI()
75
+ r = Redis(host="localhost", port=6379, decode_responses=True)
76
+
77
+ # 5 requests per 60 seconds for login only
78
+ login_limiter = RateLimiter(r, Rate(5, 60), prefix="rate:login")
79
+
80
+ @app.get("/")
81
+ def home():
82
+ return {"message": "unlimited"}
83
+
84
+ @app.post("/login", dependencies=[Depends(login_limiter.as_dependency())])
85
+ def login():
86
+ return {"message": "login successful"}
87
+ ```
88
+
89
+ ### 3. Multiple Rate Limits
90
+
91
+ You can define multiple rules (e.g. 5 req/sec burst limit and 100 req/min). If one rule fails, tokens are not deducted from the others:
92
+
93
+ ```python
94
+ rates = [
95
+ Rate(5, 1), # 5 requests per 1 second
96
+ Rate(100, 60), # 100 requests per 60 seconds
97
+ ]
98
+
99
+ # Or string shorthand:
100
+ rates = ["5/s", "100/m"]
101
+ ```
102
+
103
+ ### 4. Cache TTL
104
+
105
+ By default, Redis keys expire automatically after `max(window * 2, 60)` seconds. You can specify a custom TTL:
106
+
107
+ ```python
108
+ # Per rate rule:
109
+ Rate(40, 60, ttl=120)
110
+
111
+ # Disable TTL (persist in Redis forever):
112
+ Rate(40, 60, ttl=-1)
113
+
114
+ # Or globally on middleware / limiter:
115
+ RateLimitMiddleware(redis=r, rates=Rate(40, 60), ttl=120)
116
+ ```
117
+
118
+ ### 5. Manual Usage (Specify IP Yourself)
119
+
120
+ If you prefer to get the IP yourself, simply pass your IP variable directly:
121
+
122
+ ```python
123
+ from p_redis_limiter import RateLimiter, Rate
124
+
125
+ limiter = RateLimiter(r, Rate(5, 60))
126
+
127
+ @app.post("/login")
128
+ def login(request: Request):
129
+ user_ip = get_my_ip(request) # your own IP variable
130
+
131
+ # Simple boolean check:
132
+ if not limiter.is_allowed(user_ip):
133
+ return JSONResponse({"detail": "Too many requests"}, status_code=429)
134
+
135
+ return {"ok": True}
136
+ ```
137
+
138
+ Or get full details (`remaining`, `retry_after`):
139
+
140
+ ```python
141
+ result = limiter.check(user_ip)
142
+ if not result.allowed:
143
+ print(f"Blocked! Retry after {result.retry_after}s")
144
+ ```
145
+
146
+ ## Options
147
+
148
+ | Parameter | Type | Default | Description |
149
+ |---|---|---|---|
150
+ | `redis` | `Redis` | Required | Redis client instance (`redis.Redis` or `redis.asyncio.Redis`) |
151
+ | `rates` | `Rate` / `list` / `str` | Required | Rate rules, e.g. `Rate(40, 60)` or `["5/s", "100/m"]` |
152
+ | `ttl` | `int` | Auto | Custom Redis key TTL in seconds |
153
+ | `prefix` | `str` | `"rate"` | Prefix for Redis keys |
154
+ | `identifier` | `Callable` / `str` | Client IP | Custom function or header name (auto-detects Cloudflare, Nginx, ALB, direct IP) |
155
+ | `exclude_paths` | `list[str]` | `None` | List of paths to exclude from rate limiting |
156
+ | `error_detail` | `str` | `"Too many requests"` | Response detail message on 429 |
@@ -0,0 +1,13 @@
1
+ README.md
2
+ pyproject.toml
3
+ p_redis_limiter/__init__.py
4
+ p_redis_limiter/limiter.py
5
+ p_redis_limiter/lua.py
6
+ p_redis_limiter/middleware.py
7
+ p_redis_limiter/py.typed
8
+ p_redis_limiter/rate.py
9
+ p_redis_limiter.egg-info/PKG-INFO
10
+ p_redis_limiter.egg-info/SOURCES.txt
11
+ p_redis_limiter.egg-info/dependency_links.txt
12
+ p_redis_limiter.egg-info/requires.txt
13
+ p_redis_limiter.egg-info/top_level.txt
@@ -0,0 +1,7 @@
1
+ redis>=4.2.0
2
+
3
+ [dev]
4
+ pytest>=7.0.0
5
+ pytest-asyncio>=0.20.0
6
+ httpx>=0.23.0
7
+ fastapi>=0.70.0
@@ -0,0 +1 @@
1
+ p_redis_limiter
@@ -0,0 +1,47 @@
1
+ [build-system]
2
+ requires = ["setuptools>=61.0"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "p-redis-limiter"
7
+ version = "0.1.0"
8
+ description = "A fast, atomic, multi-tier Token Bucket rate limiter for Python and FastAPI using Redis."
9
+ readme = "README.md"
10
+ requires-python = ">=3.9"
11
+ license = { text = "MIT" }
12
+ authors = [
13
+ { name = "Mehran" }
14
+ ]
15
+ keywords = ["redis", "rate-limiter", "token-bucket", "fastapi", "throttling"]
16
+ classifiers = [
17
+ "Development Status :: 4 - Beta",
18
+ "Intended Audience :: Developers",
19
+ "Programming Language :: Python :: 3",
20
+ "Programming Language :: Python :: 3.9",
21
+ "Programming Language :: Python :: 3.10",
22
+ "Programming Language :: Python :: 3.11",
23
+ "Programming Language :: Python :: 3.12",
24
+ "Programming Language :: Python :: 3.13",
25
+ "Programming Language :: Python :: 3.14",
26
+ "License :: OSI Approved :: MIT License",
27
+ "Operating System :: OS Independent",
28
+ "Framework :: FastAPI",
29
+ ]
30
+ dependencies = [
31
+ "redis>=4.2.0",
32
+ ]
33
+
34
+ [project.optional-dependencies]
35
+ dev = [
36
+ "pytest>=7.0.0",
37
+ "pytest-asyncio>=0.20.0",
38
+ "httpx>=0.23.0",
39
+ "fastapi>=0.70.0",
40
+ ]
41
+
42
+ [tool.setuptools.packages.find]
43
+ where = ["."]
44
+ include = ["p_redis_limiter*"]
45
+
46
+ [tool.setuptools.package-data]
47
+ p_redis_limiter = ["py.typed"]
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+