diff --git a/pkg/agent/instance.go b/pkg/agent/instance.go index 880725660..ec59a8e19 100644 --- a/pkg/agent/instance.go +++ b/pkg/agent/instance.go @@ -51,6 +51,10 @@ type AgentInstance struct { // LightProvider is the concrete provider instance for the configured light model. // It is only used when routing selects the light tier for a turn. LightProvider providers.LLMProvider + // CandidateProviders maps "provider/model" keys to per-candidate LLMProvider + // instances. This allows each fallback model to use its own api_base and api_key + // from model_list, instead of inheriting the primary model's provider config. + CandidateProviders map[string]providers.LLMProvider } // NewAgentInstance creates an agent instance from config. @@ -170,6 +174,10 @@ func NewAgentInstance( // Resolve fallback candidates candidates := resolveModelCandidates(cfg, defaults.Provider, model, fallbacks) + candidateProviders := make(map[string]providers.LLMProvider) + modelIndex := buildModelIndex(cfg) + populateCandidateProviders(modelIndex, cfg.WorkspacePath(), candidates, candidateProviders) + // Model routing setup: pre-resolve light model candidates at creation time // to avoid repeated model_list lookups on every incoming message. var router *routing.Router @@ -194,6 +202,7 @@ func NewAgentInstance( }) lightCandidates = resolved lightProvider = lp + populateCandidateProviders(modelIndex, cfg.WorkspacePath(), resolved, candidateProviders) } } } else { @@ -225,9 +234,85 @@ func NewAgentInstance( Router: router, LightCandidates: lightCandidates, LightProvider: lightProvider, + CandidateProviders: candidateProviders, } } +// buildModelIndex returns a normalised "provider/model" → *ModelConfig lookup for +// cfg.ModelList. First entry wins for duplicate keys. Returns nil when the list is empty. +// +// Uses ExtractProtocol + NormalizeProvider (not cfg.GetModelConfig) because candidates +// hold resolved (Provider, Model) pairs that must be matched against the Model field, +// not the model_name alias. +func buildModelIndex(cfg *config.Config) map[string]*config.ModelConfig { + if cfg == nil || len(cfg.ModelList) == 0 { + return nil + } + index := make(map[string]*config.ModelConfig, len(cfg.ModelList)) + for _, mc := range cfg.ModelList { + entry := strings.TrimSpace(mc.Model) + if entry == "" { + continue + } + protocol, modelID := providers.ExtractProtocol(entry) + k := providers.ModelKey(providers.NormalizeProvider(protocol), modelID) + if _, exists := index[k]; !exists { + index[k] = mc + } + } + return index +} + +// populateCandidateProviders creates an LLMProvider for each candidate using the +// pre-built model index and stores it in out. Candidates absent from the index are +// skipped with a warning; they inherit the primary provider's credentials at runtime. +func populateCandidateProviders( + index map[string]*config.ModelConfig, + workspacePath string, + candidates []providers.FallbackCandidate, + out map[string]providers.LLMProvider, +) { + for _, c := range candidates { + key := providers.ModelKey(c.Provider, c.Model) + if _, exists := out[key]; exists { + continue + } + mc, found := index[key] + if !found { + logger.WarnCF( + "agent", + "fallback provider: no model_list entry found; will inherit primary provider credentials", + map[string]any{"provider": c.Provider, "model": c.Model}, + ) + continue + } + mCopy := *mc + if mCopy.Workspace == "" { + mCopy.Workspace = workspacePath + } + p, _, err := providers.CreateProviderFromConfig(&mCopy) + if err == nil { + out[key] = p + } else { + logger.WarnCF("agent", "fallback provider: failed to create provider", + map[string]any{"model": mc.Model, "error": err.Error()}) + } + } +} + +// registerCandidateProviders is the test-facing entry point that wraps +// buildModelIndex + populateCandidateProviders in a single call. +func registerCandidateProviders( + cfg *config.Config, + candidates []providers.FallbackCandidate, + out map[string]providers.LLMProvider, +) { + if cfg == nil || len(cfg.ModelList) == 0 || len(candidates) == 0 { + return + } + populateCandidateProviders(buildModelIndex(cfg), cfg.WorkspacePath(), candidates, out) +} + // resolveAgentWorkspace determines the workspace directory for an agent. func resolveAgentWorkspace(agentCfg *config.AgentConfig, defaults *config.AgentDefaults) string { if agentCfg != nil && strings.TrimSpace(agentCfg.Workspace) != "" { diff --git a/pkg/agent/instance_test.go b/pkg/agent/instance_test.go index e296a18cb..78da6ed19 100644 --- a/pkg/agent/instance_test.go +++ b/pkg/agent/instance_test.go @@ -9,6 +9,7 @@ import ( "github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/media" + "github.com/sipeed/picoclaw/pkg/providers" ) func TestNewAgentInstance_UsesDefaultsTemperatureAndMaxTokens(t *testing.T) { @@ -248,6 +249,228 @@ func TestNewAgentInstance_AllowsMediaTempDirForReadListAndExec(t *testing.T) { } } +// TestRegisterCandidateProviders_NilCfgIsNoop verifies that passing a nil +// config does not panic and leaves the output map empty. +func TestRegisterCandidateProviders_NilCfgIsNoop(t *testing.T) { + out := map[string]providers.LLMProvider{} + registerCandidateProviders(nil, []providers.FallbackCandidate{{Provider: "openai", Model: "gpt-4o"}}, out) + if len(out) != 0 { + t.Fatalf("expected empty map, got %d entries", len(out)) + } +} + +// TestRegisterCandidateProviders_SkipsExistingKeys verifies that a key already +// present in the output map is not overwritten. +func TestRegisterCandidateProviders_SkipsExistingKeys(t *testing.T) { + existing := &mockProvider{} + key := providers.ModelKey("openai", "gpt-4o") + out := map[string]providers.LLMProvider{key: existing} + + cfg := &config.Config{ + ModelList: []*config.ModelConfig{ + {Model: "openai/gpt-4o", APIKeys: config.SimpleSecureStrings("test-key")}, + }, + } + registerCandidateProviders(cfg, []providers.FallbackCandidate{{Provider: "openai", Model: "gpt-4o"}}, out) + + if out[key] != existing { + t.Fatal("existing provider entry was overwritten; expected it to be preserved") + } +} + +// TestRegisterCandidateProviders_MatchesBareModelName verifies that a +// model_list entry without a provider prefix (e.g. "gpt-4o") still matches a +// candidate whose provider is "openai" — the default protocol that +// ExtractProtocol assigns to bare names. +func TestRegisterCandidateProviders_MatchesBareModelName(t *testing.T) { + workspace := t.TempDir() + out := map[string]providers.LLMProvider{} + + cfg := &config.Config{ + ModelList: []*config.ModelConfig{ + // Bare name — no "openai/" prefix. ExtractProtocol should + // assign "openai" as the default protocol. + {Model: "gpt-4o", APIBase: "https://api.openai.com/v1", Workspace: workspace}, + }, + } + registerCandidateProviders(cfg, []providers.FallbackCandidate{{Provider: "openai", Model: "gpt-4o"}}, out) + + key := providers.ModelKey("openai", "gpt-4o") + if out[key] == nil { + t.Fatalf("expected CandidateProviders[%q] to be populated for bare model name", key) + } +} + +// TestRegisterCandidateProviders_MatchesWithProtocolPrefix verifies that a +// model_list entry using full "provider/model" notation (e.g. +// "gemini/gemma-3-27b-it") is matched correctly. +func TestRegisterCandidateProviders_MatchesWithProtocolPrefix(t *testing.T) { + workspace := t.TempDir() + out := map[string]providers.LLMProvider{} + + cfg := &config.Config{ + ModelList: []*config.ModelConfig{ + { + Model: "gemini/gemma-3-27b-it", + APIKeys: config.SimpleSecureStrings("gemini-test-key"), + Workspace: workspace, + }, + }, + } + registerCandidateProviders(cfg, []providers.FallbackCandidate{{Provider: "gemini", Model: "gemma-3-27b-it"}}, out) + + key := providers.ModelKey("gemini", "gemma-3-27b-it") + if out[key] == nil { + t.Fatalf("expected CandidateProviders[%q] to be populated for protocol-prefixed model name", key) + } +} + +// TestRegisterCandidateProviders_EmptyCandidatesIsNoop verifies the early-exit +// path when the candidates slice is empty — no index is built and the map +// remains unchanged. +func TestRegisterCandidateProviders_EmptyCandidatesIsNoop(t *testing.T) { + out := map[string]providers.LLMProvider{} + cfg := &config.Config{ + ModelList: []*config.ModelConfig{ + {Model: "openai/gpt-4o", APIKeys: config.SimpleSecureStrings("key")}, + }, + } + registerCandidateProviders(cfg, nil, out) + if len(out) != 0 { + t.Fatalf("expected empty map, got %d entries", len(out)) + } +} + +// TestRegisterCandidateProviders_EmptyModelListIsNoop verifies the early-exit +// path when model_list is empty — no provider can be created. +func TestRegisterCandidateProviders_EmptyModelListIsNoop(t *testing.T) { + out := map[string]providers.LLMProvider{} + cfg := &config.Config{} + registerCandidateProviders(cfg, []providers.FallbackCandidate{{Provider: "openai", Model: "gpt-4o"}}, out) + if len(out) != 0 { + t.Fatalf("expected empty map, got %d entries", len(out)) + } +} + +// TestRegisterCandidateProviders_FirstModelListEntryWinsForDuplicates verifies +// that when model_list contains two entries with the same normalised +// provider/model key, the first one is used (mirrors model_list precedence). +func TestRegisterCandidateProviders_FirstModelListEntryWinsForDuplicates(t *testing.T) { + workspace := t.TempDir() + out := map[string]providers.LLMProvider{} + + cfg := &config.Config{ + ModelList: []*config.ModelConfig{ + {Model: "openai/gpt-4o", APIBase: "https://first.example.com/v1", Workspace: workspace}, + {Model: "openai/gpt-4o", APIBase: "https://second.example.com/v1", Workspace: workspace}, + }, + } + registerCandidateProviders(cfg, []providers.FallbackCandidate{{Provider: "openai", Model: "gpt-4o"}}, out) + + key := providers.ModelKey("openai", "gpt-4o") + if out[key] == nil { + t.Fatalf("expected CandidateProviders[%q] to be populated", key) + } + // Only one entry should be registered despite two model_list entries. + if len(out) != 1 { + t.Fatalf("expected 1 entry, got %d", len(out)) + } +} + +// TestRegisterCandidateProviders_UnmatchedCandidateIsSkipped verifies that a +// candidate with no matching model_list entry is silently skipped and does not +// cause a panic or leave a nil entry in the map. +func TestRegisterCandidateProviders_UnmatchedCandidateIsSkipped(t *testing.T) { + out := map[string]providers.LLMProvider{} + cfg := &config.Config{ + ModelList: []*config.ModelConfig{ + {Model: "openai/gpt-4o", APIKeys: config.SimpleSecureStrings("key")}, + }, + } + // "anthropic/claude-3-opus" has no matching model_list entry. + registerCandidateProviders(cfg, []providers.FallbackCandidate{{Provider: "anthropic", Model: "claude-3-opus"}}, out) + + if len(out) != 0 { + t.Fatalf("expected empty map for unmatched candidate, got %d entries", len(out)) + } +} + +// TestNewAgentInstance_CandidateProvidersPopulatedForCrossProviderFallbacks +// mirrors the exact scenario from bug #2140: primary model on OpenRouter with +// Gemini fallbacks. Each entry must get its own provider instance so that +// fallback requests go to the correct API endpoint, not the primary's. +func TestNewAgentInstance_CandidateProvidersPopulatedForCrossProviderFallbacks(t *testing.T) { + workspace := t.TempDir() + + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{ + Workspace: workspace, + ModelName: "mistral-small-3.1", + ModelFallbacks: []string{"gemma-3-27b", "gemini-images"}, + }, + }, + ModelList: []*config.ModelConfig{ + { + ModelName: "mistral-small-3.1", + Model: "openrouter/mistralai/mistral-small-3.1-24b-instruct:free", + APIBase: "https://openrouter.ai/api/v1", + APIKeys: config.SimpleSecureStrings("sk-or-test"), + Workspace: workspace, + }, + { + ModelName: "gemma-3-27b", + Model: "gemini/gemma-3-27b-it", + APIKeys: config.SimpleSecureStrings("AIzaSy-test"), + Workspace: workspace, + }, + { + ModelName: "gemini-images", + Model: "gemini/gemini-2.5-flash-lite", + APIKeys: config.SimpleSecureStrings("AIzaSy-test"), + Workspace: workspace, + }, + }, + } + + primaryProvider := &mockProvider{} + agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, primaryProvider) + + wantKeys := []string{ + providers.ModelKey("openrouter", "mistralai/mistral-small-3.1-24b-instruct:free"), + providers.ModelKey("gemini", "gemma-3-27b-it"), + providers.ModelKey("gemini", "gemini-2.5-flash-lite"), + } + + for _, key := range wantKeys { + p, ok := agent.CandidateProviders[key] + if !ok { + t.Errorf("CandidateProviders missing key %q", key) + continue + } + if p == nil { + t.Errorf("CandidateProviders[%q] is nil", key) + } + // Each fallback must use its own provider, not the injected primary. + if p == primaryProvider { + t.Errorf( + "CandidateProviders[%q] is the same instance as the primary provider; fallback would inherit primary credentials", + key, + ) + } + } + + if t.Failed() { + t.Logf("CandidateProviders keys present: %v", func() []string { + keys := make([]string, 0, len(agent.CandidateProviders)) + for k := range agent.CandidateProviders { + keys = append(keys, k) + } + return keys + }()) + } +} + func TestNewAgentInstance_InvalidExecConfigDoesNotExit(t *testing.T) { workspace := t.TempDir() diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index 624ff261b..092f528d4 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -2002,7 +2002,11 @@ turnLoop: providerCtx, activeCandidates, func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) { - return activeProvider.Chat(ctx, messagesForCall, toolDefsForCall, model, llmOpts) + candidateProvider := activeProvider + if cp, ok := ts.agent.CandidateProviders[providers.ModelKey(provider, model)]; ok { + candidateProvider = cp + } + return candidateProvider.Chat(ctx, messagesForCall, toolDefsForCall, model, llmOpts) }, ) if fbErr != nil { diff --git a/pkg/agent/loop_test.go b/pkg/agent/loop_test.go index 9513d8aca..3d04b81cc 100644 --- a/pkg/agent/loop_test.go +++ b/pkg/agent/loop_test.go @@ -1839,6 +1839,164 @@ func TestProcessMessage_ModelRoutingUsesLightProvider(t *testing.T) { } } +// TestProcessMessage_FallbackUsesPerCandidateProvider is the loop-level test for +// bug #2140. It verifies that when the primary model returns a rate-limit error +// the fallback closure routes the retry to the fallback model's own provider +// (its own api_base), not back to the primary provider's endpoint. +func TestProcessMessage_FallbackUsesPerCandidateProvider(t *testing.T) { + workspace := t.TempDir() + + primaryCalls := 0 + primaryServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + primaryCalls++ + // Return 429 so FallbackChain classifies this as retriable and moves on. + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusTooManyRequests) + _ = json.NewEncoder(w).Encode(map[string]any{ + "error": map[string]any{ + "message": "rate limit exceeded", + "type": "rate_limit_error", + }, + }) + })) + defer primaryServer.Close() + + fallbackCalls := 0 + fallbackServer := newStrictChatCompletionTestServer( + t, "fallback", "gemma-3-27b-it", "fallback reply", &fallbackCalls, + ) + defer fallbackServer.Close() + + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{ + Workspace: workspace, + ModelName: "mistral-primary", + ModelFallbacks: []string{"gemma-fallback"}, + MaxTokens: 4096, + MaxToolIterations: 3, + }, + }, + ModelList: []*config.ModelConfig{ + { + ModelName: "mistral-primary", + Model: "openrouter/mistralai/mistral-small-3.1", + APIBase: primaryServer.URL, + APIKeys: config.SimpleSecureStrings("primary-key"), + Workspace: workspace, + }, + { + ModelName: "gemma-fallback", + Model: "gemini/gemma-3-27b-it", + APIBase: fallbackServer.URL, + APIKeys: config.SimpleSecureStrings("fallback-key"), + Workspace: workspace, + }, + }, + } + + provider, _, err := providers.CreateProvider(cfg) + if err != nil { + t.Fatalf("CreateProvider() error = %v", err) + } + msgBus := bus.NewMessageBus() + al := NewAgentLoop(cfg, msgBus, provider) + helper := testHelper{al: al} + + resp := helper.executeAndGetResponse(t, context.Background(), bus.InboundMessage{ + Channel: "telegram", + SenderID: "user1", + ChatID: "chat1", + Content: "hi", + Peer: bus.Peer{Kind: "direct", ID: "user1"}, + }) + + if resp != "fallback reply" { + t.Fatalf("response = %q, want %q (fallback provider)", resp, "fallback reply") + } + if primaryCalls == 0 { + t.Fatal("primary server was never called; expected at least one attempt") + } + if fallbackCalls != 1 { + t.Fatalf("fallback server calls = %d, want 1", fallbackCalls) + } +} + +// TestProcessMessage_FallbackUsesActiveProviderWhenCandidateNotRegistered verifies +// that when a candidate has no model_list entry it is absent from CandidateProviders +// and the fallback closure falls back to activeProvider instead of panicking. +func TestProcessMessage_FallbackUsesActiveProviderWhenCandidateNotRegistered(t *testing.T) { + workspace := t.TempDir() + + // Primary server: returns 429 on first call, succeeds on second. + // Both the primary and the unregistered fallback share this server + // (same api_base) so activeProvider routes both calls here. + callCount := 0 + primaryServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + callCount++ + w.Header().Set("Content-Type", "application/json") + if callCount == 1 { + w.WriteHeader(http.StatusTooManyRequests) + _ = json.NewEncoder(w).Encode(map[string]any{ + "error": map[string]any{"message": "rate limit", "type": "rate_limit_error"}, + }) + return + } + // Second call (fallback via activeProvider) succeeds. + _ = json.NewEncoder(w).Encode(map[string]any{ + "choices": []map[string]any{ + {"message": map[string]any{"content": "active provider reply"}, "finish_reason": "stop"}, + }, + }) + })) + defer primaryServer.Close() + + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{ + Workspace: workspace, + ModelName: "primary-model", + MaxTokens: 4096, + MaxToolIterations: 3, + // No model_list entry for this alias — absent from CandidateProviders. + ModelFallbacks: []string{"openrouter/fallback-model"}, + }, + }, + ModelList: []*config.ModelConfig{ + { + ModelName: "primary-model", + Model: "openrouter/primary-model", + APIBase: primaryServer.URL, + APIKeys: config.SimpleSecureStrings("primary-key"), + Workspace: workspace, + }, + }, + } + + provider, _, err := providers.CreateProvider(cfg) + if err != nil { + t.Fatalf("CreateProvider() error = %v", err) + } + msgBus := bus.NewMessageBus() + al := NewAgentLoop(cfg, msgBus, provider) + + helper := testHelper{al: al} + resp := helper.executeAndGetResponse(t, context.Background(), bus.InboundMessage{ + Channel: "telegram", + SenderID: "user1", + ChatID: "chat1", + Content: "hi", + Peer: bus.Peer{Kind: "direct", ID: "user1"}, + }) + + if resp != "active provider reply" { + t.Fatalf("response = %q, want %q", resp, "active provider reply") + } + if callCount < 2 { + t.Fatalf("primary server calls = %d, want >= 2 (one 429 + one success via activeProvider)", callCount) + } +} + // TestToolResult_SilentToolDoesNotSendUserMessage verifies silent tools don't trigger outbound func TestToolResult_SilentToolDoesNotSendUserMessage(t *testing.T) { tmpDir, err := os.MkdirTemp("", "agent-test-*")