fix(agent): use resolved model names for direct calls

This commit is contained in:
duomi 2026-03-15 20:41:20 +08:00
parent 0f700a6bf0
commit a1fdab0806
3 changed files with 70 additions and 5 deletions

View file

@ -1415,7 +1415,7 @@ func (al *AgentLoop) selectCandidates(
history []providers.Message, history []providers.Message,
) (candidates []providers.FallbackCandidate, model string) { ) (candidates []providers.FallbackCandidate, model string) {
if agent.Router == nil || len(agent.LightCandidates) == 0 { if agent.Router == nil || len(agent.LightCandidates) == 0 {
return agent.Candidates, agent.Model return agent.Candidates, resolvedDirectModel(agent.Model, agent.Candidates)
} }
_, usedLight, score := agent.Router.SelectModel(userMsg, history, agent.Model) _, usedLight, score := agent.Router.SelectModel(userMsg, history, agent.Model)
@ -1426,7 +1426,7 @@ func (al *AgentLoop) selectCandidates(
"score": score, "score": score,
"threshold": agent.Router.Threshold(), "threshold": agent.Router.Threshold(),
}) })
return agent.Candidates, agent.Model return agent.Candidates, resolvedDirectModel(agent.Model, agent.Candidates)
} }
logger.InfoCF("agent", "Model routing: light model selected", logger.InfoCF("agent", "Model routing: light model selected",
@ -1436,7 +1436,14 @@ func (al *AgentLoop) selectCandidates(
"score": score, "score": score,
"threshold": agent.Router.Threshold(), "threshold": agent.Router.Threshold(),
}) })
return agent.LightCandidates, agent.Router.LightModel() return agent.LightCandidates, resolvedDirectModel(agent.Router.LightModel(), agent.LightCandidates)
}
func resolvedDirectModel(raw string, candidates []providers.FallbackCandidate) string {
if len(candidates) == 1 && strings.TrimSpace(candidates[0].Model) != "" {
return candidates[0].Model
}
return raw
} }
// maybeSummarize triggers summarization if the session history exceeds thresholds. // maybeSummarize triggers summarization if the session history exceeds thresholds.
@ -1728,7 +1735,7 @@ func (al *AgentLoop) retryLLMCall(
ctx, ctx,
[]providers.Message{{Role: "user", Content: prompt}}, []providers.Message{{Role: "user", Content: prompt}},
nil, nil,
agent.Model, resolvedDirectModel(agent.Model, agent.Candidates),
map[string]any{ map[string]any{
"max_tokens": agent.MaxTokens, "max_tokens": agent.MaxTokens,
"temperature": llmTemperature, "temperature": llmTemperature,

View file

@ -220,6 +220,51 @@ func TestToolRegistry_GetDefinitions(t *testing.T) {
} }
} }
func TestSelectCandidates_UsesResolvedSingleCandidateModel(t *testing.T) {
al := &AgentLoop{}
agent := &AgentInstance{
ID: "agent-1",
Model: "nemotron-3-nano-30b-a3b",
Candidates: []providers.FallbackCandidate{
{Provider: "openrouter", Model: "nvidia/nemotron-3-nano-30b-a3b:free"},
},
}
candidates, model := al.selectCandidates(agent, "hello", []providers.Message{{Role: "user", Content: "hello"}})
if len(candidates) != 1 {
t.Fatalf("len(candidates) = %d, want 1", len(candidates))
}
if model != "nvidia/nemotron-3-nano-30b-a3b:free" {
t.Fatalf("model = %q, want resolved model", model)
}
}
func TestRetryLLMCall_UsesResolvedSingleCandidateModel(t *testing.T) {
provider := &mockProvider{}
al := &AgentLoop{}
agent := &AgentInstance{
ID: "agent-1",
Model: "nemotron-3-nano-30b-a3b",
MaxTokens: 1024,
Provider: provider,
Candidates: []providers.FallbackCandidate{
{Provider: "openrouter", Model: "nvidia/nemotron-3-nano-30b-a3b:free"},
},
}
resp, err := al.retryLLMCall(context.Background(), agent, "hello", 1)
if err != nil {
t.Fatalf("retryLLMCall() error = %v", err)
}
if resp == nil {
t.Fatal("retryLLMCall() response = nil, want non-nil")
}
if got := provider.LastModel(); got != "nvidia/nemotron-3-nano-30b-a3b:free" {
t.Fatalf("provider model = %q, want resolved model", got)
}
}
// TestAgentLoop_GetStartupInfo verifies startup info contains tools // TestAgentLoop_GetStartupInfo verifies startup info contains tools
func TestAgentLoop_GetStartupInfo(t *testing.T) { func TestAgentLoop_GetStartupInfo(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-test-*") tmpDir, err := os.MkdirTemp("", "agent-test-*")

View file

@ -2,11 +2,15 @@ package agent
import ( import (
"context" "context"
"sync"
"github.com/sipeed/picoclaw/pkg/providers" "github.com/sipeed/picoclaw/pkg/providers"
) )
type mockProvider struct{} type mockProvider struct {
mu sync.Mutex
lastModel string
}
func (m *mockProvider) Chat( func (m *mockProvider) Chat(
ctx context.Context, ctx context.Context,
@ -15,6 +19,9 @@ func (m *mockProvider) Chat(
model string, model string,
opts map[string]any, opts map[string]any,
) (*providers.LLMResponse, error) { ) (*providers.LLMResponse, error) {
m.mu.Lock()
m.lastModel = model
m.mu.Unlock()
return &providers.LLMResponse{ return &providers.LLMResponse{
Content: "Mock response", Content: "Mock response",
ToolCalls: []providers.ToolCall{}, ToolCalls: []providers.ToolCall{},
@ -24,3 +31,9 @@ func (m *mockProvider) Chat(
func (m *mockProvider) GetDefaultModel() string { func (m *mockProvider) GetDefaultModel() string {
return "mock-model" return "mock-model"
} }
func (m *mockProvider) LastModel() string {
m.mu.Lock()
defer m.mu.Unlock()
return m.lastModel
}