Files
HeXi/hexi/core/rate_limit.py
T
sansenhoshi 131b92b319 结构调整
视频解析多图/多媒体结构 消息体适配
2026-09-08 14:25:32 +08:00

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