diff --git a/pkg/agent/instance.go b/pkg/agent/instance.go index 0c7baa1ee..6634a73d0 100644 --- a/pkg/agent/instance.go +++ b/pkg/agent/instance.go @@ -39,6 +39,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 @@ -195,6 +197,20 @@ func NewAgentInstance( candidates := providers.ResolveCandidatesWithLookup(modelCfg, defaults.Provider, resolveFromModelList) + var imageCandidates []providers.FallbackCandidate + imageModel := strings.TrimSpace(defaults.ImageModel) + if imageModel != "" { + imageModelCfg := providers.ModelConfig{ + Primary: imageModel, + Fallbacks: defaults.ImageModelFallbacks, + } + imageCandidates = providers.ResolveCandidatesWithLookup(imageModelCfg, defaults.Provider, resolveFromModelList) + if len(imageCandidates) == 0 { + log.Printf("image_model %q not found in model_list — image routing disabled for agent %q", + imageModel, agentID) + } + } + // 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 @@ -234,6 +250,8 @@ func NewAgentInstance( Subagents: subagents, SkillsFilter: skillsFilter, Candidates: candidates, + ImageModel: imageModel, + ImageCandidates: imageCandidates, Router: router, LightCandidates: lightCandidates, } diff --git a/pkg/agent/instance_test.go b/pkg/agent/instance_test.go index 4f41ecd1c..156f8b893 100644 --- a/pkg/agent/instance_test.go +++ b/pkg/agent/instance_test.go @@ -160,3 +160,55 @@ func TestNewAgentInstance_ResolveCandidatesFromModelListAlias(t *testing.T) { }) } } + +func TestNewAgentInstance_ResolvesImageCandidatesFromModelList(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, + Model: "text-default", + ImageModel: "vision-primary", + ImageModelFallbacks: []string{"vision-fallback"}, + }, + }, + ModelList: []config.ModelConfig{ + { + ModelName: "text-default", + Model: "openrouter/example/text-model", + APIBase: "https://openrouter.ai/api/v1", + }, + { + ModelName: "vision-primary", + Model: "gemini/gemini-2.5-flash-lite", + APIBase: "https://generativelanguage.googleapis.com/v1beta/openai/", + }, + { + ModelName: "vision-fallback", + Model: "openrouter/example/vision-fallback", + APIBase: "https://openrouter.ai/api/v1", + }, + }, + } + + provider := &mockProvider{} + agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider) + + if agent.ImageModel != "vision-primary" { + t.Fatalf("ImageModel = %q, want %q", agent.ImageModel, "vision-primary") + } + 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 != "openrouter" || agent.ImageCandidates[1].Model != "example/vision-fallback" { + t.Fatalf("second image candidate = %+v, want openrouter/example/vision-fallback", agent.ImageCandidates[1]) + } +} diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index f20a56b9c..bbdd3de6a 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -1004,7 +1004,7 @@ func (al *AgentLoop) runLLMIteration( // selectCandidates evaluates routing once and the decision is sticky for // all tool-follow-up iterations within the same turn so that a multi-step // tool chain doesn't switch models mid-way through. - activeCandidates, activeModel := al.selectCandidates(agent, opts.UserMessage, messages) + activeCandidates, activeModel, imageTurn := al.selectCandidates(agent, opts.UserMessage, messages) for iteration < agent.MaxIterations { iteration++ @@ -1065,6 +1065,28 @@ func (al *AgentLoop) runLLMIteration( defer al.activeRequests.Done() if len(activeCandidates) > 1 && al.fallback != nil { + if imageTurn { + fbResult, fbErr := al.fallback.ExecuteImage( + ctx, + activeCandidates, + func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) { + return agent.Provider.Chat(ctx, messages, providerToolDefs, model, llmOpts) + }, + ) + if fbErr != nil { + return nil, fbErr + } + if fbResult.Provider != "" && len(fbResult.Attempts) > 0 { + logger.InfoCF( + "agent", + fmt.Sprintf("Image fallback: succeeded with %s/%s after %d attempts", + fbResult.Provider, fbResult.Model, len(fbResult.Attempts)+1), + map[string]any{"agent_id": agent.ID, "iteration": iteration}, + ) + } + return fbResult.Response, nil + } + fbResult, fbErr := al.fallback.Execute( ctx, activeCandidates, @@ -1413,9 +1435,18 @@ func (al *AgentLoop) selectCandidates( agent *AgentInstance, userMsg string, history []providers.Message, -) (candidates []providers.FallbackCandidate, model string) { +) (candidates []providers.FallbackCandidate, model string, imageTurn bool) { + if hasCurrentTurnMedia(history) && len(agent.ImageCandidates) > 0 { + logger.InfoCF("agent", "Model routing: image model selected", + map[string]any{ + "agent_id": agent.ID, + "image_model": agent.ImageModel, + }) + return agent.ImageCandidates, agent.ImageModel, true + } + if agent.Router == nil || len(agent.LightCandidates) == 0 { - return agent.Candidates, agent.Model + return agent.Candidates, agent.Model, false } _, usedLight, score := agent.Router.SelectModel(userMsg, history, agent.Model) @@ -1426,7 +1457,7 @@ func (al *AgentLoop) selectCandidates( "score": score, "threshold": agent.Router.Threshold(), }) - return agent.Candidates, agent.Model + return agent.Candidates, agent.Model, false } logger.InfoCF("agent", "Model routing: light model selected", @@ -1436,7 +1467,17 @@ func (al *AgentLoop) selectCandidates( "score": score, "threshold": agent.Router.Threshold(), }) - return agent.LightCandidates, agent.Router.LightModel() + return agent.LightCandidates, agent.Router.LightModel(), false +} + +func hasCurrentTurnMedia(messages []providers.Message) bool { + for i := len(messages) - 1; i >= 0; i-- { + if messages[i].Role != "user" { + continue + } + return len(messages[i].Media) > 0 + } + 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 a6604e87f..793fa7f39 100644 --- a/pkg/agent/loop_test.go +++ b/pkg/agent/loop_test.go @@ -220,6 +220,84 @@ func TestToolRegistry_GetDefinitions(t *testing.T) { } } +func TestSelectCandidates_PrefersImageModelForMediaTurns(t *testing.T) { + al := &AgentLoop{} + agent := &AgentInstance{ + ID: "agent-1", + Model: "primary-text", + Candidates: []providers.FallbackCandidate{ + {Provider: "openrouter", Model: "example/text"}, + }, + ImageModel: "vision-model", + ImageCandidates: []providers.FallbackCandidate{ + {Provider: "gemini", Model: "gemini-2.5-flash-lite"}, + }, + Router: routing.New(routing.RouterConfig{ + LightModel: "light-model", + Threshold: 0.95, + }), + LightCandidates: []providers.FallbackCandidate{ + {Provider: "openrouter", Model: "example/light"}, + }, + } + + messages := []providers.Message{ + {Role: "system", Content: "system"}, + {Role: "user", Content: "describe this image", Media: []string{"https://example.com/image.png"}}, + } + + candidates, model, imageTurn := al.selectCandidates(agent, "describe this image", messages) + + if !imageTurn { + t.Fatal("expected imageTurn=true for media message") + } + if model != "vision-model" { + t.Fatalf("model = %q, want %q", model, "vision-model") + } + if len(candidates) != 1 || candidates[0].Provider != "gemini" || candidates[0].Model != "gemini-2.5-flash-lite" { + t.Fatalf("unexpected image candidates: %+v", candidates) + } +} + +func TestSelectCandidates_UsesRoutingWhenNoMedia(t *testing.T) { + al := &AgentLoop{} + agent := &AgentInstance{ + ID: "agent-1", + Model: "primary-text", + Candidates: []providers.FallbackCandidate{ + {Provider: "openrouter", Model: "example/text"}, + }, + ImageModel: "vision-model", + ImageCandidates: []providers.FallbackCandidate{ + {Provider: "gemini", Model: "gemini-2.5-flash-lite"}, + }, + Router: routing.New(routing.RouterConfig{ + LightModel: "light-model", + Threshold: 0.95, + }), + LightCandidates: []providers.FallbackCandidate{ + {Provider: "openrouter", Model: "example/light"}, + }, + } + + messages := []providers.Message{ + {Role: "system", Content: "system"}, + {Role: "user", Content: "hi"}, + } + + candidates, model, imageTurn := al.selectCandidates(agent, "hi", messages) + + if imageTurn { + t.Fatal("expected imageTurn=false for text-only message") + } + if model != "light-model" { + t.Fatalf("model = %q, want %q", model, "light-model") + } + if len(candidates) != 1 || candidates[0].Model != "example/light" { + t.Fatalf("unexpected routed candidates: %+v", candidates) + } +} + // TestAgentLoop_GetStartupInfo verifies startup info contains tools func TestAgentLoop_GetStartupInfo(t *testing.T) { tmpDir, err := os.MkdirTemp("", "agent-test-*")