Files
writer-work-flow/packages/llm_gateway/tests/fakes_resilience.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

136 lines
4.5 KiB
Python
Raw Permalink 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.

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