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:
parent
46d06ff131
commit
b06d0adab5
4 changed files with 48 additions and 24 deletions
|
|
@ -74,12 +74,7 @@ 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)
|
providerCache := make(map[string]providers.LLMProvider)
|
||||||
primaryName := strings.ToLower(cfg.Agents.Defaults.Provider)
|
|
||||||
if primaryName != "" {
|
|
||||||
providerCache[primaryName] = provider
|
|
||||||
}
|
|
||||||
|
|
||||||
// Create stats tracker if enabled
|
// Create stats tracker if enabled
|
||||||
var statsTracker *stats.Tracker
|
var statsTracker *stats.Tracker
|
||||||
|
|
@ -273,24 +268,38 @@ func (al *AgentLoop) SetChannelManager(cm *channels.Manager) {
|
||||||
al.channelManager = cm
|
al.channelManager = cm
|
||||||
}
|
}
|
||||||
|
|
||||||
// resolveProvider returns the LLMProvider for the given provider name.
|
// resolveProvider returns the LLMProvider for the given provider/model pair.
|
||||||
// It caches created providers so each name is only resolved once.
|
// It caches created providers by "provider/model" key so each combination is
|
||||||
// On creation failure, it logs a warning and returns the fallback provider.
|
// only resolved once. Looks up model_list first (new format), then falls back
|
||||||
func (al *AgentLoop) resolveProvider(providerName string, fallback providers.LLMProvider) providers.LLMProvider {
|
// to the legacy providers section via CreateProviderByName.
|
||||||
name := strings.ToLower(providerName)
|
func (al *AgentLoop) resolveProvider(providerName, modelName string, fallback providers.LLMProvider) providers.LLMProvider {
|
||||||
if name == "" {
|
key := strings.ToLower(providerName + "/" + modelName)
|
||||||
|
if key == "/" {
|
||||||
return fallback
|
return fallback
|
||||||
}
|
}
|
||||||
if p, ok := al.providerCache[name]; ok {
|
if p, ok := al.providerCache[key]; ok {
|
||||||
return p
|
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 {
|
if err != nil {
|
||||||
logger.WarnCF("agent", "Failed to create provider for fallback, using primary",
|
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
|
return fallback
|
||||||
}
|
}
|
||||||
al.providerCache[name] = p
|
al.providerCache[key] = p
|
||||||
return p
|
return p
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -762,7 +771,7 @@ func (al *AgentLoop) runLLMIteration(
|
||||||
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) {
|
||||||
p := al.resolveProvider(provider, agent.Provider)
|
p := al.resolveProvider(provider, model, agent.Provider)
|
||||||
return p.Chat(ctx, messages, providerToolDefs, model, map[string]interface{}{
|
return p.Chat(ctx, messages, providerToolDefs, model, map[string]interface{}{
|
||||||
"max_tokens": agent.MaxTokens,
|
"max_tokens": agent.MaxTokens,
|
||||||
"temperature": agent.Temperature,
|
"temperature": agent.Temperature,
|
||||||
|
|
|
||||||
|
|
@ -725,14 +725,14 @@ func TestResolveProvider_CachesProviders(t *testing.T) {
|
||||||
primary := &mockProvider{}
|
primary := &mockProvider{}
|
||||||
al := NewAgentLoop(cfg, msgBus, primary)
|
al := NewAgentLoop(cfg, msgBus, primary)
|
||||||
|
|
||||||
// Primary provider should be cached under "vllm"
|
// First call creates and caches a provider for "vllm/test-model"
|
||||||
p1 := al.resolveProvider("vllm", primary)
|
p1 := al.resolveProvider("vllm", "test-model", primary)
|
||||||
if p1 != primary {
|
if p1 == primary {
|
||||||
t.Fatal("expected primary provider for 'vllm'")
|
t.Fatal("expected a new provider from legacy providers config, not the fallback")
|
||||||
}
|
}
|
||||||
|
|
||||||
// Calling again should return the same instance (cached)
|
// Calling again should return the same instance (cached)
|
||||||
p2 := al.resolveProvider("vllm", primary)
|
p2 := al.resolveProvider("vllm", "test-model", primary)
|
||||||
if p1 != p2 {
|
if p1 != p2 {
|
||||||
t.Fatal("expected same cached instance on second call")
|
t.Fatal("expected same cached instance on second call")
|
||||||
}
|
}
|
||||||
|
|
@ -762,7 +762,7 @@ func TestResolveProvider_FallsBackOnError(t *testing.T) {
|
||||||
al := NewAgentLoop(cfg, msgBus, primary)
|
al := NewAgentLoop(cfg, msgBus, primary)
|
||||||
|
|
||||||
// Request a provider that can't be created (no config for "nonexistent")
|
// 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 {
|
if p != primary {
|
||||||
t.Fatal("expected fallback to primary provider on creation error")
|
t.Fatal("expected fallback to primary provider on creation error")
|
||||||
}
|
}
|
||||||
|
|
@ -795,7 +795,7 @@ func TestResolveProvider_EmptyNameReturnsFallback(t *testing.T) {
|
||||||
primary := &mockProvider{}
|
primary := &mockProvider{}
|
||||||
al := NewAgentLoop(cfg, msgBus, primary)
|
al := NewAgentLoop(cfg, msgBus, primary)
|
||||||
|
|
||||||
p := al.resolveProvider("", primary)
|
p := al.resolveProvider("", "", primary)
|
||||||
if p != primary {
|
if p != primary {
|
||||||
t.Fatal("expected fallback provider for empty name")
|
t.Fatal("expected fallback provider for empty name")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,7 @@ import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/caarlos0/env/v11"
|
"github.com/caarlos0/env/v11"
|
||||||
|
|
@ -618,6 +619,19 @@ func (c *Config) findMatches(modelName string) []ModelConfig {
|
||||||
return matches
|
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.
|
// HasProvidersConfig checks if any provider in the old providers config has configuration.
|
||||||
func (c *Config) HasProvidersConfig() bool {
|
func (c *Config) HasProvidersConfig() bool {
|
||||||
v := c.Providers
|
v := c.Providers
|
||||||
|
|
|
||||||
|
|
@ -17,6 +17,7 @@ func DefaultConfig() *Config {
|
||||||
MaxTokens: 8192,
|
MaxTokens: 8192,
|
||||||
Temperature: nil, // nil means use provider default
|
Temperature: nil, // nil means use provider default
|
||||||
MaxToolIterations: 20,
|
MaxToolIterations: 20,
|
||||||
|
TaskReminderInterval: 5,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
Bindings: []AgentBinding{},
|
Bindings: []AgentBinding{},
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue