feat: add redis six-gate rate limiter backend

This commit is contained in:
2026-07-21 00:27:16 -04:00
parent acc1bcc18a
commit aae739cebe
4 changed files with 498 additions and 14 deletions
@@ -0,0 +1,6 @@
"""Redis 治理后端(M2): 与内存版同一契约的跨进程实现(D3 双后端)。"""
from polygateway.backends.redis.breaker import RedisGate
from polygateway.backends.redis.limiter import RedisLimiter
__all__ = ["RedisGate", "RedisLimiter"]
+307
View File
@@ -0,0 +1,307 @@
"""跨进程限流后端: CHS 六道闸 Lua 逐字移植(D3 双后端;M2 设计 §2)。
语义蓝本 `reference/CHSAnalyzer/app/coordination/{scripts,limiter}.py`:
读-判-占在单条 Lua 内原子完成;并发 = ZSET 租约(score=过期时刻 ms,惰性
清理 + PEXPIRE 双保险);RPM/TPM = 分钟固定窗口计数,窗口 id 用 **Redis
服务器时钟**(多进程口径统一);TPM 预扣 settle 多退少补且退款落 acquire
时的窗口(INCRBY 可为负,存储保留负值保证跨窗口结算正确)。
对 CHS 的已声明偏离(设计 §9 勘误): ① 限额 0 = 该闸不启用(M1 契约冻结;
CHS Lua 无此守卫,0 在其语义下是全拒),Lua 判据显式加 `limit > 0`;
② key 前缀 `cclimit:` → `pgw:limit:`;③ `LimiterError` → `GovernanceBackendError`;
④ `source_stats` 读侧 clamp ≥0(展示口径与内存版统一,存储不 clamp)。
契约量纲为秒,毫秒换算是本后端私事(ports.py 模块 docstring)。
"""
from __future__ import annotations
import asyncio
import uuid
from typing import TYPE_CHECKING
from loguru import logger
from redis.exceptions import RedisError
from polygateway.errors import GovernanceBackendError
from polygateway.types import GlobalLimits, SourceConfig, SourceStats
if TYPE_CHECKING:
from redis.asyncio import Redis
_WINDOW_TTL_S = 120 # 分钟窗口兜底过期(CHS limiter.py:16)
_PROGRESS_TTL_S = 3600 # 远大于任何 stall_window,防进度键过期造成假停滞(CHS :21)
# KEYS: g_lease s_lease g_rpm s_rpm g_tpm s_tpm(rpm/tpm key 已含 window 后缀)
# ARGV: g_conc s_conc g_rpm s_rpm g_tpm s_tpm est lease_ttl_ms window_ttl_s lease_id
# 拒绝零副作用(仅清过期 lease);闸判据 CHS scripts.py:18-23 + 库版 limit>0 守卫
ACQUIRE = """
local t = redis.call('TIME')
local now = tonumber(t[1]) * 1000 + math.floor(tonumber(t[2]) / 1000)
redis.call('ZREMRANGEBYSCORE', KEYS[1], 0, now)
redis.call('ZREMRANGEBYSCORE', KEYS[2], 0, now)
local gi = redis.call('ZCARD', KEYS[1])
local si = redis.call('ZCARD', KEYS[2])
local gr = tonumber(redis.call('GET', KEYS[3]) or '0')
local sr = tonumber(redis.call('GET', KEYS[4]) or '0')
local gt = tonumber(redis.call('GET', KEYS[5]) or '0')
local st = tonumber(redis.call('GET', KEYS[6]) or '0')
local est = tonumber(ARGV[7])
if tonumber(ARGV[1]) > 0 and gi >= tonumber(ARGV[1]) then return 0 end
if tonumber(ARGV[2]) > 0 and si >= tonumber(ARGV[2]) then return 0 end
if tonumber(ARGV[3]) > 0 and gr >= tonumber(ARGV[3]) then return 0 end
if tonumber(ARGV[4]) > 0 and sr >= tonumber(ARGV[4]) then return 0 end
if tonumber(ARGV[5]) > 0 and gt + est > tonumber(ARGV[5]) then return 0 end
if tonumber(ARGV[6]) > 0 and st + est > tonumber(ARGV[6]) then return 0 end
local expire_at = now + tonumber(ARGV[8])
redis.call('ZADD', KEYS[1], expire_at, ARGV[10])
redis.call('ZADD', KEYS[2], expire_at, ARGV[10])
redis.call('PEXPIRE', KEYS[1], ARGV[8])
redis.call('PEXPIRE', KEYS[2], ARGV[8])
redis.call('INCR', KEYS[3]); redis.call('EXPIRE', KEYS[3], ARGV[9])
redis.call('INCR', KEYS[4]); redis.call('EXPIRE', KEYS[4], ARGV[9])
redis.call('INCRBY', KEYS[5], est); redis.call('EXPIRE', KEYS[5], ARGV[9])
redis.call('INCRBY', KEYS[6], est); redis.call('EXPIRE', KEYS[6], ARGV[9])
return 1
"""
# KEYS: g_lease s_lease ; ARGV: lease_id(ZREM 幂等,CHS scripts.py:37-41)
RELEASE = """
redis.call('ZREM', KEYS[1], ARGV[1])
redis.call('ZREM', KEYS[2], ARGV[1])
return 1
"""
# KEYS: g_tpm s_tpm ; ARGV: delta window_ttl_s(delta 可负=退款;落 acquire 窗口)
SETTLE = """
redis.call('INCRBY', KEYS[1], ARGV[1]); redis.call('EXPIRE', KEYS[1], ARGV[2])
redis.call('INCRBY', KEYS[2], ARGV[1]); redis.call('EXPIRE', KEYS[2], ARGV[2])
return 1
"""
# KEYS: s_lease s_rpm s_tpm(CHS scripts.py:51-58)
STATS = """
local t = redis.call('TIME')
local now = tonumber(t[1]) * 1000 + math.floor(tonumber(t[2]) / 1000)
redis.call('ZREMRANGEBYSCORE', KEYS[1], 0, now)
return {redis.call('ZCARD', KEYS[1]),
tonumber(redis.call('GET', KEYS[2]) or '0'),
tonumber(redis.call('GET', KEYS[3]) or '0')}
"""
# KEYS: progress_key ; ARGV: ttl_s(CHS scripts.py:61-67)
PROGRESS_MARK = """
local t = redis.call('TIME')
local now = tonumber(t[1]) * 1000 + math.floor(tonumber(t[2]) / 1000)
redis.call('SET', KEYS[1], now)
redis.call('EXPIRE', KEYS[1], ARGV[1])
return 1
"""
# KEYS: progress_key(-1=从未成功;时钟回拨 clamp 0;CHS scripts.py:70-78)
PROGRESS_AGE = """
local v = redis.call('GET', KEYS[1])
if not v then return -1 end
local t = redis.call('TIME')
local now = tonumber(t[1]) * 1000 + math.floor(tonumber(t[2]) / 1000)
local d = now - tonumber(v)
if d < 0 then return 0 end
return d
"""
class _RedisPermit:
"""入场许可: 持有并发租约,负责释放与按 acquire 窗口结算 TPM(幂等)。"""
def __init__(
self, limiter: RedisLimiter, source: str, lease_id: str, est: int, window: int
) -> None:
self._limiter = limiter
self._source = source
self._lease_id = lease_id
self._est = est
self._window = window
self._released = False
self._settled = False
async def release(self) -> None:
"""归还并发租约;释放侧 Redis 失败降级 warning(设计 §10),不冒泡。"""
if self._released:
return
self._released = True
try:
await self._limiter._release_lease(self._source, self._lease_id)
except GovernanceBackendError as exc:
logger.warning("permit release 降级(租约将由 TTL 回收): {}", exc)
async def settle(self, actual_tokens: int) -> None:
"""按实际 usage 结算 TPM 差额;释放侧失败降级 warning,不冒泡。"""
if self._settled:
return
self._settled = True
delta = actual_tokens - self._est
if delta == 0:
return
try:
await self._limiter._settle_tpm(self._source, delta, self._window)
except GovernanceBackendError as exc:
logger.warning("permit settle 降级(窗口计数将随 TTL 过期): {}", exc)
class RedisLimiter:
"""六道闸的跨进程实现;同 scope 的多进程 worker 天然共享全局限额。"""
def __init__(
self,
*,
scope: str,
sources: dict[str, SourceConfig],
global_limits: GlobalLimits,
redis: Redis,
lease_ttl_s: float = 1500.0,
poll_interval_s: float = 0.05,
) -> None:
if lease_ttl_s <= 0 or poll_interval_s <= 0:
raise ValueError("lease_ttl_s 与 poll_interval_s 必须 > 0")
self._scope = scope.lower()
self._sources = dict(sources)
self._global = global_limits
self._lease_ttl_ms = int(lease_ttl_s * 1000)
self._poll_interval_s = poll_interval_s
self._redis = redis
self._owns_client = False
self._acquire_lua = redis.register_script(ACQUIRE)
self._release_lua = redis.register_script(RELEASE)
self._settle_lua = redis.register_script(SETTLE)
self._stats_lua = redis.register_script(STATS)
self._progress_mark_lua = redis.register_script(PROGRESS_MARK)
self._progress_age_lua = redis.register_script(PROGRESS_AGE)
@classmethod
def from_url(cls, url: str, **kwargs) -> RedisLimiter:
"""自建并持有 Redis 客户端(aclose 时代关);共享后端请直接注入 redis。"""
import redis.asyncio as aioredis
limiter = cls(redis=aioredis.from_url(url), **kwargs)
limiter._owns_client = True
return limiter
# —— key 布局(CHS limiter.py:100-115,前缀改 pgw:limit:)——
def _cfg(self, source_key: str) -> SourceConfig:
cfg = self._sources.get(source_key)
if cfg is None:
raise GovernanceBackendError(f"未知源 {source_key!r}(scope={self._scope})")
return cfg
def _lease_keys(self, source_key: str) -> tuple[str, str]:
return (
f"pgw:limit:GLOBAL:{self._scope}:lease",
f"pgw:limit:{self._scope}:{source_key}:lease",
)
def _window_keys(self, source_key: str, window: int) -> dict[str, str]:
gl = f"GLOBAL:{self._scope}"
return {
"g_rpm": f"pgw:limit:{gl}:rpm:{window}",
"s_rpm": f"pgw:limit:{self._scope}:{source_key}:rpm:{window}",
"g_tpm": f"pgw:limit:{gl}:tpm:{window}",
"s_tpm": f"pgw:limit:{self._scope}:{source_key}:tpm:{window}",
}
def _progress_key(self) -> str:
return f"pgw:limit:GLOBAL:{self._scope}:progress_ms"
async def _window_id(self) -> int:
"""Redis 服务器时钟生成分钟窗口 id(CHS limiter.py:95-98)。"""
sec, _ = await self._redis.time()
return int(sec) // 60
# —— 契约方法 ——
async def try_acquire(self, source_key: str, est_tokens: int) -> _RedisPermit | None:
"""非阻塞准入: 六闸单 Lua 原子判定;拒绝零副作用;Redis 失败报错不放行。"""
cfg = self._cfg(source_key)
lease_id = uuid.uuid4().hex
try:
window = await self._window_id()
gl, sl = self._lease_keys(source_key)
wk = self._window_keys(source_key, window)
ok = await self._acquire_lua(
keys=[gl, sl, wk["g_rpm"], wk["s_rpm"], wk["g_tpm"], wk["s_tpm"]],
args=[
self._global.max_concurrency,
cfg.max_concurrency,
self._global.rpm,
cfg.rpm,
self._global.tpm,
cfg.tpm,
est_tokens,
self._lease_ttl_ms,
_WINDOW_TTL_S,
lease_id,
],
)
except RedisError as exc:
raise GovernanceBackendError(f"限流后端 try_acquire 失败: {exc}") from exc
if ok != 1:
return None
return _RedisPermit(self, source_key, lease_id, est_tokens, window)
async def acquire(self, source_key: str, est_tokens: int) -> _RedisPermit:
"""阻塞准入: 轮询直到拿到 permit;sleep 可被取消(取消穿透铁律)。"""
while True:
permit = await self.try_acquire(source_key, est_tokens)
if permit is not None:
return permit
await asyncio.sleep(self._poll_interval_s)
async def _release_lease(self, source_key: str, lease_id: str) -> None:
gl, sl = self._lease_keys(source_key)
try:
await self._release_lua(keys=[gl, sl], args=[lease_id])
except RedisError as exc:
raise GovernanceBackendError(f"限流后端 release 失败: {exc}") from exc
async def _settle_tpm(self, source_key: str, delta: int, window: int) -> None:
wk = self._window_keys(source_key, window)
try:
await self._settle_lua(keys=[wk["g_tpm"], wk["s_tpm"]], args=[delta, _WINDOW_TTL_S])
except RedisError as exc:
raise GovernanceBackendError(f"限流后端 settle 失败: {exc}") from exc
async def source_stats(self, source_key: str) -> SourceStats:
"""当前窗口快照;读侧 clamp ≥0(展示口径,存储保留负值)。"""
self._cfg(source_key)
try:
window = await self._window_id()
_, sl = self._lease_keys(source_key)
wk = self._window_keys(source_key, window)
res = await self._stats_lua(keys=[sl, wk["s_rpm"], wk["s_tpm"]])
except RedisError as exc:
raise GovernanceBackendError(f"限流后端 source_stats 失败: {exc}") from exc
return SourceStats(
inflight=int(res[0]),
rpm_used=max(0, int(res[1])),
tpm_used=max(0, int(res[2])),
)
async def mark_progress(self) -> None:
"""记录 scope 级"最近出餐"时刻;失败报错,降级责任在中间件记账侧(T5)。"""
try:
await self._progress_mark_lua(keys=[self._progress_key()], args=[_PROGRESS_TTL_S])
except RedisError as exc:
raise GovernanceBackendError(f"限流后端 mark_progress 失败: {exc}") from exc
async def progress_age_s(self) -> float:
"""距上次全局成功的秒数;仅键缺失(-1)= 从未进展 → inf(CHS limiter.py:208)。"""
try:
res = await self._progress_age_lua(keys=[self._progress_key()])
except RedisError as exc:
raise GovernanceBackendError(f"限流后端 progress_age_s 失败: {exc}") from exc
return float("inf") if int(res) == -1 else int(res) / 1000.0
async def aclose(self) -> None:
"""幂等释放自建客户端;注入的客户端归注入方管理。"""
if self._owns_client:
self._owns_client = False
await self._redis.aclose()
+88 -14
View File
@@ -1,18 +1,30 @@
"""契约测试共享 fixture: 后端参数化(M1 仅 memory,M2 增 redis 零改测试)。
"""契约测试共享 fixture: 后端参数化(memory + redis,D3 双后端同一契约)。
FakeClock 仅对支持时钟注入的后端有效(memory);M2 接入 Redis 后端时,
依赖时钟推进的用例按后端能力跳过或改用真实等待。
结构(M2 设计 §2.3 定案): 单一 `backend` fixture 承载参数化;`clock` 与两工厂
依赖它——memory 用 FakeClock(时钟注入,时间语义可快进);redis 用真实实验室
Redis(db3),FakeClock 对服务器时钟不可注入,故 `clock.advance()` 是哨兵:
触发 skip,对应用例由 tests/integration/test_redis_governance_time.py 的
1:1 真实等待变体覆盖(人类拍板: 不缩放时长)。
"""
from __future__ import annotations
import asyncio
import os
from uuid import uuid4
import pytest
from dotenv import dotenv_values
from polygateway.backends.memory.breaker import InMemoryGate
from polygateway.backends.memory.limiter import InMemoryLimiter
from polygateway.types import BreakerConfig, GlobalLimits, SourceConfig
_HEADROOM_S = 10 # 分钟窗口剩余不足则等翻滚,防 RPM/TPM 用例跨窗 flake(CHS 同款)
class FakeClock:
"""确定性单调时钟;契约测试推进时间验证租约/冷却语义。"""
"""确定性单调时钟;memory 后端的契约测试推进时间验证租约/冷却语义。"""
def __init__(self, start: float = 1000.0) -> None:
self.t = start
@@ -24,6 +36,16 @@ class FakeClock:
self.t += seconds
class SkipClock:
"""redis 后端的哨兵时钟: 服务器时钟不可注入,依赖快进的用例整例跳过。"""
def __call__(self) -> float: # pragma: no cover - 不应被消费
raise AssertionError("redis 后端不消费注入时钟")
def advance(self, seconds: float) -> None:
pytest.skip("redis 时间语义由 tests/integration/test_redis_governance_time.py 变体覆盖")
def make_source(name: str = "s1", **overrides) -> SourceConfig:
base = {
"name": name,
@@ -37,32 +59,84 @@ def make_source(name: str = "s1", **overrides) -> SourceConfig:
return SourceConfig(**base)
def redis_url_from_env() -> str | None:
"""读 REDIS_URL(.env 与进程环境合并,后者优先);供契约与集成测试共用。"""
merged = {**dotenv_values(".env"), **os.environ}
return merged.get("REDIS_URL") or None
async def await_window_headroom(client, min_headroom_s: int = _HEADROOM_S) -> None:
"""按 Redis 服务器时钟等待分钟窗口翻滚防抖(CHS test_redis_limiter.py:17-30)。"""
sec, _ = await client.time()
remaining = 60 - int(sec) % 60
if remaining < min_headroom_s:
await asyncio.sleep(remaining + 0.5)
class _Backend:
def __init__(self, name: str, redis=None) -> None:
self.name = name
self.redis = redis
@pytest.fixture(params=["memory", "redis"])
async def backend(request):
if request.param == "memory":
yield _Backend("memory")
return
url = redis_url_from_env()
if url is None:
pytest.skip("REDIS_URL 未配置,跳过 redis 后端契约")
import redis.asyncio as aioredis
client = aioredis.from_url(url)
try:
await await_window_headroom(client)
yield _Backend("redis", client)
finally:
await client.aclose()
@pytest.fixture
def clock() -> FakeClock:
return FakeClock()
def clock(backend):
return FakeClock() if backend.name == "memory" else SkipClock()
@pytest.fixture(params=["memory"])
def limiter_factory(request, clock):
@pytest.fixture
def limiter_factory(backend, clock):
"""返回 (sources, global_limits, lease_ttl_s) -> RateLimiter 的工厂。"""
def make(sources: list[SourceConfig], global_limits: GlobalLimits, lease_ttl_s: float = 100.0):
return InMemoryLimiter(
scope="llm",
if backend.name == "memory":
return InMemoryLimiter(
scope="llm",
sources={s.name: s for s in sources},
global_limits=global_limits,
lease_ttl_s=lease_ttl_s,
now=clock,
)
from polygateway.backends.redis.limiter import RedisLimiter
return RedisLimiter(
scope=f"t{uuid4().hex[:8]}",
sources={s.name: s for s in sources},
global_limits=global_limits,
redis=backend.redis,
lease_ttl_s=lease_ttl_s,
now=clock,
)
return make
@pytest.fixture(params=["memory"])
def gate_factory(request, clock):
@pytest.fixture
def gate_factory(backend, clock):
"""返回 (BreakerConfig) -> ProviderGate 的工厂。"""
def make(config: BreakerConfig):
return InMemoryGate(config=config, now=clock)
if backend.name == "memory":
return InMemoryGate(config=config, now=clock)
from polygateway.backends.redis.breaker import RedisGate
return RedisGate(config=config, redis=backend.redis, scope=f"t{uuid4().hex[:8]}")
return make
+97
View File
@@ -0,0 +1,97 @@
"""RedisLimiter 纯函数部分(key 布局、秒毫秒换算)——不需要真实 Redis。"""
import pytest
from polygateway.backends.redis.limiter import RedisLimiter
from polygateway.types import GlobalLimits, SourceConfig
class _StubRedis:
"""仅满足构造期 register_script 的桩;任何执行路径不可达。"""
def register_script(self, script: str):
def _never(**kwargs):
raise AssertionError("unit 测试不应执行 Lua")
return _never
def _limiter(**overrides) -> RedisLimiter:
src = SourceConfig(
name="s1",
provider="p",
base_url="https://gw.example/v1",
api_key="sk",
model="m",
timeout_s=10.0,
)
base = {
"scope": "llm",
"sources": {"s1": src},
"global_limits": GlobalLimits(max_concurrency=0, rpm=0, tpm=0),
"redis": _StubRedis(),
"lease_ttl_s": 30.0,
}
base.update(overrides)
return RedisLimiter(**base)
class TestKeyLayout:
def test_lease_keys_prefixed_pgw(self):
gl, sl = _limiter()._lease_keys("s1")
assert gl == "pgw:limit:GLOBAL:llm:lease"
assert sl == "pgw:limit:llm:s1:lease"
def test_window_keys_carry_window_suffix(self):
wk = _limiter()._window_keys("s1", 12345)
assert wk["g_rpm"] == "pgw:limit:GLOBAL:llm:rpm:12345"
assert wk["s_rpm"] == "pgw:limit:llm:s1:rpm:12345"
assert wk["g_tpm"] == "pgw:limit:GLOBAL:llm:tpm:12345"
assert wk["s_tpm"] == "pgw:limit:llm:s1:tpm:12345"
def test_progress_key_scope_global(self):
assert _limiter()._progress_key() == "pgw:limit:GLOBAL:llm:progress_ms"
def test_scope_lowercased(self):
gl, _ = _limiter(scope="LLM")._lease_keys("s1")
assert gl == "pgw:limit:GLOBAL:llm:lease"
class TestConversions:
def test_lease_ttl_seconds_to_ms(self):
# 契约量纲为秒,Redis 内部毫秒是后端私事(ports.py docstring)
assert _limiter(lease_ttl_s=30.0)._lease_ttl_ms == 30_000
assert _limiter(lease_ttl_s=0.5)._lease_ttl_ms == 500
def test_invalid_lease_ttl_rejected(self):
with pytest.raises(ValueError):
_limiter(lease_ttl_s=0)
def test_unknown_source_rejected(self):
from polygateway.errors import GovernanceBackendError
with pytest.raises(GovernanceBackendError):
_limiter()._cfg("nope")
class TestLuaFidelity:
"""Lua 常量的移植锚点:守卫与判据语义(逐段比对 CHS scripts.py:6-34)。"""
def test_acquire_has_zero_disabled_guards(self):
"""0=闸不启用 是对 CHS 的有意偏离(设计 §9 勘误):每道闸带 limit>0 守卫。"""
from polygateway.backends.redis.limiter import ACQUIRE
assert ACQUIRE.count("> 0 and") == 6
def test_acquire_keeps_chs_gate_operators(self):
"""并发/RPM 用 >=(占后即满),TPM 用 + est >(预扣后是否超)——CHS 同款。"""
from polygateway.backends.redis.limiter import ACQUIRE
assert ACQUIRE.count(">=") == 4
assert ACQUIRE.count("+ est >") == 2
def test_settle_lands_on_acquire_window(self):
"""SETTLE 的 key 由 Python 侧按 acquire 时窗口生成;Lua 只做 INCRBY。"""
from polygateway.backends.redis.limiter import SETTLE
assert "INCRBY" in SETTLE and "TIME" not in SETTLE