diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index 8a6b307e1..cda9be25c 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -1683,10 +1683,19 @@ func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState) (turnResult, er ts.recordPersistedMessage(rootMsg) } + currentTurnResolved := resolveMediaRefs([]providers.Message{ + {Role: "user", Media: append([]string(nil), ts.media...)}, + }, al.mediaStore, maxMediaSize) + var currentTurnMedia []string + if len(currentTurnResolved) > 0 { + currentTurnMedia = currentTurnResolved[0].Media + } + activeCandidates, activeModel, usedLight, useImageFallback := al.selectCandidates( ts.agent, ts.userMessage, messages, + currentTurnMedia, ) activeProvider := ts.agent.Provider if usedLight && ts.agent.LightProvider != nil { @@ -2764,8 +2773,9 @@ func (al *AgentLoop) selectCandidates( agent *AgentInstance, userMsg string, history []providers.Message, + currentTurnMedia []string, ) (candidates []providers.FallbackCandidate, model string, usedLight bool, useImageFallback bool) { - if hasImageMedia(history) && len(agent.ImageCandidates) > 0 { + if hasImageMediaRefs(currentTurnMedia) && len(agent.ImageCandidates) > 0 { logger.InfoCF("agent", "Image model selected", map[string]any{ "agent_id": agent.ID, @@ -2798,12 +2808,10 @@ func (al *AgentLoop) selectCandidates( 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 - } +func hasImageMediaRefs(mediaRefs []string) bool { + for _, ref := range mediaRefs { + if strings.HasPrefix(strings.ToLower(ref), "data:image/") { + return true } } return false diff --git a/pkg/agent/loop_test.go b/pkg/agent/loop_test.go index 6d03954d3..6f79d21df 100644 --- a/pkg/agent/loop_test.go +++ b/pkg/agent/loop_test.go @@ -462,6 +462,7 @@ func TestSelectCandidates_UsesImageModelWhenImageMediaPresent(t *testing.T) { []providers.Message{ {Role: "user", Content: "describe this image", Media: []string{"data:image/png;base64,AAAA"}}, }, + []string{"data:image/png;base64,AAAA"}, ) if usedLight { @@ -481,6 +482,48 @@ func TestSelectCandidates_UsesImageModelWhenImageMediaPresent(t *testing.T) { } } +func TestSelectCandidates_DoesNotUseImageModelForHistoricalImages(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, useImageFallback := (&AgentLoop{}).selectCandidates( + agent, + "text-only follow-up", + []providers.Message{ + {Role: "user", Content: "earlier image", Media: []string{"data:image/png;base64,AAAA"}}, + {Role: "assistant", Content: "previous response"}, + {Role: "user", Content: "current text-only turn"}, + }, + nil, + ) + + if useImageFallback { + t.Fatal("expected image fallback to be disabled for text-only current turn") + } + if model != "gpt-5.4" { + t.Fatalf("model = %q, want %q", model, "gpt-5.4") + } + if len(candidates) != 1 { + t.Fatalf("len(candidates) = %d, want 1", len(candidates)) + } + if candidates[0].Provider != "openai" || candidates[0].Model != "gpt-5.4" { + t.Fatalf("candidate = %+v, want openai/gpt-5.4", candidates[0]) + } +} + func TestRunAgentLoop_UsesImageModelForImageMessages(t *testing.T) { tmpDir := t.TempDir() cfg := &config.Config{ @@ -521,6 +564,62 @@ func TestRunAgentLoop_UsesImageModelForImageMessages(t *testing.T) { } } +func TestRunAgentLoop_DoesNotUseImageModelWhenOnlyHistoryHasImages(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() + + sessionKey := "image-history-session" + _, err := al.runAgentLoop(context.Background(), agent, processOptions{ + SessionKey: sessionKey, + 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("first runAgentLoop returned error: %v", err) + } + if provider.lastModel != "gemini-2.5-flash-lite" { + t.Fatalf("first run provider lastModel = %q, want %q", provider.lastModel, "gemini-2.5-flash-lite") + } + + _, err = al.runAgentLoop(context.Background(), agent, processOptions{ + SessionKey: sessionKey, + Channel: "telegram", + ChatID: "chat-1", + UserMessage: "now answer text only", + DefaultResponse: "fallback", + SendResponse: false, + }) + if err != nil { + t.Fatalf("second runAgentLoop returned error: %v", err) + } + if provider.lastModel != "gpt-5.4" { + t.Fatalf("second run provider lastModel = %q, want %q", provider.lastModel, "gpt-5.4") + } +} + func TestNewAgentLoop_StateInitialized(t *testing.T) { // Create temp workspace tmpDir, err := os.MkdirTemp("", "agent-test-*")