Files
writer-work-flow/packages/llm_gateway/tests/test_fallback_chain.py

295 lines
11 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 回退链 + 重试 + 熔断ARCH §4.5)。
主模型 transient 失败 → 退避重试 → 仍失败切回退链下一个;回退服务时标
`served_by.fell_back=True`;记账落实际服务方;链耗尽抛 LLM_UNAVAILABLE。
熔断:某 provider 连续失败超阈值后短时熔断、直接跳过走回退。
"""
from __future__ import annotations
import uuid
import pytest
import structlog
from fakes_resilience import (
AuthError,
FakeLedger,
ScriptedAdapter,
auth_error,
chain,
chain_resolver,
transient,
)
from ww_llm_gateway.gateway import CircuitBreaker, Gateway
from ww_llm_gateway.types import LlmRequest, Scope
from ww_shared import AppError, ErrorCode
_SCOPE = Scope(user_id=uuid.UUID(int=1))
def _failing_gw() -> Gateway:
"""单 provider 链、每次 transient 失败:耗尽重试后记 `llm_provider_failed` 再抛。"""
primary = ScriptedAdapter("deepseek", failures=[transient() for _ in range(10)])
return Gateway(
{"deepseek": primary},
FakeLedger(),
chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"))),
max_retries=1,
)
async def test_primary_success_no_fallback(req: LlmRequest) -> None:
primary = ScriptedAdapter("deepseek", text="主模型")
backup = ScriptedAdapter("openai", text="备用")
ledger = FakeLedger()
gw = Gateway(
{"deepseek": primary, "openai": backup},
ledger,
chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"), ("openai", "gpt-4o"))),
)
resp = await gw.run(req)
assert resp.text == "主模型"
assert resp.served_by.provider == "deepseek"
assert resp.served_by.fell_back is False
assert backup.complete_calls == 0
assert ledger.records[0].provider == "deepseek"
async def test_falls_through_to_next_provider_on_transient(req: LlmRequest) -> None:
# 主模型每次 complete 都 transient 失败(足够耗尽重试) → 切回退。
primary = ScriptedAdapter("deepseek", failures=[transient() for _ in range(10)])
backup = ScriptedAdapter("openai", text="备用结果")
ledger = FakeLedger()
gw = Gateway(
{"deepseek": primary, "openai": backup},
ledger,
chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"), ("openai", "gpt-4o"))),
max_retries=2,
)
resp = await gw.run(req)
assert resp.text == "备用结果"
assert resp.served_by.provider == "openai"
assert resp.served_by.fell_back is True
# 记账记实际服务方 openai不是失败的 deepseek。
assert len(ledger.records) == 1
assert ledger.records[0].provider == "openai"
async def test_retries_then_succeeds_on_same_provider(req: LlmRequest) -> None:
# 前两次 transient第三次成功 → 不应切回退max_retries=2 即最多 3 次尝试)。
primary = ScriptedAdapter("deepseek", text="重试后成功", failures=[transient(), transient()])
backup = ScriptedAdapter("openai", text="备用")
ledger = FakeLedger()
gw = Gateway(
{"deepseek": primary, "openai": backup},
ledger,
chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"), ("openai", "gpt-4o"))),
max_retries=2,
)
resp = await gw.run(req)
assert resp.text == "重试后成功"
assert resp.served_by.provider == "deepseek"
assert resp.served_by.fell_back is False
assert backup.complete_calls == 0
async def test_rate_limited_triggers_fallback(req: LlmRequest) -> None:
rate_limited = AppError(ErrorCode.RATE_LIMITED, "429")
primary = ScriptedAdapter("deepseek", failures=[rate_limited for _ in range(10)])
backup = ScriptedAdapter("openai", text="降级备用")
ledger = FakeLedger()
gw = Gateway(
{"deepseek": primary, "openai": backup},
ledger,
chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"), ("openai", "gpt-4o"))),
max_retries=1,
)
resp = await gw.run(req)
assert resp.served_by.provider == "openai"
assert resp.served_by.fell_back is True
async def test_chain_exhausted_raises_llm_unavailable(req: LlmRequest) -> None:
p1 = ScriptedAdapter("deepseek", failures=[transient() for _ in range(10)])
p2 = ScriptedAdapter("openai", failures=[transient() for _ in range(10)])
gw = Gateway(
{"deepseek": p1, "openai": p2},
FakeLedger(),
chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"), ("openai", "gpt-4o"))),
max_retries=1,
)
with pytest.raises(AppError) as exc:
await gw.run(req)
assert exc.value.code == ErrorCode.LLM_UNAVAILABLE
async def test_missing_adapter_in_chain_skipped(req: LlmRequest) -> None:
# 链上首个 provider 没注册适配器 → 跳过、走下一个(不硬失败)。
backup = ScriptedAdapter("openai", text="可用")
gw = Gateway(
{"openai": backup},
FakeLedger(),
chain_resolver=chain_resolver(chain(("missing", "m"), ("openai", "gpt-4o"))),
)
resp = await gw.run(req)
assert resp.served_by.provider == "openai"
assert resp.served_by.fell_back is True
async def test_stream_retries_before_first_token_then_succeeds(req: LlmRequest) -> None:
# CR-H6首 token 前的瞬时失败应 per-provider 退避重试(对称 run 的 _complete_with_retry
# 而非直接烧掉一个回退名额切备用。max_retries=2 → 最多 3 次尝试;前两次 transient、
# 第三次成功 → 主 provider 服务,备用零调用。
primary = ScriptedAdapter(
"deepseek", text="流式重试后成功", failures=[transient(), transient()]
)
backup = ScriptedAdapter("openai", text="备用")
ledger = FakeLedger()
gw = Gateway(
{"deepseek": primary, "openai": backup},
ledger,
chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"), ("openai", "gpt-4o"))),
max_retries=2,
)
collected = [d.text async for d in gw.stream(req)]
assert "".join(collected) == "流式重试后成功"
assert primary.stream_calls == 3 # 2 次失败 + 1 次成功(重试在首 token 前)
assert backup.stream_calls == 0 # 未切备用
assert len(ledger.records) == 1
assert ledger.records[0].provider == "deepseek"
# ---- 熔断器 ----
def test_circuit_breaker_trips_after_threshold() -> None:
cb = CircuitBreaker(threshold=3, reset_seconds=60.0)
assert cb.is_open("deepseek") is False
cb.record_failure("deepseek")
cb.record_failure("deepseek")
assert cb.is_open("deepseek") is False # 未到阈值
cb.record_failure("deepseek")
assert cb.is_open("deepseek") is True # 第 3 次 → 熔断
def test_circuit_breaker_success_resets() -> None:
cb = CircuitBreaker(threshold=2, reset_seconds=60.0)
cb.record_failure("deepseek")
cb.record_success("deepseek")
cb.record_failure("deepseek")
assert cb.is_open("deepseek") is False # 成功清零计数
def test_circuit_breaker_reopens_after_cooldown() -> None:
now = [1000.0]
cb = CircuitBreaker(threshold=1, reset_seconds=30.0, clock=lambda: now[0])
cb.record_failure("deepseek")
assert cb.is_open("deepseek") is True
now[0] += 31.0 # 冷却窗口过 → 半开(放行试探)
assert cb.is_open("deepseek") is False
async def test_persistent_auth_error_counts_toward_breaker(req: LlmRequest) -> None:
"""P1-2持续性 401坏 key虽不可重试但应计入熔断——连续 N 次后熔断打开。
每次 `run` 命中 401 立即上抛(不重试、不回退),但 raise 前 `record_failure`
达到阈值后熔断打开,后续请求直接跳过该 provider → 链耗尽抛 LLM_UNAVAILABLE。
"""
threshold = 3
primary = ScriptedAdapter("deepseek", failures=[auth_error() for _ in range(threshold + 2)])
cb = CircuitBreaker(threshold=threshold, reset_seconds=60.0)
gw = Gateway(
{"deepseek": primary},
FakeLedger(),
chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"))),
breaker=cb,
)
# 前 threshold 次401 原样上抛(不可重试),但每次计一次熔断失败。
for _ in range(threshold):
assert cb.is_open("deepseek") is False
with pytest.raises(AuthError):
await gw.run(req)
# 第 threshold 次失败后熔断打开。
assert cb.is_open("deepseek") is True
# 后续请求provider 被熔断跳过 → 不再调 complete链耗尽 → LLM_UNAVAILABLE。
calls_before = primary.complete_calls
with pytest.raises(AppError) as exc:
await gw.run(req)
assert exc.value.code == ErrorCode.LLM_UNAVAILABLE
assert primary.complete_calls == calls_before # 熔断后未再触达坏 provider
async def test_open_circuit_skips_provider(req: LlmRequest) -> None:
# 熔断已打开的主 provider 被直接跳过,连 complete 都不调,直接走回退。
primary = ScriptedAdapter("deepseek", text="不该被调")
backup = ScriptedAdapter("openai", text="回退服务")
cb = CircuitBreaker(threshold=1, reset_seconds=60.0)
cb.record_failure("deepseek") # 预先熔断
gw = Gateway(
{"deepseek": primary, "openai": backup},
FakeLedger(),
chain_resolver=chain_resolver(chain(("deepseek", "deepseek-chat"), ("openai", "gpt-4o"))),
breaker=cb,
)
resp = await gw.run(req)
assert resp.served_by.provider == "openai"
assert resp.served_by.fell_back is True
assert primary.complete_calls == 0
# CR-H4 收尾:`llm_provider_failed`(失败/回退日志)也须条件透传 request_id——
# 设了则带上、未设则不发键(否则 None 覆盖 merge_contextvars 供的 id回退 sync/SSE 追踪)。
async def test_run_failure_log_carries_request_id_conditionally() -> None:
with structlog.testing.capture_logs() as logs:
with pytest.raises(AppError):
await _failing_gw().run(
LlmRequest(tier="writer", input="x", scope=_SCOPE, request_id="rid-run")
)
failed = next(e for e in logs if e["event"] == "llm_provider_failed")
assert failed["request_id"] == "rid-run"
with structlog.testing.capture_logs() as logs:
with pytest.raises(AppError):
await _failing_gw().run(LlmRequest(tier="writer", input="x", scope=_SCOPE))
failed = next(e for e in logs if e["event"] == "llm_provider_failed")
assert "request_id" not in failed
async def test_stream_failure_log_carries_request_id_conditionally() -> None:
with structlog.testing.capture_logs() as logs:
with pytest.raises(AppError):
async for _ in _failing_gw().stream(
LlmRequest(tier="writer", input="x", scope=_SCOPE, request_id="rid-stream")
):
pass
failed = next(e for e in logs if e["event"] == "llm_provider_failed")
assert failed["request_id"] == "rid-stream"
with structlog.testing.capture_logs() as logs:
with pytest.raises(AppError):
async for _ in _failing_gw().stream(LlmRequest(tier="writer", input="x", scope=_SCOPE)):
pass
failed = next(e for e in logs if e["event"] == "llm_provider_failed")
assert "request_id" not in failed