"""网关核心:路由 → 回退链 → 重试/熔断 → 调用适配器 → 记账 → 返回。 ARCH §4.3(路由)/§4.4(能力协商降级)/§4.5(回退/重试/熔断)/§4.8(记账)。 T5.4 多 provider 韧性: - **回退链**:`chain_resolver(tier)` 返 [(provider, model), ...];主失败→退避重试→ 仍失败切下一个;首个成功者服务并标 `served_by.fell_back`(非链首即 True)。 - **熔断**:`CircuitBreaker` 按 provider 连续失败计数,超阈值短时熔断、直接跳过。 - **能力降级**:结构化输出请求优先选声明支持的 provider;链上无支持者则降级用 首个可用(适配器自走 instructor JSON 提示),标 `served_by.degraded`。 - **重试归网关**(不变量):瞬时错误(TransientProviderError / RATE_LIMITED)经 tenacity 指数退避重试 R 次;节点只感知干净的最终失败,绝不自循环。 日志脱敏:只记长度,绝不记原文/api key(不变量、§9.3)。 """ from __future__ import annotations import time from collections.abc import AsyncIterator, Callable import structlog from tenacity import ( AsyncRetrying, retry_if_exception, stop_after_attempt, wait_exponential, ) from ww_shared import AppError, ErrorCode from .adapters.base import ProviderAdapter, ProviderResult, ProviderUsage from .errors import TransientProviderError from .ledger import LedgerSink from .pricing import cost_minor from .routing import Route, resolve_chain from .types import Delta, LlmRequest, LlmResponse, ServedBy, Tier, Usage log = structlog.get_logger(__name__) # 退避:min 0.05s、max 2s(确定性测试通过 max_retries 控制次数,退避不影响断言)。 _RETRY_MIN_SECONDS = 0.05 _RETRY_MAX_SECONDS = 2.0 _DEFAULT_MAX_RETRIES = 2 _DEFAULT_BREAKER_THRESHOLD = 5 _DEFAULT_BREAKER_RESET_SECONDS = 30.0 def _is_retryable(exc: BaseException) -> bool: """瞬时、可退避重试 + 可回退的错误:TransientProviderError 或 RATE_LIMITED。""" if isinstance(exc, TransientProviderError): return True return isinstance(exc, AppError) and exc.code == ErrorCode.RATE_LIMITED # 持续性鉴权错误状态码:坏 key / 被禁用 → 每次重打都失败,应计入熔断(P1-2)。 _PERSISTENT_AUTH_STATUSES = frozenset({401, 403}) def _is_persistent_auth_error(exc: BaseException) -> bool: """持续性鉴权失败(401/403):错误 key / 账号被禁,重打无意义 → 计入熔断。 按 `status_code` 属性识别(不硬依赖任何厂商 SDK 异常类型)。 """ status = getattr(exc, "status_code", None) return isinstance(status, int) and status in _PERSISTENT_AUTH_STATUSES def _input_len(req: LlmRequest) -> int: if isinstance(req.input, str): return len(req.input) return sum(len(b.text) for b in req.input) class CircuitBreaker: """每 provider 连续失败计数 → 超阈值短时熔断;成功清零;冷却后半开放行。""" def __init__( self, *, threshold: int = _DEFAULT_BREAKER_THRESHOLD, reset_seconds: float = _DEFAULT_BREAKER_RESET_SECONDS, clock: Callable[[], float] = time.monotonic, ) -> None: self._threshold = threshold self._reset_seconds = reset_seconds self._clock = clock self._failures: dict[str, int] = {} self._opened_at: dict[str, float] = {} def is_open(self, provider: str) -> bool: opened = self._opened_at.get(provider) if opened is None: return False if self._clock() - opened >= self._reset_seconds: # 冷却窗口过 → 半开:清状态、放行一次试探。 self._opened_at.pop(provider, None) self._failures.pop(provider, None) return False return True def record_failure(self, provider: str) -> None: count = self._failures.get(provider, 0) + 1 self._failures[provider] = count if count >= self._threshold: self._opened_at[provider] = self._clock() def record_success(self, provider: str) -> None: self._failures.pop(provider, None) self._opened_at.pop(provider, None) class Gateway: def __init__( self, adapters: dict[str, ProviderAdapter], ledger: LedgerSink, *, chain_resolver: Callable[[Tier], list[Route]] | None = None, resolver: Callable[[Tier], Route] | None = None, max_retries: int = _DEFAULT_MAX_RETRIES, breaker: CircuitBreaker | None = None, ) -> None: self._adapters = adapters self._ledger = ledger # `resolver`(单路由,C1/M1 兼容)会被包成单元素链;优先用 `chain_resolver`(回退链)。 if chain_resolver is not None: self._resolve_chain = chain_resolver elif resolver is not None: single = resolver def _as_chain(tier: Tier) -> list[Route]: return [single(tier)] self._resolve_chain = _as_chain else: self._resolve_chain = resolve_chain self._max_retries = max_retries self._breaker = breaker or CircuitBreaker() # ---- 路由 / 链选择 ---- def _ordered_chain(self, req: LlmRequest) -> list[Route]: """据能力协商重排回退链:结构化输出请求优先把声明支持的 provider 提前。 仅在请求带 `output_schema` 时重排——把支持结构化输出的(且已注册适配器、未熔断) provider 提到前面,降低降级概率(§4.4)。保持相对顺序稳定。 """ chain = self._resolve_chain(req.tier) if req.output_schema is None: return chain capable: list[Route] = [] rest: list[Route] = [] for route in chain: adapter = self._adapters.get(route.provider) if adapter is not None and adapter.capabilities().structured_output: capable.append(route) else: rest.append(route) return capable + rest def _usage(self, route: Route, pu: ProviderUsage) -> Usage: cost, currency = cost_minor(route.provider, route.model, pu.input_tokens, pu.output_tokens) return Usage( provider=route.provider, model=route.model, input_tokens=pu.input_tokens, output_tokens=pu.output_tokens, cache_read_tokens=pu.cache_read_tokens, cost_minor=cost, currency=currency, ) def _served_by( self, route: Route, req: LlmRequest, adapter: ProviderAdapter, *, fell_back: bool ) -> ServedBy: degraded = bool( req.output_schema is not None and not adapter.capabilities().structured_output ) return ServedBy( provider=route.provider, model=route.model, fell_back=fell_back, degraded=degraded ) def _log_call( self, req: LlmRequest, usage: Usage, served_by: ServedBy, *, stream: bool ) -> None: fields: dict[str, object] = { "provider": usage.provider, "model": usage.model, "tier": req.tier, "input_tokens": usage.input_tokens, "output_tokens": usage.output_tokens, "cache_read_tokens": usage.cache_read_tokens, "cost_minor": usage.cost_minor, "currency": usage.currency, "stream": stream, "fell_back": served_by.fell_back, "degraded": served_by.degraded, "input_chars": _input_len(req), "project_id": str(req.scope.project_id) if req.scope.project_id else None, } # 仅在调用方显式设了 request_id 时才发——否则 None 会覆盖 merge_contextvars # 供的 id,反倒回退 sync/SSE 路径的追踪(§9.3)。 if req.request_id is not None: fields["request_id"] = req.request_id log.info("llm_call", **fields) def _retrying(self) -> AsyncRetrying: return AsyncRetrying( stop=stop_after_attempt(self._max_retries + 1), wait=wait_exponential(min=_RETRY_MIN_SECONDS, max=_RETRY_MAX_SECONDS), retry=retry_if_exception(_is_retryable), reraise=True, ) # ---- run(非流式)---- async def run(self, req: LlmRequest) -> LlmResponse: chain = self._ordered_chain(req) last_error: Exception | None = None for idx, route in enumerate(chain): adapter = self._adapters.get(route.provider) if adapter is None: log.warning("llm_route_no_adapter", provider=route.provider, tier=req.tier) continue if self._breaker.is_open(route.provider): log.warning("llm_circuit_open_skip", provider=route.provider, tier=req.tier) continue try: result = await self._complete_with_retry(adapter, req, route.model) except Exception as exc: # noqa: BLE001 — 链内逐 provider 兜底,最终统一上抛 if not _is_retryable(exc): # 持续性鉴权失败(401/403)虽不可重试,但坏 key 应触发熔断(P1-2), # 否则每次都白打同一坏 provider。其它不可重试错误直接上抛。 if _is_persistent_auth_error(exc): self._breaker.record_failure(route.provider) raise self._breaker.record_failure(route.provider) last_error = exc log.warning( "llm_provider_failed", provider=route.provider, tier=req.tier, error=type(exc).__name__, **({"request_id": req.request_id} if req.request_id is not None else {}), ) continue self._breaker.record_success(route.provider) fell_back = idx > 0 served_by = self._served_by(route, req, adapter, fell_back=fell_back) usage = self._usage(route, result.usage) await self._ledger.record(req.scope, usage) self._log_call(req, usage, served_by, stream=False) return LlmResponse( text=result.text, parsed=result.parsed, usage=usage, served_by=served_by ) raise AppError( ErrorCode.LLM_UNAVAILABLE, "所有 provider 回退链耗尽,请稍后重试或更换档位", {"tier": req.tier, "last_error": type(last_error).__name__ if last_error else None}, ) async def _complete_with_retry( self, adapter: ProviderAdapter, req: LlmRequest, model: str ) -> ProviderResult: retrying = self._retrying() async for attempt in retrying: with attempt: return await adapter.complete(req, model) raise AssertionError("unreachable: reraise=True") # ---- stream(流式)---- async def stream(self, req: LlmRequest) -> AsyncIterator[Delta]: chain = self._ordered_chain(req) last_error: Exception | None = None for idx, route in enumerate(chain): adapter = self._adapters.get(route.provider) if adapter is None: continue if self._breaker.is_open(route.provider): continue # 流式回退:仅在「尚未吐出任何 token」前可切。首块产出后失败属中途失败, # 不静默重连(§4.5:已存部分留 draft,节点报错停在 write 前 checkpoint)。 try: stream_iter = adapter.stream(req, route.model) final = ProviderUsage(input_tokens=0, output_tokens=0) started = False async for chunk in stream_iter: if chunk.text: started = True yield Delta(text=chunk.text) if chunk.usage is not None: final = chunk.usage except Exception as exc: # noqa: BLE001 if started or not _is_retryable(exc): raise self._breaker.record_failure(route.provider) last_error = exc log.warning( "llm_provider_failed", provider=route.provider, tier=req.tier, error=type(exc).__name__, stream=True, **({"request_id": req.request_id} if req.request_id is not None else {}), ) continue self._breaker.record_success(route.provider) fell_back = idx > 0 served_by = self._served_by(route, req, adapter, fell_back=fell_back) usage = self._usage(route, final) await self._ledger.record(req.scope, usage) self._log_call(req, usage, served_by, stream=True) return raise AppError( ErrorCode.LLM_UNAVAILABLE, "所有 provider 回退链耗尽,请稍后重试或更换档位", {"tier": req.tier, "last_error": type(last_error).__name__ if last_error else None}, )