237 lines
8.4 KiB
Python
237 lines
8.4 KiB
Python
|
|
import asyncio
|
|||
|
|
import contextlib
|
|||
|
|
import time
|
|||
|
|
from collections import defaultdict
|
|||
|
|
from collections.abc import Callable
|
|||
|
|
from dataclasses import dataclass, field
|
|||
|
|
from enum import StrEnum
|
|||
|
|
from functools import wraps
|
|||
|
|
|
|||
|
|
|
|||
|
|
class RateLimitScope(StrEnum):
|
|||
|
|
LOGIN = "login"
|
|||
|
|
API = "api"
|
|||
|
|
SIGNUP = "signup"
|
|||
|
|
PASSWORD_RESET = "password_reset"
|
|||
|
|
TOKEN_REFRESH = "token_refresh"
|
|||
|
|
|
|||
|
|
|
|||
|
|
@dataclass
|
|||
|
|
class ScopeConfig:
|
|||
|
|
max_attempts: int = 10
|
|||
|
|
window_seconds: int = 60
|
|||
|
|
cooldown_seconds: int = 300
|
|||
|
|
loopback_exempt: bool = True
|
|||
|
|
|
|||
|
|
|
|||
|
|
DEFAULT_SCOPE_CONFIGS = {
|
|||
|
|
RateLimitScope.LOGIN: ScopeConfig(max_attempts=10, window_seconds=60, cooldown_seconds=300),
|
|||
|
|
RateLimitScope.API: ScopeConfig(max_attempts=100, window_seconds=60, cooldown_seconds=60),
|
|||
|
|
RateLimitScope.SIGNUP: ScopeConfig(max_attempts=3, window_seconds=3600, cooldown_seconds=7200),
|
|||
|
|
RateLimitScope.PASSWORD_RESET: ScopeConfig(max_attempts=5, window_seconds=300, cooldown_seconds=600),
|
|||
|
|
RateLimitScope.TOKEN_REFRESH: ScopeConfig(max_attempts=30, window_seconds=60, cooldown_seconds=120),
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
_LOCAL_IPS = frozenset({"127.0.0.1", "::1", "localhost"})
|
|||
|
|
|
|||
|
|
|
|||
|
|
@dataclass
|
|||
|
|
class _BucketState:
|
|||
|
|
timestamps: list[float] = field(default_factory=list)
|
|||
|
|
cooldown_until: float = 0.0
|
|||
|
|
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
|||
|
|
|
|||
|
|
|
|||
|
|
class AuthRateLimiter:
|
|||
|
|
"""认证级限流器:内存滑动窗口 + 多 Scope + 本地回环豁免 + 锁定冷却。
|
|||
|
|
|
|||
|
|
特性:
|
|||
|
|
- 多 Scope 独立配置(login/api/signup/password_reset/token_refresh)
|
|||
|
|
- 滑动窗口计数,窗口外记录自动淘汰
|
|||
|
|
- 本地回环 (127.0.0.1, ::1) 可豁免
|
|||
|
|
- 锁定冷却期 5min(超限后拒绝所有尝试)
|
|||
|
|
- asyncio.Lock 串行化防竞态
|
|||
|
|
- 后台定期清理失效 bucket,防止 DDoS 下内存持续增长
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
_CLEANUP_INTERVAL = 300 # seconds between stale bucket cleanups
|
|||
|
|
|
|||
|
|
def __init__(self, scope_configs: dict[RateLimitScope, ScopeConfig] | None = None):
|
|||
|
|
self._configs = scope_configs or DEFAULT_SCOPE_CONFIGS
|
|||
|
|
self._buckets: dict[str, dict[RateLimitScope, _BucketState]] = defaultdict(dict)
|
|||
|
|
self._lock = asyncio.Lock()
|
|||
|
|
self._last_cleanup: float = 0.0
|
|||
|
|
self._cleanup_task: asyncio.Task | None = None
|
|||
|
|
self._start_background_cleanup()
|
|||
|
|
|
|||
|
|
def _is_loopback(self, identifier: str) -> bool:
|
|||
|
|
if identifier in _LOCAL_IPS:
|
|||
|
|
return True
|
|||
|
|
host = identifier.split(":")[0]
|
|||
|
|
if host.startswith("127."):
|
|||
|
|
return True
|
|||
|
|
try:
|
|||
|
|
import ipaddress
|
|||
|
|
|
|||
|
|
ip = ipaddress.ip_address(host)
|
|||
|
|
return ip.is_loopback
|
|||
|
|
except ValueError:
|
|||
|
|
return False
|
|||
|
|
|
|||
|
|
async def _get_bucket(self, identifier: str, scope: RateLimitScope) -> _BucketState:
|
|||
|
|
async with self._lock:
|
|||
|
|
scope_buckets = self._buckets[identifier]
|
|||
|
|
if scope not in scope_buckets:
|
|||
|
|
scope_buckets[scope] = _BucketState()
|
|||
|
|
return scope_buckets[scope]
|
|||
|
|
|
|||
|
|
async def _cleanup_stale_buckets(self, now: float) -> None:
|
|||
|
|
async with self._lock:
|
|||
|
|
stale_identifiers: list[str] = []
|
|||
|
|
for identifier, scopes in self._buckets.items():
|
|||
|
|
active_scopes: dict[RateLimitScope, _BucketState] = {}
|
|||
|
|
for scope, bucket in scopes.items():
|
|||
|
|
config = self._configs.get(scope)
|
|||
|
|
window = config.window_seconds if config else 60
|
|||
|
|
bucket.timestamps = [t for t in bucket.timestamps if t > now - window]
|
|||
|
|
has_activity = bucket.timestamps or bucket.cooldown_until > now
|
|||
|
|
if has_activity:
|
|||
|
|
active_scopes[scope] = bucket
|
|||
|
|
if active_scopes:
|
|||
|
|
self._buckets[identifier] = active_scopes
|
|||
|
|
else:
|
|||
|
|
stale_identifiers.append(identifier)
|
|||
|
|
for identifier in stale_identifiers:
|
|||
|
|
del self._buckets[identifier]
|
|||
|
|
|
|||
|
|
def _start_background_cleanup(self) -> None:
|
|||
|
|
async def _loop() -> None:
|
|||
|
|
while True:
|
|||
|
|
await asyncio.sleep(self._CLEANUP_INTERVAL)
|
|||
|
|
await self._cleanup_stale_buckets(time.monotonic())
|
|||
|
|
|
|||
|
|
try:
|
|||
|
|
loop = asyncio.get_running_loop()
|
|||
|
|
self._cleanup_task = loop.create_task(_loop())
|
|||
|
|
except RuntimeError:
|
|||
|
|
pass
|
|||
|
|
|
|||
|
|
async def shutdown(self) -> None:
|
|||
|
|
if self._cleanup_task and not self._cleanup_task.done():
|
|||
|
|
self._cleanup_task.cancel()
|
|||
|
|
with contextlib.suppress(asyncio.CancelledError):
|
|||
|
|
await self._cleanup_task
|
|||
|
|
|
|||
|
|
async def check_and_record(
|
|||
|
|
self,
|
|||
|
|
identifier: str,
|
|||
|
|
scope: RateLimitScope,
|
|||
|
|
) -> bool:
|
|||
|
|
"""检查是否允许请求,并记录本次尝试。
|
|||
|
|
|
|||
|
|
返回 True 表示允许,False 表示被限流。
|
|||
|
|
"""
|
|||
|
|
config = self._configs.get(scope)
|
|||
|
|
if config is None:
|
|||
|
|
return True
|
|||
|
|
|
|||
|
|
if config.loopback_exempt and self._is_loopback(identifier):
|
|||
|
|
return True
|
|||
|
|
|
|||
|
|
now = time.monotonic()
|
|||
|
|
|
|||
|
|
if now - self._last_cleanup > self._CLEANUP_INTERVAL:
|
|||
|
|
self._last_cleanup = now
|
|||
|
|
await self._cleanup_stale_buckets(now)
|
|||
|
|
|
|||
|
|
bucket = await self._get_bucket(identifier, scope)
|
|||
|
|
now = time.monotonic()
|
|||
|
|
|
|||
|
|
async with bucket.lock:
|
|||
|
|
if bucket.cooldown_until > now:
|
|||
|
|
return False
|
|||
|
|
|
|||
|
|
window_start = now - config.window_seconds
|
|||
|
|
bucket.timestamps = [t for t in bucket.timestamps if t > window_start]
|
|||
|
|
|
|||
|
|
if len(bucket.timestamps) >= config.max_attempts:
|
|||
|
|
bucket.cooldown_until = now + config.cooldown_seconds
|
|||
|
|
return False
|
|||
|
|
|
|||
|
|
bucket.timestamps.append(now)
|
|||
|
|
return True
|
|||
|
|
|
|||
|
|
async def check_and_record_or_raise(
|
|||
|
|
self,
|
|||
|
|
identifier: str,
|
|||
|
|
scope: RateLimitScope,
|
|||
|
|
) -> None:
|
|||
|
|
if not await self.check_and_record(identifier, scope):
|
|||
|
|
config = self._configs.get(scope)
|
|||
|
|
retry_after = config.cooldown_seconds if config else 60
|
|||
|
|
raise RateLimitExceeded(scope=scope, retry_after=retry_after)
|
|||
|
|
|
|||
|
|
async def clear_cooldown(self, identifier: str, scope: RateLimitScope):
|
|||
|
|
bucket = await self._get_bucket(identifier, scope)
|
|||
|
|
async with bucket.lock:
|
|||
|
|
bucket.cooldown_until = 0.0
|
|||
|
|
|
|||
|
|
async def reset(self, identifier: str, scope: RateLimitScope | None = None):
|
|||
|
|
async with self._lock:
|
|||
|
|
if scope:
|
|||
|
|
if identifier in self._buckets and scope in self._buckets[identifier]:
|
|||
|
|
async with self._buckets[identifier][scope].lock:
|
|||
|
|
self._buckets[identifier].pop(scope, None)
|
|||
|
|
else:
|
|||
|
|
self._buckets.pop(identifier, None)
|
|||
|
|
|
|||
|
|
def get_stats(self) -> dict:
|
|||
|
|
now = time.monotonic()
|
|||
|
|
stats: dict[str, dict] = {}
|
|||
|
|
for identifier, scopes in self._buckets.items():
|
|||
|
|
stats[identifier] = {}
|
|||
|
|
for scope, bucket in scopes.items():
|
|||
|
|
config = self._configs.get(scope)
|
|||
|
|
window = config.window_seconds if config else 60
|
|||
|
|
active = len([t for t in bucket.timestamps if t > now - window])
|
|||
|
|
stats[identifier][scope.value] = {
|
|||
|
|
"attempts_in_window": active,
|
|||
|
|
"cooldown_active": bucket.cooldown_until > now,
|
|||
|
|
"cooldown_remaining": max(0, bucket.cooldown_until - now),
|
|||
|
|
}
|
|||
|
|
return stats
|
|||
|
|
|
|||
|
|
def wrap(
|
|||
|
|
self,
|
|||
|
|
scope: RateLimitScope,
|
|||
|
|
identifier_fn: Callable[..., str] | None = None,
|
|||
|
|
):
|
|||
|
|
"""装饰器:将限流包装到异步函数上。
|
|||
|
|
|
|||
|
|
使用方式:
|
|||
|
|
@auth_rate_limiter.wrap(RateLimitScope.LOGIN)
|
|||
|
|
async def login(request: Request):
|
|||
|
|
...
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
def decorator(fn):
|
|||
|
|
@wraps(fn)
|
|||
|
|
async def wrapper(*args, **kwargs):
|
|||
|
|
ident = identifier_fn(*args, **kwargs) if identifier_fn else "default"
|
|||
|
|
await self.check_and_record_or_raise(ident, scope)
|
|||
|
|
return await fn(*args, **kwargs)
|
|||
|
|
|
|||
|
|
return wrapper
|
|||
|
|
|
|||
|
|
return decorator
|
|||
|
|
|
|||
|
|
|
|||
|
|
class RateLimitExceeded(Exception):
|
|||
|
|
def __init__(self, scope: RateLimitScope, retry_after: int = 60):
|
|||
|
|
self.scope = scope
|
|||
|
|
self.retry_after = retry_after
|
|||
|
|
super().__init__(f"Rate limit exceeded for scope '{scope.value}', retry after {retry_after}s")
|
|||
|
|
|
|||
|
|
|
|||
|
|
auth_rate_limiter = AuthRateLimiter()
|