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()