Files
writer-work-flow/packages/llm_gateway/ww_llm_gateway/gateway.py
Yaojia Wang 016509c5c6 fix(gateway): 熔断计入持续 4xx + 去 assert + Protocol/transient 去重
P1-2 非可重试错误(持续 401/403)也 record_failure,坏 key 可触发熔断。
P1-9 _complete_structured 的 assert 改显式 raise ValueError(-O 安全)。
P2 GatewayRun 抽到 orchestrator/_protocols.py 单点(去 4 处重复);
  _is_transient 抽到 adapters/base.py is_transient_by_name(去 3 处重复);
  Gemini Protocol 改 async def;gateway._retrying 去无用 async。
2026-06-21 19:32:49 +02:00

314 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""网关核心:路由 → 回退链 → 重试/熔断 → 调用适配器 → 记账 → 返回。
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:
log.info(
"llm_call",
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,
)
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__,
)
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,
)
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},
)