feat: add dual-condition stall detection and quiet accounting degradation
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user