From aae739cebeb9d41d7395d26cb2590689b6716425 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Tue, 21 Jul 2026 00:27:16 -0400 Subject: [PATCH] feat: add redis six-gate rate limiter backend --- src/polygateway/backends/redis/__init__.py | 6 + src/polygateway/backends/redis/limiter.py | 307 +++++++++++++++++++++ tests/contracts/conftest.py | 102 ++++++- tests/unit/test_redis_key_layout.py | 97 +++++++ 4 files changed, 498 insertions(+), 14 deletions(-) create mode 100644 src/polygateway/backends/redis/__init__.py create mode 100644 src/polygateway/backends/redis/limiter.py create mode 100644 tests/unit/test_redis_key_layout.py diff --git a/src/polygateway/backends/redis/__init__.py b/src/polygateway/backends/redis/__init__.py new file mode 100644 index 0000000..914696e --- /dev/null +++ b/src/polygateway/backends/redis/__init__.py @@ -0,0 +1,6 @@ +"""Redis 治理后端(M2): 与内存版同一契约的跨进程实现(D3 双后端)。""" + +from polygateway.backends.redis.breaker import RedisGate +from polygateway.backends.redis.limiter import RedisLimiter + +__all__ = ["RedisGate", "RedisLimiter"] diff --git a/src/polygateway/backends/redis/limiter.py b/src/polygateway/backends/redis/limiter.py new file mode 100644 index 0000000..367405b --- /dev/null +++ b/src/polygateway/backends/redis/limiter.py @@ -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() diff --git a/tests/contracts/conftest.py b/tests/contracts/conftest.py index cb0de99..2b67a97 100644 --- a/tests/contracts/conftest.py +++ b/tests/contracts/conftest.py @@ -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 diff --git a/tests/unit/test_redis_key_layout.py b/tests/unit/test_redis_key_layout.py new file mode 100644 index 0000000..20c5132 --- /dev/null +++ b/tests/unit/test_redis_key_layout.py @@ -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