Terminal / tests /test_ai_client_provider_unavailability.py
Baida-A's picture
deploy: disable Qwen reasoning budget for concise responses (#4)
509c85e
Raw
History Blame Contribute Delete
3.71 kB
import asyncio
import unittest
from unittest.mock import patch
from models.ai_client import AIClient, ProviderConfig, ProviderUnavailableError
class _FailingCompletions:
def create(self, **_kwargs):
raise RuntimeError("quota exhausted")
class _FailingChat:
completions = _FailingCompletions()
class _FailingClient:
chat = _FailingChat()
class _ClientWithFailingProviders(AIClient):
def __init__(self):
self.providers = [
ProviderConfig(name="primary", api_key="x", base_url="https://example.invalid", default_model="model-a"),
ProviderConfig(name="fallback", api_key="y", base_url="https://example.invalid", default_model="model-b"),
]
self._client_cache = {}
self._rr_indices = {}
# Stato minimo richiesto dai percorsi chat/stream dopo l’introduzione
# del circuit breaker per profilo. Non chiama AIClient.__init__ e non
# carica provider o segreti dall’ambiente.
self._breaker = {}
self._breaker_threshold = 2
self._breaker_cooldown_s = 60.0
def _client_for(self, _provider):
return _FailingClient()
class ProviderUnavailableTests(unittest.IsolatedAsyncioTestCase):
async def test_chat_raises_structured_error_when_every_provider_fails(self):
client = _ClientWithFailingProviders()
with self.assertRaises(ProviderUnavailableError) as raised:
await client.chat([{"role": "user", "content": "hello"}], max_tokens=8)
self.assertCountEqual(raised.exception.providers, ("primary", "fallback"))
self.assertNotIn("api_key", str(raised.exception).lower())
async def test_stream_chat_raises_structured_error_when_every_provider_fails(self):
client = _ClientWithFailingProviders()
with self.assertRaises(ProviderUnavailableError) as raised:
async for _ in client.stream_chat([{"role": "user", "content": "hello"}], max_tokens=8):
pass
self.assertCountEqual(raised.exception.providers, ("primary", "fallback"))
self.assertNotIn("api_key", str(raised.exception).lower())
async def test_stream_chat_expands_a_role_specific_provider_pool(self):
client = _ClientWithFailingProviders()
client.providers = [
ProviderConfig(name="gemini-role", api_key="x", base_url="https://example.invalid", default_model="gemini")
]
runtime_fallback = ProviderConfig(
name="nvidia", api_key="y", base_url="https://fallback.invalid", default_model="nemotron"
)
with patch.object(client, "_load_providers", return_value=[runtime_fallback]):
with self.assertRaises(ProviderUnavailableError) as raised:
async for _ in client.stream_chat([{"role": "user", "content": "hello"}], max_tokens=8):
pass
self.assertCountEqual(raised.exception.providers, ("gemini-role", "nvidia"))
class RuntimeModelOverrideTests(unittest.TestCase):
def test_groq_runtime_model_overrides_database_model(self):
row = {
"base_url": "https://api.groq.com/openai/v1",
"default_model": "llama-3.3-70b-versatile",
}
with patch.dict("os.environ", {"GROQ_MODEL": "openai/gpt-oss-120b"}, clear=False):
self.assertEqual(
AIClient._runtime_model_override(row),
"openai/gpt-oss-120b",
)
def test_unknown_provider_keeps_database_model(self):
row = {"base_url": "https://example.invalid/v1", "default_model": "custom-model"}
self.assertEqual(AIClient._runtime_model_override(row), "custom-model")
if __name__ == "__main__":
unittest.main()