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
|
summarizing sync.Map
|
||||||
fallback *providers.FallbackChain
|
fallback *providers.FallbackChain
|
||||||
channelManager *channels.Manager
|
channelManager *channels.Manager
|
||||||
|
providerCache map[string]providers.LLMProvider
|
||||||
}
|
}
|
||||||
|
|
||||||
// processOptions configures how a message is processed
|
// 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)
|
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{
|
return &AgentLoop{
|
||||||
bus: msgBus,
|
bus: msgBus,
|
||||||
cfg: cfg,
|
cfg: cfg,
|
||||||
registry: registry,
|
registry: registry,
|
||||||
state: stateManager,
|
state: stateManager,
|
||||||
summarizing: sync.Map{},
|
summarizing: sync.Map{},
|
||||||
fallback: fallbackChain,
|
fallback: fallbackChain,
|
||||||
|
providerCache: providerCache,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -193,6 +202,27 @@ func (al *AgentLoop) SetChannelManager(cm *channels.Manager) {
|
||||||
al.channelManager = cm
|
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.
|
// RecordLastChannel records the last active channel for this workspace.
|
||||||
// This uses the atomic state save mechanism to prevent data loss on crash.
|
// This uses the atomic state save mechanism to prevent data loss on crash.
|
||||||
func (al *AgentLoop) RecordLastChannel(channel string) error {
|
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 {
|
if len(agent.Candidates) > 1 && al.fallback != nil {
|
||||||
fbResult, fbErr := al.fallback.Execute(ctx, agent.Candidates,
|
fbResult, fbErr := al.fallback.Execute(ctx, agent.Candidates,
|
||||||
func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) {
|
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,
|
"max_tokens": 8192,
|
||||||
"temperature": 0.7,
|
"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) {
|
func TestBuildTaskReminder_Truncation(t *testing.T) {
|
||||||
// Build a long message (1000 runes)
|
// Build a long message (1000 runes)
|
||||||
longMsg := strings.Repeat("あ", 1000)
|
longMsg := strings.Repeat("あ", 1000)
|
||||||
|
|
|
||||||
|
|
@ -63,8 +63,11 @@ func createCodexAuthProvider(enableWebSearch bool) (LLMProvider, error) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func resolveProviderSelection(cfg *config.Config) (providerSelection, 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
|
model := cfg.Agents.Defaults.Model
|
||||||
providerName := strings.ToLower(cfg.Agents.Defaults.Provider)
|
|
||||||
lowerModel := strings.ToLower(model)
|
lowerModel := strings.ToLower(model)
|
||||||
|
|
||||||
sel := providerSelection{
|
sel := providerSelection{
|
||||||
|
|
@ -358,3 +361,31 @@ func CreateProvider(cfg *config.Config) (LLMProvider, error) {
|
||||||
return NewHTTPProvider(sel.apiKey, sel.apiBase, sel.proxy), nil
|
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)
|
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