diff --git a/backend/utils/limiter.py b/backend/utils/limiter.py index 25aa4568..25280771 100644 --- a/backend/utils/limiter.py +++ b/backend/utils/limiter.py @@ -1,11 +1,17 @@ +from asyncio import Lock +from collections import OrderedDict from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from hashlib import sha256 +from inspect import isawaitable from math import ceil -from typing import TypeAlias +from typing import TypeAlias, TypeVar from fastapi import Request, Response from fastapi_pagination.utils import is_async_callable -from pyrate_limiter import AbstractBucket, Limiter, Rate +from pyrate_limiter import AbstractBucket, BucketFactory, Limiter, Rate, RateItem from pyrate_limiter.buckets import RedisBucket +from redis.asyncio import Redis from starlette.concurrency import run_in_threadpool from backend.common.exception import errors @@ -18,6 +24,174 @@ IdentifierCallable: TypeAlias = Callable[[Request], str] | Callable[[Request], A CallbackCallable: TypeAlias = ( Callable[[Request, Response, int], None] | Callable[[Request, Response, int], Awaitable[None]] ) +T = TypeVar('T') + +REQUEST_LIMITER_BUCKET_CACHE_MAX_SIZE = 4096 +REQUEST_LIMITER_BUCKET_CACHE_BUFFER_MS = 10_000 + + +@dataclass(slots=True) +class RedisBucketState: + """Redis bucket 缓存状态""" + + bucket: RedisBucket + last_seen: int + + +async def _maybe_await(value: T | Awaitable[T]) -> T: + """ + 兼容同步值和异步值 + + :param value: 同步值或 Awaitable 对象 + :return: + """ + if isawaitable(value): + return await value + return value + + +async def _redis_time_ms(redis: Redis) -> int: + """ + 获取 Redis 服务端当前时间 + + :return: + """ + seconds, microseconds = await redis.time() + return seconds * 1000 + microseconds // 1000 + + +class RedisTimeBucket(RedisBucket): + """使用 Redis 服务端时间的 Redis bucket""" + + async def now(self) -> int: + """获取 Redis 服务端当前时间""" + return await _redis_time_ms(self.redis) + + +class RedisBucketFactory(BucketFactory): + """按请求标识符路由到独立 Redis bucket""" + + def __init__( + self, + rates: list[Rate], + bucket_key: str, + max_cache_size: int = REQUEST_LIMITER_BUCKET_CACHE_MAX_SIZE, + ) -> None: + """ + 初始化 Redis bucket 工厂 + + :param rates: pyrate_limiter Rate 对象列表 + :param bucket_key: Redis key 前缀 + :param max_cache_size: 本地 bucket 缓存最大数量 + :return: + """ + self.rates = rates + self.bucket_key = f'{bucket_key}:{self._rate_key(rates)}' + self.max_cache_size = max(1, max_cache_size) + self.cache_ttl = max(rate.interval for rate in rates) + REQUEST_LIMITER_BUCKET_CACHE_BUFFER_MS + self.lock = Lock() + self.buckets: OrderedDict[str, RedisBucketState] = OrderedDict() + + async def wrap_item(self, name: str, weight: int = 1) -> RateItem: + """ + 包装限流项 + + :param name: 限流标识符 + :param weight: 请求权重 + :return: + """ + return RateItem(name, await _redis_time_ms(redis_client), weight=weight) + + async def get(self, item: RateItem) -> RedisBucket: + """ + 获取标识符对应的 Redis bucket + + :param item: 限流项 + :return: + """ + bucket_key = self._bucket_key(item.name) + now = await _redis_time_ms(redis_client) + + async with self.lock: + state = self.buckets.get(bucket_key) + if state is not None: + state.last_seen = now + self.buckets.move_to_end(bucket_key) + return state.bucket + + bucket_result = RedisTimeBucket.init( + rates=self.rates, + redis=redis_client, + bucket_key=bucket_key, + ) + bucket = await _maybe_await(bucket_result) + self.buckets[bucket_key] = RedisBucketState(bucket=bucket, last_seen=now) + self.schedule_leak(bucket) + await self._evict(now) + return bucket + + async def get_bucket(self, name: str) -> RedisBucket: + """ + 获取标识符对应的 Redis bucket + + :param name: 限流标识符 + :return: + """ + return await self.get(await self.wrap_item(name)) + + async def _evict(self, now: int) -> None: + """ + 淘汰本地 bucket 缓存 + + :param now: 当前时间戳,单位毫秒 + :return: + """ + for bucket_key, state in list(self.buckets.items()): + if now - state.last_seen <= self.cache_ttl: + continue + await self._dispose(bucket_key, state, now, cleanup=True) + + while len(self.buckets) > self.max_cache_size: + bucket_key, state = next(iter(self.buckets.items())) + await self._dispose(bucket_key, state, now, cleanup=False) + + async def _dispose(self, bucket_key: str, state: RedisBucketState, now: int, *, cleanup: bool) -> None: + """ + 移除本地 bucket 并按需清理 Redis 过期数据 + + :param bucket_key: Redis bucket key + :param state: Redis bucket 缓存状态 + :param now: 当前时间戳,单位毫秒 + :param cleanup: 是否执行 Redis 过期数据清理 + :return: + """ + self.buckets.pop(bucket_key, None) + self.dispose(state.bucket) + if cleanup: + await _maybe_await(state.bucket.leak(now)) + if await _maybe_await(state.bucket.count()) == 0: + await _maybe_await(state.bucket.flush()) + + def _bucket_key(self, name: str) -> str: + """ + 生成标识符对应的 Redis bucket key + + :param name: 限流标识符 + :return: + """ + digest = sha256(name.encode()).hexdigest() + return f'{self.bucket_key}:{digest}' + + @staticmethod + def _rate_key(rates: list[Rate]) -> str: + """ + 生成限流策略对应的 Redis key 片段 + + :param rates: pyrate_limiter Rate 对象列表 + :return: + """ + value = ':'.join(f'{rate.limit}:{rate.interval}' for rate in sorted(rates, key=lambda rate: rate.interval)) + return sha256(value.encode()).hexdigest() def default_identifier(request: Request) -> str: @@ -68,23 +242,32 @@ class RateLimiter: :param callback: 自定义限流回调函数 :return: """ - if not rates and bucket is None: - raise errors.ServerError(msg='至少需要传入一个 Rate 或 bucket 实例') + if limiter is None and not rates and bucket is None: + raise errors.ServerError(msg='至少需要传入一个 Rate、bucket 或 limiter 实例') self.rates = list(rates) self.identifier = identifier self.bucket = bucket self.limiter = limiter self.callback = callback + self.bucket_factory: RedisBucketFactory | None = None async def __call__(self, request: Request, response: Response) -> None: + """ + 执行请求限流检查 + + :param request: FastAPI 请求对象 + :param response: FastAPI 响应对象 + :return: + """ if self.limiter is None: if self.bucket is None: - self.bucket = await RedisBucket.init( # type: ignore + self.bucket_factory = RedisBucketFactory( rates=self.rates, - redis=redis_client, bucket_key=f'{settings.REQUEST_LIMITER_REDIS_PREFIX}', ) - self.limiter = Limiter(self.bucket) + self.limiter = Limiter(self.bucket_factory) + else: + self.limiter = Limiter(self.bucket) if is_async_callable(self.identifier): identifier = await self.identifier(request) @@ -93,8 +276,35 @@ class RateLimiter: acquired = await self.limiter.try_acquire_async(identifier, blocking=False) if not acquired: - retry_after = ceil(self.bucket.failing_rate.interval / 1000) + retry_after = await self._retry_after(identifier) if is_async_callable(self.callback): await self.callback(request, response, retry_after) else: await run_in_threadpool(self.callback, request, response, retry_after) + + async def _retry_after(self, identifier: str) -> int: + """ + 计算限流重试等待时间 + + :param identifier: 限流标识符 + :return: + """ + if self.bucket_factory is not None: + failing_rate = (await self.bucket_factory.get_bucket(identifier)).failing_rate + elif self.bucket is not None: + failing_rate = self.bucket.failing_rate + else: + failing_rate = None + + if failing_rate is not None: + return ceil(failing_rate.interval / 1000) + + if self.limiter is not None: + for bucket in self.limiter.buckets(): + if bucket.failing_rate is not None: + return ceil(bucket.failing_rate.interval / 1000) + + if self.rates: + return ceil(max(rate.interval for rate in self.rates) / 1000) + + return 1