feat(runner): support structured fallback models
Bind fallback model chains to the active model configuration so defaults and presets do not inherit or merge fallback behavior implicitly. Require explicit fallback providers while preserving per-fallback generation overrides and context-window safety. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -7,6 +7,7 @@ from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.config.schema import ModelFallbackConfig
|
||||
from nanobot.providers.base import LLMProvider, LLMResponse
|
||||
from nanobot.providers.fallback_provider import FallbackProvider
|
||||
|
||||
@@ -24,6 +25,25 @@ def _error_response(content: str = "api error") -> LLMResponse:
|
||||
return _make_response(content, finish_reason="error", error_kind="server_error")
|
||||
|
||||
|
||||
def _fallback(
|
||||
model: str,
|
||||
provider: str = "fallback",
|
||||
*,
|
||||
max_tokens: int | None = None,
|
||||
context_window_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
reasoning_effort: str | None = None,
|
||||
) -> ModelFallbackConfig:
|
||||
return ModelFallbackConfig(
|
||||
model=model,
|
||||
provider=provider,
|
||||
max_tokens=max_tokens,
|
||||
context_window_tokens=context_window_tokens,
|
||||
temperature=temperature,
|
||||
reasoning_effort=reasoning_effort,
|
||||
)
|
||||
|
||||
|
||||
class _FakeProvider(LLMProvider):
|
||||
"""Fake provider for testing."""
|
||||
|
||||
@@ -60,17 +80,113 @@ def test_fallback_models_default_empty() -> None:
|
||||
|
||||
def test_fallback_models_accepts_list() -> None:
|
||||
from nanobot.config.schema import ModelPresetConfig
|
||||
p = ModelPresetConfig(model="test/primary", fallback_models=["test/a", "test/b"])
|
||||
assert p.fallback_models == ["test/a", "test/b"]
|
||||
p = ModelPresetConfig(
|
||||
model="test/primary",
|
||||
fallback_models=[{"provider": "test", "model": "test/a"}],
|
||||
)
|
||||
assert p.fallback_models == [_fallback("test/a", provider="test")]
|
||||
|
||||
|
||||
def test_fallback_models_from_camel_case() -> None:
|
||||
from nanobot.config.schema import ModelPresetConfig
|
||||
p = ModelPresetConfig.model_validate({
|
||||
"model": "test/primary",
|
||||
"fallbackModels": ["test/a"],
|
||||
"fallbackModels": [{"provider": "test", "model": "test/a"}],
|
||||
})
|
||||
assert p.fallback_models == ["test/a"]
|
||||
assert p.fallback_models == [_fallback("test/a", provider="test")]
|
||||
|
||||
|
||||
def test_provider_signature_tracks_fallback_models_and_provider_config() -> None:
|
||||
from nanobot.config.schema import Config
|
||||
from nanobot.providers.factory import provider_signature
|
||||
|
||||
base = {
|
||||
"modelPresets": {
|
||||
"prod": {
|
||||
"model": "openai/gpt-4.1",
|
||||
"fallbackModels": [
|
||||
{"provider": "anthropic", "model": "anthropic/claude-sonnet-4-6"}
|
||||
],
|
||||
}
|
||||
},
|
||||
"providers": {
|
||||
"openai": {"apiKey": "primary-key"},
|
||||
"anthropic": {"apiKey": "fallback-key"},
|
||||
},
|
||||
}
|
||||
changed_fallback = {
|
||||
**base,
|
||||
"modelPresets": {
|
||||
"prod": {
|
||||
"model": "openai/gpt-4.1",
|
||||
"fallbackModels": [{"provider": "deepseek", "model": "deepseek/deepseek-chat"}],
|
||||
}
|
||||
},
|
||||
"providers": {
|
||||
**base["providers"],
|
||||
"deepseek": {"apiKey": "deepseek-key"},
|
||||
},
|
||||
}
|
||||
changed_key = {
|
||||
**base,
|
||||
"providers": {
|
||||
"openai": {"apiKey": "primary-key"},
|
||||
"anthropic": {"apiKey": "new-fallback-key"},
|
||||
},
|
||||
}
|
||||
|
||||
signature = provider_signature(Config.model_validate(base), preset_name="prod")
|
||||
|
||||
assert signature != provider_signature(Config.model_validate(changed_fallback), preset_name="prod")
|
||||
assert signature != provider_signature(Config.model_validate(changed_key), preset_name="prod")
|
||||
|
||||
|
||||
def test_agent_defaults_can_define_fallback_models() -> None:
|
||||
from nanobot.config.schema import Config
|
||||
|
||||
config = Config.model_validate({
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"model": "primary-model",
|
||||
"provider": "custom",
|
||||
"fallbackModels": [{"provider": "deepseek", "model": "deepseek-v4-pro"}],
|
||||
}
|
||||
}
|
||||
})
|
||||
|
||||
assert config.resolve_preset().fallback_models == [
|
||||
_fallback("deepseek-v4-pro", provider="deepseek")
|
||||
]
|
||||
|
||||
|
||||
def test_provider_snapshot_uses_smallest_fallback_context_window() -> None:
|
||||
from nanobot.config.schema import Config
|
||||
from nanobot.providers.factory import build_provider_snapshot
|
||||
|
||||
config = Config.model_validate({
|
||||
"modelPresets": {
|
||||
"prod": {
|
||||
"model": "openai/gpt-4.1",
|
||||
"provider": "openai",
|
||||
"contextWindowTokens": 128000,
|
||||
"fallbackModels": [
|
||||
{
|
||||
"provider": "deepseek",
|
||||
"model": "deepseek/deepseek-chat",
|
||||
"contextWindowTokens": 64000,
|
||||
}
|
||||
],
|
||||
}
|
||||
},
|
||||
"providers": {
|
||||
"openai": {"apiKey": "primary-key"},
|
||||
"deepseek": {"apiKey": "fallback-key"},
|
||||
},
|
||||
})
|
||||
|
||||
snapshot = build_provider_snapshot(config, preset_name="prod")
|
||||
|
||||
assert snapshot.context_window_tokens == 64000
|
||||
|
||||
|
||||
# -- FallbackProvider tests --
|
||||
@@ -83,7 +199,7 @@ class TestNoFallbackWhenPrimarySucceeds:
|
||||
factory = MagicMock()
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_models=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -102,14 +218,14 @@ class TestFallbackOnPrimaryError:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_models=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
result = await fb.chat(messages=[{"role": "user", "content": "hi"}], model="primary-model")
|
||||
assert result.content == "fallback ok"
|
||||
assert result.finish_reason == "stop"
|
||||
factory.assert_called_once_with("fallback-a")
|
||||
factory.assert_called_once_with(_fallback("fallback-a"))
|
||||
assert primary.chat_calls[0]["model"] == "primary-model"
|
||||
assert fallback.chat_calls[0]["model"] == "fallback-a"
|
||||
|
||||
@@ -121,7 +237,7 @@ class TestNoFallbackWhenContentStreamed:
|
||||
factory = MagicMock()
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_models=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -146,14 +262,14 @@ class TestFailoverOnTransientError:
|
||||
factory = MagicMock(return_value=fallback)
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_models=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
result = await fb.chat(messages=[{"role": "user", "content": "hi"}])
|
||||
assert result.content == "fallback ok"
|
||||
assert result.finish_reason == "stop"
|
||||
factory.assert_called_once_with("fallback-a")
|
||||
factory.assert_called_once_with(_fallback("fallback-a"))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_timeout(self) -> None:
|
||||
@@ -165,14 +281,14 @@ class TestFailoverOnTransientError:
|
||||
factory = MagicMock(return_value=fallback)
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_models=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
result = await fb.chat(messages=[{"role": "user", "content": "hi"}])
|
||||
assert result.content == "fallback ok"
|
||||
assert result.finish_reason == "stop"
|
||||
factory.assert_called_once_with("fallback-a")
|
||||
factory.assert_called_once_with(_fallback("fallback-a"))
|
||||
|
||||
|
||||
class TestFallbackTriesModelsInOrder:
|
||||
@@ -185,15 +301,15 @@ class TestFallbackTriesModelsInOrder:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a", "fallback-b"],
|
||||
fallback_models=[_fallback("fallback-a"), _fallback("fallback-b")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
result = await fb.chat(messages=[{"role": "user", "content": "hi"}])
|
||||
assert result.content == "b ok"
|
||||
assert factory.call_count == 2
|
||||
factory.assert_any_call("fallback-a")
|
||||
factory.assert_any_call("fallback-b")
|
||||
factory.assert_any_call(_fallback("fallback-a"))
|
||||
factory.assert_any_call(_fallback("fallback-b"))
|
||||
|
||||
|
||||
class TestAllFallbacksFail:
|
||||
@@ -205,7 +321,7 @@ class TestAllFallbacksFail:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_models=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -223,7 +339,7 @@ class TestFactoryExceptionSkipsModel:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a", "fallback-b"],
|
||||
fallback_models=[_fallback("fallback-a"), _fallback("fallback-b")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -242,13 +358,43 @@ class TestFallbackModelParameter:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-model"],
|
||||
fallback_models=[_fallback("fallback-model")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
await fb.chat(messages=[{"role": "user", "content": "hi"}], model="primary-model")
|
||||
assert fallback.chat_calls[0]["model"] == "fallback-model"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_overrides_generation_fields_when_configured(self) -> None:
|
||||
primary = _FakeProvider("primary", _error_response())
|
||||
fallback = _FakeProvider("fallback", _make_response("ok"))
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=[
|
||||
_fallback(
|
||||
"fallback-model",
|
||||
max_tokens=1234,
|
||||
temperature=0.4,
|
||||
reasoning_effort="low",
|
||||
)
|
||||
],
|
||||
provider_factory=MagicMock(return_value=fallback),
|
||||
)
|
||||
|
||||
await fb.chat(
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
model="primary-model",
|
||||
max_tokens=8192,
|
||||
temperature=0.1,
|
||||
reasoning_effort="high",
|
||||
)
|
||||
|
||||
assert fallback.chat_calls[0]["model"] == "fallback-model"
|
||||
assert fallback.chat_calls[0]["max_tokens"] == 1234
|
||||
assert fallback.chat_calls[0]["temperature"] == 0.4
|
||||
assert fallback.chat_calls[0]["reasoning_effort"] == "low"
|
||||
|
||||
|
||||
class TestNoFallbackWhenEmptyList:
|
||||
@pytest.mark.asyncio
|
||||
@@ -277,7 +423,7 @@ class TestChatStreamFailover:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_models=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -291,7 +437,7 @@ class TestGetDefaultModel:
|
||||
primary = _FakeProvider("primary")
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["a"],
|
||||
fallback_models=[_fallback("a")],
|
||||
provider_factory=MagicMock(),
|
||||
)
|
||||
assert fb.get_default_model() == "primary/model"
|
||||
@@ -305,7 +451,7 @@ class TestCircuitBreaker:
|
||||
factory = MagicMock(return_value=fallback)
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_models=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -329,7 +475,7 @@ class TestCircuitBreaker:
|
||||
factory = MagicMock(return_value=fallback)
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_models=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -357,7 +503,7 @@ class TestGenerationForwarded:
|
||||
primary.generation = GenerationSettings(temperature=0.5, max_tokens=1024)
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["a"],
|
||||
fallback_models=[_fallback("a")],
|
||||
provider_factory=MagicMock(),
|
||||
)
|
||||
assert fb.generation.temperature == 0.5
|
||||
|
||||
Reference in New Issue
Block a user