feat(runner): support fallback candidates
Resolve fallbackModels as preset references or explicit inline provider configs so failover uses complete model settings without exposing fallback logic to the agent loop. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -3,10 +3,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.config.schema import ModelPresetConfig
|
||||
from nanobot.providers.base import LLMProvider, LLMResponse
|
||||
from nanobot.providers.fallback_provider import FallbackProvider
|
||||
|
||||
@@ -16,14 +17,45 @@ def _make_response(
|
||||
finish_reason: str = "stop",
|
||||
*,
|
||||
error_kind: str | None = None,
|
||||
error_status_code: int | None = None,
|
||||
error_type: str | None = None,
|
||||
error_code: str | None = None,
|
||||
error_should_retry: bool | None = None,
|
||||
) -> LLMResponse:
|
||||
return LLMResponse(content=content, finish_reason=finish_reason, error_kind=error_kind)
|
||||
return LLMResponse(
|
||||
content=content,
|
||||
finish_reason=finish_reason,
|
||||
error_kind=error_kind,
|
||||
error_status_code=error_status_code,
|
||||
error_type=error_type,
|
||||
error_code=error_code,
|
||||
error_should_retry=error_should_retry,
|
||||
)
|
||||
|
||||
|
||||
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 = "custom",
|
||||
*,
|
||||
max_tokens: int = 8192,
|
||||
context_window_tokens: int = 65_536,
|
||||
temperature: float = 0.1,
|
||||
reasoning_effort: str | None = None,
|
||||
) -> ModelPresetConfig:
|
||||
return ModelPresetConfig(
|
||||
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."""
|
||||
|
||||
@@ -53,24 +85,163 @@ class _FakeProvider(LLMProvider):
|
||||
|
||||
|
||||
def test_fallback_models_default_empty() -> None:
|
||||
from nanobot.config.schema import ModelPresetConfig
|
||||
p = ModelPresetConfig(model="test/model")
|
||||
assert p.fallback_models == []
|
||||
from nanobot.config.schema import AgentDefaults
|
||||
|
||||
defaults = AgentDefaults()
|
||||
|
||||
assert defaults.fallback_models == []
|
||||
|
||||
|
||||
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"]
|
||||
def test_fallback_models_accept_preset_refs_and_inline_configs() -> None:
|
||||
from nanobot.config.schema import Config, InlineFallbackConfig
|
||||
|
||||
|
||||
def test_fallback_models_from_camel_case() -> None:
|
||||
from nanobot.config.schema import ModelPresetConfig
|
||||
p = ModelPresetConfig.model_validate({
|
||||
"model": "test/primary",
|
||||
"fallbackModels": ["test/a"],
|
||||
config = Config.model_validate({
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"fallbackModels": [
|
||||
"deep",
|
||||
{
|
||||
"provider": "openai",
|
||||
"model": "gpt-4.1",
|
||||
"maxTokens": 4096,
|
||||
},
|
||||
]
|
||||
}
|
||||
},
|
||||
"modelPresets": {
|
||||
"deep": {"provider": "anthropic", "model": "claude-opus-4-7"}
|
||||
},
|
||||
})
|
||||
assert p.fallback_models == ["test/a"]
|
||||
|
||||
assert config.agents.defaults.fallback_models[0] == "deep"
|
||||
assert config.agents.defaults.fallback_models[1] == InlineFallbackConfig(
|
||||
provider="openai",
|
||||
model="gpt-4.1",
|
||||
max_tokens=4096,
|
||||
)
|
||||
|
||||
|
||||
def test_fallback_model_preset_ref_must_exist() -> None:
|
||||
from nanobot.config.schema import Config
|
||||
|
||||
with pytest.raises(ValueError, match="fallback_models.*not found"):
|
||||
Config.model_validate({
|
||||
"agents": {"defaults": {"fallbackModels": ["missing"]}},
|
||||
"modelPresets": {},
|
||||
})
|
||||
|
||||
|
||||
def test_provider_signature_tracks_fallback_presets_and_provider_config() -> None:
|
||||
from nanobot.config.schema import Config
|
||||
from nanobot.providers.factory import provider_signature
|
||||
|
||||
base = {
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"modelPreset": "fast",
|
||||
"fallbackModels": ["deep"],
|
||||
}
|
||||
},
|
||||
"modelPresets": {
|
||||
"fast": {"model": "openai/gpt-4.1", "provider": "openai"},
|
||||
"deep": {"model": "anthropic/claude-sonnet-4-6", "provider": "anthropic"},
|
||||
},
|
||||
"providers": {
|
||||
"openai": {"apiKey": "primary-key"},
|
||||
"anthropic": {"apiKey": "fallback-key"},
|
||||
},
|
||||
}
|
||||
changed_fallback = {
|
||||
**base,
|
||||
"agents": {"defaults": {"modelPreset": "fast", "fallbackModels": ["backup"]}},
|
||||
"modelPresets": {
|
||||
**base["modelPresets"],
|
||||
"backup": {"model": "deepseek/deepseek-chat", "provider": "deepseek"},
|
||||
},
|
||||
"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))
|
||||
|
||||
assert signature != provider_signature(Config.model_validate(changed_fallback))
|
||||
assert signature != provider_signature(Config.model_validate(changed_key))
|
||||
|
||||
|
||||
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({
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"modelPreset": "fast",
|
||||
"fallbackModels": ["deep"],
|
||||
}
|
||||
},
|
||||
"modelPresets": {
|
||||
"fast": {
|
||||
"model": "openai/gpt-4.1",
|
||||
"provider": "openai",
|
||||
"contextWindowTokens": 128000,
|
||||
},
|
||||
"deep": {
|
||||
"model": "deepseek/deepseek-chat",
|
||||
"provider": "deepseek",
|
||||
"contextWindowTokens": 64000,
|
||||
},
|
||||
},
|
||||
"providers": {
|
||||
"openai": {"apiKey": "primary-key"},
|
||||
"deepseek": {"apiKey": "fallback-key"},
|
||||
},
|
||||
})
|
||||
|
||||
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
|
||||
snapshot = build_provider_snapshot(config)
|
||||
|
||||
assert snapshot.context_window_tokens == 64000
|
||||
|
||||
|
||||
def test_inline_fallback_reasoning_effort_does_not_inherit_primary() -> None:
|
||||
from nanobot.config.schema import Config
|
||||
from nanobot.providers.factory import provider_signature
|
||||
|
||||
config = Config.model_validate({
|
||||
"agents": {
|
||||
"defaults": {
|
||||
"modelPreset": "fast",
|
||||
"fallbackModels": [
|
||||
{"provider": "openai", "model": "gpt-4.1"}
|
||||
],
|
||||
}
|
||||
},
|
||||
"modelPresets": {
|
||||
"fast": {
|
||||
"model": "anthropic/claude-opus-4-5",
|
||||
"provider": "anthropic",
|
||||
"reasoningEffort": "high",
|
||||
}
|
||||
},
|
||||
"providers": {
|
||||
"anthropic": {"apiKey": "primary-key"},
|
||||
"openai": {"apiKey": "fallback-key"},
|
||||
},
|
||||
})
|
||||
|
||||
signature = provider_signature(config)
|
||||
fallback_signatures = signature[-1]
|
||||
|
||||
assert fallback_signatures[0][11] is None
|
||||
|
||||
|
||||
# -- FallbackProvider tests --
|
||||
@@ -83,7 +254,7 @@ class TestNoFallbackWhenPrimarySucceeds:
|
||||
factory = MagicMock()
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_presets=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -102,14 +273,14 @@ class TestFallbackOnPrimaryError:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_presets=[_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 +292,7 @@ class TestNoFallbackWhenContentStreamed:
|
||||
factory = MagicMock()
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_presets=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -146,14 +317,62 @@ class TestFailoverOnTransientError:
|
||||
factory = MagicMock(return_value=fallback)
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_presets=[_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 TestNoFallbackOnNonRetryableError:
|
||||
@pytest.mark.asyncio
|
||||
async def test_bad_request(self) -> None:
|
||||
primary = _FakeProvider(
|
||||
"primary",
|
||||
_make_response(
|
||||
"invalid request",
|
||||
finish_reason="error",
|
||||
error_status_code=400,
|
||||
error_kind="invalid_request",
|
||||
),
|
||||
)
|
||||
factory = MagicMock()
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_presets=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
result = await fb.chat(messages=[{"role": "user", "content": "hi"}])
|
||||
|
||||
assert result.finish_reason == "error"
|
||||
factory.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auth_error(self) -> None:
|
||||
primary = _FakeProvider(
|
||||
"primary",
|
||||
_make_response(
|
||||
"unauthorized",
|
||||
finish_reason="error",
|
||||
error_status_code=401,
|
||||
error_kind="authentication",
|
||||
),
|
||||
)
|
||||
factory = MagicMock()
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_presets=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
result = await fb.chat(messages=[{"role": "user", "content": "hi"}])
|
||||
|
||||
assert result.finish_reason == "error"
|
||||
factory.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_timeout(self) -> None:
|
||||
@@ -165,14 +384,14 @@ class TestFailoverOnTransientError:
|
||||
factory = MagicMock(return_value=fallback)
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_presets=[_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 +404,15 @@ class TestFallbackTriesModelsInOrder:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a", "fallback-b"],
|
||||
fallback_presets=[_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 +424,7 @@ class TestAllFallbacksFail:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_presets=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -223,7 +442,7 @@ class TestFactoryExceptionSkipsModel:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a", "fallback-b"],
|
||||
fallback_presets=[_fallback("fallback-a"), _fallback("fallback-b")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -242,13 +461,43 @@ class TestFallbackModelParameter:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-model"],
|
||||
fallback_presets=[_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_uses_fallback_generation_fields(self) -> None:
|
||||
primary = _FakeProvider("primary", _error_response())
|
||||
fallback = _FakeProvider("fallback", _make_response("ok"))
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_presets=[
|
||||
_fallback(
|
||||
"fallback-model",
|
||||
max_tokens=1234,
|
||||
temperature=0.4,
|
||||
reasoning_effort=None,
|
||||
)
|
||||
],
|
||||
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 "reasoning_effort" not in fallback.chat_calls[0]
|
||||
|
||||
|
||||
class TestNoFallbackWhenEmptyList:
|
||||
@pytest.mark.asyncio
|
||||
@@ -258,7 +507,7 @@ class TestNoFallbackWhenEmptyList:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=[],
|
||||
fallback_presets=[],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -277,7 +526,7 @@ class TestChatStreamFailover:
|
||||
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_presets=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -291,7 +540,7 @@ class TestGetDefaultModel:
|
||||
primary = _FakeProvider("primary")
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["a"],
|
||||
fallback_presets=[_fallback("a")],
|
||||
provider_factory=MagicMock(),
|
||||
)
|
||||
assert fb.get_default_model() == "primary/model"
|
||||
@@ -305,7 +554,7 @@ class TestCircuitBreaker:
|
||||
factory = MagicMock(return_value=fallback)
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_presets=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -329,7 +578,7 @@ class TestCircuitBreaker:
|
||||
factory = MagicMock(return_value=fallback)
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["fallback-a"],
|
||||
fallback_presets=[_fallback("fallback-a")],
|
||||
provider_factory=factory,
|
||||
)
|
||||
|
||||
@@ -357,7 +606,7 @@ class TestGenerationForwarded:
|
||||
primary.generation = GenerationSettings(temperature=0.5, max_tokens=1024)
|
||||
fb = FallbackProvider(
|
||||
primary=primary,
|
||||
fallback_models=["a"],
|
||||
fallback_presets=[_fallback("a")],
|
||||
provider_factory=MagicMock(),
|
||||
)
|
||||
assert fb.generation.temperature == 0.5
|
||||
|
||||
Reference in New Issue
Block a user