diff --git a/src/polygateway/middleware/ratelimit.py b/src/polygateway/middleware/ratelimit.py index 6035dcf..f9ee398 100644 --- a/src/polygateway/middleware/ratelimit.py +++ b/src/polygateway/middleware/ratelimit.py @@ -44,3 +44,11 @@ class QuotaGate: raise except Exception as exc: raise GovernanceBackendError(f"限流后端故障(mark_progress): {exc}") from exc + + async def progress_age_s(self) -> float: + try: + return await self._limiter.progress_age_s() + except GovernanceBackendError: + raise + except Exception as exc: + raise GovernanceBackendError(f"限流后端故障(progress_age_s): {exc}") from exc diff --git a/src/polygateway/middleware/retry.py b/src/polygateway/middleware/retry.py index 6a681ce..b892ac2 100644 --- a/src/polygateway/middleware/retry.py +++ b/src/polygateway/middleware/retry.py @@ -23,6 +23,7 @@ from loguru import logger from polygateway.errors import ( AllSourcesExhausted, CircuitOpenError, + GovernanceBackendError, PolyGatewayError, RequestRejectedError, ResultInvalidError, @@ -55,6 +56,22 @@ if TYPE_CHECKING: ) +def backoff_delay( + policy: RetryPolicy, + fails: int, + exc: BaseException | None, + rng: Callable[[], float], +) -> float: + """指数退避+jitter,与 Retry-After 提示取大(ARCH §7.2;VT jitter 系数)。 + + 模块级纯函数: RetryMW 与 EmbeddingClient(M2 §7)共用同一公式。 + """ + base = min(policy.backoff_base_s * (2 ** (fails - 1)), policy.backoff_max_s) + delay = base * (0.5 + rng()) + retry_after = getattr(exc, "retry_after_s", None) or 0.0 + return max(delay, retry_after) + + def _failure_reason(exc: PolyGatewayError) -> str: """失败原因归类(CHS governance.py:169 同款)。""" if isinstance(exc, SourceDeadError): @@ -118,10 +135,11 @@ class RetryMW: raise AllSourcesExhausted(scope=self._scope, reason="no_sources", retry_after_s=0.0) fails = 0 reasons: dict[str, str] = {} + entered_at = self._now() # 调用级累计计时,循环内不重置(CHS governance.py:207) while True: picked, gate_rejections = await self._pick_runnable(reasons) if picked is None: - await self._on_no_runnable(gate_rejections, reasons) + await self._on_no_runnable(gate_rejections, reasons, entered_at) continue outcome = await self._attempt(request, *picked, reasons) if isinstance(outcome, LLMResponse): @@ -170,7 +188,9 @@ class RetryMW: await self._settle_and_release(permit, 0) return None, gate_rejections - async def _on_no_runnable(self, gate_rejections: int, reasons: dict[str, str]) -> None: + async def _on_no_runnable( + self, gate_rejections: int, reasons: dict[str, str], entered_at: float + ) -> None: if gate_rejections == len(self._sources): names = tuple(s.name for s in self._sources) raise CircuitOpenError( @@ -185,7 +205,19 @@ class RetryMW: retry_after_s=self._bp.poll_interval_s, per_source_reasons=reasons, ) - await self._sleep(self._bp.poll_interval_s) + # 双条件 stall 判死(CHS governance.py:270-281): 本地累计等待与全局 + # 无进展**同时**超窗才判死——本地 monotonic 与后端时钟刻意不混用。 + stall = self._bp.stall_window_s + if self._now() - entered_at > stall and await self._quota.progress_age_s() > stall: + names = tuple(s.name for s in self._sources) + raise AllSourcesExhausted( + scope=self._scope, + reason="stalled", + retry_after_s=await self._breaker.retry_after_s(names), + per_source_reasons=reasons, + ) + # jitter ∈ [0.5p, 1.0p] 防惊群(CHS governance.py:283-285) + await self._sleep(self._bp.poll_interval_s * (0.5 + 0.5 * self._rng())) # —— 单次尝试(CHS run 200-268)—— @@ -209,8 +241,8 @@ class RetryMW: call_id=call_id, ) actual = result.prompt_tokens + result.completion_tokens - await self._breaker.record_success(entry) - await self._quota.mark_progress() + await self._record_quietly(self._breaker.record_success(entry)) + await self._record_quietly(self._quota.mark_progress()) response = self._build_response(source, result, call_id, started) await self._emit(request, source, call_id, started, response=response) return response @@ -220,19 +252,19 @@ class RetryMW: raise except ResultInvalidError as exc: # 坏结果 ≠ 坏服务: 熔断记成功,异常上抛消耗业务失败预算(§6.3) - await self._breaker.record_success(entry) + await self._record_quietly(self._breaker.record_success(entry)) await self._emit(request, source, call_id, started, error=exc) raise except asyncio.CancelledError: if entry.is_probe: - await self._breaker.release_probe(entry) + await self._record_quietly(self._breaker.release_probe(entry)) await self._emit(request, source, call_id, started, error="cancelled") raise except (SourceDeadError, TransientError) as exc: dead = isinstance(exc, SourceDeadError) reason = _failure_reason(exc) reasons[source.name] = reason - await self._breaker.record_failure(entry, reason, dead) + await self._record_quietly(self._breaker.record_failure(entry, reason, dead)) if not dead: actual = source.est_tokens # 保守: 失败请求可能已被网关计费(CHS 同款) await self._emit(request, source, call_id, started, error=exc) @@ -245,18 +277,29 @@ class RetryMW: ) -> None: provider_responded = exc.source_name == source.name and exc.status_code is not None if provider_responded: - await self._breaker.record_success(entry) # 网关健康地拒了坏请求 + # 网关健康地拒了坏请求 + await self._record_quietly(self._breaker.record_success(entry)) elif entry.is_probe: - await self._breaker.release_probe(entry) + await self._record_quietly(self._breaker.release_probe(entry)) + + async def _record_quietly(self, write_back: Awaitable[object]) -> None: + """记账侧写回(record_*/mark_progress/release_probe)降级执行(设计 §10)。 + + 调用已真实完成: 后端失败若冒泡会丢弃真实成功响应或掩盖原始尝试 + 异常,故 warning 降级(ARCH §7.3 勘误,CHS 全 fail-closed 的有意反转); + 取消照常穿透。 + """ + try: + await write_back + except asyncio.CancelledError: + raise + except GovernanceBackendError as exc: + logger.warning("治理记账写回降级(不冒泡): {}", exc) # —— 辅助 —— def _backoff_delay(self, fails: int, exc: PolyGatewayError) -> float: - """指数退避+jitter,与 Retry-After 提示取大(ARCH §7.2;VT jitter 系数)。""" - base = min(self._retry.backoff_base_s * (2 ** (fails - 1)), self._retry.backoff_max_s) - delay = base * (0.5 + self._rng()) - retry_after = getattr(exc, "retry_after_s", None) or 0.0 - return max(delay, retry_after) + return backoff_delay(self._retry, fails, exc, self._rng) def _build_response( self, source: SourceConfig, result: TransportResult, call_id: str, started: float diff --git a/tests/unit/test_backpressure.py b/tests/unit/test_backpressure.py new file mode 100644 index 0000000..a9065ed --- /dev/null +++ b/tests/unit/test_backpressure.py @@ -0,0 +1,257 @@ +"""背压 stall 双条件判死 + poll jitter + 记账侧降级(M2 设计 §4/§10)。 + +保真蓝本 CHS governance.py:200-285: local_waited(本地 monotonic,调用级 +累计不重置)与 progress_age(全局活性)**同时**超 stall_window 才判死; +poll 间隔带 [0.5p, 1.0p] jitter 防惊群。 +""" + +import asyncio + +import pytest + +from polygateway.backends.memory.breaker import InMemoryGate +from polygateway.backends.memory.limiter import InMemoryLimiter +from polygateway.errors import AllSourcesExhausted, GovernanceBackendError, TransientError +from polygateway.middleware.retry import RetryMW, backoff_delay +from polygateway.sources import RoundRobinSelector, SourceCooldownMemo +from polygateway.types import ( + BackpressurePolicy, + BreakerConfig, + ChatRequest, + GlobalLimits, + RetryPolicy, +) +from tests.contracts.conftest import FakeClock, make_source +from tests.unit.test_retry import FakeTransport, _ok + +_BREAKER = BreakerConfig(fail_threshold=3, cooldown_s=60.0, probe_ttl_s=120.0) +_NO_GLOBAL = GlobalLimits(max_concurrency=0, rpm=0, tpm=0) +_REQ = ChatRequest(messages=[{"role": "user", "content": "hi"}]) +_STALL = 300.0 + + +class BoundedSleep: + """每次 sleep 执行注入的副作用;超过上限仍在轮询 = 判死逻辑失效,炸出而非死循环。""" + + def __init__(self, side_effect=None, limit: int = 10): + self.delays: list[float] = [] + self._side_effect = side_effect + self._limit = limit + + async def __call__(self, seconds: float) -> None: + self.delays.append(seconds) + if len(self.delays) > self._limit: + raise RuntimeError(f"超过 {self._limit} 次轮询仍未判死/未获 permit") + if self._side_effect is not None: + await self._side_effect(len(self.delays)) + + +def _mw(sources, limiter, script, *, clock, sleep, rng=lambda: 0.0, quota_full="wait", gate=None): + return RetryMW( + scope="llm", + sources=sources, + selector=RoundRobinSelector(), + limiter=limiter, + gate=gate or InMemoryGate(config=_BREAKER, now=clock), + transport=FakeTransport(script), + retry=RetryPolicy(max_attempts=3, backoff_base_s=2.0, backoff_max_s=30.0), + backpressure=BackpressurePolicy(stall_window_s=_STALL, poll_interval_s=0.01), + quota_full=quota_full, + cooldown_memo=SourceCooldownMemo(now=clock), + emitter=None, + now=clock, + sleep=sleep, + rng=rng, + ) + + +def _blocked_limiter(clock): + """单并发源被外部占满 → RetryMW 进入 wait 轮询分支。""" + src = make_source(max_concurrency=1) + limiter = InMemoryLimiter( + scope="llm", + sources={"s1": src}, + global_limits=_NO_GLOBAL, + lease_ttl_s=10_000.0, + now=clock, + ) + return src, limiter + + +class TestStallQuadrants: + async def test_both_windows_exceeded_raises_stalled(self): + """双超窗: 本地等待与全局无进展同时 > stall_window → stalled。""" + clock = FakeClock() + src, limiter = _blocked_limiter(clock) + _held = await limiter.try_acquire("s1", 0) + + async def advance(_n): + clock.advance(_STALL + 100) + + mw = _mw([src], limiter, [], clock=clock, sleep=BoundedSleep(advance)) + with pytest.raises(AllSourcesExhausted) as ei: + await mw(_REQ) + assert ei.value.reason == "stalled" + assert ei.value.per_source_reasons == {"s1": "rate_limited"} + + async def test_local_exceeded_but_global_fresh_keeps_waiting(self): + """仅本地超窗: 别人一直在出餐 → 不判死,等到 permit 后正常成功。""" + clock = FakeClock() + src, limiter = _blocked_limiter(clock) + held = await limiter.try_acquire("s1", 0) + + async def advance_and_feed(n): + clock.advance(_STALL + 100) # 本地远超窗 + await limiter.mark_progress() # 但全局刚出过餐 + if n >= 3: + await held.release() # 第 3 轮让出 permit + + mw = _mw([src], limiter, [_ok()], clock=clock, sleep=BoundedSleep(advance_and_feed)) + resp = await mw(_REQ) + assert resp.content == "ok" + + async def test_global_stale_but_local_fresh_keeps_waiting(self): + """仅全局超窗(从未出餐 age=inf): 本地才刚开始等 → 不判死。""" + clock = FakeClock() + src, limiter = _blocked_limiter(clock) + held = await limiter.try_acquire("s1", 0) + + async def release_soon(n): + if n >= 2: # 本地累计 poll 极短,未超窗 + await held.release() + + mw = _mw([src], limiter, [_ok()], clock=clock, sleep=BoundedSleep(release_soon)) + resp = await mw(_REQ) + assert resp.content == "ok" + + async def test_poll_jitter_bounds(self): + """wait 轮询间隔 ∈ [0.5p, 1.0p](CHS governance.py:283-285)。""" + clock = FakeClock() + src, limiter = _blocked_limiter(clock) + held = await limiter.try_acquire("s1", 0) + + async def release_soon(n): + if n >= 2: + await held.release() + + for rng_v, expected in ((0.0, 0.005), (1.0, 0.01)): + clock2 = FakeClock() + src2, limiter2 = _blocked_limiter(clock2) + held2 = await limiter2.try_acquire("s1", 0) + + async def release2(n, _h=held2): + if n >= 2: + await _h.release() + + sleep = BoundedSleep(release2) + mw = _mw([src2], limiter2, [_ok()], clock=clock2, sleep=sleep, rng=lambda v=rng_v: v) + await mw(_REQ) + assert sleep.delays[0] == pytest.approx(expected) + + async def test_fail_fast_unaffected(self): + """fail_fast 路径回归: 不进入 stall 判定,立即 quota_exhausted。""" + clock = FakeClock() + src, limiter = _blocked_limiter(clock) + _held = await limiter.try_acquire("s1", 0) + mw = _mw([src], limiter, [], clock=clock, sleep=BoundedSleep(), quota_full="fail_fast") + with pytest.raises(AllSourcesExhausted) as ei: + await mw(_REQ) + assert ei.value.reason == "quota_exhausted" + + async def test_cancellation_pierces_wait_loop(self): + """stall 等待中的取消穿透(铁律): sleep 可取消,任务立即终止。""" + clock = FakeClock() + src, limiter = _blocked_limiter(clock) + _held = await limiter.try_acquire("s1", 0) + mw = _mw([src], limiter, [], clock=clock, sleep=asyncio.sleep) + task = asyncio.create_task(mw(_REQ)) + await asyncio.sleep(0.03) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + +class _GateSuccessBroken(InMemoryGate): + async def record_success(self, entry): + raise GovernanceBackendError("redis 抖动") + + +class _GateFailureBroken(InMemoryGate): + async def record_failure(self, entry, reason, force_open): + raise GovernanceBackendError("redis 抖动") + + +class _LimiterProgressBroken(InMemoryLimiter): + async def mark_progress(self): + raise GovernanceBackendError("redis 抖动") + + +class TestAccountingDegradation: + """记账侧降级(设计 §10,ARCH §7.3 勘误): 调用已完成,写回失败不冒泡。""" + + async def test_record_success_failure_does_not_lose_response(self): + clock = FakeClock() + src = make_source() + limiter = InMemoryLimiter( + scope="llm", sources={"s1": src}, global_limits=_NO_GLOBAL, now=clock + ) + gate = _GateSuccessBroken(config=_BREAKER, now=clock) + mw = _mw([src], limiter, [_ok()], clock=clock, sleep=BoundedSleep(), gate=gate) + resp = await mw(_REQ) + assert resp.content == "ok" # 真实成功响应不因记账失败被丢弃 + + async def test_mark_progress_failure_does_not_lose_response(self): + clock = FakeClock() + src = make_source() + limiter = _LimiterProgressBroken( + scope="llm", sources={"s1": src}, global_limits=_NO_GLOBAL, now=clock + ) + mw = _mw([src], limiter, [_ok()], clock=clock, sleep=BoundedSleep()) + resp = await mw(_REQ) + assert resp.content == "ok" + + async def test_record_failure_failure_does_not_mask_retry(self): + """失败记账挂掉 → 原始尝试异常不被掩盖,重试照常换发并成功。""" + clock = FakeClock() + src = make_source() + limiter = InMemoryLimiter( + scope="llm", sources={"s1": src}, global_limits=_NO_GLOBAL, now=clock + ) + gate = _GateFailureBroken(config=_BREAKER, now=clock) + script = [TransientError("boom", status_code=500), _ok()] + mw = _mw([src], limiter, script, clock=clock, sleep=BoundedSleep(), gate=gate) + resp = await mw(_REQ) + assert resp.content == "ok" + + +class TestBackoffExtraction: + """T5 提取的模块级纯函数(T9 Embedding 复用)。""" + + def test_formula_matches_policy(self): + policy = RetryPolicy(max_attempts=3, backoff_base_s=2.0, backoff_max_s=30.0) + exc = TransientError("x", status_code=500) + assert backoff_delay(policy, 1, exc, lambda: 0.0) == pytest.approx(1.0) # 2*0.5 + assert backoff_delay(policy, 1, exc, lambda: 1.0) == pytest.approx(3.0) # 2*1.5 + assert backoff_delay(policy, 10, exc, lambda: 1.0) == pytest.approx(45.0) # 封顶 30*1.5 + + def test_retry_after_hint_wins_when_larger(self): + policy = RetryPolicy(max_attempts=3, backoff_base_s=2.0, backoff_max_s=30.0) + exc = TransientError("x", status_code=429, retry_after_s=7.5) + assert backoff_delay(policy, 1, exc, lambda: 0.0) == pytest.approx(7.5) + + +class TestQuotaGateProgressAge: + async def test_passthrough_and_wrap(self): + from polygateway.middleware.ratelimit import QuotaGate + + class _L: + async def progress_age_s(self): + return 12.5 + + class _Broken: + async def progress_age_s(self): + raise OSError("down") + + assert await QuotaGate(_L()).progress_age_s() == 12.5 + with pytest.raises(GovernanceBackendError): + await QuotaGate(_Broken()).progress_age_s()