feat(fallback): enable cross-provider fallback in model_fallbacks chain

Previously the fallback callback always called agent.Provider.Chat(),
ignoring the candidate's provider name. This adds a provider cache to
AgentLoop with lazy creation via CreateProviderByName, so fallback
candidates like "openai/gpt-4o" correctly resolve to a different
LLMProvider instance (e.g. CodexProvider) instead of reusing the
primary vllm HTTPProvider.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
dj-oyu 2026-02-20 14:53:48 +09:00
parent f165fb8c83
commit ae2541e9b8
4 changed files with 241 additions and 8 deletions

View file

@ -37,6 +37,7 @@ type AgentLoop struct {
summarizing sync.Map
fallback *providers.FallbackChain
channelManager *channels.Manager
providerCache map[string]providers.LLMProvider
}
// processOptions configures how a message is processed
@ -68,6 +69,13 @@ func NewAgentLoop(cfg *config.Config, msgBus *bus.MessageBus, provider providers
stateManager = state.NewManager(defaultAgent.Workspace)
}
// Initialize provider cache and seed with primary provider
providerCache := make(map[string]providers.LLMProvider)
primaryName := strings.ToLower(cfg.Agents.Defaults.Provider)
if primaryName != "" {
providerCache[primaryName] = provider
}
return &AgentLoop{
bus: msgBus,
cfg: cfg,
@ -75,6 +83,7 @@ func NewAgentLoop(cfg *config.Config, msgBus *bus.MessageBus, provider providers
state: stateManager,
summarizing: sync.Map{},
fallback: fallbackChain,
providerCache: providerCache,
}
}
@ -193,6 +202,27 @@ func (al *AgentLoop) SetChannelManager(cm *channels.Manager) {
al.channelManager = cm
}
// resolveProvider returns the LLMProvider for the given provider name.
// It caches created providers so each name is only resolved once.
// On creation failure, it logs a warning and returns the fallback provider.
func (al *AgentLoop) resolveProvider(providerName string, fallback providers.LLMProvider) providers.LLMProvider {
name := strings.ToLower(providerName)
if name == "" {
return fallback
}
if p, ok := al.providerCache[name]; ok {
return p
}
p, err := providers.CreateProviderByName(al.cfg, name)
if err != nil {
logger.WarnCF("agent", "Failed to create provider for fallback, using primary",
map[string]interface{}{"provider": name, "error": err.Error()})
return fallback
}
al.providerCache[name] = p
return p
}
// RecordLastChannel records the last active channel for this workspace.
// This uses the atomic state save mechanism to prevent data loss on crash.
func (al *AgentLoop) RecordLastChannel(channel string) error {
@ -526,7 +556,8 @@ func (al *AgentLoop) runLLMIteration(ctx context.Context, agent *AgentInstance,
if len(agent.Candidates) > 1 && al.fallback != nil {
fbResult, fbErr := al.fallback.Execute(ctx, agent.Candidates,
func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) {
return agent.Provider.Chat(ctx, messages, providerToolDefs, model, map[string]interface{}{
p := al.resolveProvider(provider, agent.Provider)
return p.Chat(ctx, messages, providerToolDefs, model, map[string]interface{}{
"max_tokens": 8192,
"temperature": 0.7,
})

View file

@ -693,6 +693,111 @@ func TestBuildTaskReminder_WithBlocker(t *testing.T) {
}
}
func TestResolveProvider_CachesProviders(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-test-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Model: "test-model",
Provider: "vllm",
MaxTokens: 4096,
MaxToolIterations: 10,
},
},
Providers: config.ProvidersConfig{
VLLM: config.ProviderConfig{
APIKey: "test-key",
APIBase: "https://example.com/v1",
},
},
}
msgBus := bus.NewMessageBus()
primary := &mockProvider{}
al := NewAgentLoop(cfg, msgBus, primary)
// Primary provider should be cached under "vllm"
p1 := al.resolveProvider("vllm", primary)
if p1 != primary {
t.Fatal("expected primary provider for 'vllm'")
}
// Calling again should return the same instance (cached)
p2 := al.resolveProvider("vllm", primary)
if p1 != p2 {
t.Fatal("expected same cached instance on second call")
}
}
func TestResolveProvider_FallsBackOnError(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-test-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Model: "test-model",
Provider: "vllm",
MaxTokens: 4096,
MaxToolIterations: 10,
},
},
}
msgBus := bus.NewMessageBus()
primary := &mockProvider{}
al := NewAgentLoop(cfg, msgBus, primary)
// Request a provider that can't be created (no config for "nonexistent")
p := al.resolveProvider("nonexistent", primary)
if p != primary {
t.Fatal("expected fallback to primary provider on creation error")
}
// Ensure the failed provider is NOT cached
if _, ok := al.providerCache["nonexistent"]; ok {
t.Fatal("failed provider should not be cached")
}
}
func TestResolveProvider_EmptyNameReturnsFallback(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-test-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Model: "test-model",
MaxTokens: 4096,
MaxToolIterations: 10,
},
},
}
msgBus := bus.NewMessageBus()
primary := &mockProvider{}
al := NewAgentLoop(cfg, msgBus, primary)
p := al.resolveProvider("", primary)
if p != primary {
t.Fatal("expected fallback provider for empty name")
}
}
func TestBuildTaskReminder_Truncation(t *testing.T) {
// Build a long message (1000 runes)
longMsg := strings.Repeat("あ", 1000)

View file

@ -63,8 +63,11 @@ func createCodexAuthProvider(enableWebSearch bool) (LLMProvider, error) {
}
func resolveProviderSelection(cfg *config.Config) (providerSelection, error) {
return resolveProviderSelectionByName(cfg, strings.ToLower(cfg.Agents.Defaults.Provider))
}
func resolveProviderSelectionByName(cfg *config.Config, providerName string) (providerSelection, error) {
model := cfg.Agents.Defaults.Model
providerName := strings.ToLower(cfg.Agents.Defaults.Provider)
lowerModel := strings.ToLower(model)
sel := providerSelection{
@ -358,3 +361,31 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) {
return NewHTTPProvider(sel.apiKey, sel.apiBase, sel.proxy), nil
}
}
// CreateProviderByName creates a provider for the given explicit provider name.
// Used by the fallback chain to resolve cross-provider candidates.
func CreateProviderByName(cfg *config.Config, providerName string) (LLMProvider, error) {
sel, err := resolveProviderSelectionByName(cfg, strings.ToLower(providerName))
if err != nil {
return nil, err
}
switch sel.providerType {
case providerTypeClaudeAuth:
return createClaudeAuthProvider(sel.apiBase)
case providerTypeCodexAuth:
return createCodexAuthProvider(sel.enableWebSearch)
case providerTypeCodexCLIToken:
c := NewCodexProviderWithTokenSource("", "", CreateCodexCliTokenSource())
c.enableWebSearch = sel.enableWebSearch
return c, nil
case providerTypeClaudeCLI:
return NewClaudeCliProvider(sel.workspace), nil
case providerTypeCodexCLI:
return NewCodexCliProvider(sel.workspace), nil
case providerTypeGitHubCopilot:
return NewGitHubCopilotProvider(sel.apiBase, sel.connectMode, sel.model)
default:
return NewHTTPProvider(sel.apiKey, sel.apiBase, sel.proxy), nil
}
}

View file

@ -297,3 +297,69 @@ func TestCreateProviderReturnsCodexProviderForOpenAIOAuth(t *testing.T) {
t.Fatalf("provider type = %T, want *CodexProvider", provider)
}
}
func TestCreateProviderByName_OpenAI_OAuth(t *testing.T) {
originalGetCredential := getCredential
t.Cleanup(func() { getCredential = originalGetCredential })
getCredential = func(provider string) (*auth.AuthCredential, error) {
if provider != "openai" {
t.Fatalf("provider = %q, want openai", provider)
}
return &auth.AuthCredential{
AccessToken: "openai-token",
AccountID: "acct_test",
}, nil
}
cfg := config.DefaultConfig()
cfg.Providers.OpenAI.AuthMethod = "oauth"
provider, err := CreateProviderByName(cfg, "openai")
if err != nil {
t.Fatalf("CreateProviderByName() error = %v", err)
}
if _, ok := provider.(*CodexProvider); !ok {
t.Fatalf("provider type = %T, want *CodexProvider", provider)
}
}
func TestCreateProviderByName_VLLM(t *testing.T) {
cfg := config.DefaultConfig()
cfg.Providers.VLLM.APIKey = "minimax-key"
cfg.Providers.VLLM.APIBase = "https://api.minimax.io/v1"
provider, err := CreateProviderByName(cfg, "vllm")
if err != nil {
t.Fatalf("CreateProviderByName() error = %v", err)
}
if _, ok := provider.(*HTTPProvider); !ok {
t.Fatalf("provider type = %T, want *HTTPProvider", provider)
}
}
func TestCreateProviderByName_Unknown(t *testing.T) {
cfg := config.DefaultConfig()
_, err := CreateProviderByName(cfg, "nonexistent-provider")
if err == nil {
t.Fatal("expected error for unknown provider, got nil")
}
}
func TestCreateProviderByName_CaseInsensitive(t *testing.T) {
cfg := config.DefaultConfig()
cfg.Providers.VLLM.APIKey = "test-key"
cfg.Providers.VLLM.APIBase = "https://example.com/v1"
provider, err := CreateProviderByName(cfg, "VLLM")
if err != nil {
t.Fatalf("CreateProviderByName() error = %v", err)
}
if _, ok := provider.(*HTTPProvider); !ok {
t.Fatalf("provider type = %T, want *HTTPProvider", provider)
}
}