"""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