Skip to content

Commit dd103e4

Browse files
fix: map models across provider fallbacks
1 parent c26f738 commit dd103e4

3 files changed

Lines changed: 135 additions & 14 deletions

File tree

‎README.md‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -576,6 +576,15 @@ The local CLI is intentionally a high-permission agent, similar to Codex or Open
576576
}
577577
```
578578

579+
When automatic fallback is used, `auto`, an empty model, and the primary
580+
provider's configured simple/complex model names are mapped to the equivalent
581+
simple/complex model configured for each fallback provider. An explicit model
582+
name that is not recognised as one of the primary provider's names is passed
583+
through unchanged; use that only when the fallback endpoint supports the same
584+
name. The mapping is identical for streaming and non-streaming requests, and
585+
an all-provider failure retains every provider error with the primary failure
586+
as its cause.
587+
579588
The file is auto-managed. Use `/provider` or `/api_key` in-chat to update it interactively.
580589

581590
---

‎providers.py‎

Lines changed: 57 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -443,30 +443,73 @@ def __init__(self, primary_config: ProviderConfig) -> None:
443443
def config(self) -> ProviderConfig:
444444
return self._primary.config
445445

446+
def _model_for(self, provider: LLMProvider, requested_model: str | None) -> str:
447+
"""Choose a model that has the same simple/complex meaning for *provider*.
448+
449+
The model names in a provider configuration are provider-specific. A
450+
model selected from the primary provider must therefore not be sent to
451+
a fallback merely because it is a non-empty string. Names configured
452+
for the primary (including its built-in defaults) are treated as
453+
semantic simple/complex slots and mapped to the fallback's slots.
454+
An unrecognised explicit name is forwarded as-is so users can use a
455+
model name shared by several compatible endpoints, with the fallback
456+
provider responsible for accepting or rejecting it.
457+
"""
458+
primary_config = self._primary.config
459+
target_config = provider.config
460+
requested = (requested_model or "").strip()
461+
462+
if not requested or requested.lower() == "auto":
463+
return target_config.model_simple
464+
465+
primary_simple = {
466+
primary_config.model_simple,
467+
PROVIDER_DEFAULT_MODELS.get(primary_config.provider, ("", ""))[0],
468+
}
469+
primary_complex = {
470+
primary_config.model_complex,
471+
PROVIDER_DEFAULT_MODELS.get(primary_config.provider, ("", ""))[1],
472+
}
473+
if requested in primary_complex and requested not in primary_simple:
474+
return target_config.model_complex
475+
if requested in primary_simple:
476+
return target_config.model_simple
477+
return requested
478+
479+
@staticmethod
480+
def _raise_all_failed(attempts: list[tuple[str, Exception]]) -> None:
481+
"""Raise a useful error without discarding the primary failure."""
482+
if not attempts:
483+
raise RuntimeError("All providers failed")
484+
if len(attempts) == 1:
485+
raise attempts[0][1]
486+
details = "; ".join(
487+
f"{provider}: {type(error).__name__}: {error}"
488+
for provider, error in attempts
489+
)
490+
error = RuntimeError(f"All providers failed ({details})")
491+
raise error from attempts[0][1]
492+
446493
def chat(self, messages: list[dict[str, str]], model: str | None = None) -> tuple[str, dict | None]:
447494
providers = [self._primary] + self._fallbacks
448-
last_error = None
449-
for i, prov in enumerate(providers):
495+
attempts: list[tuple[str, Exception]] = []
496+
for prov in providers:
450497
try:
451-
return prov.chat(messages, model)
498+
return prov.chat(messages, self._model_for(prov, model))
452499
except Exception as e:
453-
last_error = e
454-
if i < len(providers) - 1:
455-
continue # try next
456-
raise last_error or RuntimeError("All providers failed")
500+
attempts.append((prov.name, e))
501+
self._raise_all_failed(attempts)
457502

458503
def chat_stream(self, messages: list[dict[str, str]], model: str | None = None) -> Iterator[str]:
459504
providers = [self._primary] + self._fallbacks
460-
last_error = None
461-
for i, prov in enumerate(providers):
505+
attempts: list[tuple[str, Exception]] = []
506+
for prov in providers:
462507
try:
463-
yield from prov.chat_stream(messages, model)
508+
yield from prov.chat_stream(messages, self._model_for(prov, model))
464509
return
465510
except Exception as e:
466-
last_error = e
467-
if i < len(providers) - 1:
468-
continue
469-
raise last_error or RuntimeError("All providers failed")
511+
attempts.append((prov.name, e))
512+
self._raise_all_failed(attempts)
470513

471514
@property
472515
def name(self) -> str:

‎tests/test_providers.py‎

Lines changed: 69 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,8 @@
66
from pathlib import Path
77

88
from providers import (
9+
FallbackProvider,
10+
LLMProvider,
911
ProviderConfig,
1012
_get_encryption_key,
1113
decrypt_api_key,
@@ -14,6 +16,29 @@
1416
)
1517

1618

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+
1742
class ProviderConfigTests(unittest.TestCase):
1843
def setUp(self):
1944
self._home = os.environ.get("HOME")
@@ -48,5 +73,49 @@ def test_legacy_xor_ciphertext_remains_readable(self):
4873
self.assertEqual(decrypt_api_key(ciphertext), plaintext)
4974

5075

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+
51120
if __name__ == "__main__":
52121
unittest.main()

0 commit comments

Comments
 (0)