"""RetryMW: 重试循环拥有"尝试"的全部编排(D13 自研;蓝本 CHS governance.py:107-268)。 每次尝试 = 选源(跳过冷却源)→ 该源熔断门 → 限流 permit → transport; 换源、退避、逐次遥测、熔断写回、permit 结算全部在循环内。作为洋葱的 终端(compose 的 terminal)被调用: `await retry_mw(request)`。 例外说明(P7): 本模块 import httpx 仅用于失败原因归类(CHS 同款判定), httpx 是库的核心依赖而非实现层内部件,不违反"middleware 只依赖端口"。 """ from __future__ import annotations import asyncio import contextlib import random import time import uuid from dataclasses import dataclass from typing import TYPE_CHECKING import httpx from loguru import logger from polygateway.errors import ( AllSourcesExhausted, CircuitOpenError, GovernanceBackendError, PolyGatewayError, RequestRejectedError, ResultInvalidError, SourceDeadError, SourceNotConfiguredError, TransientError, ) from polygateway.middleware.breaker import BreakerGate from polygateway.middleware.ratelimit import QuotaGate from polygateway.ports import OutcomeAwareSelector from polygateway.sources import AdaptivePacer, SourceCooldownMemo from polygateway.streaming import StreamLivenessTimeout from polygateway.types import LLMResponse if TYPE_CHECKING: from collections.abc import AsyncIterator, Awaitable, Callable from polygateway.ports import ( GateDecision, Permit, ProviderGate, RateLimiter, SourceSelector, Transport, ) from polygateway.types import ( BackpressurePolicy, ChatRequest, RetryPolicy, SourceConfig, TransportResult, ) def backoff_delay( policy: RetryPolicy, fails: int, exc: BaseException | None, rng: Callable[[], float], ) -> float: """指数退避+jitter,与 Retry-After 提示取大(ARCH §7.2;VT jitter 系数)。 模块级纯函数: RetryMW 与 EmbeddingClient(M2 §7)共用同一公式。 """ base = min(policy.backoff_base_s * (2 ** (fails - 1)), policy.backoff_max_s) delay = base * (0.5 + rng()) retry_after = getattr(exc, "retry_after_s", None) or 0.0 return max(delay, retry_after) class _Attempt: """一次尝试的计时句柄;`refund()` 把它退还给 stall 账(见 `StallClock`)。""" __slots__ = ("productive",) def __init__(self) -> None: self.productive = True def refund(self) -> None: """该次尝试不消耗重试预算(429),故其耗时归 stall 治理而非重试治理。""" self.productive = False class StallClock: """调用级 stall 计时器: 只累计非生产性等待(issue #8 设计 §3.1)。 **划分依据是"谁消耗重试预算"**,不是"是否发出了请求"。消耗 `max_attempts` 的时间已被重试预算治理,从 stall 账扣除;不消耗它的时间无人治理,归 stall。 两者重叠计费正是 issue #8 的根因: stall 预算(默认 300s)小于重试预算 (3 × timeout_s),必然先耗尽,于是重试预算在超时场景下永远用不上。 "生产性"的边界即 `_attempt` 的边界,含该次尝试的记账与遥测收尾——它们是 "尝试已有结论"之后的动作,不是在等待重试机会;把它们计入 stall 会让遥测 抖动参与判死。 **例外: 429 尝试须 `refund()`**。429 免重试预算(饱和期等待而非死亡),若其 耗时又算生产性,就掉进两个预算的缝隙——排队型网关持满 timeout 才回 429 时, 每轮只有退避那一两秒进 stall 账,调用可挂满 `stall_window/backoff_base` 轮 (实测 timeout=300/base=2 时达 25 小时)。退还后缝隙闭合。 每次调用创建一个实例。严禁提升为实例属性: `_entered_at` 会固定在进程启动 时刻,使 `stalled_s()` 随进程运行时长单调增长,最终所有调用被误判 stalled。 模块级共享单元, EmbeddingClient 与 OcrClient 复用(同 `backoff_delay`)。 """ __slots__ = ("_now", "_entered_at", "_productive_s") def __init__(self, now: Callable[[], float]) -> None: self._now = now self._entered_at = now() self._productive_s = 0.0 def stalled_s(self) -> float: """非生产性等待累计秒数 = 调用总耗时 - 消耗重试预算的时间。""" return self._now() - self._entered_at - self._productive_s @contextlib.asynccontextmanager async def attempting(self) -> AsyncIterator[_Attempt]: """包裹一次真实尝试,其耗时默认记为生产性(除非被 `refund()`)。""" handle = _Attempt() started = self._now() try: yield handle finally: # 只做算术与取值, 不吞任何异常——CancelledError 逐字穿透(库铁律) if handle.productive: self._productive_s += self._now() - started def _demote_call_failures( ordered: list[SourceConfig], attempt_fails: dict[str, int], health: Callable[[str], float] | None, ) -> list[SourceConfig]: """调用内降权(设计 §3.3/§3.36): 失败 ≥2 次且存在可信替代才让位。 可信替代 = 某未失败候选 health ≥ 0.5 × 失败源 health——异构池里健康源 偶发失败不该被推向已知坏源(第三轮教训: 期望成功率 83% vs 10%)。 无健康视图(round_robin 等)保持无条件降权(冷启动保护)。 """ demoted = [s for s in ordered if attempt_fails.get(s.name, 0) >= 2] if not demoted or len(demoted) == len(ordered): return ordered if health is None: return _move_to_tail(ordered, demoted) return _health_gated_reorder(ordered, demoted, attempt_fails, health) def _move_to_tail(ordered: list[SourceConfig], demoted: list[SourceConfig]) -> list[SourceConfig]: """无健康视图: 无条件移尾(冷启动保护原语义)。""" names = {d.name for d in demoted} return [s for s in ordered if s.name not in names] + demoted def _health_gated_reorder( ordered: list[SourceConfig], demoted: list[SourceConfig], attempt_fails: dict[str, int], health: Callable[[str], float], ) -> list[SourceConfig]: """健康门槛降权: 无可信替代则原地重试;有则插到可信替代之后。""" demoted = _credible_demotions(ordered, demoted, attempt_fails, health) if not demoted: return ordered names = {d.name for d in demoted} rest = [s for s in ordered if s.name not in names] return _insert_after_credible(rest, demoted, health) def _insert_after_credible( rest: list[SourceConfig], demoted: list[SourceConfig], health: Callable[[str], float], ) -> list[SourceConfig]: """插入位置(第四轮教训): 被降权源排在可信替代之后、不可信源之前—— 可信替代被限流闸/熔断跳过时,下一候选是失败源本身而非垃圾源。""" bar = 0.5 * max(health(d.name) for d in demoted) credible = [s for s in rest if health(s.name) >= bar] junk = [s for s in rest if health(s.name) < bar] return credible + demoted + junk def _credible_demotions( ordered: list[SourceConfig], demoted: list[SourceConfig], attempt_fails: dict[str, int], health: Callable[[str], float], ) -> list[SourceConfig]: """健康门槛过滤: 仅当存在"健康分 ≥ 失败源一半"的未失败候选,让位才有意义。""" alts = [o for o in ordered if attempt_fails.get(o.name, 0) < 2] return [s for s in demoted if any(health(o.name) >= 0.5 * health(s.name) for o in alts)] def _failure_reason(exc: PolyGatewayError) -> str: """失败原因归类(CHS governance.py:169 同款)。""" if isinstance(exc, SourceDeadError): return "source_dead" if exc.status_code == 429: return "rate_limited" if isinstance(exc.__cause__, httpx.TimeoutException | StreamLivenessTimeout): return "timeout" return "network_error" @dataclass(frozen=True) class _Failed: """一次可重试失败的结果(SourceDead 立即换源,Transient 退避)。""" exc: PolyGatewayError immediate: bool def _is_rate_limited(outcome: LLMResponse | _Failed) -> bool: """429 = 服务端调度指令(gRPC pushback 语义,迭代 5): 按 Retry-After 退避但 **不消耗重试预算**——饱和窗口里等待而非死亡;其余失败照常计数。 因其免重试预算,该次尝试的耗时必须归 stall 治理(`StallClock` 的 refund)。 """ return isinstance(outcome, _Failed) and _failure_reason(outcome.exc) == "rate_limited" class RetryMW: """尝试编排器;时钟/睡眠/随机全部注入,纯确定性可测(P6)。""" def __init__( self, *, scope: str, sources: list[SourceConfig], selector: SourceSelector, limiter: RateLimiter, gate: ProviderGate, transport: Transport, retry: RetryPolicy, backpressure: BackpressurePolicy, quota_full: str = "wait", cooldown_memo: SourceCooldownMemo | None = None, pacer: AdaptivePacer | None = None, emitter: object | None = None, now: Callable[[], float] = time.monotonic, sleep: Callable[[float], Awaitable[None]] = asyncio.sleep, rng: Callable[[], float] = random.random, ) -> None: if quota_full not in ("wait", "fail_fast"): raise ValueError(f"quota_full 必须是 wait|fail_fast: {quota_full!r}") self._scope = scope self._sources = list(sources) self._selector = selector self._quota = QuotaGate(limiter, scope=self._scope) self._breaker = BreakerGate(gate, scope=self._scope) self._transport = transport self._retry = retry self._bp = backpressure self._quota_full = quota_full self._memo = cooldown_memo or SourceCooldownMemo(now=now) # M2.5: 选源器可选健康喂数端口,构造期 isinstance 判定一次(设计 §3.2) self._outcome_sink = selector if isinstance(selector, OutcomeAwareSelector) else None self._health_view = self._outcome_sink.health if self._outcome_sink else None # M2.5 §3.35: AIMD 自适应并发——429 收紧、成功回涨,超限调用排队不烧预算 self._pacer = pacer or AdaptivePacer(ceiling=64.0) self._emitter = emitter self._now = now self._sleep = sleep self._rng = rng async def __call__(self, request: ChatRequest) -> LLMResponse: """执行治理调用;scope 级失败按 §6.1 携结构化字段上抛。""" if not self._sources: raise AllSourcesExhausted(scope=self._scope, reason="no_sources", retry_after_s=0.0) fails = 0 reasons: dict[str, str] = {} # 调用内失败计数(设计 §3.3): 局部状态,调用结束即弃;严禁实例属性(并发共享) attempt_fails: dict[str, int] = {} # 调用级累计计时,循环内不重置(CHS governance.py:207);issue #8 起只计 # 非生产性等待——真实尝试由重试预算治理,不再重复烧 stall 预算 clock = StallClock(self._now) while True: # 调用级时间上限(迭代 5): 429 免预算后的兜底,防饱和期无限循环 if await self._stalled(clock): raise AllSourcesExhausted( scope=self._scope, reason="stalled", retry_after_s=self._retry.backoff_base_s, per_source_reasons=reasons, ) picked, gate_rejections = await self._pick_runnable(reasons, attempt_fails) if picked is None: await self._on_no_runnable(gate_rejections, reasons, clock) continue async with clock.attempting() as attempt: outcome = await self._attempt(request, *picked, reasons, attempt_fails) rate_limited = _is_rate_limited(outcome) if rate_limited: # 免了重试预算就得进 stall 账,否则这段耗时无人治理(见 StallClock) attempt.refund() if isinstance(outcome, LLMResponse): return outcome if not rate_limited: fails += 1 if fails >= self._retry.max_attempts: raise AllSourcesExhausted( scope=self._scope, reason="retry_exhausted", retry_after_s=self._retry.backoff_base_s, per_source_reasons=reasons, ) from outcome.exc if not outcome.immediate: await self._sleep(self._backoff_delay(max(fails, 1), outcome.exc)) # —— 选源与准入(CHS _pick_runnable 120-167)—— async def _pick_runnable( self, reasons: dict[str, str], attempt_fails: dict[str, int] ) -> tuple[tuple[SourceConfig, Permit, GateDecision] | None, int]: 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 self._memo.active(cand.name): # 冷却备忘跳过也计入拒绝数,保住 circuit_open 判据(CHS 同款) gate_rejections += 1 reasons[cand.name] = "cooldown" continue if not self._pacer.admit(cand.name): # AIMD 超限: 不计 gate_rejections → 走 quota-wait 排队,不误判熔断 reasons.setdefault(cand.name, "adaptive_paced") continue permit = await self._quota.try_acquire(cand) if permit is None: reasons.setdefault(cand.name, "rate_limited") continue entry = None try: entry = await self._breaker.try_enter(cand, uuid.uuid4().hex) finally: # try_enter 未归还 entry(异常/取消)→ 释放已占 permit,不吞任何异常 if entry is None: await self._settle_and_release(permit, 0) if entry.allowed: self._pacer.enter(cand.name) return (cand, permit, entry), gate_rejections gate_rejections += 1 reasons[cand.name] = "circuit_open" # 开路源本地记冷却,避免每轮白烧 RPM 探测(CHS governance.py:107) self._memo.set_until(cand.name, self._now() + entry.retry_after_s) await self._settle_and_release(permit, 0) return None, gate_rejections # —— 背压与 stall 判死(CHS governance.py:270-285)—— async def _stalled(self, clock: StallClock) -> bool: """双条件 stall 判死(CHS governance.py:270-281): 本地累计等待与全局 无进展**同时**超窗才判死——本地 monotonic 与后端时钟刻意不混用。 本地一侧只计非生产性等待(issue #8,见 `StallClock`)。短路顺序有意为之: 本地未超窗就不问后端,省一次 Redis 往返。 """ stall = self._bp.stall_window_s return clock.stalled_s() > stall and await self._quota.progress_age_s() > stall async def _on_no_runnable( self, gate_rejections: int, reasons: dict[str, str], clock: StallClock ) -> None: if gate_rejections == len(self._sources): names = tuple(s.name for s in self._sources) raise CircuitOpenError( scope=self._scope, retry_after_s=await self._breaker.retry_after_s(names), per_source_reasons=reasons, ) if self._quota_full == "fail_fast": raise AllSourcesExhausted( scope=self._scope, reason="quota_exhausted", retry_after_s=self._bp.poll_interval_s, per_source_reasons=reasons, ) if await self._stalled(clock): names = tuple(s.name for s in self._sources) raise AllSourcesExhausted( scope=self._scope, reason="stalled", retry_after_s=await self._breaker.retry_after_s(names), per_source_reasons=reasons, ) # jitter ∈ [0.5p, 1.0p] 防惊群(CHS governance.py:283-285) await self._sleep(self._bp.poll_interval_s * (0.5 + 0.5 * self._rng())) # —— 单次尝试(CHS run 200-268)—— async def _attempt( self, request: ChatRequest, source: SourceConfig, permit: Permit, entry: GateDecision, reasons: dict[str, str], attempt_fails: dict[str, int], ) -> LLMResponse | _Failed: call_id = str(uuid.uuid4()) started = self._now() actual = 0 try: result = await self._transport.complete( messages=request.messages, source=source, stream=request.stream, overlay=request.overlay, call_id=call_id, ) if result.usage_source == "unavailable": # 用量不可得时按入场预扣量结算(delta==0),否则押金会被整笔退回, # 对"从不返回 usage 帧"的源等于 TPM 闸失效(设计 §3.2 #9) actual = source.effective_est_tokens() else: actual = result.prompt_tokens + result.completion_tokens await self._record_quietly(self._breaker.record_success(entry)) await self._record_quietly(self._quota.mark_progress()) self._feed_outcome(source.name, ok=True) self._pacer.on_success(source.name) response = self._build_response(source, result, call_id, started) await self._emit(request, source, call_id, started, response=response) return response except RequestRejectedError as exc: await self._on_rejected(exc, source, entry) await self._emit(request, source, call_id, started, error=exc) raise except ResultInvalidError as exc: # 坏结果 ≠ 坏服务: 熔断记成功但不计窗口样本,亦不喂健康分(M2.5 §3.1) await self._record_quietly(self._breaker.record_success(entry, count_attempt=False)) await self._emit(request, source, call_id, started, error=exc) raise except asyncio.CancelledError: if entry.is_probe: await self._record_quietly(self._breaker.release_probe(entry)) await self._emit(request, source, call_id, started, error="cancelled") raise except (SourceDeadError, TransientError) as exc: dead = isinstance(exc, SourceDeadError) reason = _failure_reason(exc) reasons[source.name] = reason attempt_fails[source.name] = attempt_fails.get(source.name, 0) + 1 self._feed_outcome(source.name, ok=False) if reason == "rate_limited": self._pacer.on_backpressure(source.name) await self._record_quietly(self._breaker.record_failure(entry, reason, dead)) if not dead: # 保守: 失败请求可能已被网关计费(CHS 同款);与入场预扣同源取值 actual = source.effective_est_tokens() await self._emit(request, source, call_id, started, error=exc) return _Failed(exc, immediate=dead) finally: self._pacer.leave(source.name) await self._settle_and_release(permit, actual) async def _on_rejected( self, exc: RequestRejectedError, source: SourceConfig, entry: GateDecision ) -> None: provider_responded = exc.source_name == source.name and exc.status_code is not None if provider_responded: # 网关健康地拒了坏请求: 记成功但不计窗口样本(M2.5 §3.1) await self._record_quietly(self._breaker.record_success(entry, count_attempt=False)) elif entry.is_probe: await self._record_quietly(self._breaker.release_probe(entry)) async def _record_quietly(self, write_back: Awaitable[object]) -> None: """记账侧写回(record_*/mark_progress/release_probe)降级执行(设计 §10)。 调用已真实完成: 后端失败若冒泡会丢弃真实成功响应或掩盖原始尝试 异常,故 warning 降级(ARCH §7.3 勘误,CHS 全 fail-closed 的有意反转); 取消照常穿透。 """ try: await write_back except asyncio.CancelledError: raise except (GovernanceBackendError, SourceNotConfiguredError) as exc: logger.warning("治理记账写回降级(不冒泡): {}", exc) def _feed_outcome(self, source_name: str, ok: bool) -> None: """健康喂数降级执行: 选源器异常不得打断真实成功/失败的主路径(设计 §4)。""" if self._outcome_sink is None: return try: self._outcome_sink.record_outcome(source_name, ok) except Exception as exc: logger.warning("选源健康喂数失败(降级不冒泡): {}", exc) # —— 辅助 —— def _backoff_delay(self, fails: int, exc: PolyGatewayError) -> float: return backoff_delay(self._retry, fails, exc, self._rng) def _build_response( self, source: SourceConfig, result: TransportResult, call_id: str, started: float ) -> LLMResponse: return LLMResponse( content=result.content, thinking=result.thinking, model=source.model, provider=source.provider, prompt_tokens=result.prompt_tokens, completion_tokens=result.completion_tokens, latency_ms=int((self._now() - started) * 1000), ttft_ms=result.ttft_ms, max_inter_token_ms=result.max_inter_token_ms, cache_hit=False, call_id=call_id, source_name=source.name, cost=None, usage_source=result.usage_source, cached_prompt_tokens=result.cached_prompt_tokens, model_reported=result.model_reported, reasoning_tokens=result.reasoning_tokens, ) async def _settle_and_release(self, permit: Permit, actual: int) -> None: """finally 专用: settle 后必 release;失败降级 warning,绝不掩盖主异常/取消。""" try: try: await permit.settle(actual) finally: await permit.release() except asyncio.CancelledError: raise except Exception as exc: logger.warning("permit 结算/释放失败(不掩盖主异常): {}", exc) async def _emit( self, request: ChatRequest, source: SourceConfig, call_id: str, started: float, *, response: LLMResponse | None = None, error: object | None = None, ) -> None: """逐次遥测(经注入的单一 Emitter);遥测失败不得影响调用(铁律)。""" if self._emitter is None: return try: await self._emitter.emit_attempt( request=request, source=source, call_id=call_id, latency_ms=int((self._now() - started) * 1000), response=response, error=None if error is None else str(error), ) except asyncio.CancelledError: raise except Exception as exc: logger.warning("逐次遥测记录失败(降级不冒泡): {}", exc)