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。
136 lines
4.5 KiB
Python
136 lines
4.5 KiB
Python
"""T5.4 韧性测试替身(多 provider 适配器 + 失败模拟 + 记账嗅探)——绝不联网。
|
||
|
||
放独立模块(非 conftest),测试走绝对导入 `from fakes_resilience import ...`。
|
||
本目录无 __init__.py(见 fakes.py 注释)。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from collections.abc import AsyncIterator, Callable
|
||
|
||
from ww_llm_gateway.adapters.base import (
|
||
Capabilities,
|
||
ProviderResult,
|
||
ProviderUsage,
|
||
StreamChunk,
|
||
)
|
||
from ww_llm_gateway.errors import TransientProviderError
|
||
from ww_llm_gateway.routing import Route
|
||
from ww_llm_gateway.types import LlmRequest, Scope, Tier, Usage
|
||
|
||
|
||
class ScriptedAdapter:
|
||
"""可编排成功/失败序列的假适配器。
|
||
|
||
`failures` 为开头要抛的异常列表(每次 complete/stream 消费一个);耗尽后正常返回。
|
||
`capabilities_` 控制能力矩阵(测降级)。记录调用次数以断言回退/重试行为。
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
provider: str,
|
||
*,
|
||
text: str = "ok",
|
||
failures: list[Exception] | None = None,
|
||
capabilities_: Capabilities | None = None,
|
||
input_tokens: int = 100,
|
||
output_tokens: int = 50,
|
||
structured_text: str | None = None,
|
||
) -> None:
|
||
self.provider = provider
|
||
self.text = text
|
||
self._failures = list(failures or [])
|
||
self._caps = capabilities_ or Capabilities(structured_output=True, prefix_cache=True)
|
||
self.input_tokens = input_tokens
|
||
self.output_tokens = output_tokens
|
||
self.structured_text = structured_text
|
||
self.complete_calls = 0
|
||
self.stream_calls = 0
|
||
|
||
def capabilities(self) -> Capabilities:
|
||
return self._caps
|
||
|
||
def _maybe_fail(self) -> None:
|
||
if self._failures:
|
||
raise self._failures.pop(0)
|
||
|
||
async def complete(self, req: LlmRequest, model: str) -> ProviderResult:
|
||
self.complete_calls += 1
|
||
self._maybe_fail()
|
||
if req.output_schema is not None:
|
||
# 模拟原生结构化输出:仅当能力声明支持时才会被网关派到这里。
|
||
parsed = req.output_schema.model_validate({} if not _has_fields(req) else _stub(req))
|
||
return ProviderResult(
|
||
text=parsed.model_dump_json(),
|
||
usage=ProviderUsage(
|
||
input_tokens=self.input_tokens, output_tokens=self.output_tokens
|
||
),
|
||
parsed=parsed,
|
||
)
|
||
return ProviderResult(
|
||
text=self.structured_text or self.text,
|
||
usage=ProviderUsage(input_tokens=self.input_tokens, output_tokens=self.output_tokens),
|
||
)
|
||
|
||
async def stream(self, req: LlmRequest, model: str) -> AsyncIterator[StreamChunk]:
|
||
self.stream_calls += 1
|
||
self._maybe_fail()
|
||
for ch in self.text:
|
||
yield StreamChunk(text=ch)
|
||
yield StreamChunk(
|
||
usage=ProviderUsage(input_tokens=self.input_tokens, output_tokens=self.output_tokens)
|
||
)
|
||
|
||
|
||
def _has_fields(req: LlmRequest) -> bool:
|
||
return bool(req.output_schema and req.output_schema.model_fields)
|
||
|
||
|
||
def _stub(req: LlmRequest) -> dict[str, object]:
|
||
# 用各字段默认/最小填充——测试 schema 仅需可构造。
|
||
assert req.output_schema is not None
|
||
out: dict[str, object] = {}
|
||
for name, field in req.output_schema.model_fields.items():
|
||
if field.is_required():
|
||
out[name] = "x"
|
||
return out
|
||
|
||
|
||
class FakeLedger:
|
||
def __init__(self) -> None:
|
||
self.records: list[Usage] = []
|
||
|
||
async def record(self, scope: Scope, usage: Usage) -> None:
|
||
self.records.append(usage)
|
||
|
||
|
||
def transient(msg: str = "boom") -> TransientProviderError:
|
||
return TransientProviderError(msg)
|
||
|
||
|
||
class AuthError(Exception):
|
||
"""模拟持续性鉴权失败(坏 key / 账号被禁):带 `status_code`,**非**瞬时不可重试。
|
||
|
||
适配器对 401/403 不翻译为 `TransientProviderError`,原样上抛;网关须对其计入熔断
|
||
(P1-2)而非每次白打同一坏 provider。
|
||
"""
|
||
|
||
def __init__(self, msg: str = "unauthorized", *, status_code: int = 401) -> None:
|
||
super().__init__(msg)
|
||
self.status_code = status_code
|
||
|
||
|
||
def auth_error(status_code: int = 401) -> AuthError:
|
||
return AuthError(status_code=status_code)
|
||
|
||
|
||
def chain(*routes: tuple[str, str]) -> list[Route]:
|
||
return [Route(provider=p, model=m) for p, m in routes]
|
||
|
||
|
||
def chain_resolver(routes: list[Route]) -> Callable[[Tier], list[Route]]:
|
||
def _resolve(tier: Tier) -> list[Route]:
|
||
return routes
|
||
|
||
return _resolve
|