Files

97 lines
3.9 KiB
Python
Raw Permalink Normal View History

"""限流计数后端:Redis 固定窗口计数,**Redis 不可用时一律放行**。
取舍(为什么是固定窗口,而不是 ZSET 滑动窗口或令牌桶):
- 固定窗口用一条 `INCR` + `TTL` 就能原子地完成"计数 + 取重试等待时间",不需要 Lua
或事务,故障面最小;滑动窗口要 `ZREMRANGEBYSCORE`/`ZADD`/`ZCARD`/`EXPIRE` 四条命令
才能近似原子,令牌桶还要在服务端保存补充速率的状态。本平台的限流目的是**保护**
底座不被单个客户端打爆,不是做精确计费,窗口边界最多放过一个窗口的量可以接受。
- 计数键带 TTL,Redis 自己回收,不需要额外的清理任务,也不会留下永久脏键。
降级语义(用户要求,也是 `MemoryCacheAdapter` 的同一原则):限流是保护措施,不能因为
Redis 故障把正常请求全部拒掉。因此后端**只返回 `None` 表示"无法判定"**,由调用方放行;
这里绝不抛异常、绝不返回"计数超限"。
Redis 客户端是**懒建**的:进程启动时 Redis 不可用不应该让应用起不来,第一次真正需要
限流判定时才建连(与 `bootstrap.py` 的记忆缓存适配器一致)。
"""
import logging
from typing import Any, Protocol
from app.core.config import get_settings
logger = logging.getLogger(__name__)
class CounterBackend(Protocol):
"""限流计数后端。
`increment` 返回 `(窗口内计数, 剩余秒数)`;返回 `None` 表示后端不可用,
调用方必须**放行**(fail-open)。
"""
async def increment(self, key: str, window_seconds: int) -> tuple[int, int] | None: ...
class RedisCounterBackend:
def __init__(
self,
redis_url: str,
*,
connect_timeout_seconds: float,
client_factory: Any | None = None,
) -> None:
self._redis_url = redis_url
self._connect_timeout_seconds = connect_timeout_seconds
# 注入点仅用于测试:默认走 `redis.asyncio.Redis.from_url`。
self._client_factory = client_factory
self._client: Any | None = None
async def _client_or_none(self) -> Any | None:
if self._client is not None:
return self._client
try:
if self._client_factory is not None:
client = self._client_factory()
else:
from redis.asyncio import Redis
client = Redis.from_url(
self._redis_url,
socket_connect_timeout=self._connect_timeout_seconds,
socket_timeout=self._connect_timeout_seconds,
decode_responses=True,
)
self._client = client
except Exception:
logger.warning("限流后端不可用:Redis 客户端构造失败,本次降级放行", exc_info=True)
return None
return self._client
async def increment(self, key: str, window_seconds: int) -> tuple[int, int] | None:
client = await self._client_or_none()
if client is None:
return None
try:
async with client.pipeline(transaction=False) as pipe:
pipe.incr(key)
pipe.ttl(key)
count, ttl = await pipe.execute()
remaining = int(ttl)
if remaining < 0:
# 键首次创建(或历史上丢过 TTL):补一次过期为窗口长度。
await client.expire(key, window_seconds)
remaining = window_seconds
return int(count), max(1, remaining)
except Exception:
logger.warning("限流后端不可用:Redis 计数失败,本次降级放行", exc_info=True)
return None
def default_counter_backend() -> CounterBackend:
settings = get_settings()
return RedisCounterBackend(
settings.redis_url, connect_timeout_seconds=settings.redis_connect_timeout_seconds
)