feat: add opt-in cross-source hedged requests for chat
This commit is contained in:
@@ -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
@@ -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 专用键。
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)。
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,456 @@
|
||||
"""对冲编排测试(issue #24 设计 §4.6 / 计划批次 G)。
|
||||
|
||||
设施纪律(计划 §5): 两源 scope、事件驱动假 transport(每源一对 entered/release
|
||||
Event + 可脚本化"先置首 token 再挂起")、**真实 loop 钟**(不注入 FakeClock)、
|
||||
`hedge_after_s=0.05`、断言容差 4–10×;取消窗口用 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"] # 恰一轮两路: 计一次预算即耗尽
|
||||
Reference in New Issue
Block a user