Files
PolyGateway/tests/unit/test_retry.py
T
iomgaa ab496bb298 feat: let tpm be configured without an est_tokens companion
The gate check forced operators to guess a per-call token size before
they could enable the TPM gate at all; est_tokens is now an optional
tuning override and effective_est_tokens() derives the reservation from
the provider quota. Reservation and settlement already read the same
derived value, so the deposit still nets to zero on both the success
path and the non-dead transient failure path.

The rest of _validate_gates is untouched, and the est_tokens field plus
its EST_TOKENS env key stay put for migration compatibility.
2026-07-30 10:57:40 -04:00

657 lines
24 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""RetryMW 尝试编排测试(保真蓝本 CHS governance.py:107-268)。
用真实内存后端 + FakeClock + 可编程 fake transport,验证换源/退避/熔断
写回/permit 结算/取消穿透等治理行为。
"""
import asyncio
import pytest
from polygateway.backends.memory.breaker import InMemoryGate
from polygateway.backends.memory.limiter import InMemoryLimiter
from polygateway.errors import (
AllSourcesExhausted,
CircuitOpenError,
RequestRejectedError,
ResultInvalidError,
SourceDeadError,
TransientError,
)
from polygateway.middleware.retry import RetryMW
from polygateway.sources import RoundRobinSelector, SourceCooldownMemo
from polygateway.types import (
BackpressurePolicy,
BreakerConfig,
ChatRequest,
GlobalLimits,
RetryPolicy,
SourceConfig,
TransportResult,
)
from tests.contracts.conftest import FakeClock
_BREAKER = BreakerConfig(fail_threshold=3, cooldown_s=60.0, probe_ttl_s=120.0)
_NO_GLOBAL = GlobalLimits(max_concurrency=0, rpm=0, tpm=0)
def _src(name, **overrides):
base = {
"name": name,
"provider": "openai",
"base_url": "https://gw.example/v1",
"api_key": "sk",
"model": "m",
"timeout_s": 10.0,
}
base.update(overrides)
return SourceConfig(**base)
def _ok(content="ok"):
return TransportResult(
content=content,
thinking="",
prompt_tokens=10,
completion_tokens=5,
usage_source="measured",
ttft_ms=12.0,
max_inter_token_ms=3.0,
raw={},
)
class FakeTransport:
"""按脚本逐次返回结果或抛异常;记录每次 (source_name, call_id)。"""
def __init__(self, script):
self.script = list(script)
self.calls = []
async def complete(self, *, messages, source, stream, overlay, call_id):
self.calls.append((source.name, call_id))
action = self.script.pop(0)
if isinstance(action, Exception):
raise action
if action == "hang":
await asyncio.Event().wait()
return action
class FakeSleep:
"""记录退避时长,立即返回(不真等)。"""
def __init__(self):
self.delays = []
async def __call__(self, seconds):
self.delays.append(seconds)
def _harness(
sources,
script,
*,
clock=None,
max_attempts=3,
quota_full="wait",
global_limits=_NO_GLOBAL,
rng=lambda: 0.0,
selector=None,
pacer=None,
):
clock = clock or FakeClock()
limiter = InMemoryLimiter(
scope="llm",
sources={s.name: s for s in sources},
global_limits=global_limits,
lease_ttl_s=100.0,
now=clock,
)
gate = InMemoryGate(config=_BREAKER, now=clock)
transport = FakeTransport(script)
sleep = FakeSleep()
mw = RetryMW(
scope="llm",
sources=sources,
selector=selector if selector is not None else RoundRobinSelector(),
limiter=limiter,
gate=gate,
transport=transport,
retry=RetryPolicy(max_attempts=max_attempts, backoff_base_s=2.0, backoff_max_s=30.0),
backpressure=BackpressurePolicy(stall_window_s=300.0, poll_interval_s=0.01),
quota_full=quota_full,
cooldown_memo=SourceCooldownMemo(now=clock),
emitter=None,
now=clock,
sleep=sleep,
rng=rng,
pacer=pacer,
)
return mw, limiter, gate, transport, sleep, clock
_REQ = ChatRequest(messages=[{"role": "user", "content": "hi"}])
class TestSuccessPath:
async def test_first_attempt_success_builds_response(self):
mw, limiter, gate, transport, sleep, _ = _harness([_src("a")], [_ok("hello")])
resp = await mw(_REQ)
assert resp.content == "hello"
assert resp.source_name == "a" and resp.provider == "openai"
assert resp.cache_hit is False and resp.call_id
assert (await limiter.source_stats("a")).inflight == 0 # permit 已释放
assert await limiter.progress_age_s() < 5.0 # mark_progress 已调用
assert sleep.delays == []
async def test_settle_uses_actual_usage(self):
src = _src("a", tpm=1000, est_tokens=400)
mw, limiter, *_ = _harness([src], [_ok()])
await mw(_REQ)
# 预扣 400,实际 15 → settle 后窗口只记 15
assert (await limiter.source_stats("a")).tpm_used == 15
@pytest.mark.parametrize("usage_source", ["measured", "estimated"])
async def test_settle_uses_measured_sum_when_usage_available(self, usage_source):
"""用量可得(含打捞降级的 estimated)时结算恒取实测之和,不落派生兜底分支。"""
src = _src("a", tpm=1000, est_tokens=400)
result = TransportResult(
content="ok",
thinking="",
prompt_tokens=40,
completion_tokens=60,
usage_source=usage_source,
ttft_ms=12.0,
max_inter_token_ms=3.0,
raw={},
)
mw, limiter, *_ = _harness([src], [result])
await mw(_REQ)
# 预扣 400,实测 40+60 → settle 后窗口记 100(而非派生兜底的 400)
assert (await limiter.source_stats("a")).tpm_used == 100
async def test_settle_keeps_derived_deposit_when_usage_unavailable(self):
"""未填 est_tokens + usage 帧缺失的**成功**调用: 押金留存而非整笔退回。
入场预扣与结算须同取 `effective_est_tokens()`(delta==0),否则对
"从不返回 usage 帧"的源等于 TPM 闸进门即放行、出门即清账(设计 §3.2 #9)。
"""
src = _src("a", tpm=1000, est_tokens=0) # 派生预扣量 = max(1, 1000 // 60) = 16
result = TransportResult(
content="ok",
thinking="",
prompt_tokens=0,
completion_tokens=0,
usage_source="unavailable",
ttft_ms=12.0,
max_inter_token_ms=3.0,
raw={},
)
mw, limiter, *_ = _harness([src], [result])
await mw(_REQ)
assert src.effective_est_tokens() == 16
assert (await limiter.source_stats("a")).tpm_used == 16
class TestRetryAndFailover:
async def test_transient_switches_source_then_succeeds(self):
mw, _, _, transport, sleep, _ = _harness(
[_src("a"), _src("b")], [TransientError("boom"), _ok()]
)
resp = await mw(_REQ)
assert [name for name, _ in transport.calls] == ["a", "b"]
assert resp.source_name == "b"
assert len(sleep.delays) == 1 # 瞬时错误退避一次
async def test_each_attempt_gets_fresh_call_id(self):
mw, _, _, transport, _, _ = _harness([_src("a")], [TransientError("x"), _ok()])
await mw(_REQ)
ids = [cid for _, cid in transport.calls]
assert len(ids) == 2 and ids[0] != ids[1]
async def test_max_attempts_is_total_attempts(self):
mw, _, _, transport, _, _ = _harness(
[_src("a")],
[TransientError("1"), TransientError("2"), TransientError("3")],
max_attempts=3,
)
with pytest.raises(AllSourcesExhausted) as ei:
await mw(_REQ)
assert len(transport.calls) == 3 # 恰 3 次总尝试(含首次)
assert ei.value.reason == "retry_exhausted"
assert ei.value.per_source_reasons.get("a") == "network_error"
async def test_backoff_formula_and_retry_after_max(self):
# rng=0 → jitter 因子 0.5;第一次退避 = 2*2^0*0.5 = 1.0
mw, _, _, _, sleep, _ = _harness([_src("a")], [TransientError("x"), _ok()])
await mw(_REQ)
assert sleep.delays == [1.0]
# Retry-After 提示更大时取提示值
mw2, _, _, _, sleep2, _ = _harness(
[_src("a")], [TransientError("x", retry_after_s=7.5), _ok()]
)
await mw2(_REQ)
assert sleep2.delays == [7.5]
async def test_source_dead_switches_immediately_and_force_opens(self):
mw, _, gate, transport, sleep, _ = _harness(
[_src("a"), _src("b")], [SourceDeadError("401"), _ok()]
)
resp = await mw(_REQ)
assert resp.source_name == "b"
assert sleep.delays == [] # 源死亡不退避
assert not (await gate.try_enter("a", "w")).allowed # a 已 force_open
async def test_transient_failure_keeps_derived_deposit(self):
"""未填 est_tokens 的**非 dead 瞬时失败**同样按派生预扣量保守结算。
失败请求可能已被网关计费,退掉押金会低估用量(设计 §3.2 #8);
max_attempts=1 保证恰一次尝试,窗口残留量即单次预扣量。
"""
src = _src("a", tpm=1000, est_tokens=0) # 派生预扣量 = 16
mw, limiter, *_ = _harness([src], [TransientError("boom")], max_attempts=1)
with pytest.raises(AllSourcesExhausted):
await mw(_REQ)
assert (await limiter.source_stats("a")).tpm_used == 16
class TestNonRetryableOutcomes:
async def test_request_rejected_propagates_without_retry(self):
exc = RequestRejectedError("400", source_name="a", status_code=400)
mw, _, gate, transport, _, _ = _harness([_src("a")], [exc])
with pytest.raises(RequestRejectedError):
await mw(_REQ)
assert len(transport.calls) == 1
# 网关已应答 → 记成功,熔断计数未增长
assert (await gate.try_enter("a", "w")).allowed
async def test_result_invalid_records_success_and_propagates(self):
mw, limiter, gate, transport, _, _ = _harness(
[_src("a")], [ResultInvalidError("bad json", raw_text="{oops")]
)
with pytest.raises(ResultInvalidError):
await mw(_REQ)
assert len(transport.calls) == 1 # 坏结果不重试
assert (await gate.try_enter("a", "w")).allowed # 熔断记成功
assert (await limiter.source_stats("a")).inflight == 0
class TestScopeUnavailable:
async def test_all_sources_circuit_open(self):
clock = FakeClock()
script = [TransientError(str(i)) for i in range(9)]
mw, _, gate, _, _, _ = _harness([_src("a")], script, clock=clock, max_attempts=99)
# 3 次失败后 a 开路 → 第 4 次尝试选不到源且 gate_rejections==全部 → CircuitOpen
with pytest.raises(CircuitOpenError) as ei:
await mw(_REQ)
assert ei.value.reason == "circuit_open"
assert ei.value.retry_after_s > 0
assert ei.value.per_source_reasons # 携逐源原因
async def test_no_sources_configured(self):
mw, *_ = _harness([], [])
with pytest.raises(AllSourcesExhausted) as ei:
await mw(_REQ)
assert ei.value.reason == "no_sources"
async def test_quota_fail_fast(self):
src = _src("a", max_concurrency=1)
mw, limiter, _, _, _, _ = _harness([src], [_ok()], quota_full="fail_fast")
held = await limiter.try_acquire("a", 0) # 外部占满并发
assert held is not None
with pytest.raises(AllSourcesExhausted) as ei:
await mw(_REQ)
assert ei.value.reason == "quota_exhausted"
async def test_quota_wait_polls_until_slot_frees(self):
src = _src("a", max_concurrency=1)
clock = FakeClock()
limiter = InMemoryLimiter(
scope="llm",
sources={"a": src},
global_limits=_NO_GLOBAL,
lease_ttl_s=100.0,
now=clock,
)
held = await limiter.try_acquire("a", 0)
released = {"done": False}
async def sleep_and_release(seconds):
if not released["done"]:
released["done"] = True
await held.release()
gate = InMemoryGate(config=_BREAKER, now=clock)
transport = FakeTransport([_ok()])
mw = RetryMW(
scope="llm",
sources=[src],
selector=RoundRobinSelector(),
limiter=limiter,
gate=gate,
transport=transport,
retry=RetryPolicy(max_attempts=3, backoff_base_s=2.0, backoff_max_s=30.0),
backpressure=BackpressurePolicy(stall_window_s=300.0, poll_interval_s=0.01),
quota_full="wait",
cooldown_memo=SourceCooldownMemo(now=clock),
emitter=None,
now=clock,
sleep=sleep_and_release,
rng=lambda: 0.0,
)
resp = await mw(_REQ)
assert resp.content == "ok" and released["done"]
class TestCancellation:
async def test_cancel_mid_flight_releases_permit(self):
mw, limiter, _, _, _, _ = _harness([_src("a", max_concurrency=1)], ["hang"])
task = asyncio.ensure_future(mw(_REQ))
await asyncio.sleep(0.05)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert (await limiter.source_stats("a")).inflight == 0 # finally 释放
async def test_cancel_mid_backoff_propagates_with_no_held_permit(self):
"""退避 sleep 中取消: CancelledError 穿透,且 permit 早已在 finally 释放。"""
clock = FakeClock()
src = _src("a", max_concurrency=1)
limiter = InMemoryLimiter(
scope="llm",
sources={"a": src},
global_limits=_NO_GLOBAL,
lease_ttl_s=100.0,
now=clock,
)
mw = RetryMW(
scope="llm",
sources=[src],
selector=RoundRobinSelector(),
limiter=limiter,
gate=InMemoryGate(config=_BREAKER, now=clock),
transport=FakeTransport([TransientError("x"), _ok()]),
retry=RetryPolicy(max_attempts=3, backoff_base_s=30.0, backoff_max_s=60.0),
backpressure=BackpressurePolicy(300.0, 0.01),
cooldown_memo=SourceCooldownMemo(now=clock),
emitter=None,
now=clock,
sleep=asyncio.sleep,
rng=lambda: 0.5,
)
task = asyncio.ensure_future(mw(_REQ))
await asyncio.sleep(0.05) # 第一次失败后进入 30s 真实退避
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert (await limiter.source_stats("a")).inflight == 0 # 退避期不占并发槽
async def test_cancel_probe_releases_probe_lease(self):
clock = FakeClock()
mw, _, gate, _, _, _ = _harness(
[_src("a")],
[TransientError("1"), TransientError("2"), TransientError("3"), "hang"],
clock=clock,
max_attempts=99,
)
# 三连失败开路
with pytest.raises(CircuitOpenError):
await mw(_REQ)
clock.advance(_BREAKER.cooldown_s + 1)
task = asyncio.ensure_future(mw(_REQ)) # 半开探针 → hang
await asyncio.sleep(0.05)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
# 探针租约已归还: 下一 caller 立即拿到探针而非等租约过期
nxt = await gate.try_enter("a", "w2")
assert nxt.allowed and nxt.is_probe
class StaticSelector:
"""固定配置序,隔离测试调用内降权(不带 record_outcome)。"""
def order(self, sources, stats):
return list(sources)
class RecordingSelector(StaticSelector):
def __init__(self):
self.outcomes = []
def record_outcome(self, source_name, ok):
self.outcomes.append((source_name, ok))
def health(self, source_name):
return 1.0
class ExplodingSelector(StaticSelector):
def record_outcome(self, source_name, ok):
raise RuntimeError("sink boom")
def health(self, source_name):
return 1.0
class TestM25Orchestration:
"""M2.5 设计 §3.3: 调用内失败降权 + 健康喂数(先红后绿)。"""
async def test_failed_source_demoted_after_two_strikes(self):
# 失败 1 次仍首选(原地退避重试);失败 2 次让位次优源
mw, _, _, transport, _, _ = _harness(
[_src("a"), _src("b")],
[TransientError("1"), TransientError("2"), _ok()],
selector=StaticSelector(),
)
resp = await mw(_REQ)
assert [n for n, _ in transport.calls] == ["a", "a", "b"]
assert resp.source_name == "b"
async def test_attempt_fails_reset_between_calls(self):
mw, _, _, transport, _, _ = _harness(
[_src("a"), _src("b")],
[TransientError("1"), TransientError("2"), _ok(), _ok()],
selector=StaticSelector(),
)
await mw(_REQ)
await mw(_REQ) # 新调用状态清零: 回到首选 a
assert [n for n, _ in transport.calls] == ["a", "a", "b", "a"]
async def test_outcome_feeding_success_and_transient(self):
sel = RecordingSelector()
mw, _, _, _, _, _ = _harness(
[_src("a"), _src("b")], [TransientError("1"), _ok()], selector=sel
)
await mw(_REQ)
assert sel.outcomes == [("a", False), ("a", True)]
async def test_outcome_feeding_source_dead(self):
sel = RecordingSelector()
mw, _, _, _, _, _ = _harness(
[_src("a"), _src("b")], [SourceDeadError("401"), _ok()], selector=sel
)
await mw(_REQ)
assert sel.outcomes == [("a", False), ("b", True)]
async def test_result_invalid_and_rejected_not_fed(self):
sel = RecordingSelector()
mw, _, _, _, _, _ = _harness(
[_src("a")], [ResultInvalidError("bad", raw_text="x")], selector=sel
)
with pytest.raises(ResultInvalidError):
await mw(_REQ)
sel2 = RecordingSelector()
mw2, _, _, _, _, _ = _harness(
[_src("a")],
[RequestRejectedError("400", source_name="a", status_code=400)],
selector=sel2,
)
with pytest.raises(RequestRejectedError):
await mw2(_REQ)
assert sel.outcomes == [] and sel2.outcomes == []
async def test_outcome_sink_exception_swallowed(self):
mw, _, _, _, _, _ = _harness([_src("a")], [_ok()], selector=ExplodingSelector())
resp = await mw(_REQ)
assert resp.content == "ok" # 喂数异常不得打断真实成功返回
async def test_result_invalid_gate_window_untouched(self):
# 坏结果 ≠ 坏服务: count_attempt=False,失败率窗口 attempts 不得增长
mw, _, gate, _, _, _ = _harness([_src("a")], [ResultInvalidError("bad", raw_text="x")])
with pytest.raises(ResultInvalidError):
await mw(_REQ)
g = gate._gates["a"]
assert g.a0 + g.a1 == 0
class TestAdaptivePacing:
"""AIMD 接线(设计 §3.35): 429 收紧准入,paced 源等待而非烧预算。"""
async def test_paced_source_waits_without_consuming_budget(self):
from polygateway.sources import AdaptivePacer
pacer = AdaptivePacer(ceiling=32.0)
for _ in range(200):
pacer.on_backpressure("a") # limit → 1
pacer.enter("a") # 模拟一个在途占满名额
mw, _, _, transport, _, _ = _harness(
[_src("a")], [_ok()], quota_full="fail_fast", pacer=pacer
)
with pytest.raises(AllSourcesExhausted) as ei:
await mw(_REQ)
assert ei.value.reason == "quota_exhausted" # 走配额等待通道,非 CircuitOpen
assert ei.value.per_source_reasons.get("a") == "adaptive_paced"
assert transport.calls == [] # 未发起尝试 → 不烧重试预算
async def test_429_cuts_limit_success_grows_it(self):
from polygateway.sources import AdaptivePacer
pacer = AdaptivePacer(ceiling=32.0)
mw, _, _, _, _, _ = _harness(
[_src("a")],
[TransientError("throttled", status_code=429), _ok()],
pacer=pacer,
)
await mw(_REQ)
# 429 削减一次(8→4),随后成功加性增长(4 + 1/4)
assert pacer.limit("a") == pytest.approx(8.0 * 0.5 + 1.0 / (8.0 * 0.5))
async def test_inflight_returns_to_zero_after_call(self):
from polygateway.sources import AdaptivePacer
pacer = AdaptivePacer(ceiling=32.0)
mw, _, _, _, _, _ = _harness(
[_src("a"), _src("b")],
[TransientError("x"), _ok()],
pacer=pacer,
)
await mw(_REQ)
assert pacer._inflight.get("a", 0) == 0
assert pacer._inflight.get("b", 0) == 0
class HealthySink(StaticSelector):
"""带健康视图的选源器桩(OutcomeAwareSelector 全量实现)。"""
def __init__(self, health):
self._health = health
self.outcomes = []
def record_outcome(self, source_name, ok):
self.outcomes.append((source_name, ok))
def health(self, source_name):
return self._health.get(source_name, 1.0)
class TestHealthGatedDemotion:
"""迭代 2(设计 §3.36): 降权需可信替代,否则原地重试。"""
async def test_no_credible_alternative_stays_on_healthy(self):
# 替补健康分 0.08 < 0.5×0.9 → 不让位,第三次仍打 a
sel = HealthySink({"a": 0.9, "b": 0.08})
mw, _, _, transport, _, _ = _harness(
[_src("a"), _src("b")],
[TransientError("1"), TransientError("2"), _ok()],
selector=sel,
)
resp = await mw(_REQ)
assert [n for n, _ in transport.calls] == ["a", "a", "a"]
assert resp.source_name == "a"
async def test_credible_alternative_still_yields(self):
sel = HealthySink({"a": 0.9, "b": 0.9})
mw, _, _, transport, _, _ = _harness(
[_src("a"), _src("b")],
[TransientError("1"), TransientError("2"), _ok()],
selector=sel,
)
await mw(_REQ)
assert [n for n, _ in transport.calls] == ["a", "a", "b"]
class TestDemotionInsertPosition:
"""迭代 3(设计 §3.36 补): 被降权源插在可信替代之后、不可信源之前。"""
async def test_demoted_lands_before_junk_sources(self):
# a 失败 2 次;b 可信(0.9)但会被跳过时,第三候选应是 a 而非垃圾源 c
from polygateway.middleware.retry import _demote_call_failures
srcs = [_src("a"), _src("b"), _src("c")]
health = {"a": 0.9, "b": 0.9, "c": 0.05}.__getitem__
out = _demote_call_failures(srcs, {"a": 2}, health)
assert [s.name for s in out] == ["b", "a", "c"]
async def test_health_blind_demotion_still_tail(self):
from polygateway.middleware.retry import _demote_call_failures
srcs = [_src("a"), _src("b"), _src("c")]
out = _demote_call_failures(srcs, {"a": 2}, None)
assert [s.name for s in out] == ["b", "c", "a"]
class TestRateLimitPushback:
"""迭代 5(设计 §3.38): 429 是服务端调度指令,不耗重试预算;时间上限兜底。"""
async def test_429_does_not_consume_retry_budget(self):
# 3 连 429 后成功——若 429 计预算,max_attempts=3 时第 4 次不会发生
mw, _, _, transport, sleep, _ = _harness(
[_src("a")],
[
TransientError("t1", status_code=429, retry_after_s=1.0),
TransientError("t2", status_code=429, retry_after_s=1.0),
TransientError("t3", status_code=429, retry_after_s=1.0),
_ok(),
],
)
resp = await mw(_REQ)
assert resp.content == "ok"
assert len(transport.calls) == 4
assert len(sleep.delays) == 3 # 每次 429 仍按 Retry-After 退避
async def test_429_storm_bounded_by_stall_window(self):
# 持续 429 且时钟推进超 stall_window → stalled 兜底,不无限循环
clock = FakeClock()
script = [TransientError(str(i), status_code=429, retry_after_s=30.0) for i in range(99)]
mw, _, _, _, _, _ = _harness([_src("a")], script, clock=clock)
async def advancing_sleep(seconds):
clock.advance(seconds)
mw._sleep = advancing_sleep
with pytest.raises(AllSourcesExhausted) as ei:
await mw(_REQ)
assert ei.value.reason == "stalled"
async def test_non_429_transient_still_consumes_budget(self):
mw, _, _, transport, _, _ = _harness(
[_src("a")],
[TransientError("1"), TransientError("2"), TransientError("3")],
)
with pytest.raises(AllSourcesExhausted) as ei:
await mw(_REQ)
assert ei.value.reason == "retry_exhausted"
assert len(transport.calls) == 3