diff --git a/pkg/agent/instance.go b/pkg/agent/instance.go index 880725660..5a9fc0836 100644 --- a/pkg/agent/instance.go +++ b/pkg/agent/instance.go @@ -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, diff --git a/pkg/agent/instance_test.go b/pkg/agent/instance_test.go index e296a18cb..030411405 100644 --- a/pkg/agent/instance_test.go +++ b/pkg/agent/instance_test.go @@ -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() diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index ef2951365..8a6b307e1 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -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. diff --git a/pkg/agent/loop_test.go b/pkg/agent/loop_test.go index 25d20c689..6d03954d3 100644 --- a/pkg/agent/loop_test.go +++ b/pkg/agent/loop_test.go @@ -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-*") diff --git a/pkg/agent/mock_provider_test.go b/pkg/agent/mock_provider_test.go index 4962810dc..07facc4c3 100644 --- a/pkg/agent/mock_provider_test.go +++ b/pkg/agent/mock_provider_test.go @@ -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{},