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.
- p_redis_limiter-0.1.0/PKG-INFO +156 -0
- p_redis_limiter-0.1.0/README.md +128 -0
- p_redis_limiter-0.1.0/p_redis_limiter/__init__.py +16 -0
- p_redis_limiter-0.1.0/p_redis_limiter/limiter.py +264 -0
- p_redis_limiter-0.1.0/p_redis_limiter/lua.py +99 -0
- p_redis_limiter-0.1.0/p_redis_limiter/middleware.py +148 -0
- p_redis_limiter-0.1.0/p_redis_limiter/py.typed +0 -0
- p_redis_limiter-0.1.0/p_redis_limiter/rate.py +121 -0
- p_redis_limiter-0.1.0/p_redis_limiter.egg-info/PKG-INFO +156 -0
- p_redis_limiter-0.1.0/p_redis_limiter.egg-info/SOURCES.txt +13 -0
- p_redis_limiter-0.1.0/p_redis_limiter.egg-info/dependency_links.txt +1 -0
- p_redis_limiter-0.1.0/p_redis_limiter.egg-info/requires.txt +7 -0
- p_redis_limiter-0.1.0/p_redis_limiter.egg-info/top_level.txt +1 -0
- p_redis_limiter-0.1.0/pyproject.toml +47 -0
- p_redis_limiter-0.1.0/setup.cfg +4 -0
|
@@ -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 @@
|
|
|
1
|
+
|
|
@@ -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"]
|