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:
parent
8ee13e53c5
commit
da664fd4a5
4 changed files with 241 additions and 8 deletions
|
|
@ -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,13 +69,21 @@ 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,
|
||||
registry: registry,
|
||||
state: stateManager,
|
||||
summarizing: sync.Map{},
|
||||
fallback: fallbackChain,
|
||||
bus: msgBus,
|
||||
cfg: cfg,
|
||||
registry: registry,
|
||||
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,
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue