320 lines
13 KiB
Python
320 lines
13 KiB
Python
"""网关核心:路由 → 回退链 → 重试/熔断 → 调用适配器 → 记账 → 返回。
|
||
|
||
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},
|
||
)
|