feat: auto committed
This commit is contained in:
129
services/ai/tests/test_failover.py
Normal file
129
services/ai/tests/test_failover.py
Normal file
@@ -0,0 +1,129 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user