62 lines
1.9 KiB
Python
62 lines
1.9 KiB
Python
"""内存限流(MVP: 单进程 TokenBucket)。
|
|
|
|
未来切换 Redis: 实现 RedisRateLimiter,接口不变。
|
|
"""
|
|
|
|
import time
|
|
from collections import defaultdict
|
|
|
|
from app.core.config import get_settings
|
|
|
|
|
|
class _Bucket:
|
|
__slots__ = ("tokens", "max_tokens", "refill_rate", "last_refill")
|
|
|
|
def __init__(self, max_tokens: int, per_seconds: int) -> None:
|
|
self.max_tokens = max_tokens
|
|
self.refill_rate = max_tokens / per_seconds
|
|
self.tokens = float(max_tokens)
|
|
self.last_refill = time.monotonic()
|
|
|
|
def allow(self) -> bool:
|
|
now = time.monotonic()
|
|
elapsed = now - self.last_refill
|
|
self.tokens = min(self.max_tokens, self.tokens + elapsed * self.refill_rate)
|
|
self.last_refill = now
|
|
if self.tokens >= 1.0:
|
|
self.tokens -= 1.0
|
|
return True
|
|
return False
|
|
|
|
|
|
# key → Bucket
|
|
_token_buckets: dict[str, _Bucket] = defaultdict(lambda: _Bucket(_token_limit, 60))
|
|
_ip_buckets: dict[str, _Bucket] = defaultdict(lambda: _Bucket(_ip_limit, 60))
|
|
|
|
# 延迟初始化
|
|
_token_limit = 60
|
|
_ip_limit = 30
|
|
_initialized = False
|
|
|
|
|
|
def _ensure_init() -> None:
|
|
global _token_limit, _ip_limit, _initialized
|
|
if _initialized:
|
|
return
|
|
settings = get_settings()
|
|
_token_limit = settings.rate_limit_per_token_per_min
|
|
_ip_limit = settings.rate_limit_per_ip_per_min
|
|
# 重建 default dict factories
|
|
_token_buckets.default_factory = lambda: _Bucket(_token_limit, 60) # type: ignore[assignment]
|
|
_ip_buckets.default_factory = lambda: _Bucket(_ip_limit, 60) # type: ignore[assignment]
|
|
_initialized = True
|
|
|
|
|
|
def check_rate_limit(token_key: str | None = None, ip_key: str | None = None) -> bool:
|
|
"""检查是否允许请求。返回 True 表示允许。"""
|
|
_ensure_init()
|
|
if token_key and not _token_buckets[token_key].allow():
|
|
return False
|
|
if ip_key and not _ip_buckets[ip_key].allow():
|
|
return False
|
|
return True |