Spaces:
Paused
Paused
| 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): | |
| 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() | |