Terminal / tests /test_provider_profile_pool.py
Baida-A's picture
deploy: disable Qwen reasoning budget for concise responses (#4)
509c85e
Raw
History Blame Contribute Delete
9.44 kB
import json
import os
import types
import time
import unittest
from unittest.mock import AsyncMock, patch
from models.ai_client import AIClient, ProviderConfig, _PROVIDER_DEFS
class ProviderProfilePoolTests(unittest.TestCase):
def _profiles(self):
return [
ProviderConfig(name="openrouter", api_key="test-key-a", base_url="https://openrouter.ai/api/v1", profile="a", purpose="coding"),
ProviderConfig(name="openrouter", api_key="test-key-b", base_url="https://openrouter.ai/api/v1", profile="b", purpose="coding"),
ProviderConfig(name="openrouter", api_key="test-key-c", base_url="https://openrouter.ai/api/v1", profile="c", purpose="coding"),
]
def test_profile_json_loads_alongside_legacy_key(self):
raw = json.dumps([
{"profile": "primary", "api_key": "profile-key-1"},
{"profile": "backup", "api_key": "profile-key-2", "model": "openai/gpt-oss-20b:free"},
])
with patch.dict(os.environ, {"OPENROUTER_PROFILES_JSON": raw, "OPENROUTER_API_KEY": "legacy-key"}, clear=True):
client = AIClient()
profiles = [p for p in client.providers if p.name == "openrouter"]
self.assertEqual([p.profile for p in profiles], ["primary", "backup"])
self.assertEqual([p.api_key for p in profiles], ["profile-key-1", "profile-key-2"])
def test_profile_json_is_supported_for_every_provider(self):
env = {
f"{definition['name'].upper()}_PROFILES_JSON": json.dumps([
{"profile": "primary", "api_key": f"{definition['name']}-key"},
{"profile": "backup", "api_key": f"{definition['name']}-backup"},
])
for definition in _PROVIDER_DEFS
}
with patch.dict(os.environ, env, clear=True):
client = AIClient()
for definition in _PROVIDER_DEFS:
profiles = [p for p in client.providers if p.name == definition["name"]]
self.assertEqual([p.profile for p in profiles], ["primary", "backup"])
def test_environment_pool_overrides_same_provider_database_row(self):
raw = json.dumps([
{"profile": "primary", "api_key": "profile-key-1"},
{"profile": "backup", "api_key": "profile-key-2"},
])
database_row = ProviderConfig(
name="openrouter", api_key="database-key",
base_url="https://openrouter.ai/api/v1", profile="db-1",
)
with patch.dict(os.environ, {"OPENROUTER_PROFILES_JSON": raw}, clear=True), \
patch.object(AIClient, "_try_load_from_supabase", return_value=[database_row]):
client = AIClient()
profiles = [p for p in client.providers if p.name == "openrouter"]
self.assertEqual([p.profile for p in profiles], ["primary", "backup"])
self.assertNotIn("database-key", [p.api_key for p in profiles])
def test_inter_provider_pool_excludes_exhausted_provider(self):
client = AIClient()
openrouter = ProviderConfig(name="openrouter", api_key="or", base_url="https://openrouter.ai/api/v1", profile="a", purpose="coding", tier=1)
groq = ProviderConfig(name="groq", api_key="groq", base_url="https://api.groq.com/openai/v1", profile="a", purpose="reasoning", tier=0)
client.providers = [openrouter, groq]
client._breaker[openrouter.identity] = {"failures": 2, "open_until": 10**12}
selected = client._inter_provider_fallback_pool("coding", {"openrouter"})
self.assertEqual([provider.name for provider in selected], ["groq"])
def test_profiles_have_distinct_client_cache_entries(self):
client = AIClient()
first, second = self._profiles()[:2]
first_client = client._client_for(first)
second_client = client._client_for(second)
self.assertIsNot(first_client, second_client)
self.assertEqual(len(client._client_cache), 2)
def test_round_robin_rotates_profiles_and_skips_open_circuit(self):
client = AIClient()
profiles = self._profiles()
self.assertEqual(client._execution_pool(profiles, "coding")[0].profile, "a")
self.assertEqual(client._execution_pool(profiles, "coding")[0].profile, "b")
client._record_failure(profiles[1], RuntimeError("HTTP 429 rate limit"))
self.assertFalse(client._is_available(profiles[1]))
selected = client._execution_pool(profiles, "coding")
self.assertNotEqual(selected[0].profile, "b")
def test_rate_limit_reset_opens_profile_on_first_error(self):
client = AIClient()
profile = self._profiles()[0]
client._record_failure(
profile,
RuntimeError("429 free-models-per-day X-RateLimit-Reset: 4102444800000"),
)
self.assertFalse(client._is_available(profile))
self.assertGreater(
client._breaker[profile.identity]["open_until"],
time.monotonic() + 900,
)
def test_all_openrouter_profiles_are_removed_from_execution_pool(self):
client = AIClient()
profiles = self._profiles()
for profile in profiles:
client._record_failure(profile, RuntimeError("HTTP 429 free-models-per-day"))
self.assertEqual(client._execution_pool(profiles, "coding"), [])
class InterProviderFallbackChatTests(unittest.IsolatedAsyncioTestCase):
async def test_chat_falls_back_when_primary_pool_returns_errors(self):
client = AIClient()
openrouter = ProviderConfig(name="openrouter", api_key="or", base_url="https://openrouter.ai/api/v1", profile="a", purpose="coding", tier=1)
groq = ProviderConfig(name="groq", api_key="groq", base_url="https://api.groq.com/openai/v1", profile="a", purpose="reasoning", tier=0)
client.providers = [openrouter, groq]
async def fake_fetch(provider, messages, temperature, max_tokens):
if provider.name == "openrouter":
return provider, "ERROR: HTTP 429 rate limit", 0.0
return provider, "fallback answer from healthy provider", 0.2
with patch("models.ai_client.get_cached_response", new=AsyncMock(return_value=None)), \
patch("models.ai_client.set_cached_response", new=AsyncMock()), \
patch.object(client, "_fetch_one", side_effect=fake_fetch):
answer = await client.chat([{"role": "user", "content": "write code"}])
self.assertEqual(answer, "fallback answer from healthy provider")
class StreamingFallbackTests(unittest.IsolatedAsyncioTestCase):
@staticmethod
def _chunk(text):
return types.SimpleNamespace(
choices=[types.SimpleNamespace(
delta=types.SimpleNamespace(content=text),
)],
)
async def test_stream_retries_next_provider_before_first_chunk(self):
client = AIClient()
first = ProviderConfig(name="openrouter", api_key="or", base_url="https://openrouter.ai/api/v1", profile="a", purpose="coding", tier=1)
second = ProviderConfig(name="groq", api_key="groq", base_url="https://api.groq.com/openai/v1", profile="a", purpose="reasoning", tier=0)
client.providers = [first, second]
class FakeCompletions:
def __init__(self, provider):
self.provider = provider
def create(self, **_kwargs):
if self.provider == "openrouter":
raise RuntimeError("HTTP 429 rate limit")
return iter([StreamingFallbackTests._chunk("healthy "), StreamingFallbackTests._chunk("stream")])
def fake_client(provider):
return types.SimpleNamespace(chat=types.SimpleNamespace(completions=FakeCompletions(provider.name)))
with patch.object(client, "_client_for", side_effect=fake_client):
output = [part async for part in client.stream_chat([{"role": "user", "content": "hello"}])]
self.assertEqual(output, ["healthy ", "stream"])
async def test_stream_does_not_retry_after_partial_output(self):
client = AIClient()
first = ProviderConfig(name="openrouter", api_key="or", base_url="https://openrouter.ai/api/v1", profile="a", purpose="coding", tier=1)
second = ProviderConfig(name="groq", api_key="groq", base_url="https://api.groq.com/openai/v1", profile="a", purpose="reasoning", tier=2)
client.providers = [first, second]
calls = []
class FakeCompletions:
def __init__(self, provider):
self.provider = provider
def create(self, **_kwargs):
calls.append(self.provider)
if self.provider == "openrouter":
def broken_stream():
yield StreamingFallbackTests._chunk("partial")
raise RuntimeError("stream disconnected")
return broken_stream()
return iter([StreamingFallbackTests._chunk("should not run")])
def fake_client(provider):
return types.SimpleNamespace(chat=types.SimpleNamespace(completions=FakeCompletions(provider.name)))
with patch.object(client, "_client_for", side_effect=fake_client):
with self.assertRaises(RuntimeError):
_ = [part async for part in client.stream_chat([{"role": "user", "content": "hello"}])]
self.assertEqual(calls, ["openrouter"])
if __name__ == "__main__":
unittest.main()