Files
writer-work-flow/packages/llm_gateway/ww_llm_gateway/gateway.py
Yaojia Wang 765dbdfbd4 feat: M4 文风 + M5 生成/多provider/Skill + Kimi Code 订阅接入 + 本地联调修复
M4(文风): style-auditor 双轨(提取指纹/漂移第四审)+ jobs 长任务框架(zombie reaper) + 回炉 refine + GET /style read-back。
M5(生成+扩展): worldbuilder/character-gen(入库 continuity 409 gate + partition_writes 白名单 + schema→JSONB 形变);
  网关多 provider 回退链/熔断/能力降级(Anthropic/Gemini 适配器);Skill registry + 表权限沙箱 + 规则;
  前端 角色生成器/世界观/Codex/规则页/技能库/⌘K 命令面板。
K1(Kimi Code 订阅接入): OAuth device-flow(kimi-code)+ 静态 Console key(kimi-code-key)两路径;
  coding 端点 KimiCLI 伪造头(实测 UA allow-list 门禁,缺则 403)+ JSON 模式结构化(thinking ⊥ tool_choice)。
本地联调修复: CORS 中间件;assemble 注入 premise+「写第N章」指令(修空 prompt 400);
  GET /outline·/draft read-back + 大纲/工作台/审稿页重载;写页 client/server 常量边界 + notFound 健壮化;
  字数 toLocaleString locale 水合;审稿页终稿从已存草稿 seed(修 accept 422)。
门禁: backend ruff/mypy(157)/alembic 无漂移/pytest 451 · frontend lint/tsc/vitest/build。

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-20 10:39:58 +02:00

297 lines
12 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
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,
)
async 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):
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 = await 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},
)