feat: add opt-in cross-source hedged requests for chat

This commit is contained in:
2026-09-10 14:00:50 -04:00
parent 463eca380d
commit adc069447a
7 changed files with 1018 additions and 13 deletions
+16 -1
View File
@@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, Any, Literal
from polygateway.backends.memory.breaker import InMemoryGate
from polygateway.backends.memory.cache import InMemoryCache
from polygateway.backends.memory.limiter import InMemoryLimiter
from polygateway.config import GatewaySettings
from polygateway.config import GatewaySettings, check_hedge_assembly
from polygateway.deadline import ensure_call_deadline, with_call_deadline
from polygateway.errors import PolyGatewayError
from polygateway.middleware.base import compose
@@ -237,6 +237,8 @@ class GatewayClient:
structured_escalation: StructuredOutputStrategy | None = None,
structured_max_retries: int = 1,
call_deadline_s: float | None = None,
hedge_after_s: float | None = None,
hedge_max_extra: int = 1,
now: Any = time.monotonic,
sleep: Any = asyncio.sleep,
rng: Any = random.random,
@@ -245,6 +247,15 @@ class GatewayClient:
self._call_deadline_s = ensure_call_deadline(
call_deadline_s, "GatewayClient(call_deadline_s=...)"
)
# 对冲守卫与 GatewaySettings 共用同一份(issue #24 H4): 直接构造这条路
# 不经过 settings,值域/交叉守卫若只挂在 settings 上就会被它绕过
self._hedge_after_s = check_hedge_assembly(
hedge_after_s=hedge_after_s,
hedge_max_extra=hedge_max_extra,
sources=sources,
call_deadline_s=self._call_deadline_s,
origin="GatewayClient(hedge_after_s=...)",
)
emitter = (
TelemetryEmitter(telemetry, scope=scope, pricing=pricing, text_cap=text_cap)
if telemetry is not None
@@ -267,6 +278,8 @@ class GatewayClient:
ceiling=float(max([64, *(s.max_concurrency for s in sources if s.max_concurrency)]))
),
emitter=emitter,
# max_extra 不下传(issue #24 H5): v1 编排固定单路对冲,无消费者
hedge_after_s=self._hedge_after_s,
now=now,
sleep=sleep,
rng=rng,
@@ -509,6 +522,8 @@ class GatewayClient:
structured_escalation=escalation,
structured_max_retries=settings.structured_max_retries,
call_deadline_s=settings.call_deadline_s,
hedge_after_s=settings.hedge_after_s,
hedge_max_extra=settings.hedge_max_extra,
)
_mark_owned_components(client, limiter=limiter, breaker=breaker, telemetry=telemetry)
client._owns_cache = cache is None # 缓存后端可以是 None(backend=none),helper 会跳过
+127 -2
View File
@@ -30,7 +30,7 @@ from polygateway.types import (
)
if TYPE_CHECKING:
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
# FIELD → (SourceConfig 属性, 类型);CHS config.py:95-104 全集 + M1 新增
_SOURCE_FIELDS: dict[str, tuple[str, str]] = {
@@ -54,7 +54,7 @@ _SOURCE_FIELDS: dict[str, tuple[str, str]] = {
"TRUST_ENV": ("trust_env", "bool"),
"EXTRA_BODY": ("extra_body", "json"),
}
_RESERVED_SEGMENTS = frozenset({"GLOBAL", "RETRY", "BREAKER", "BACKPRESSURE"})
_RESERVED_SEGMENTS = frozenset({"GLOBAL", "RETRY", "BREAKER", "BACKPRESSURE", "HEDGE"})
_SELECTORS = frozenset({"round_robin", "least_inflight", "health_aware"})
_QUOTA_FULL = frozenset({"wait", "fail_fast"})
# 熔断全拒时的处置(issue #14);值域与 _QUOTA_FULL 相同但语义不同——配额满是
@@ -188,6 +188,12 @@ class GatewaySettings:
# 与 `OcrSettings.gateway` 自动继承。值域由 `_validate_call_deadline` 把关,
# 直接构造、`dataclasses.replace` 与 env 三条路一致
call_deadline_s: float | None = None
# 长尾对冲(issue #24 H4): 挂起超阈值时并发向异源再发一次,先回者赢、输家取消。
# 缺省 None = 关闭,行为逐字等于 1.3.6。`hedge_max_extra` 值域 [1,3] 为 H5 梯次
# 预留,v1 仅单路生效(>1 装配期 warning);**不下传 RetryMW**。值域与交叉守卫
# 由 `check_hedge_assembly` 单一定义点把关,直接构造/replace/env 三路一致
hedge_after_s: float | None = None
hedge_max_extra: int = 1
def __post_init__(self) -> None:
self._normalize()
@@ -199,6 +205,7 @@ class GatewaySettings:
self._validate_stall()
self._validate_probe()
self._validate_call_deadline()
self._validate_hedge()
def _normalize(self) -> None:
"""把 `from_env` 一直在做的规范化补到构造路上,两条路必须产出同一个值。
@@ -365,6 +372,24 @@ class GatewaySettings:
ensure_call_deadline(self.call_deadline_s, "GatewaySettings.call_deadline_s"),
)
def _validate_hedge(self) -> None:
"""对冲装配守卫(issue #24 H4): 与期限同款,盖住直接构造与 replace 两条路。
env 路的值域错误已在 `_load_hedge` 里带真实键名报过;此处对合法值是幂等
空操作,交叉守卫(阈值 vs timeout/deadline/ttft、单源)只在这里有一处。
"""
object.__setattr__(
self,
"hedge_after_s",
check_hedge_assembly(
hedge_after_s=self.hedge_after_s,
hedge_max_extra=self.hedge_max_extra,
sources=self.sources,
call_deadline_s=self.call_deadline_s,
origin="GatewaySettings.hedge_after_s",
),
)
@classmethod
def from_env(
cls,
@@ -394,6 +419,7 @@ class GatewaySettings:
quota_full=_load_choice(env, f"{scope_u}__QUOTA_FULL", _QUOTA_FULL, "wait"),
circuit_open=_load_choice(env, f"{scope_u}__CIRCUIT_OPEN", _CIRCUIT_OPEN, "fail_fast"),
call_deadline_s=_load_call_deadline(scope_u, env),
**_load_hedge(scope_u, env),
**_load_pgw(env),
)
@@ -716,6 +742,105 @@ def _load_call_deadline(scope: str, env: Mapping[str, str]) -> float | None:
return ensure_call_deadline(_cast(found[1], "float", found[0]), found[0])
def _load_hedge(scope: str, env: Mapping[str, str]) -> dict[str, object]:
"""读 `{SCOPE}__HEDGE__AFTER_S`/`{SCOPE}__HEDGE__MAX_EXTRA`(issue #24 H4)。
两键均为 3 段键(`split("__")` 长度 3 ≠ 4),`_load_sources` 的段数判据天然
跳过它们;`HEDGE` 已进 `_RESERVED_SEGMENTS`,4 段的 `{SCOPE}__HEDGE__{N}__*`
也不会被当成 provider 段造出源。`AFTER_S` 未设 = 关闭(缺省逐字等于 1.3.6);
`MAX_EXTRA` 未设 = 1。origin 传实际命中键名(同 `_load_call_deadline` 纪律);
值域的交叉守卫(ttft/单源/max_extra 生效口径)归 `check_hedge_assembly` 一处。
Args:
scope: 已大写的 scope 名。
env: 已合并的环境映射。
Returns:
`{"hedge_after_s": float | None, "hedge_max_extra": int}`,直传构造器。
"""
found_after = _first(env, f"{scope}__HEDGE__AFTER_S")
after = (
ensure_call_deadline(_cast(found_after[1], "float", found_after[0]), found_after[0])
if found_after
else None
)
found_extra = _first(env, f"{scope}__HEDGE__MAX_EXTRA")
extra = int(_cast(found_extra[1], "int", found_extra[0])) if found_extra else 1
return {"hedge_after_s": after, "hedge_max_extra": extra}
def check_hedge_assembly(
*,
hedge_after_s: float | None,
hedge_max_extra: int,
sources: Sequence[SourceConfig],
call_deadline_s: float | None,
origin: str,
) -> float | None:
"""对冲装配守卫的唯一事实源(issue #24 设计 §5 全表,H4 批准)。
`GatewaySettings.__post_init__` 与 `GatewayClient.__init__` 调同一份,两条装配
路的值域/交叉守卫不漂移。返回归一化后的 `hedge_after_s`(None 或有限正数,
复用 `ensure_call_deadline` 的值域校验);`hedge_after_s is None`(未启用)时
值域归一化后直接返回,交叉守卫不查——它们没有可校验的对象。
Args:
hedge_after_s: 对冲触发阈值(秒);None = 关闭。
hedge_max_extra: 每次逻辑调用最多并发对冲路数;v1 仅单路生效(H5)。
sources: 本 scope 的源集合(交叉守卫要读 timeout_s/ttft_timeout_s)。
call_deadline_s: 调用期限(秒);与对冲的组合守卫见设计 §6。
origin: after_s 值域报错的定位串(env 键名 / `GatewaySettings.hedge_after_s` /
`GatewayClient(hedge_after_s=...)`);跨字段守卫与 warning 沿用本模块先例,
消息自带字段名与 env 键型,不挂 origin。
Raises:
ValueError: max_extra 非 int/bool 或出 [1,3];阈值 ≥ 最小源 timeout_s(永不
可能触发);阈值 ≥ call_deadline_s(期限先于对冲触发,对冲形同虚设)。
"""
if isinstance(hedge_max_extra, bool) or not isinstance(hedge_max_extra, int):
raise ValueError(
f"hedge_max_extra({{SCOPE}}__HEDGE__MAX_EXTRA)必须是 int: {hedge_max_extra!r}"
)
if not 1 <= hedge_max_extra <= 3:
raise ValueError(
f"hedge_max_extra({{SCOPE}}__HEDGE__MAX_EXTRA)须在 [1,3]: {hedge_max_extra};"
"每次逻辑调用最多并发对冲路数,v1 仅单路生效"
)
after = ensure_call_deadline(hedge_after_s, origin)
if after is None:
return None
min_timeout = min(s.timeout_s for s in sources)
if after >= min_timeout:
raise ValueError(
f"hedge_after_s({{SCOPE}}__HEDGE__AFTER_S={after})须 < 最小源 timeout_s"
f"({min_timeout});对冲永不可能触发,配置即错误"
)
if call_deadline_s is not None and after >= call_deadline_s:
raise ValueError(
f"hedge_after_s({after})须 < call_deadline_s({call_deadline_s});"
"期限会先于对冲触发,对冲形同虚设(设计 §6)"
)
ttfts = [s.ttft_timeout_s for s in sources if s.ttft_timeout_s is not None]
if ttfts and after >= min(ttfts):
logger.warning(
"hedge_after_s({}) ≥ 最小源 ttft_timeout_s({}): 流式挂起会被 TTFT 看门狗"
"先行切断,对冲对流式形同虚设(非流式仍有效)",
after,
min(ttfts),
)
if len(sources) == 1:
logger.warning(
"单源 scope 配置了对冲阈值 hedge_after_s={};运行期拿不到异源候选,对冲自然静默",
after,
)
if hedge_max_extra > 1:
logger.warning(
"hedge_max_extra={} 已接受,但 v1 仅单路对冲生效(梯次追加为 H5 预留)",
hedge_max_extra,
)
return after
@dataclass(frozen=True)
class EmbeddingSettings:
"""Embedding scope 装配配置(M2 §7): 复用 GatewaySettings + embedding 专用键。
+15 -2
View File
@@ -169,15 +169,28 @@ class SourceAdmission:
# —— 选源与准入(CHS _pick_runnable 120-167)——
async def pick(
self, reasons: dict[str, str], attempt_fails: dict[str, int]
self,
reasons: dict[str, str],
attempt_fails: dict[str, int],
*,
exclude: frozenset[str] | None = None,
) -> tuple[tuple[SourceConfig, Permit, GateDecision] | None, int]:
"""挑出第一个过闸的候选;返回 (选中三元组 | None, 熔断类拒绝计数)。"""
"""挑出第一个过闸的候选;返回 (选中三元组 | None, 熔断类拒绝计数)。
`exclude` 是对冲编排的私有排除参数(issue #24): 以在途源名为排除集,
保证对冲路落在**异源**。被排除不是源的拒绝——不计 gate_rejections、
不写 reasons,否则会污染 `on_no_runnable` 的分派判据与 per_source_reasons
对账。拿不到候选时返回 None,是否放弃由调用方决定(对冲方静默等原路,
**严禁**对这个 None 调 `on_no_runnable`)。
"""
stats = {s.name: await self._quota.stats(s) for s in self._sources}
gate_rejections = 0
ordered = _demote_call_failures(
self._selector.order(self._sources, stats), attempt_fails, self._health_view
)
for cand in ordered:
if exclude and cand.name in exclude:
continue # 对冲排除在途源: 不是拒绝, 不计数不写原因(见 docstring)
if self._memo.active(cand.name):
# 冷却备忘跳过也计入拒绝数,保住 circuit_open 判据(CHS 同款)
gate_rejections += 1
+214 -6
View File
@@ -164,6 +164,17 @@ def _is_rate_limited(outcome: LLMResponse | _Failed) -> bool:
return isinstance(outcome, _Failed) and _failure_reason(outcome.exc) == "rate_limited"
_HEDGE_LOSER_ATTR = "_polygateway_hedge_loser"
"""编排在 cancel() 之前给输家任务置位的标记;_attempt 读它选遥测标签。"""
def _combine_failures(primary_f: _Failed, hedge_f: _Failed) -> _Failed:
"""任一非 429 优先(计预算);两路皆 429 才按 429 免预算退还 stall 账;同类取原路。"""
if _failure_reason(primary_f.exc) == "rate_limited" != _failure_reason(hedge_f.exc):
return hedge_f
return primary_f
class RetryMW:
"""尝试编排器;时钟/睡眠/随机全部注入,纯确定性可测(P6)。"""
@@ -183,6 +194,7 @@ class RetryMW:
cooldown_memo: SourceCooldownMemo | None = None,
pacer: AdaptivePacer | None = None,
emitter: TelemetryEmitter | None = None,
hedge_after_s: float | None = None,
now: Callable[[], float] = time.monotonic,
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
rng: Callable[[], float] = random.random,
@@ -200,6 +212,10 @@ class RetryMW:
# M2.5 §3.35: AIMD 自适应并发——429 收紧、成功回涨,超限调用排队不烧预算
self._pacer = pacer or AdaptivePacer(ceiling=64.0)
self._emitter = emitter
# 对冲触发阈值(issue #24): None = 关闭(__call__ 逐字走 1.3.6 单路路径)。
# 值域/交叉守卫在装配层(config.check_hedge_assembly),本类不重复校验;
# hedge_max_extra 不下传——v1 编排固定单路对冲(H5)
self._hedge_after_s = hedge_after_s
self._now = now
self._sleep = sleep
self._rng = rng
@@ -249,16 +265,27 @@ class RetryMW:
# 每轮一个 sink: 裸生成时间由编排裁定归属(T2 单路 = 成功那轮;
# T3 对冲 = 赢家那一路),`_attempt` 只负责把本次 transport 耗时投进来
generation_sink: list[int] = []
if self._hedge_after_s is None:
# 默认关闭: 逐字 1.3.6 单路路径(默认关闭回归门据此成立)
outcome = await self._attempt(
request, *picked, reasons, attempt_fails, generation_sink=generation_sink
request,
*picked,
reasons,
attempt_fails,
first_token_event=None,
generation_sink=generation_sink,
)
else:
# 对冲轮 = 一次"超级尝试": 编排内部自建 sink 并按赢家归属登记
outcome = await self._attempt_hedged(request, picked, reasons, attempt_fails)
rate_limited = _is_rate_limited(outcome)
if rate_limited:
# 免了重试预算就得进 stall 账,否则这段耗时无人治理(见 StallClock)
attempt.refund()
if isinstance(outcome, LLMResponse):
# 与 register_attempt 同款 None 守卫: 库内现场构造的请求跳过登记
if request.call_context is not None:
# 与 register_attempt 同款 None 守卫: 库内现场构造的请求跳过登记;
# 对冲轮由 _attempt_hedged 按赢家归属登记,此处外层 sink 恒为空
if request.call_context is not None and generation_sink:
request.call_context.record_generation(generation_sink[0], accumulate=False)
return outcome
if not rate_limited:
@@ -284,6 +311,7 @@ class RetryMW:
reasons: dict[str, str],
attempt_fails: dict[str, int],
*,
first_token_event: asyncio.Event | None,
generation_sink: list[int],
) -> LLMResponse | _Failed:
call_id = str(uuid.uuid4())
@@ -311,8 +339,9 @@ class RetryMW:
# 逐次尝试原样重传: 换源不改变调用方要的档位(源级默认由 transport
# 自己按选中的源解析,两者在 effective_effort 里汇合)
reasoning_effort=request.reasoning_effort,
# T3 对冲编排接线前恒为 None: 调用方不观测首 token(计划 §3.1)
first_token_event=None,
# 对冲编排(计划 §3.4): 原路携带首 token 事件,对冲路恒 None(v1 单路,
# 不再梯次);未启用对冲时 __call__ 传 None = 调用方不观测首 token
first_token_event=first_token_event,
)
generation_sink.append(int((self._now() - gen_started) * 1000))
if result.usage_source == "unavailable":
@@ -347,7 +376,15 @@ class RetryMW:
actual = source.effective_est_tokens()
if entry.is_probe:
await self._record_quietly(self._breaker.release_probe(entry))
await self._emit(request, source, call_id, started, error="cancelled")
# 对冲输家在 cancel() 前被编排置位标记(计划 §3.4): 据此区分"被对冲
# 淘汰"与"外部取消",零新遥测列(设计 §4.5);外部取消与对冲取消同时
# 到达的竞速可能误贴——记账方向一致(est 保留),属已批准的可接受残留
label = (
"hedge_cancelled"
if getattr(asyncio.current_task(), _HEDGE_LOSER_ATTR, False)
else "cancelled"
)
await self._emit(request, source, call_id, started, error=label)
raise
except (SourceDeadError, TransientError) as exc:
dead = isinstance(exc, SourceDeadError)
@@ -369,6 +406,177 @@ class RetryMW:
self._pacer.leave(source.name)
await settle_and_release(permit, actual)
# —— 对冲编排(issue #24 设计 §4.6;默认关闭,`hedge_after_s is None` 不进这里)——
async def _attempt_hedged(
self,
request: ChatRequest,
picked: tuple[SourceConfig, Permit, GateDecision],
reasons: dict[str, str],
attempt_fails: dict[str, int],
) -> LLMResponse | _Failed:
"""一次"超级尝试": 原路 + (触发后)异源对冲路,FIRST_COMPLETED 竞速。
计时只用事件循环相对时长(`asyncio.wait` timeout),绝不读注入 `now`
(设计 §4.1 时钟纪律);两路各持各的 permit,结算/遥测/熔断写回全部沿用
`_attempt` 既有路径,本方法只做编排与赢家裁定。
"""
# Phase 1 启动原路: 事件与 sink 每轮新建(局部状态,严禁实例属性)
source, permit, entry = picked
first_token: asyncio.Event = asyncio.Event()
sink_p: list[int] = []
primary = asyncio.create_task(
self._attempt(
request,
source,
permit,
entry,
reasons,
attempt_fails,
first_token_event=first_token,
generation_sink=sink_p,
)
)
hedge: asyncio.Task[LLMResponse | _Failed] | None = None
try:
# Phase 2 触发窗: 只认"阈值到 + 首 token 未至 + 原路在途"(loop 相对时长)
if not await self._past_hedge_window(primary, first_token):
# 原路已了结/首 token 已至: 等价于未配置对冲
outcome = await primary
if isinstance(outcome, LLMResponse):
self._record_generation(request, sink_p)
return outcome
# Phase 3 异源准入(完整 pick 路径,无旁路): 拿不到候选 = 静默等原路
# (对冲是优化不是权利;严禁对这个 None 调 on_no_runnable,见 §3.5)
hedge_picked, _ = await self._admission.pick(
reasons, attempt_fails, exclude=frozenset({source.name})
)
if hedge_picked is None:
outcome = await primary
if isinstance(outcome, LLMResponse):
self._record_generation(request, sink_p)
return outcome
# Phase 4 启动对冲路(v1 单路,H5): 对冲路恒传 first_token_event=None,
# 不再触发梯次对冲
sink_h: list[int] = []
hedge = asyncio.create_task(
self._attempt(
request,
*hedge_picked,
reasons,
attempt_fails,
first_token_event=None,
generation_sink=sink_h,
)
)
# Phase 5 赢家裁定与收口
done, pending = await asyncio.wait(
{primary, hedge}, return_when=asyncio.FIRST_COMPLETED
)
if pending and not any(self._succeeded(t, done) for t in done):
# 先了结的是失败: 等另一路的结论再裁定(它可能后发先至;
# 原路先败、对冲在途时不重试,设计 §4.6)
more, pending = await asyncio.wait(pending)
done |= more
winner = (
primary
if self._succeeded(primary, done)
else hedge
if self._succeeded(hedge, done)
else None
)
if winner is None:
primary_exc = primary.exception()
hedge_exc = hedge.exception() # 两路都取回,不留 never-retrieved 告警
# 非可重试族(RequestRejected/ResultInvalid)穿透,与单路同口径
if primary_exc is not None:
raise primary_exc
if hedge_exc is not None:
raise hedge_exc
# 两败(H6): 汇合为一次失败,重试预算只计一次;路数照登(无赢家)
outcome = _combine_failures(primary.result(), hedge.result())
self._register_hedge(request, hedge_won=False, winner_sink=None)
return outcome
# 两路同时成功的竞速 → 原路优先(保守不弃原路成果),由上面 primary
# 先判实现;竞速落选者已完成,不置标记——它没被取消,行是正常行
loser = hedge if winner is primary else primary
if not loser.done():
# 先置标记再 cancel: _attempt 的取消分支据标记选 hedge_cancelled
setattr(loser, _HEDGE_LOSER_ATTR, True)
loser.cancel()
# 收口: gather 等输家 finally 的结算/遥测跑完才返回(快照含输家,
# attempts==2);对已完成的输家顺带取回结果,不留 never-retrieved 告警
await asyncio.gather(loser, return_exceptions=True)
self._register_hedge(
request,
hedge_won=winner is hedge,
winner_sink=sink_h if winner is hedge else sink_p,
)
return winner.result()
except BaseException:
# 取消穿透与准入冒泡(GovernanceBackendError)同路: 两任务(存在且未完者)
# 同消,尽力收口后原异常上抛;不 shield,收口 await 允许被再取消(136 清理纪律)
tasks = [t for t in (primary, hedge) if t is not None and not t.done()]
for t in tasks:
t.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
raise
async def _past_hedge_window(self, primary: asyncio.Task, first_token: asyncio.Event) -> bool:
"""对冲触发窗: 阈值到 + 首 token 未至 + 原路在途,三者齐备才放行对冲。
只用事件循环相对时长(`asyncio.wait` timeout),绝不读注入 `now`——测试
伪造注入钟跳变不得触发对冲(设计 §4.1)。waiter 任务必须收口,且**不得吞
外部取消**(finally 里 await 已取消的 waiter 会接住一个 CancelledError,
须靠 cancelling() 区分它是 waiter 自己的还是外面打进来的)。
"""
waiter = asyncio.create_task(first_token.wait())
try:
await asyncio.wait(
{primary, waiter}, timeout=self._hedge_after_s, return_when=asyncio.FIRST_COMPLETED
)
finally:
waiter.cancel()
try:
await waiter
except asyncio.CancelledError:
if asyncio.current_task().cancelling(): # 外部取消,穿透
raise
return not primary.done() and not first_token.is_set()
@staticmethod
def _succeeded(task: asyncio.Task, done: set[asyncio.Task]) -> bool:
"""赢家裁定: 已完成、未被取消、无异常且结果为 LLMResponse。"""
return (
task in done
and not task.cancelled()
and task.exception() is None
and isinstance(task.result(), LLMResponse)
)
@staticmethod
def _record_generation(request: ChatRequest, sink: list[int]) -> None:
"""record_generation 的共享守卫: 上下文缺失(库内现场构造)或空 sink 均跳过。"""
if request.call_context is not None and sink:
request.call_context.record_generation(sink[0], accumulate=False)
@staticmethod
def _register_hedge(
request: ChatRequest, *, hedge_won: bool, winner_sink: list[int] | None
) -> None:
"""对冲登记(赢家裁定后一次性;同 register_attempt 的 None 守卫)。
hedges 计"实际并发发出的对冲路数"——触发但准入失败的静默不计(没走到
这里);两败轮次无赢家: 路数照登(hedge_won=False),裸生成时间无归属不记。
"""
context = request.call_context
if context is None:
return
if winner_sink:
context.record_generation(winner_sink[0], accumulate=False)
context.register_hedge(hedge_won=hedge_won)
async def _on_rejected(
self, exc: RequestRejectedError, source: SourceConfig, entry: GateDecision
) -> None:
+26
View File
@@ -144,6 +144,32 @@ class TestChatEndToEnd:
await client.chat([{"role": "user", "content": "hi"}], structured="json")
class TestHedgeParams:
"""对冲两个 keyword-only 参数的 client 入口校验(issue #24 H4;计划批次 H)。
编排行为本身见 tests/unit/test_hedge.py;这里只钉装配面: 入口即校、
值域/交叉守卫与 settings 共用同一份 `check_hedge_assembly`。
"""
def test_client_hedge_params_entry_validation(self):
# 值域: 0/负/非有限当场 ValueError(不经 settings 那道守卫)
with pytest.raises(ValueError, match=r"GatewayClient\(hedge_after_s"):
_client(hedge_after_s=0)
with pytest.raises(ValueError, match="hedge_max_extra"):
_client(hedge_max_extra=0)
with pytest.raises(ValueError, match="hedge_max_extra"):
_client(hedge_max_extra=True) # bool 不是 int 档位
# 交叉: 阈值 ≥ 源 timeout_s(10)= 对冲永不可能触发
with pytest.raises(ValueError, match="timeout_s"):
_client(hedge_after_s=99.0)
# 合法值透传到 RetryMW(单源 scope 的装配 warning 是预期噪音,不断言)
client = _client(hedge_after_s=0.05)
assert client._hedge_after_s == 0.05
assert client._terminal._hedge_after_s == 0.05
# 未启用(缺省): 行为逐字等于 1.3.6
assert _client()._terminal._hedge_after_s is None
class TestGenerationMsClient:
"""裸生成时间的 client 级口径(1.3.7 批次 C2/F)。
+162
View File
@@ -1096,3 +1096,165 @@ class TestCallDeadlineConfig:
with pytest.raises(ValueError, match=r"GatewayClient\(call_deadline_s"):
_client(call_deadline_s=0)
class TestHedgeConfig:
"""`{SCOPE}__HEDGE__AFTER_S`/`{SCOPE}__HEDGE__MAX_EXTRA` 两键与装配守卫(issue #24 H4)。
对冲默认关闭: 键未设 = None/1,行为逐字等于 1.3.6守卫四路覆盖
(env/直接构造/dataclasses.replace/client 直传),单一定义点是
`config.check_hedge_assembly`
"""
def _two_source_env(self, **overrides):
"""双源 env(同 provider 避免注册表依赖): 隔离单源 warning 的干扰。"""
return _env(
**{
"LLM__QWEN__2__BASE_URL": "https://gw-b.example/v1",
"LLM__QWEN__2__API_KEY": "sk-b",
"LLM__QWEN__2__MODEL": "qwen-plus",
"LLM__QWEN__2__TIMEOUT_S": "90",
**overrides,
}
)
def test_hedge_keys_from_env_skip_source_loader(self):
"""两键为 3 段键,天然不被 `_load_sources` 当源字段;`HEDGE` 进保留段防 4 段撞名。"""
env = self._two_source_env(**{"LLM__HEDGE__AFTER_S": "8", "LLM__HEDGE__MAX_EXTRA": "2"})
with _captured_warnings(): # max_extra>1 的 v1 单路 warning, 本用例不断言它
s = GatewaySettings.from_env("LLM", env=env)
assert s.hedge_after_s == 8.0
assert s.hedge_max_extra == 2
assert {src.name for src in s.sources} == {"qwen_1", "qwen_2"} # HEDGE 键未造源
# `HEDGE` 在保留段: `LLM__HEDGE__1__*` 四段键不得造出一个名为 hedge_1 的源
env_collision = _env(
**{
"LLM__HEDGE__1__BASE_URL": "https://gw-c.example/v1",
"LLM__HEDGE__1__API_KEY": "sk-c",
"LLM__HEDGE__1__MODEL": "m-c",
"LLM__HEDGE__1__TIMEOUT_S": "60",
}
)
s2 = GatewaySettings.from_env("LLM", env=env_collision)
assert [src.name for src in s2.sources] == ["qwen_1"]
def test_hedge_keys_unset_mean_disabled(self):
"""默认关闭: 两键未设 = None/1,且不产生任何 warning。"""
with _captured_warnings() as warnings:
s = GatewaySettings.from_env("LLM", env=_env())
assert s.hedge_after_s is None and s.hedge_max_extra == 1
assert not warnings
def test_hedge_after_s_domain_four_paths(self):
"""非法值四条装配路全部当场 ValueError(消息须定位得到是哪个键/参数)。"""
from tests.unit.test_client import _client
# 路 1: env(origin 是实际命中的键名)
with pytest.raises(ValueError, match="LLM__HEDGE__AFTER_S"):
GatewaySettings.from_env("LLM", env=_env(**{"LLM__HEDGE__AFTER_S": "0"}))
with pytest.raises(ValueError, match="LLM__HEDGE__AFTER_S"):
GatewaySettings.from_env("LLM", env=_env(**{"LLM__HEDGE__AFTER_S": "abc"}))
base = GatewaySettings.from_env("LLM", env=_env())
# 路 2: 直接构造
fields = {f.name: getattr(base, f.name) for f in dataclasses.fields(base)}
with pytest.raises(ValueError, match="hedge_after_s"):
GatewaySettings(**{**fields, "hedge_after_s": float("nan")})
# 路 3: dataclasses.replace
with pytest.raises(ValueError, match="hedge_after_s"):
dataclasses.replace(base, hedge_after_s=-1)
# 路 4: client 直传(不经 settings 那道守卫)
with pytest.raises(ValueError, match=r"GatewayClient\(hedge_after_s"):
_client(hedge_after_s=0)
def test_hedge_guard_below_min_timeout_raises(self):
"""阈值 ≥ 最小源 timeout_s = 对冲永不可能触发,装配期炸掉(ValueError)。"""
env = self._two_source_env(**{"LLM__HEDGE__AFTER_S": "90"}) # min(timeout)=90
with pytest.raises(ValueError, match="timeout_s"):
GatewaySettings.from_env("LLM", env=env)
# 边界内侧合法(89 < 90)
s = GatewaySettings.from_env(
"LLM", env=self._two_source_env(**{"LLM__HEDGE__AFTER_S": "89"})
)
assert s.hedge_after_s == 89.0
def test_hedge_guard_ttft_warns(self):
"""阈值 ≥ 最小已设 ttft_timeout_s: 流式被看门狗先切,装配期 warning 而非 ValueError。"""
env = self._two_source_env(
**{
"LLM__QWEN__1__TTFT_TIMEOUT_S": "30",
"LLM__QWEN__1__INTER_TOKEN_TIMEOUT_S": "15",
"LLM__HEDGE__AFTER_S": "35", # ≥ ttft 30, < timeout 90
}
)
with _captured_warnings() as warnings:
s = GatewaySettings.from_env("LLM", env=env)
assert s.hedge_after_s == 35.0 # warning 不是拒绝: 非流式仍有效
assert any("ttft_timeout_s" in m for m in warnings)
# 阈值低于看门狗时不告警
with _captured_warnings() as warnings2:
GatewaySettings.from_env(
"LLM",
env=self._two_source_env(
**{
"LLM__QWEN__1__TTFT_TIMEOUT_S": "30",
"LLM__QWEN__1__INTER_TOKEN_TIMEOUT_S": "15",
"LLM__HEDGE__AFTER_S": "25",
}
),
)
assert not warnings2
def test_hedge_guard_single_source_warns(self):
"""单源 scope 设阈值: 装配期 warning 放行,运行期拿不到候选自然静默。"""
with _captured_warnings() as warnings:
s = GatewaySettings.from_env("LLM", env=_env(**{"LLM__HEDGE__AFTER_S": "8"}))
assert s.hedge_after_s == 8.0
assert any("单源" in m for m in warnings)
def test_hedge_guard_deadline_conflict_raises(self):
"""阈值 ≥ call_deadline_s: 期限先于对冲触发,对冲形同虚设 → ValueError(§6)。"""
env = self._two_source_env(**{"LLM__HEDGE__AFTER_S": "35", "LLM__CALL_DEADLINE_S": "30"})
with pytest.raises(ValueError, match="call_deadline_s"):
GatewaySettings.from_env("LLM", env=env)
# 边界值(恰好相等)同样拒绝
with pytest.raises(ValueError, match="call_deadline_s"):
GatewaySettings.from_env(
"LLM",
env=self._two_source_env(
**{"LLM__HEDGE__AFTER_S": "30", "LLM__CALL_DEADLINE_S": "30"}
),
)
# 阈值 < 期限是合法组合
s = GatewaySettings.from_env(
"LLM",
env=self._two_source_env(**{"LLM__HEDGE__AFTER_S": "29", "LLM__CALL_DEADLINE_S": "30"}),
)
assert s.hedge_after_s == 29.0 and s.call_deadline_s == 30.0
def test_hedge_max_extra_v1_cap(self):
"""max_extra 值域 [1,3] 的四路校验;>1 已接受但 warning 声明 v1 仅单路生效(H5)。"""
# 域外值无条件拒绝(即使对冲未启用: 非法值没有"惰性"豁免)
with pytest.raises(ValueError, match="hedge_max_extra"):
GatewaySettings.from_env("LLM", env=_env(**{"LLM__HEDGE__MAX_EXTRA": "0"}))
with pytest.raises(ValueError, match="LLM__HEDGE__MAX_EXTRA"):
GatewaySettings.from_env("LLM", env=_env(**{"LLM__HEDGE__MAX_EXTRA": "x"}))
with pytest.raises(ValueError, match="hedge_max_extra"):
GatewaySettings.from_env("LLM", env=_env(**{"LLM__HEDGE__MAX_EXTRA": "4"}))
base = GatewaySettings.from_env("LLM", env=_env())
with pytest.raises(ValueError, match="hedge_max_extra"):
dataclasses.replace(base, hedge_max_extra=0)
# 2/3 接受 + warning: v1 运行期恒单路(对冲任务不再携带首 token 观测,
# 行为面由 test_hedge.py 的 ft_events 断言钉住),梯次追加为 H5 预留
env = self._two_source_env(**{"LLM__HEDGE__AFTER_S": "8", "LLM__HEDGE__MAX_EXTRA": "2"})
with _captured_warnings() as warnings:
s = GatewaySettings.from_env("LLM", env=env)
assert s.hedge_max_extra == 2
assert any("单路" in m for m in warnings)
def test_from_settings_propagates_hedge_to_client(self):
"""from_settings 透传: RetryMW 拿到归一化阈值;max_extra 不下传(v1 无消费者)。"""
env = self._two_source_env(**{"LLM__HEDGE__AFTER_S": "8"})
s = GatewaySettings.from_env("LLM", env=env)
client = GatewayClient.from_settings(s)
assert client._hedge_after_s == 8.0
assert client._terminal._hedge_after_s == 8.0
+456
View File
@@ -0,0 +1,456 @@
"""对冲编排测试(issue #24 设计 §4.6 / 计划批次 G)。
设施纪律(计划 §5): 两源 scope事件驱动假 transport(每源一对 entered/release
Event + 可脚本化"先置首 token 再挂起")**真实 loop **(不注入 FakeClock)
`hedge_after_s=0.05`断言容差 410×;取消窗口用 entered Event 栅栏钉死,
sleep 撞窗口;计时只断言下界与相对比较,不断言精确值
"""
import asyncio
import time
import pytest
from polygateway import CallDeadlineExceeded
from polygateway.backends.memory.breaker import InMemoryGate
from polygateway.backends.memory.limiter import InMemoryLimiter
from polygateway.errors import AllSourcesExhausted, TransientError
from polygateway.middleware.retry import RetryMW
from polygateway.sources import SourceCooldownMemo
from polygateway.types import (
BackpressurePolicy,
BreakerConfig,
ChatRequest,
GlobalLimits,
RetryPolicy,
_CallContext,
)
from tests.unit.test_retry import RecordingSelector, StaticSelector, _ok, _src
_BREAKER = BreakerConfig(fail_threshold=3, cooldown_s=60.0, probe_ttl_s=120.0)
_NO_GLOBAL = GlobalLimits(max_concurrency=0, rpm=0, tpm=0)
_HEDGE_AFTER_S = 0.05 # 真实 loop 钟阈值;断言上界取 10×(0.5s)
class HedgeTransport:
"""事件驱动假 transport: 按源剧本精确控制首 token 置位与完成时刻。
剧本动作(每源一条队列,耗尽后重复最后一项429 风暴用例需无限供应):
("hang",) 置位该源 entered,挂起直到该源 release(或被取消)
("token_then_hang",) 先置 first_token_event hang( token 已至)
("succeed", content) 立即成功
("succeed_after", delay, content) 真实 loop 钟睡 delay 后成功
("fail_after", delay, factory) delay 后抛 `factory()` 新造的异常
"""
def __init__(self, scripts: dict[str, list[tuple]]):
self._scripts = {name: list(actions) for name, actions in scripts.items()}
self.calls: list[str] = []
# 逐次记录收到的 first_token_event 身份(H5: 对冲路恒为 None,不再梯次)
self.ft_events: list[object] = []
self.entered = {name: asyncio.Event() for name in scripts}
self.release = {name: asyncio.Event() for name in scripts}
async def complete(
self, *, messages, source, stream, overlay, call_id, reasoning_effort, first_token_event
):
self.calls.append(source.name)
self.ft_events.append(first_token_event)
actions = self._scripts[source.name]
action = actions.pop(0) if len(actions) > 1 else actions[0]
kind = action[0]
if kind == "hang":
self.entered[source.name].set()
await self.release[source.name].wait()
return _ok(f"ok-{source.name}")
if kind == "token_then_hang":
if first_token_event is not None:
first_token_event.set()
self.entered[source.name].set()
await self.release[source.name].wait()
return _ok(f"ok-{source.name}")
if kind == "succeed":
return _ok(action[1])
if kind == "succeed_after":
await asyncio.sleep(action[1])
return _ok(action[2])
if kind == "fail_after":
await asyncio.sleep(action[1])
raise action[2]()
raise AssertionError(f"未知剧本动作: {action!r}")
class RecordingEmitter:
"""逐次遥测假 emitter: 记录每行的源/错误标签/逻辑调用 ID/attempt call_id。"""
def __init__(self):
self.rows: list[dict] = []
async def emit_attempt(
self,
*,
request,
source,
call_id,
latency_ms,
response,
error,
reasoning_applies,
operation,
):
self.rows.append(
{
"source": source.name,
"call_id": call_id,
"logical_call_id": request.call_context.logical_call_id
if request.call_context is not None
else None,
"error": error,
}
)
def _harness(
sources,
transport,
*,
hedge_after_s=_HEDGE_AFTER_S,
max_attempts=3,
emitter=None,
selector=None,
limiter=None,
gate=None,
stall_window_s=300.0,
):
"""真实 loop 钟装配(对冲计时纪律: 只用 loop 相对时长,不注入 FakeClock)。"""
limiter = limiter or InMemoryLimiter(
scope="llm",
sources={s.name: s for s in sources},
global_limits=_NO_GLOBAL,
lease_ttl_s=100.0,
)
gate = gate or InMemoryGate(config=_BREAKER)
mw = RetryMW(
scope="llm",
sources=sources,
# 固定配置序: 原路恒为 s1、对冲路恒为 s2,断言不依赖选源器内部状态
selector=selector if selector is not None else StaticSelector(),
limiter=limiter,
gate=gate,
transport=transport,
retry=RetryPolicy(max_attempts=max_attempts, backoff_base_s=0.01, backoff_max_s=0.05),
backpressure=BackpressurePolicy(stall_window_s=stall_window_s, poll_interval_s=0.01),
quota_full="wait",
cooldown_memo=SourceCooldownMemo(),
emitter=emitter,
hedge_after_s=hedge_after_s,
)
return mw, limiter, gate
def _req(*, stream=False, ctx=None):
return ChatRequest(
messages=[{"role": "user", "content": "hi"}], stream=stream, call_context=ctx
)
def _ctx():
return _CallContext(now=time.monotonic)
class TestHedgeTrigger:
"""触发两形态(验收矩阵 ①): 非流式纯时间阈值;流式以首 token 未至为判据。"""
async def test_non_stream_triggers_hedge_and_fast_leg_wins(self):
"""s1 挂起、s2 即时成功: 对冲截断长尾,赢家为对冲路(⑩: 计时不含触发前等待)。"""
s1, s2 = _src("s1"), _src("s2")
transport = HedgeTransport({"s1": [("hang",)], "s2": [("succeed_after", 0.02, "fast")]})
mw, _, _ = _harness([s1, s2], transport)
ctx = _ctx()
started = time.monotonic()
# wait_for 是防挂安全带(红相位无实现时 5s 判负),不是计时断言
resp = await asyncio.wait_for(mw(_req(stream=False, ctx=ctx)), timeout=5)
elapsed = time.monotonic() - started
assert resp.content == "fast" and resp.source_name == "s2"
assert elapsed < 10 * _HEDGE_AFTER_S # 挂起路被对冲截断,而非等到释放
stats = ctx.snapshot()
assert stats.hedges == 1 and stats.hedge_won is True
# 赢家裸生成时间: 下界 = s2 实际 transport 耗时(0.02s 留截断余量),
# 且严格小于总时长(不含 0.05s 触发窗等待)
assert stats.generation_ms >= 15
assert stats.generation_ms < stats.total_latency_ms
async def test_stream_triggers_only_when_first_token_absent(self):
"""两例(①): 首 token 未至 → 触发;先置首 token 再挂起 → 不触发。"""
# 例一: 流式但首 token 未至,阈值到 → 对冲触发
s1, s2 = _src("s1"), _src("s2")
transport = HedgeTransport({"s1": [("hang",)], "s2": [("succeed", "fast")]})
mw, _, _ = _harness([s1, s2], transport)
ctx = _ctx()
resp = await asyncio.wait_for(mw(_req(stream=True, ctx=ctx)), timeout=5)
assert resp.source_name == "s2"
stats = ctx.snapshot()
assert stats.hedges == 1 and stats.hedge_won is True and stats.attempts == 2
# 例二: 首 token 已至(假 transport 先 set 再挂起)→ 不触发,误杀慢生成即此处
transport2 = HedgeTransport({"s1": [("token_then_hang",)], "s2": [("succeed", "other")]})
mw2, _, _ = _harness([_src("s1"), _src("s2")], transport2)
ctx2 = _ctx()
async def release_later():
await asyncio.sleep(4 * _HEDGE_AFTER_S) # 4× 余量确认窗口已过
transport2.release["s1"].set()
releaser = asyncio.create_task(release_later())
resp2 = await asyncio.wait_for(mw2(_req(stream=True, ctx=ctx2)), timeout=5)
await releaser
assert resp2.source_name == "s1"
assert transport2.calls == ["s1"] # 对冲从未发出
stats2 = ctx2.snapshot()
assert stats2.hedges == 0 and stats2.hedge_won is False and stats2.attempts == 1
class TestHedgeRouting:
"""异源排除与静默放弃(验收矩阵 ②③)。"""
async def test_hedge_goes_to_other_source(self):
"""对冲请求落在另一源;两 attempt 行共享同一 logical_call_id(②⑤)。"""
emitter = RecordingEmitter()
transport = HedgeTransport({"s1": [("hang",)], "s2": [("succeed", "hedged")]})
mw, _, _ = _harness([_src("s1"), _src("s2")], transport, emitter=emitter)
ctx = _ctx()
resp = await asyncio.wait_for(mw(_req(ctx=ctx)), timeout=5)
assert resp.source_name == "s2"
assert transport.calls == ["s1", "s2"] # 第二请求落在异源
# H5 接缝: 原路携带首 token 观测,对冲路恒 None(v1 单路,不再梯次)
assert [e is not None for e in transport.ft_events] == [True, False]
assert len(emitter.rows) == 2
assert {r["logical_call_id"] for r in emitter.rows} == {ctx.logical_call_id}
assert emitter.rows[0]["call_id"] != emitter.rows[1]["call_id"] # 各 attempt 独立 ID
async def test_hedge_silent_when_no_candidate(self):
"""异源配额被占满 → 准入失败静默放弃: 不对冲、不抛错、原请求照等(③)。"""
s1 = _src("s1")
s2 = _src("s2", max_concurrency=1)
limiter = InMemoryLimiter(
scope="llm", sources={"s1": s1, "s2": s2}, global_limits=_NO_GLOBAL, lease_ttl_s=100.0
)
held = await limiter.try_acquire("s2", 0) # 外部预占满 s2 并发
assert held is not None
try:
transport = HedgeTransport({"s1": [("hang",)], "s2": [("succeed", "x")]})
mw, _, _ = _harness([s1, s2], transport, limiter=limiter)
ctx = _ctx()
task = asyncio.ensure_future(mw(_req(ctx=ctx)))
await transport.entered["s1"].wait()
# 4× 余量: 给对冲窗与那次注定失败的准入留足发生时间
await asyncio.sleep(4 * _HEDGE_AFTER_S)
assert transport.calls == ["s1"] # 对冲静默未发出
transport.release["s1"].set()
resp = await asyncio.wait_for(task, timeout=5)
assert resp.source_name == "s1"
stats = ctx.snapshot()
assert stats.hedges == 0 and stats.attempts == 1 and stats.hedge_won is False
finally:
await held.release()
async def test_hedge_silent_when_single_source(self):
"""单源 scope: 运行期拿不到异源候选自然静默,行为与不配阈值逐字相同(②)。"""
transport = HedgeTransport({"s1": [("hang",)]})
mw, _, _ = _harness([_src("s1")], transport)
ctx = _ctx()
task = asyncio.ensure_future(mw(_req(ctx=ctx)))
await transport.entered["s1"].wait()
await asyncio.sleep(4 * _HEDGE_AFTER_S) # 窗口已过,仍无候选
assert transport.calls == ["s1"]
transport.release["s1"].set()
resp = await asyncio.wait_for(task, timeout=5)
assert resp.content == "ok-s1"
stats = ctx.snapshot()
assert stats.hedges == 0 and stats.attempts == 1 and stats.hedge_won is False
class TestHedgeSettlementAndSignals:
"""赢输记账与熔断/健康信号(验收矩阵 ④⑤;设计 §3 关键判断: 挂起 ≠ 源死亡)。"""
async def test_winner_settles_actual_loser_keeps_est(self):
"""赢家按真实 usage 结算;输家取消落 1.3.6 S3 格: est 预扣保留(④)。"""
s1 = _src("s1", tpm=1000, est_tokens=400)
s2 = _src("s2", tpm=1000, est_tokens=400)
transport = HedgeTransport({"s1": [("hang",)], "s2": [("succeed", "win")]})
mw, limiter, _ = _harness([s1, s2], transport)
resp = await asyncio.wait_for(mw(_req()), timeout=5)
assert resp.source_name == "s2"
winner_stats = await limiter.source_stats("s2")
loser_stats = await limiter.source_stats("s1")
assert winner_stats.tpm_used == 15 # 预扣 400,实测 10+5 → settle 后只记 15
assert loser_stats.tpm_used == 400 # 输家 est 保留(可能被上游计费,保守下限)
assert winner_stats.inflight == 0 and loser_stats.inflight == 0
async def test_loser_row_labelled_hedge_cancelled(self):
"""输家 attempt 行 error=='hedge_cancelled',赢家行无 error,同行逻辑调用(④⑤)。
终态行是 client 级语义且成功调用本就不写终态行(emit_terminal_once 只在
异常/取消路径触发),MW 级可观测面即这两条 attempt
"""
emitter = RecordingEmitter()
transport = HedgeTransport({"s1": [("hang",)], "s2": [("succeed", "win")]})
mw, _, _ = _harness([_src("s1"), _src("s2")], transport, emitter=emitter)
ctx = _ctx()
await asyncio.wait_for(mw(_req(ctx=ctx)), timeout=5)
assert len(emitter.rows) == 2
loser = next(r for r in emitter.rows if r["source"] == "s1")
winner = next(r for r in emitter.rows if r["source"] == "s2")
assert loser["error"] == "hedge_cancelled"
assert winner["error"] is None
assert loser["logical_call_id"] == winner["logical_call_id"] == ctx.logical_call_id
async def test_loser_does_not_feed_breaker(self):
"""输家取消不喂熔断失败计数、不喂健康分;赢家照常 record_success(④)。"""
selector = RecordingSelector()
gate = InMemoryGate(config=_BREAKER)
transport = HedgeTransport({"s1": [("hang",)], "s2": [("succeed", "win")]})
mw, _, _ = _harness([_src("s1"), _src("s2")], transport, selector=selector, gate=gate)
await asyncio.wait_for(mw(_req()), timeout=5)
gate_s1 = gate._gates["s1"]
assert gate_s1.a0 + gate_s1.a1 == 0 # 熔断失败率窗口无样本
assert selector.outcomes == [("s2", True)] # 健康喂数只有赢家的成功
assert (await gate.try_enter("s1", "w")).allowed # 挂起源未被标记
async def test_attempts_two_and_no_task_leak(self):
"""快照 attempts==2(含输家);返回后无本调用残留任务(⑤)。"""
before = asyncio.all_tasks()
transport = HedgeTransport({"s1": [("hang",)], "s2": [("succeed", "win")]})
mw, _, _ = _harness([_src("s1"), _src("s2")], transport)
ctx = _ctx()
await asyncio.wait_for(mw(_req(ctx=ctx)), timeout=5)
assert ctx.snapshot().attempts == 2
assert asyncio.all_tasks() == before
class TestHedgeCancellation:
"""取消穿透(⑦)与期限组合(⑧): 两任务同消、不 shield、不留后台任务。"""
async def test_external_cancel_cancels_both_legs(self):
"""两路均在途时外部取消: CancelledError 上抛,两 permit 释放。
输家标记只在赢家产生后才置位外部取消下没有赢家,两行都是普通
"cancelled"(;竞速误贴属设计 §4.5 已批准残留)
"""
emitter = RecordingEmitter()
s1, s2 = _src("s1"), _src("s2")
transport = HedgeTransport({"s1": [("hang",)], "s2": [("hang",)]})
mw, limiter, _ = _harness([s1, s2], transport, emitter=emitter)
task = asyncio.ensure_future(mw(_req()))
# 双 Event 栅栏: 确认对冲已触发、两路均在途,再取消(禁 sleep 猜窗口)
await transport.entered["s1"].wait()
await transport.entered["s2"].wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert {r["source"]: r["error"] for r in emitter.rows} == {
"s1": "cancelled",
"s2": "cancelled",
}
assert (await limiter.source_stats("s1")).inflight == 0
assert (await limiter.source_stats("s2")).inflight == 0
async def test_deadline_cuts_hedged_tree(self):
"""client 级 call_deadline_s=0.2 + 两路挂起 → CallDeadlineExceeded,permit 全释放(⑧)。"""
from tests.unit.test_client import _client
s1, s2 = _src("s1"), _src("s2")
limiter = InMemoryLimiter(
scope="llm", sources={"s1": s1, "s2": s2}, global_limits=_NO_GLOBAL
)
transport = HedgeTransport({"s1": [("hang",)], "s2": [("hang",)]})
client = _client(
sources=[s1, s2],
transport=transport,
limiter=limiter,
call_deadline_s=0.2,
hedge_after_s=_HEDGE_AFTER_S,
)
async with client:
with pytest.raises(CallDeadlineExceeded):
await client.chat([{"role": "user", "content": "hi"}], stream=False)
assert transport.calls == ["s1", "s2"] # 期限截止前对冲确已触发
assert (await limiter.source_stats("s1")).inflight == 0
assert (await limiter.source_stats("s2")).inflight == 0
class TestHedgeWinnerAdjudication:
"""赢家裁定(⑩): 取快者,含原路后发先至的对称面。"""
async def test_primary_late_success_wins_back(self):
"""s1 挂 0.3s(6× 阈值)后成功、s2 对冲路在途: 原路先完成 → 原路赢。"""
emitter = RecordingEmitter()
transport = HedgeTransport({"s1": [("succeed_after", 0.3, "late")], "s2": [("hang",)]})
mw, _, _ = _harness([_src("s1"), _src("s2")], transport, emitter=emitter)
ctx = _ctx()
resp = await asyncio.wait_for(mw(_req(ctx=ctx)), timeout=5)
assert resp.content == "late" and resp.source_name == "s1"
stats = ctx.snapshot()
assert stats.hedges == 1 and stats.hedge_won is False
# 裸生成时间为原路那次 transport 时长(≈300ms,只断言下界与相对关系)
assert 250 <= stats.generation_ms <= stats.total_latency_ms
loser = next(r for r in emitter.rows if r["source"] == "s2")
assert loser["error"] == "hedge_cancelled" # 在途对冲路被裁为输家
class TestHedgeFailureCombination:
"""两败汇合(H6): 只计一次重试预算;429 分账逐字沿用 attempt 级机制。"""
async def test_both_fail_counts_budget_once(self):
"""两路 Transient: max_attempts=2 时恰进第二轮(两败只计一次),第二轮两败后才耗尽。"""
transport = HedgeTransport(
{
"s1": [("fail_after", 0.1, lambda: TransientError("p", source_name="s1"))],
"s2": [("fail_after", 0.12, lambda: TransientError("h", source_name="s2"))],
}
)
mw, _, _ = _harness([_src("s1"), _src("s2")], transport, max_attempts=2)
ctx = _ctx()
with pytest.raises(AllSourcesExhausted) as ei:
await mw(_req(ctx=ctx))
assert ei.value.reason == "retry_exhausted"
# 若两败计两次预算,第一轮即耗尽,这些调用根本不会发生
assert transport.calls == ["s1", "s2", "s1", "s2"]
# 两败轮次同样登记对冲路数(设计 §4.5: 实际并发发出即计)
assert ctx.snapshot().hedges == 2
async def test_both_429_refund_no_budget(self):
"""两路皆 429: 免预算且耗时退 stall 账——小 stall 窗下终局 stalled 而非耗尽。"""
def _429():
return TransientError("throttled", status_code=429, retry_after_s=0.01)
transport = HedgeTransport(
{"s1": [("fail_after", 0.1, _429)], "s2": [("fail_after", 0.12, _429)]}
)
mw, _, _ = _harness([_src("s1"), _src("s2")], transport, max_attempts=1, stall_window_s=0.3)
with pytest.raises(AllSourcesExhausted) as ei:
await mw(_req())
# max_attempts=1: 任一路计预算都会当场 retry_exhausted;
# 两 429 免预算 → 循环到 stall 窗口判死
assert ei.value.reason == "stalled"
async def test_mixed_429_and_failure_counts_budget(self):
"""一路 429 一路 Transient → 计一次预算、不退还 stall 账(_combine_failures)。"""
transport = HedgeTransport(
{
"s1": [
(
"fail_after",
0.1,
lambda: TransientError("rl", status_code=429, retry_after_s=0.01),
)
],
"s2": [("fail_after", 0.12, lambda: TransientError("boom", source_name="s2"))],
}
)
mw, _, _ = _harness([_src("s1"), _src("s2")], transport, max_attempts=1, stall_window_s=0.3)
with pytest.raises(AllSourcesExhausted) as ei:
await mw(_req())
assert ei.value.reason == "retry_exhausted"
assert transport.calls == ["s1", "s2"] # 恰一轮两路: 计一次预算即耗尽