From b06d0adab5e4ac7ca05cc4d7ec1c173c89b86d30 Mon Sep 17 00:00:00 2001 From: dj-oyu <68707227+dj-oyu@users.noreply.github.com> Date: Sun, 22 Feb 2026 01:41:13 +0900 Subject: [PATCH] fix: resolveProvider now supports model_list for cross-provider fallback MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit resolveProvider tries FindModelConfigByRef → CreateProviderFromConfig (model_list path) first, then falls back to legacy CreateProviderByName (providers section). This ensures cross-provider fallback works correctly when config uses the new model_list format. - Add FindModelConfigByRef to config.go for model_list lookup - Update resolveProvider signature to accept both provider and model name - Use provider/model composite key for provider cache - Remove eager cache seeding in NewAgentLoop (lazy creation sufficient) Co-Authored-By: Claude Opus 4.6 --- pkg/agent/loop.go | 41 +++++++++++++++++++++++++---------------- pkg/agent/loop_test.go | 14 +++++++------- pkg/config/config.go | 14 ++++++++++++++ pkg/config/defaults.go | 3 ++- 4 files changed, 48 insertions(+), 24 deletions(-) diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index c27b163a4..4fe64e96e 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -74,12 +74,7 @@ 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 - } // Create stats tracker if enabled var statsTracker *stats.Tracker @@ -273,24 +268,38 @@ 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 == "" { +// resolveProvider returns the LLMProvider for the given provider/model pair. +// It caches created providers by "provider/model" key so each combination is +// only resolved once. Looks up model_list first (new format), then falls back +// to the legacy providers section via CreateProviderByName. +func (al *AgentLoop) resolveProvider(providerName, modelName string, fallback providers.LLMProvider) providers.LLMProvider { + key := strings.ToLower(providerName + "/" + modelName) + if key == "/" { return fallback } - if p, ok := al.providerCache[name]; ok { + if p, ok := al.providerCache[key]; ok { return p } - p, err := providers.CreateProviderByName(al.cfg, name) + + // Try model_list first (new config format). + if mc := al.cfg.FindModelConfigByRef(providerName, modelName); mc != nil { + p, _, err := providers.CreateProviderFromConfig(mc) + if err == nil { + al.providerCache[key] = p + return p + } + logger.WarnCF("agent", "Failed to create provider from model_list, trying legacy", + map[string]interface{}{"provider": providerName, "model": modelName, "error": err.Error()}) + } + + // Fall back to legacy providers section. + p, err := providers.CreateProviderByName(al.cfg, providerName) if err != nil { logger.WarnCF("agent", "Failed to create provider for fallback, using primary", - map[string]interface{}{"provider": name, "error": err.Error()}) + map[string]interface{}{"provider": providerName, "error": err.Error()}) return fallback } - al.providerCache[name] = p + al.providerCache[key] = p return p } @@ -762,7 +771,7 @@ func (al *AgentLoop) runLLMIteration( 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) { - p := al.resolveProvider(provider, agent.Provider) + p := al.resolveProvider(provider, model, agent.Provider) return p.Chat(ctx, messages, providerToolDefs, model, map[string]interface{}{ "max_tokens": agent.MaxTokens, "temperature": agent.Temperature, diff --git a/pkg/agent/loop_test.go b/pkg/agent/loop_test.go index 2dd8ab9f2..84a3b4866 100644 --- a/pkg/agent/loop_test.go +++ b/pkg/agent/loop_test.go @@ -725,14 +725,14 @@ func TestResolveProvider_CachesProviders(t *testing.T) { 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'") + // First call creates and caches a provider for "vllm/test-model" + p1 := al.resolveProvider("vllm", "test-model", primary) + if p1 == primary { + t.Fatal("expected a new provider from legacy providers config, not the fallback") } // Calling again should return the same instance (cached) - p2 := al.resolveProvider("vllm", primary) + p2 := al.resolveProvider("vllm", "test-model", primary) if p1 != p2 { t.Fatal("expected same cached instance on second call") } @@ -762,7 +762,7 @@ func TestResolveProvider_FallsBackOnError(t *testing.T) { al := NewAgentLoop(cfg, msgBus, primary) // Request a provider that can't be created (no config for "nonexistent") - p := al.resolveProvider("nonexistent", primary) + p := al.resolveProvider("nonexistent", "unknown-model", primary) if p != primary { t.Fatal("expected fallback to primary provider on creation error") } @@ -795,7 +795,7 @@ func TestResolveProvider_EmptyNameReturnsFallback(t *testing.T) { primary := &mockProvider{} al := NewAgentLoop(cfg, msgBus, primary) - p := al.resolveProvider("", primary) + p := al.resolveProvider("", "", primary) if p != primary { t.Fatal("expected fallback provider for empty name") } diff --git a/pkg/config/config.go b/pkg/config/config.go index 5554ffd20..7ea46a512 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path/filepath" + "strings" "sync/atomic" "github.com/caarlos0/env/v11" @@ -618,6 +619,19 @@ func (c *Config) findMatches(modelName string) []ModelConfig { return matches } +// FindModelConfigByRef finds a ModelConfig entry whose Model field matches +// "protocol/modelID" (case-insensitive). Used by the fallback chain to look up +// cross-provider candidates in model_list. +func (c *Config) FindModelConfigByRef(protocol, modelID string) *ModelConfig { + target := strings.ToLower(protocol + "/" + modelID) + for i := range c.ModelList { + if strings.ToLower(c.ModelList[i].Model) == target { + return &c.ModelList[i] + } + } + return nil +} + // HasProvidersConfig checks if any provider in the old providers config has configuration. func (c *Config) HasProvidersConfig() bool { v := c.Providers diff --git a/pkg/config/defaults.go b/pkg/config/defaults.go index 7654326e7..aa15d6d89 100644 --- a/pkg/config/defaults.go +++ b/pkg/config/defaults.go @@ -16,7 +16,8 @@ func DefaultConfig() *Config { Model: "glm-4.7", MaxTokens: 8192, Temperature: nil, // nil means use provider default - MaxToolIterations: 20, + MaxToolIterations: 20, + TaskReminderInterval: 5, }, }, Bindings: []AgentBinding{},