182 lines
6.3 KiB
Python
182 lines
6.3 KiB
Python
"""限频器(令牌桶):为依赖外部 API 的插件限速,避免请求过频被限流返回 403/429
|
|
|
|
纯逻辑实现,不依赖 NoneBot 运行时;同步(requests)与异步(httpx/aiohttp)均可用。
|
|
|
|
用法:
|
|
from hexi.core import rate_limit
|
|
|
|
# 异步请求前取令牌(无令牌时等待,最多等 timeout 秒,超时返回 False)
|
|
if not await rate_limit.acquire("steam", rate=1, capacity=2):
|
|
await UniMessage.text("查询太频繁了,歇会儿再来~").send()
|
|
return
|
|
|
|
# 同步请求前取令牌(不等待,超限立即失败)
|
|
if not rate_limit.acquire_sync("bf_api", rate=0.5, wait=False):
|
|
return "请求过于频繁,请稍后再试"
|
|
|
|
# 拿到响应后检测是否被限流(403/429),可据此告警或退避
|
|
if rate_limit.is_ratelimited(resp.status_code):
|
|
...
|
|
"""
|
|
|
|
import asyncio
|
|
import threading
|
|
import time
|
|
|
|
from nonebot.log import logger
|
|
|
|
# 常见限流状态码:403(服务端按频率拒绝)与 429(Too Many Requests)
|
|
_RATE_LIMIT_CODES = frozenset({403, 429})
|
|
|
|
|
|
class RateLimiter:
|
|
"""令牌桶限频器
|
|
|
|
- rate: 每秒补充的令牌数;capacity: 桶容量(最大突发),默认取 max(rate, 1)
|
|
- 惰性补充:按距上次取令牌的时间差计算补充量,无需后台任务
|
|
- 线程安全:临界区用 threading.Lock 保护,同步/异步调用可并发使用
|
|
"""
|
|
|
|
def __init__(self, rate: float, capacity: int | None = None):
|
|
if rate <= 0:
|
|
raise ValueError(f"rate 必须大于 0,收到 {rate!r}")
|
|
if capacity is not None and capacity < 1:
|
|
raise ValueError(f"capacity 必须 >= 1,收到 {capacity!r}")
|
|
self._rate = float(rate)
|
|
self._capacity = float(capacity if capacity is not None else max(rate, 1))
|
|
self._tokens = self._capacity
|
|
self._last_refill = time.monotonic()
|
|
self._lock = threading.Lock()
|
|
|
|
def _refill(self) -> None:
|
|
"""按经过的时间补充令牌(调用方需持有 _lock)"""
|
|
now = time.monotonic()
|
|
self._tokens = min(
|
|
self._capacity, self._tokens + (now - self._last_refill) * self._rate
|
|
)
|
|
self._last_refill = now
|
|
|
|
def try_acquire(self) -> bool:
|
|
"""尝试取 1 个令牌:有则立即成功,无则返回 False(不阻塞)"""
|
|
with self._lock:
|
|
self._refill()
|
|
if self._tokens >= 1:
|
|
self._tokens -= 1
|
|
return True
|
|
return False
|
|
|
|
def remaining_secs(self) -> float:
|
|
"""距下一个令牌可用的秒数;0 表示当前有令牌可用(不消费令牌)"""
|
|
with self._lock:
|
|
self._refill()
|
|
return max(0.0, (1.0 - self._tokens) / self._rate)
|
|
|
|
def _wait_secs(self, deadline: float | None) -> float:
|
|
"""计算距下次尝试还需等待的秒数;已超时返回 -1 表示放弃"""
|
|
with self._lock:
|
|
self._refill()
|
|
need = max(1.0 - self._tokens, 0.0) / self._rate
|
|
if deadline is not None:
|
|
remaining = deadline - time.monotonic()
|
|
if remaining <= 0:
|
|
return -1.0
|
|
need = min(need, remaining)
|
|
return need
|
|
|
|
async def acquire(self, timeout: float | None = None) -> bool:
|
|
"""异步取 1 个令牌:无令牌时 asyncio.sleep 等待补充(不阻塞事件循环)
|
|
|
|
timeout 为 None 时无限等待;超时返回 False
|
|
"""
|
|
deadline = None if timeout is None else time.monotonic() + timeout
|
|
while True:
|
|
if self.try_acquire():
|
|
return True
|
|
wait = self._wait_secs(deadline)
|
|
if wait < 0:
|
|
return False
|
|
await asyncio.sleep(wait)
|
|
|
|
def acquire_sync(self, timeout: float | None = None) -> bool:
|
|
"""同步取 1 个令牌:无令牌时 time.sleep 等待补充(阻塞当前线程)
|
|
|
|
timeout 为 None 时无限等待;超时返回 False
|
|
"""
|
|
deadline = None if timeout is None else time.monotonic() + timeout
|
|
while True:
|
|
if self.try_acquire():
|
|
return True
|
|
wait = self._wait_secs(deadline)
|
|
if wait < 0:
|
|
return False
|
|
time.sleep(wait)
|
|
|
|
|
|
# 按名称共享的限频器注册表:同名(如同一外部 API)的所有插件共用同一配额
|
|
_limiters: dict[str, RateLimiter] = {}
|
|
_registry_lock = threading.Lock()
|
|
|
|
|
|
def get_limiter(name: str, rate: float, capacity: int | None = None) -> RateLimiter:
|
|
"""获取(或创建)指定名称的限频器,同名共享同一配额
|
|
|
|
首次创建后,后续调用若 rate/capacity 与已有实例不一致会告警(以先创建者为准)。
|
|
"""
|
|
with _registry_lock:
|
|
limiter = _limiters.get(name)
|
|
if limiter is None:
|
|
limiter = RateLimiter(rate, capacity)
|
|
_limiters[name] = limiter
|
|
logger.debug(
|
|
f"限频器已创建: {name}"
|
|
f"(rate={limiter._rate}, capacity={limiter._capacity})"
|
|
)
|
|
elif (limiter._rate, limiter._capacity) != (
|
|
float(rate),
|
|
float(capacity if capacity is not None else max(rate, 1)),
|
|
):
|
|
logger.warning(
|
|
f"限频器 {name} 已存在(rate={limiter._rate}, "
|
|
f"capacity={limiter._capacity}),本次参数被忽略"
|
|
)
|
|
return limiter
|
|
|
|
|
|
async def acquire(
|
|
name: str,
|
|
rate: float,
|
|
capacity: int | None = None,
|
|
*,
|
|
wait: bool = True,
|
|
timeout: float | None = None,
|
|
) -> bool:
|
|
"""异步限频入口:获取(或创建)名为 name 的限频器并取令牌
|
|
|
|
wait=True 时等待令牌(最多 timeout 秒,None 为无限);
|
|
wait=False 时超限立即返回 False。
|
|
"""
|
|
limiter = get_limiter(name, rate, capacity)
|
|
if not wait:
|
|
return limiter.try_acquire()
|
|
return await limiter.acquire(timeout)
|
|
|
|
|
|
def acquire_sync(
|
|
name: str,
|
|
rate: float,
|
|
capacity: int | None = None,
|
|
*,
|
|
wait: bool = True,
|
|
timeout: float | None = None,
|
|
) -> bool:
|
|
"""同步限频入口(requests 等阻塞调用):语义同 acquire()"""
|
|
limiter = get_limiter(name, rate, capacity)
|
|
if not wait:
|
|
return limiter.try_acquire()
|
|
return limiter.acquire_sync(timeout)
|
|
|
|
|
|
def is_ratelimited(status_code: int | None) -> bool:
|
|
"""判断 HTTP 状态码是否为限流响应(403 或 429)"""
|
|
return status_code in _RATE_LIMIT_CODES
|