Fix fallback provider resolution for model aliases

This commit is contained in:
toocoldcc 2026-03-17 00:30:44 +08:00
parent f2addff099
commit feaf65911d
3 changed files with 235 additions and 26 deletions

View file

@ -39,6 +39,7 @@ type AgentInstance struct {
Subagents *config.SubagentsConfig
SkillsFilter []string
Candidates []providers.FallbackCandidate
candidateProviders map[string]providers.LLMProvider
// Router is non-nil when model routing is configured and the light model
// was successfully resolved. It scores each incoming message and decides
@ -153,6 +154,39 @@ func NewAgentInstance(
Primary: model,
Fallbacks: fallbacks,
}
resolvedModelConfigs := make(map[string]*config.ModelConfig)
resolveModelConfig := func(raw string) (*config.ModelConfig, bool) {
raw = strings.TrimSpace(raw)
if raw == "" || cfg == nil {
return nil, false
}
if mc, ok := resolvedModelConfigs[raw]; ok && mc != nil {
return mc, true
}
if mc, err := cfg.GetModelConfig(raw); err == nil && mc != nil && strings.TrimSpace(mc.Model) != "" {
resolvedModelConfigs[raw] = mc
return mc, true
}
for i := range cfg.ModelList {
fullModel := strings.TrimSpace(cfg.ModelList[i].Model)
if fullModel == "" {
continue
}
if fullModel == raw {
resolvedModelConfigs[raw] = &cfg.ModelList[i]
return &cfg.ModelList[i], true
}
_, modelID := providers.ExtractProtocol(fullModel)
if modelID == raw {
resolvedModelConfigs[raw] = &cfg.ModelList[i]
return &cfg.ModelList[i], true
}
}
return nil, false
}
resolveFromModelList := func(raw string) (string, bool) {
ensureProtocol := func(model string) string {
model = strings.TrimSpace(model)
@ -165,35 +199,38 @@ func NewAgentInstance(
return "openai/" + model
}
raw = strings.TrimSpace(raw)
if raw == "" {
return "", false
}
if cfg != nil {
if mc, err := cfg.GetModelConfig(raw); err == nil && mc != nil && strings.TrimSpace(mc.Model) != "" {
return ensureProtocol(mc.Model), true
}
for i := range cfg.ModelList {
fullModel := strings.TrimSpace(cfg.ModelList[i].Model)
if fullModel == "" {
continue
}
if fullModel == raw {
return ensureProtocol(fullModel), true
}
_, modelID := providers.ExtractProtocol(fullModel)
if modelID == raw {
return ensureProtocol(fullModel), true
}
}
if mc, ok := resolveModelConfig(raw); ok {
return ensureProtocol(mc.Model), true
}
return "", false
}
candidates := providers.ResolveCandidatesWithLookup(modelCfg, defaults.Provider, resolveFromModelList)
candidateProviders := make(map[string]providers.LLMProvider)
registerCandidateProvider := func(raw string) {
mc, ok := resolveModelConfig(raw)
if !ok {
return
}
protocol, modelID := providers.ExtractProtocol(strings.TrimSpace(mc.Model))
key := providers.ModelKey(protocol, modelID)
if _, exists := candidateProviders[key]; exists {
return
}
candidateProvider, candidateModel, err := providers.CreateProviderFromConfig(mc)
if err != nil || candidateProvider == nil {
return
}
key = providers.ModelKey(protocol, candidateModel)
candidateProviders[key] = candidateProvider
}
for _, raw := range fallbacks {
registerCandidateProvider(raw)
}
// Model routing setup: pre-resolve light model candidates at creation time
// to avoid repeated model_list lookups on every incoming message.
@ -203,6 +240,7 @@ func NewAgentInstance(
lightModelCfg := providers.ModelConfig{Primary: rc.LightModel}
resolved := providers.ResolveCandidatesWithLookup(lightModelCfg, defaults.Provider, resolveFromModelList)
if len(resolved) > 0 {
registerCandidateProvider(rc.LightModel)
router = routing.New(routing.RouterConfig{
LightModel: rc.LightModel,
Threshold: rc.Threshold,
@ -234,11 +272,20 @@ func NewAgentInstance(
Subagents: subagents,
SkillsFilter: skillsFilter,
Candidates: candidates,
candidateProviders: candidateProviders,
Router: router,
LightCandidates: lightCandidates,
}
}
func (a *AgentInstance) providerForCandidate(provider, model string) providers.LLMProvider {
key := providers.ModelKey(provider, model)
if candidateProvider, ok := a.candidateProviders[key]; ok && candidateProvider != nil {
return candidateProvider
}
return a.Provider
}
// resolveAgentWorkspace determines the workspace directory for an agent.
func resolveAgentWorkspace(agentCfg *config.AgentConfig, defaults *config.AgentDefaults) string {
if agentCfg != nil && strings.TrimSpace(agentCfg.Workspace) != "" {
@ -284,6 +331,14 @@ func compilePatterns(patterns []string) []*regexp.Regexp {
// Close releases resources held by the agent's session store.
func (a *AgentInstance) Close() error {
for _, candidateProvider := range a.candidateProviders {
if candidateProvider == nil || candidateProvider == a.Provider {
continue
}
if stateful, ok := candidateProvider.(providers.StatefulProvider); ok {
stateful.Close()
}
}
if a.Sessions != nil {
return a.Sessions.Close()
}

View file

@ -1049,10 +1049,15 @@ func (al *AgentLoop) runLLMIteration(
"temperature": agent.Temperature,
"prompt_cache_key": agent.ID,
}
activeProvider := agent.Provider
if len(activeCandidates) > 0 {
activeProvider = agent.providerForCandidate(activeCandidates[0].Provider, activeCandidates[0].Model)
}
// parseThinkingLevel guarantees ThinkingOff for empty/unknown values,
// so checking != ThinkingOff is sufficient.
if agent.ThinkingLevel != ThinkingOff {
if tc, ok := agent.Provider.(providers.ThinkingCapable); ok && tc.SupportsThinking() {
if tc, ok := activeProvider.(providers.ThinkingCapable); ok && tc.SupportsThinking() {
llmOpts["thinking_level"] = string(agent.ThinkingLevel)
} else {
logger.WarnCF("agent", "thinking_level is set but current provider does not support it, ignoring",
@ -1069,7 +1074,8 @@ func (al *AgentLoop) runLLMIteration(
ctx,
activeCandidates,
func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) {
return agent.Provider.Chat(ctx, messages, providerToolDefs, model, llmOpts)
candidateProvider := agent.providerForCandidate(provider, model)
return candidateProvider.Chat(ctx, messages, providerToolDefs, model, llmOpts)
},
)
if fbErr != nil {
@ -1085,7 +1091,7 @@ func (al *AgentLoop) runLLMIteration(
}
return fbResult.Response, nil
}
return agent.Provider.Chat(ctx, messages, providerToolDefs, activeModel, llmOpts)
return activeProvider.Chat(ctx, messages, providerToolDefs, activeModel, llmOpts)
}
// Retry loop for context/token errors

View file

@ -0,0 +1,148 @@
package agent
import (
"os"
"testing"
"github.com/sipeed/picoclaw/pkg/config"
)
func TestNewAgentInstance_UsesDedicatedProviderForFallbackAlias(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-provider-resolution-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Provider: "openai",
Model: "gpt-4.1-mini",
ModelFallbacks: []string{"gemini-fallback"},
},
},
ModelList: []config.ModelConfig{
{
ModelName: "gemini-fallback",
Model: "gemini-3.1-flash-lite-preview",
APIKey: "test-key",
APIBase: "https://customproxy.example/v1",
},
},
}
primaryProvider := &mockProvider{}
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, primaryProvider)
if len(agent.Candidates) != 2 {
t.Fatalf("len(Candidates) = %d, want 2", len(agent.Candidates))
}
primary := agent.providerForCandidate(agent.Candidates[0].Provider, agent.Candidates[0].Model)
if primary != primaryProvider {
t.Fatalf("primary candidate should use the primary provider")
}
fallback := agent.providerForCandidate(agent.Candidates[1].Provider, agent.Candidates[1].Model)
if fallback == nil {
t.Fatal("fallback provider is nil")
}
if fallback == primaryProvider {
t.Fatal("fallback alias reused the primary provider")
}
}
func TestNewAgentInstance_UsesConsistentProviderForLoadBalancedFallbackAlias(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-provider-resolution-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Provider: "openai",
Model: "gpt-4.1-mini",
ModelFallbacks: []string{"shared-fallback"},
},
},
ModelList: []config.ModelConfig{
{
ModelName: "shared-fallback",
Model: "gemini/gemini-2.5-flash-lite-preview-06-17",
APIKey: "gemini-key",
},
{
ModelName: "shared-fallback",
Model: "openrouter/google/gemini-2.5-flash-preview",
APIKey: "openrouter-key",
APIBase: "https://openrouter.ai/api/v1",
},
},
}
primaryProvider := &mockProvider{}
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, primaryProvider)
if len(agent.Candidates) != 2 {
t.Fatalf("len(Candidates) = %d, want 2", len(agent.Candidates))
}
fallback := agent.providerForCandidate(agent.Candidates[1].Provider, agent.Candidates[1].Model)
if fallback == nil {
t.Fatal("fallback provider is nil")
}
if fallback == primaryProvider {
t.Fatal("load-balanced fallback alias reused the primary provider")
}
}
func TestNewAgentInstance_UsesDedicatedProviderForLightModelAlias(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-provider-resolution-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Provider: "openai",
Model: "gpt-4.1-mini",
Routing: &config.RoutingConfig{
Enabled: true,
LightModel: "gemini-light",
Threshold: 0.5,
},
},
},
ModelList: []config.ModelConfig{
{
ModelName: "gemini-light",
Model: "gemini/gemini-2.5-flash-lite-preview-06-17",
APIKey: "test-key",
APIBase: "https://customproxy.example/v1",
},
},
}
primaryProvider := &mockProvider{}
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, primaryProvider)
if len(agent.LightCandidates) != 1 {
t.Fatalf("len(LightCandidates) = %d, want 1", len(agent.LightCandidates))
}
lightProvider := agent.providerForCandidate(agent.LightCandidates[0].Provider, agent.LightCandidates[0].Model)
if lightProvider == nil {
t.Fatal("light model provider is nil")
}
if lightProvider == primaryProvider {
t.Fatal("light model alias reused the primary provider")
}
}