diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index 6733af05d..8b9e5c9bc 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -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, }) diff --git a/pkg/agent/loop_test.go b/pkg/agent/loop_test.go index bbf74fc1f..8abc2f142 100644 --- a/pkg/agent/loop_test.go +++ b/pkg/agent/loop_test.go @@ -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) diff --git a/pkg/providers/factory.go b/pkg/providers/factory.go index e39cfe32b..b90d7a645 100644 --- a/pkg/providers/factory.go +++ b/pkg/providers/factory.go @@ -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 + } +} diff --git a/pkg/providers/factory_test.go b/pkg/providers/factory_test.go index e31737eb9..584460fa4 100644 --- a/pkg/providers/factory_test.go +++ b/pkg/providers/factory_test.go @@ -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) + } +}