|
6 | 6 | from pathlib import Path |
7 | 7 |
|
8 | 8 | from providers import ( |
| 9 | + FallbackProvider, |
| 10 | + LLMProvider, |
9 | 11 | ProviderConfig, |
10 | 12 | _get_encryption_key, |
11 | 13 | decrypt_api_key, |
|
14 | 16 | ) |
15 | 17 |
|
16 | 18 |
|
| 19 | +class _ProbeProvider(LLMProvider): |
| 20 | + def __init__(self, provider: str, simple: str, complex_model: str, |
| 21 | + response: str = "ok", error: Exception | None = None) -> None: |
| 22 | + super().__init__(ProviderConfig(provider=provider, model_simple=simple, |
| 23 | + model_complex=complex_model)) |
| 24 | + self.models: list[str | None] = [] |
| 25 | + self.stream_models: list[str | None] = [] |
| 26 | + self.response = response |
| 27 | + self.error = error |
| 28 | + |
| 29 | + def chat(self, messages, model=None): |
| 30 | + self.models.append(model) |
| 31 | + if self.error: |
| 32 | + raise self.error |
| 33 | + return self.response, None |
| 34 | + |
| 35 | + def chat_stream(self, messages, model=None): |
| 36 | + self.stream_models.append(model) |
| 37 | + if self.error: |
| 38 | + raise self.error |
| 39 | + yield self.response |
| 40 | + |
| 41 | + |
17 | 42 | class ProviderConfigTests(unittest.TestCase): |
18 | 43 | def setUp(self): |
19 | 44 | self._home = os.environ.get("HOME") |
@@ -48,5 +73,49 @@ def test_legacy_xor_ciphertext_remains_readable(self): |
48 | 73 | self.assertEqual(decrypt_api_key(ciphertext), plaintext) |
49 | 74 |
|
50 | 75 |
|
| 76 | +class FallbackModelTests(unittest.TestCase): |
| 77 | + def _fallback(self, primary, fallback): |
| 78 | + wrapper = FallbackProvider.__new__(FallbackProvider) |
| 79 | + wrapper._primary = primary |
| 80 | + wrapper._fallbacks = [fallback] |
| 81 | + return wrapper |
| 82 | + |
| 83 | + def test_sync_maps_deepseek_complex_model_to_openai_default(self): |
| 84 | + primary = _ProbeProvider("deepseek", "deepseek-chat", "deepseek-reasoner", error=RuntimeError("primary down")) |
| 85 | + fallback = _ProbeProvider("openai", "gpt-4o", "gpt-4o", response="fallback") |
| 86 | + result = self._fallback(primary, fallback).chat([], "deepseek-reasoner") |
| 87 | + self.assertEqual(result[0], "fallback") |
| 88 | + self.assertEqual(primary.models, ["deepseek-reasoner"]) |
| 89 | + self.assertEqual(fallback.models, ["gpt-4o"]) |
| 90 | + |
| 91 | + def test_stream_maps_anthropic_simple_model_to_deepseek_default(self): |
| 92 | + primary = _ProbeProvider("anthropic", "claude-sonnet-4-20250514", "claude-sonnet-4-20250514", |
| 93 | + error=RuntimeError("primary down")) |
| 94 | + fallback = _ProbeProvider("deepseek", "deepseek-chat", "deepseek-reasoner", response="stream fallback") |
| 95 | + result = list(self._fallback(primary, fallback).chat_stream([], "claude-sonnet-4-20250514")) |
| 96 | + self.assertEqual(result, ["stream fallback"]) |
| 97 | + self.assertEqual(primary.stream_models, ["claude-sonnet-4-20250514"]) |
| 98 | + self.assertEqual(fallback.stream_models, ["deepseek-chat"]) |
| 99 | + |
| 100 | + def test_auto_and_unknown_explicit_models_have_documented_behavior(self): |
| 101 | + primary = _ProbeProvider("deepseek", "deepseek-chat", "deepseek-reasoner", error=RuntimeError("primary down")) |
| 102 | + fallback = _ProbeProvider("openai", "gpt-4o", "gpt-4o", response="fallback") |
| 103 | + wrapper = self._fallback(primary, fallback) |
| 104 | + wrapper.chat([], "auto") |
| 105 | + wrapper.chat([], "shared-model") |
| 106 | + self.assertEqual(fallback.models, ["gpt-4o", "shared-model"]) |
| 107 | + |
| 108 | + def test_all_provider_failures_keep_primary_error_as_cause_and_include_both(self): |
| 109 | + primary_error = RuntimeError("primary outage") |
| 110 | + fallback_error = RuntimeError("fallback outage") |
| 111 | + primary = _ProbeProvider("deepseek", "deepseek-chat", "deepseek-reasoner", error=primary_error) |
| 112 | + fallback = _ProbeProvider("openai", "gpt-4o", "gpt-4o", error=fallback_error) |
| 113 | + with self.assertRaises(RuntimeError) as context: |
| 114 | + self._fallback(primary, fallback).chat([], "auto") |
| 115 | + self.assertIs(context.exception.__cause__, primary_error) |
| 116 | + self.assertIn("primary outage", str(context.exception)) |
| 117 | + self.assertIn("fallback outage", str(context.exception)) |
| 118 | + |
| 119 | + |
51 | 120 | if __name__ == "__main__": |
52 | 121 | unittest.main() |
0 commit comments