diff --git a/pkg/agent/agent_media.go b/pkg/agent/agent_media.go index c02c7392..c38a38d4 100644 --- a/pkg/agent/agent_media.go +++ b/pkg/agent/agent_media.go @@ -52,22 +52,14 @@ func resolveMediaRefs(messages []providers.Message, store media.MediaStore, maxS // When leaving a tool-message block, flush any accumulated images // as a synthetic user message. if m.Role != "tool" && len(pendingToolImages) > 0 { - result = append(result, providers.Message{ - Role: "user", - Content: "[Loaded image from tool result above]", - Media: pendingToolImages, - }) + result = append(result, toolImageFollowUpPromptMessage(pendingToolImages)) pendingToolImages = nil } if len(m.Media) == 0 { result = append(result, m) if idx == len(messages)-1 && len(pendingToolImages) > 0 { - result = append(result, providers.Message{ - Role: "user", - Content: "[Loaded image from tool result above]", - Media: pendingToolImages, - }) + result = append(result, toolImageFollowUpPromptMessage(pendingToolImages)) pendingToolImages = nil } continue @@ -104,7 +96,7 @@ func resolveMediaRefs(messages []providers.Message, store media.MediaStore, maxS mime := detectMIME(localPath, meta) pathTags = append(pathTags, buildPathTag(mime, localPath)) - if m.Role == "tool" && strings.HasPrefix(mime, "image/") { + if m.Role == "tool" && isTurnPromptMessage(m) && strings.HasPrefix(mime, "image/") { dataURL := encodeImageToDataURL(localPath, mime, info, maxSize) if dataURL != "" { pendingToolImages = append(pendingToolImages, dataURL) @@ -120,11 +112,7 @@ func resolveMediaRefs(messages []providers.Message, store media.MediaStore, maxS // If this is the last message and we have pending images, flush them. if idx == len(messages)-1 && len(pendingToolImages) > 0 { - result = append(result, providers.Message{ - Role: "user", - Content: "[Loaded image from tool result above]", - Media: pendingToolImages, - }) + result = append(result, toolImageFollowUpPromptMessage(pendingToolImages)) pendingToolImages = nil } } diff --git a/pkg/agent/agent_test.go b/pkg/agent/agent_test.go index a4f651e6..40c75ade 100644 --- a/pkg/agent/agent_test.go +++ b/pkg/agent/agent_test.go @@ -4558,6 +4558,71 @@ func (p *unexpectedTextAttachmentProvider) GetDefaultModel() string { return "unexpected-text-attachment-model" } +type loadImageThenTextFollowUpProvider struct { + path string + calls int + models []string + mediaSeen []bool + syntheticMsgSeen []bool +} + +func (p *loadImageThenTextFollowUpProvider) Chat( + ctx context.Context, + messages []providers.Message, + tools []providers.ToolDefinition, + model string, + opts map[string]any, +) (*providers.LLMResponse, error) { + p.calls++ + p.models = append(p.models, model) + + hasMedia := false + hasSynthetic := false + for _, msg := range messages { + if msg.Content == "[Loaded image from tool result above]" { + hasSynthetic = true + } + for _, ref := range msg.Media { + if strings.TrimSpace(ref) != "" { + hasMedia = true + break + } + } + } + p.mediaSeen = append(p.mediaSeen, hasMedia) + p.syntheticMsgSeen = append(p.syntheticMsgSeen, hasSynthetic) + + switch p.calls { + case 1: + return &providers.LLMResponse{ + Content: "Let me inspect the image.", + ToolCalls: []providers.ToolCall{{ + ID: "call_load_image_regression", + Type: "function", + Name: "load_image", + Arguments: map[string]any{"path": p.path}, + }}, + }, nil + case 2: + if hasMedia { + return nil, fmt.Errorf("text-only follow-up unexpectedly retained media from prior turn") + } + if hasSynthetic { + return nil, fmt.Errorf("text-only follow-up unexpectedly rebuilt synthetic tool image message from history") + } + return &providers.LLMResponse{ + Content: "text follow-up", + ToolCalls: []providers.ToolCall{}, + }, nil + default: + return nil, fmt.Errorf("unexpected extra text-model call %d", p.calls) + } +} + +func (p *loadImageThenTextFollowUpProvider) GetDefaultModel() string { + return "load-image-then-text-follow-up-model" +} + func TestAgentLoop_VisionUnsupportedErrorReturnsClearFailure(t *testing.T) { workspace := t.TempDir() @@ -4711,6 +4776,113 @@ func TestAgentLoop_UserAttachmentRoutesToImageModelAfterMediaResolution(t *testi } } +func TestAgentLoop_TextFollowUpAfterLoadImageStaysOnTextModel(t *testing.T) { + workspace := t.TempDir() + pngPath := filepath.Join(workspace, "sample.png") + pngBytes := []byte{ + 0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, + 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, 0x44, 0x52, + 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x02, + 0x00, 0x00, 0x00, 0x90, 0x77, 0x53, 0xDE, + } + if err := os.WriteFile(pngPath, pngBytes, 0o644); err != nil { + t.Fatalf("WriteFile() error = %v", err) + } + + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{ + Workspace: workspace, + ModelName: "text-model", + ImageModel: "vision-model", + MaxTokens: 4096, + MaxToolIterations: 3, + }, + }, + Tools: config.ToolsConfig{ + LoadImage: config.ToolConfig{Enabled: true}, + }, + ModelList: []*config.ModelConfig{ + {ModelName: "text-model", Model: "openai/text-model"}, + {ModelName: "vision-model", Model: "openai/vision-model"}, + }, + } + + msgBus := bus.NewMessageBus() + textProvider := &loadImageThenTextFollowUpProvider{path: pngPath} + al := NewAgentLoop(cfg, msgBus, textProvider) + al.SetMediaStore(media.NewFileMediaStore()) + + agent := al.registry.GetDefaultAgent() + if agent == nil { + t.Fatal("expected default agent") + } + if len(agent.ImageCandidates) != 1 { + t.Fatalf("len(ImageCandidates) = %d, want 1", len(agent.ImageCandidates)) + } + + visionProvider := &visionAnswerProvider{} + agent.CandidateProviders[providers.ModelKey("openai", "vision-model")] = visionProvider + + sessionKey := "agent:main:telegram:direct:user1" + timeoutCtx, cancel := context.WithTimeout(context.Background(), responseTimeout) + defer cancel() + + resp1, err := al.processMessage(timeoutCtx, testInboundMessage(bus.InboundMessage{ + Context: bus.InboundContext{ + Channel: "telegram", + ChatID: "chat1", + ChatType: "direct", + SenderID: "user1", + MessageID: "m1", + }, + Content: "describe the image you load", + SessionKey: sessionKey, + })) + if err != nil { + t.Fatalf("first processMessage() error = %v", err) + } + if resp1 != "vision answer" { + t.Fatalf("first response = %q, want %q", resp1, "vision answer") + } + + timeoutCtx2, cancel2 := context.WithTimeout(context.Background(), responseTimeout) + defer cancel2() + + resp2, err := al.processMessage(timeoutCtx2, testInboundMessage(bus.InboundMessage{ + Context: bus.InboundContext{ + Channel: "telegram", + ChatID: "chat1", + ChatType: "direct", + SenderID: "user1", + MessageID: "m2", + }, + Content: "now summarize it in one sentence", + SessionKey: sessionKey, + })) + if err != nil { + t.Fatalf("second processMessage() error = %v", err) + } + if resp2 != "text follow-up" { + t.Fatalf("second response = %q, want %q", resp2, "text follow-up") + } + if textProvider.calls != 2 { + t.Fatalf("textProvider calls = %d, want %d", textProvider.calls, 2) + } + if !slices.Equal(textProvider.models, []string{"text-model", "text-model"}) { + t.Fatalf("textProvider models = %v, want %v", textProvider.models, []string{"text-model", "text-model"}) + } + if !slices.Equal(textProvider.mediaSeen, []bool{false, false}) { + t.Fatalf("textProvider mediaSeen = %v, want %v", textProvider.mediaSeen, []bool{false, false}) + } + if !slices.Equal(textProvider.syntheticMsgSeen, []bool{false, false}) { + t.Fatalf("textProvider syntheticMsgSeen = %v, want %v", textProvider.syntheticMsgSeen, []bool{false, false}) + } + if visionProvider.calls != 1 { + t.Fatalf("visionProvider calls = %d, want %d", visionProvider.calls, 1) + } +} + func TestAgentLoop_LoadImageFollowUpRoutesToImageModel(t *testing.T) { workspace := t.TempDir() pngPath := filepath.Join(workspace, "sample.png") @@ -6123,7 +6295,7 @@ func TestResolveMediaRefs_ToolRoleImageAppendedAsUserMessage(t *testing.T) { ref, _ := store.Store(pngPath, media.MediaMeta{}, "test") messages := []providers.Message{ - {Role: "tool", Content: "Image loaded", Media: []string{ref}}, + toolResultPromptMessage("Image loaded", "call_tool_result_image", []string{ref}), } result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) @@ -6151,6 +6323,41 @@ func TestResolveMediaRefs_ToolRoleImageAppendedAsUserMessage(t *testing.T) { } } +func TestResolveMediaRefs_HistoricalToolRoleImageDoesNotAppendAsUserMessage(t *testing.T) { + store := media.NewFileMediaStore() + dir := t.TempDir() + + pngPath := filepath.Join(dir, "historical-tool-result.png") + pngHeader := []byte{ + 0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, + 0x00, 0x00, 0x00, 0x0D, + 0x49, 0x48, 0x44, 0x52, + 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x02, + 0x00, 0x00, 0x00, + 0x90, 0x77, 0x53, 0xDE, + } + if err := os.WriteFile(pngPath, pngHeader, 0o644); err != nil { + t.Fatal(err) + } + ref, _ := store.Store(pngPath, media.MediaMeta{}, "test") + + messages := []providers.Message{ + {Role: "tool", Content: "Image loaded", Media: []string{ref}}, + } + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + + if len(result) != 1 { + t.Fatalf("expected only historical tool message to remain, got %d messages", len(result)) + } + if len(result[0].Media) != 0 { + t.Fatalf("expected 0 media in historical tool message, got %d", len(result[0].Media)) + } + localPath, _, _ := store.ResolveWithMeta(ref) + if !strings.Contains(result[0].Content, "[image:"+localPath+"]") { + t.Fatalf("expected image path tag in historical tool content, got %q", result[0].Content) + } +} + func TestResolveMediaRefs_MultiToolCallPreservesOrdering(t *testing.T) { store := media.NewFileMediaStore() dir := t.TempDir() @@ -6169,8 +6376,8 @@ func TestResolveMediaRefs_MultiToolCallPreservesOrdering(t *testing.T) { // Simulate: assistant called load_image + read_file, two tool results follow messages := []providers.Message{ {Role: "assistant", Content: "Let me load the image and read the file."}, - {Role: "tool", Content: "Image loaded [image: photo]", Media: []string{imgRef}}, - {Role: "tool", Content: "file contents here"}, + toolResultPromptMessage("Image loaded [image: photo]", "call_load_image_multi_tool", []string{imgRef}), + toolResultPromptMessage("file contents here", "call_read_file_multi_tool", nil), } result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) diff --git a/pkg/agent/llm_media.go b/pkg/agent/llm_media.go index 54e1802f..92410677 100644 --- a/pkg/agent/llm_media.go +++ b/pkg/agent/llm_media.go @@ -107,8 +107,11 @@ func sameCandidateSet(a, b []providers.FallbackCandidate) bool { return true } -func messagesContainMediaTurn(messages []providers.Message) bool { +func messagesContainCurrentTurnMediaTurn(messages []providers.Message) bool { for _, msg := range messages { + if !isTurnPromptMessage(msg) { + continue + } if len(msg.Media) > 0 { return true } @@ -120,7 +123,8 @@ func messagesContainMediaTurn(messages []providers.Message) bool { } func (p *Pipeline) routeMediaTurn(ts *turnState, exec *turnExecution) error { - if p == nil || ts == nil || ts.agent == nil || exec == nil || !messagesContainMediaTurn(exec.callMessages) { + if p == nil || ts == nil || ts.agent == nil || exec == nil || + !messagesContainCurrentTurnMediaTurn(exec.callMessages) { return nil } diff --git a/pkg/agent/pipeline_execute.go b/pkg/agent/pipeline_execute.go index 81fc06ec..59e33317 100644 --- a/pkg/agent/pipeline_execute.go +++ b/pkg/agent/pipeline_execute.go @@ -295,21 +295,16 @@ toolLoop: contentForLLM = al.cfg.FilterSensitiveData(contentForLLM) } - toolResultMsg := providers.Message{ - Role: "tool", - Content: contentForLLM, - ToolCallID: tc.ID, - } - + var toolResultMedia []string if len(hookResult.Media) > 0 && !hookResult.ResponseHandled { hookResult.ArtifactTags = buildArtifactTags(al.mediaStore, hookResult.Media) contentForLLM = hookResult.ContentForLLM() if al.cfg.Tools.IsFilterSensitiveDataEnabled() { contentForLLM = al.cfg.FilterSensitiveData(contentForLLM) } - toolResultMsg.Content = contentForLLM - toolResultMsg.Media = append(toolResultMsg.Media, hookResult.Media...) + toolResultMedia = append(toolResultMedia, hookResult.Media...) } + toolResultMsg := toolResultPromptMessage(contentForLLM, tc.ID, toolResultMedia) al.emitEvent( runtimeevents.KindAgentToolExecEnd, @@ -695,14 +690,11 @@ toolLoop: contentForLLM = al.cfg.FilterSensitiveData(contentForLLM) } - toolResultMsg := providers.Message{ - Role: "tool", - Content: contentForLLM, - ToolCallID: toolCallID, - } + var toolResultMedia []string if len(toolResult.Media) > 0 && !toolResult.ResponseHandled { - toolResultMsg.Media = append(toolResultMsg.Media, toolResult.Media...) + toolResultMedia = append(toolResultMedia, toolResult.Media...) } + toolResultMsg := toolResultPromptMessage(contentForLLM, toolCallID, toolResultMedia) al.emitEvent( runtimeevents.KindAgentToolExecEnd, ts.eventMeta("runTurn", "turn.tool.end"), diff --git a/pkg/agent/prompt.go b/pkg/agent/prompt.go index 1c66fe69..6270dc32 100644 --- a/pkg/agent/prompt.go +++ b/pkg/agent/prompt.go @@ -37,6 +37,7 @@ const ( PromptSlotMessage PromptSlot = "message" PromptSlotSteering PromptSlot = "steering" PromptSlotSubTurn PromptSlot = "subturn" + PromptSlotToolResult PromptSlot = "tool_result" PromptSlotInterrupt PromptSlot = "interrupt" PromptSlotOutput PromptSlot = "output" ) @@ -60,6 +61,7 @@ const ( PromptSourceUserMessage PromptSourceID = "turn:user_message" PromptSourceSteering PromptSourceID = "turn:steering" PromptSourceSubTurnResult PromptSourceID = "turn:subturn_result" + PromptSourceToolResult PromptSourceID = "turn:tool_result" PromptSourceInterrupt PromptSourceID = "turn:interrupt" ) diff --git a/pkg/agent/prompt_turn.go b/pkg/agent/prompt_turn.go index 053eb846..abf58d63 100644 --- a/pkg/agent/prompt_turn.go +++ b/pkg/agent/prompt_turn.go @@ -194,6 +194,33 @@ func userPromptMessage(content string, media []string) providers.Message { return promptMessageWithMetadata(msg, PromptLayerTurn, PromptSlotMessage, PromptSourceUserMessage) } +func toolResultPromptMessage(content, toolCallID string, media []string) providers.Message { + msg := providers.Message{ + Role: "tool", + Content: content, + ToolCallID: toolCallID, + } + if len(media) > 0 { + msg.Media = append([]string(nil), media...) + } + return promptMessageWithMetadata(msg, PromptLayerTurn, PromptSlotToolResult, PromptSourceToolResult) +} + +func toolImageFollowUpPromptMessage(media []string) providers.Message { + msg := providers.Message{ + Role: "user", + Content: "[Loaded image from tool result above]", + } + if len(media) > 0 { + msg.Media = append([]string(nil), media...) + } + return promptMessageWithMetadata(msg, PromptLayerTurn, PromptSlotToolResult, PromptSourceToolResult) +} + +func isTurnPromptMessage(msg providers.Message) bool { + return msg.PromptLayer == string(PromptLayerTurn) +} + func steeringPromptMessage(msg providers.Message) providers.Message { return promptMessageWithDefaultMetadata(msg, PromptLayerTurn, PromptSlotSteering, PromptSourceSteering) }