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
ab99e64d38
commit
7f4a1853bd
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)
|
||||
}
|
||||
|
||||
// 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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ func DefaultConfig() *Config {
|
|||
MaxTokens: 8192,
|
||||
Temperature: nil, // nil means use provider default
|
||||
MaxToolIterations: 20,
|
||||
TaskReminderInterval: 5,
|
||||
},
|
||||
},
|
||||
Bindings: []AgentBinding{},
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue