Files
amb_rag/backend/app/core/rate_limit.py
T
2026-09-01 11:53:59 +08:00

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