fix(agent): route image turns through image-model fallback
This commit is contained in:
parent
e70928cc6f
commit
d750bfcca5
5 changed files with 180 additions and 12 deletions
|
|
@ -40,6 +40,8 @@ type AgentInstance struct {
|
|||
Subagents *config.SubagentsConfig
|
||||
SkillsFilter []string
|
||||
Candidates []providers.FallbackCandidate
|
||||
ImageModel string
|
||||
ImageCandidates []providers.FallbackCandidate
|
||||
|
||||
// Router is non-nil when model routing is configured and the light model
|
||||
// was successfully resolved. It scores each incoming message and decides
|
||||
|
|
@ -170,6 +172,12 @@ func NewAgentInstance(
|
|||
// Resolve fallback candidates
|
||||
candidates := resolveModelCandidates(cfg, defaults.Provider, model, fallbacks)
|
||||
|
||||
imageModel := strings.TrimSpace(defaults.ImageModel)
|
||||
var imageCandidates []providers.FallbackCandidate
|
||||
if imageModel != "" {
|
||||
imageCandidates = resolveModelCandidates(cfg, defaults.Provider, imageModel, defaults.ImageModelFallbacks)
|
||||
}
|
||||
|
||||
// Model routing setup: pre-resolve light model candidates at creation time
|
||||
// to avoid repeated model_list lookups on every incoming message.
|
||||
var router *routing.Router
|
||||
|
|
@ -222,6 +230,8 @@ func NewAgentInstance(
|
|||
Subagents: subagents,
|
||||
SkillsFilter: skillsFilter,
|
||||
Candidates: candidates,
|
||||
ImageModel: imageModel,
|
||||
ImageCandidates: imageCandidates,
|
||||
Router: router,
|
||||
LightCandidates: lightCandidates,
|
||||
LightProvider: lightProvider,
|
||||
|
|
|
|||
|
|
@ -165,6 +165,50 @@ func TestNewAgentInstance_ResolveCandidatesFromModelListAlias(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestNewAgentInstance_ResolveImageCandidatesFromModelListAlias(t *testing.T) {
|
||||
tmpDir, err := os.MkdirTemp("", "agent-instance-test-*")
|
||||
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,
|
||||
ModelName: "text-main",
|
||||
ImageModel: "vision-main",
|
||||
ImageModelFallbacks: []string{"vision-backup"},
|
||||
},
|
||||
},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{
|
||||
ModelName: "vision-main",
|
||||
Model: "gemini/gemini-2.5-flash-lite",
|
||||
},
|
||||
{
|
||||
ModelName: "vision-backup",
|
||||
Model: "anthropic/claude-3-7-sonnet",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, &mockProvider{})
|
||||
|
||||
if agent.ImageModel != "vision-main" {
|
||||
t.Fatalf("ImageModel = %q, want %q", agent.ImageModel, "vision-main")
|
||||
}
|
||||
if len(agent.ImageCandidates) != 2 {
|
||||
t.Fatalf("len(ImageCandidates) = %d, want 2", len(agent.ImageCandidates))
|
||||
}
|
||||
if agent.ImageCandidates[0].Provider != "gemini" || agent.ImageCandidates[0].Model != "gemini-2.5-flash-lite" {
|
||||
t.Fatalf("first image candidate = %+v, want gemini/gemini-2.5-flash-lite", agent.ImageCandidates[0])
|
||||
}
|
||||
if agent.ImageCandidates[1].Provider != "anthropic" || agent.ImageCandidates[1].Model != "claude-3-7-sonnet" {
|
||||
t.Fatalf("second image candidate = %+v, want anthropic/claude-3-7-sonnet", agent.ImageCandidates[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewAgentInstance_AllowsMediaTempDirForReadListAndExec(t *testing.T) {
|
||||
workspace := t.TempDir()
|
||||
mediaDir := media.TempDir()
|
||||
|
|
|
|||
|
|
@ -1683,7 +1683,11 @@ func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState) (turnResult, er
|
|||
ts.recordPersistedMessage(rootMsg)
|
||||
}
|
||||
|
||||
activeCandidates, activeModel, usedLight := al.selectCandidates(ts.agent, ts.userMessage, messages)
|
||||
activeCandidates, activeModel, usedLight, useImageFallback := al.selectCandidates(
|
||||
ts.agent,
|
||||
ts.userMessage,
|
||||
messages,
|
||||
)
|
||||
activeProvider := ts.agent.Provider
|
||||
if usedLight && ts.agent.LightProvider != nil {
|
||||
activeProvider = ts.agent.LightProvider
|
||||
|
|
@ -1905,13 +1909,19 @@ turnLoop:
|
|||
defer al.activeRequests.Done()
|
||||
|
||||
if len(activeCandidates) > 1 && al.fallback != nil {
|
||||
fbResult, fbErr := al.fallback.Execute(
|
||||
providerCtx,
|
||||
activeCandidates,
|
||||
func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) {
|
||||
return activeProvider.Chat(ctx, messagesForCall, toolDefsForCall, model, llmOpts)
|
||||
},
|
||||
runCandidate := func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) {
|
||||
return activeProvider.Chat(ctx, messagesForCall, toolDefsForCall, model, llmOpts)
|
||||
}
|
||||
|
||||
var (
|
||||
fbResult *providers.FallbackResult
|
||||
fbErr error
|
||||
)
|
||||
if useImageFallback {
|
||||
fbResult, fbErr = al.fallback.ExecuteImage(providerCtx, activeCandidates, runCandidate)
|
||||
} else {
|
||||
fbResult, fbErr = al.fallback.Execute(providerCtx, activeCandidates, runCandidate)
|
||||
}
|
||||
if fbErr != nil {
|
||||
return nil, fbErr
|
||||
}
|
||||
|
|
@ -2754,9 +2764,17 @@ func (al *AgentLoop) selectCandidates(
|
|||
agent *AgentInstance,
|
||||
userMsg string,
|
||||
history []providers.Message,
|
||||
) (candidates []providers.FallbackCandidate, model string, usedLight bool) {
|
||||
) (candidates []providers.FallbackCandidate, model string, usedLight bool, useImageFallback bool) {
|
||||
if hasImageMedia(history) && len(agent.ImageCandidates) > 0 {
|
||||
logger.InfoCF("agent", "Image model selected",
|
||||
map[string]any{
|
||||
"agent_id": agent.ID,
|
||||
"image_model": agent.ImageModel,
|
||||
})
|
||||
return agent.ImageCandidates, resolvedCandidateModel(agent.ImageCandidates, agent.ImageModel), false, true
|
||||
}
|
||||
if agent.Router == nil || len(agent.LightCandidates) == 0 {
|
||||
return agent.Candidates, resolvedCandidateModel(agent.Candidates, agent.Model), false
|
||||
return agent.Candidates, resolvedCandidateModel(agent.Candidates, agent.Model), false, false
|
||||
}
|
||||
|
||||
_, usedLight, score := agent.Router.SelectModel(userMsg, history, agent.Model)
|
||||
|
|
@ -2767,7 +2785,7 @@ func (al *AgentLoop) selectCandidates(
|
|||
"score": score,
|
||||
"threshold": agent.Router.Threshold(),
|
||||
})
|
||||
return agent.Candidates, resolvedCandidateModel(agent.Candidates, agent.Model), false
|
||||
return agent.Candidates, resolvedCandidateModel(agent.Candidates, agent.Model), false, false
|
||||
}
|
||||
|
||||
logger.InfoCF("agent", "Model routing: light model selected",
|
||||
|
|
@ -2777,7 +2795,18 @@ func (al *AgentLoop) selectCandidates(
|
|||
"score": score,
|
||||
"threshold": agent.Router.Threshold(),
|
||||
})
|
||||
return agent.LightCandidates, resolvedCandidateModel(agent.LightCandidates, agent.Router.LightModel()), true
|
||||
return agent.LightCandidates, resolvedCandidateModel(agent.LightCandidates, agent.Router.LightModel()), true, false
|
||||
}
|
||||
|
||||
func hasImageMedia(messages []providers.Message) bool {
|
||||
for _, msg := range messages {
|
||||
for _, ref := range msg.Media {
|
||||
if strings.HasPrefix(strings.ToLower(ref), "data:image/") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// maybeSummarize triggers summarization if the session history exceeds thresholds.
|
||||
|
|
|
|||
|
|
@ -439,6 +439,88 @@ func TestRecordLastChatID(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestSelectCandidates_UsesImageModelWhenImageMediaPresent(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
cfg := &config.Config{
|
||||
Agents: config.AgentsConfig{
|
||||
Defaults: config.AgentDefaults{
|
||||
Workspace: tmpDir,
|
||||
ModelName: "text-main",
|
||||
ImageModel: "vision-main",
|
||||
},
|
||||
},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{ModelName: "text-main", Model: "openai/gpt-5.4"},
|
||||
{ModelName: "vision-main", Model: "gemini/gemini-2.5-flash-lite"},
|
||||
},
|
||||
}
|
||||
|
||||
agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, &mockProvider{})
|
||||
candidates, model, usedLight, useImageFallback := (&AgentLoop{}).selectCandidates(
|
||||
agent,
|
||||
"describe this image",
|
||||
[]providers.Message{
|
||||
{Role: "user", Content: "describe this image", Media: []string{"data:image/png;base64,AAAA"}},
|
||||
},
|
||||
)
|
||||
|
||||
if usedLight {
|
||||
t.Fatal("did not expect light-model routing for image-only selection")
|
||||
}
|
||||
if !useImageFallback {
|
||||
t.Fatal("expected image fallback to be selected")
|
||||
}
|
||||
if model != "gemini-2.5-flash-lite" {
|
||||
t.Fatalf("model = %q, want %q", model, "gemini-2.5-flash-lite")
|
||||
}
|
||||
if len(candidates) != 1 {
|
||||
t.Fatalf("len(candidates) = %d, want 1", len(candidates))
|
||||
}
|
||||
if candidates[0].Provider != "gemini" || candidates[0].Model != "gemini-2.5-flash-lite" {
|
||||
t.Fatalf("candidate = %+v, want gemini/gemini-2.5-flash-lite", candidates[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunAgentLoop_UsesImageModelForImageMessages(t *testing.T) {
|
||||
tmpDir := t.TempDir()
|
||||
cfg := &config.Config{
|
||||
Agents: config.AgentsConfig{
|
||||
Defaults: config.AgentDefaults{
|
||||
Workspace: tmpDir,
|
||||
ModelName: "text-main",
|
||||
ImageModel: "vision-main",
|
||||
MaxTokens: 4096,
|
||||
MaxToolIterations: 5,
|
||||
},
|
||||
},
|
||||
ModelList: []*config.ModelConfig{
|
||||
{ModelName: "text-main", Model: "openai/gpt-5.4"},
|
||||
{ModelName: "vision-main", Model: "gemini/gemini-2.5-flash-lite"},
|
||||
},
|
||||
}
|
||||
|
||||
msgBus := bus.NewMessageBus()
|
||||
provider := &mockProvider{}
|
||||
al := NewAgentLoop(cfg, msgBus, provider)
|
||||
agent := al.registry.GetDefaultAgent()
|
||||
|
||||
_, err := al.runAgentLoop(context.Background(), agent, processOptions{
|
||||
SessionKey: "image-session",
|
||||
Channel: "telegram",
|
||||
ChatID: "chat-1",
|
||||
UserMessage: "describe this image",
|
||||
Media: []string{"data:image/png;base64,AAAA"},
|
||||
DefaultResponse: "fallback",
|
||||
SendResponse: false,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("runAgentLoop returned error: %v", err)
|
||||
}
|
||||
if provider.lastModel != "gemini-2.5-flash-lite" {
|
||||
t.Fatalf("provider lastModel = %q, want %q", provider.lastModel, "gemini-2.5-flash-lite")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewAgentLoop_StateInitialized(t *testing.T) {
|
||||
// Create temp workspace
|
||||
tmpDir, err := os.MkdirTemp("", "agent-test-*")
|
||||
|
|
|
|||
|
|
@ -6,7 +6,9 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
type mockProvider struct{}
|
||||
type mockProvider struct {
|
||||
lastModel string
|
||||
}
|
||||
|
||||
func (m *mockProvider) Chat(
|
||||
ctx context.Context,
|
||||
|
|
@ -15,6 +17,7 @@ func (m *mockProvider) Chat(
|
|||
model string,
|
||||
opts map[string]any,
|
||||
) (*providers.LLMResponse, error) {
|
||||
m.lastModel = model
|
||||
return &providers.LLMResponse{
|
||||
Content: "Mock response",
|
||||
ToolCalls: []providers.ToolCall{},
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue