"""Provider 故障切换链测试.""" import pytest from src.ai.errors import AILLMUnavailableError from src.ai.providers import ProviderFailoverChain from src.ai.providers.circuit_breaker import CircuitBreaker from .conftest import MockProvider class TestProviderFailoverChain: """ProviderFailoverChain 测试.""" def _make_chain(self, providers: list) -> ProviderFailoverChain: return ProviderFailoverChain(providers, CircuitBreaker()) async def test_single_provider_success(self) -> None: """单 Provider 成功.""" chain = self._make_chain([MockProvider(name="p1", response_content="ok")]) resp = await chain.chat([], "model") assert resp.content == "ok" assert resp.provider == "p1" async def test_failover_to_second(self) -> None: """第一个失败自动切换第二个.""" chain = self._make_chain([ MockProvider(name="p1", fail=True), MockProvider(name="p2", response_content="from-p2"), ]) resp = await chain.chat([], "model") assert resp.content == "from-p2" assert resp.provider == "p2" async def test_all_fail_raises(self) -> None: """全部失败抛 AILLMUnavailableError.""" chain = self._make_chain([ MockProvider(name="p1", fail=True), MockProvider(name="p2", fail=True), ]) with pytest.raises(AILLMUnavailableError): await chain.chat([], "model") async def test_skip_unavailable(self) -> None: """跳过未配置的 Provider.""" chain = self._make_chain([ MockProvider(name="p1", available=False), MockProvider(name="p2", response_content="ok"), ]) resp = await chain.chat([], "model") assert resp.provider == "p2" async def test_empty_providers_raises(self) -> None: """空 Provider 列表抛 ValueError.""" with pytest.raises(ValueError): ProviderFailoverChain([], CircuitBreaker()) async def test_stream_success(self) -> None: """流式成功.""" chain = self._make_chain([ MockProvider(name="p1", stream_chunks=["a", "b", "c"]), ]) chunks = [] async for chunk in chain.stream_chat([], "model"): chunks.append(chunk.delta) assert chunks == ["a", "b", "c"] async def test_stream_failover(self) -> None: """流式 failover.""" chain = self._make_chain([ MockProvider(name="p1", fail=True), MockProvider(name="p2", stream_chunks=["x", "y"]), ]) chunks = [] async for chunk in chain.stream_chat([], "model"): chunks.append(chunk.delta) assert chunks == ["x", "y"] async def test_stream_all_fail(self) -> None: """流式全部失败.""" chain = self._make_chain([ MockProvider(name="p1", fail=True), MockProvider(name="p2", fail=True), ]) with pytest.raises(AILLMUnavailableError): async for _ in chain.stream_chat([], "model"): pass async def test_available_providers(self) -> None: """available_providers 过滤未配置.""" chain = self._make_chain([ MockProvider(name="p1", available=True), MockProvider(name="p2", available=False), ]) avail = chain.available_providers() assert len(avail) == 1 assert avail[0].name == "p1" async def test_circuit_open_skips_provider(self) -> None: """熔断的 Provider 被跳过.""" cb = CircuitBreaker(failure_threshold=1) chain = ProviderFailoverChain( [ MockProvider(name="p1", fail=True), MockProvider(name="p2", response_content="ok"), ], cb, ) # 第一次 p1 失败,切换 p2 resp = await chain.chat([], "m") assert resp.provider == "p2" # p1 熔断,第二次直接用 p2 resp2 = await chain.chat([], "m") assert resp2.provider == "p2" def test_providers_property(self) -> None: """providers 属性返回列表副本.""" p1 = MockProvider(name="p1") chain = self._make_chain([p1]) assert len(chain.providers) == 1 # 修改返回列表不影响内部 chain.providers.clear() assert len(chain.providers) == 1 def test_circuit_breaker_property(self) -> None: """circuit_breaker 属性可访问.""" cb = CircuitBreaker() chain = ProviderFailoverChain([MockProvider()], cb) assert chain.circuit_breaker is cb