fix: resolveProvider now supports model_list for cross-provider fallback

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 <noreply@anthropic.com>
This commit is contained in:
dj-oyu 2026-02-22 01:41:13 +09:00
parent 46d06ff131
commit b06d0adab5
4 changed files with 48 additions and 24 deletions

View file

@ -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,

View file

@ -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")
}

View file

@ -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

View file

@ -17,6 +17,7 @@ func DefaultConfig() *Config {
MaxTokens: 8192,
Temperature: nil, // nil means use provider default
MaxToolIterations: 20,
TaskReminderInterval: 5,
},
},
Bindings: []AgentBinding{},