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:
Xubin Ren
2026-05-13 13:57:30 +00:00
co-authored by Cursor
parent eaa8ebd5d3
commit 02b059a616
5 changed files with 325 additions and 42 deletions
+169 -23
View File
@@ -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