diff --git a/pkg/agent/agent_init.go b/pkg/agent/agent_init.go index 17629892..23fd21b5 100644 --- a/pkg/agent/agent_init.go +++ b/pkg/agent/agent_init.go @@ -296,7 +296,7 @@ func registerSharedTools( // This keeps subagent vision support working even when the optimized // sub-turn spawner path is unavailable. subagentManager.SetMediaResolver(func(msgs []providers.Message) []providers.Message { - return resolveMediaRefs(msgs, al.mediaStore, cfg.Agents.Defaults.GetMaxMediaSize()) + return resolveMediaRefs(msgs, al.mediaStore, cfg.Agents.Defaults.GetMaxMediaSize(), 0) }) // Set the spawner that links into AgentLoop's turnState diff --git a/pkg/agent/agent_media.go b/pkg/agent/agent_media.go index c38a38d4..c15e8c16 100644 --- a/pkg/agent/agent_media.go +++ b/pkg/agent/agent_media.go @@ -31,6 +31,21 @@ var ( filePlaceholderRegex = regexp.MustCompile(`\[file(:\s+[^\]]*)?\]`) ) +func normalizeCurrentTurnStart(messages []providers.Message, currentTurnStart int) int { + if currentTurnStart < 0 { + return 0 + } + if currentTurnStart > len(messages) { + return len(messages) + } + return currentTurnStart +} + +func currentTurnMessages(messages []providers.Message, currentTurnStart int) []providers.Message { + currentTurnStart = normalizeCurrentTurnStart(messages, currentTurnStart) + return messages[currentTurnStart:] +} + // resolveMediaRefs resolves media:// refs in messages. // For user messages: images get path tags only ([image:/path]) so the LLM // can decide whether to view them via load_image or operate on the file. @@ -38,12 +53,20 @@ var ( // user message only after the contiguous tool-message block ends, so we don't // break the tool-results-must-immediately-follow-assistant constraint that // LLM APIs enforce. +// Only tool messages from the current turn may emit the synthetic user +// follow-up; historical tool results stay as plain path-tagged history. // Non-image files always get path tags regardless of role. // Returns a new slice; original messages are not mutated. -func resolveMediaRefs(messages []providers.Message, store media.MediaStore, maxSize int) []providers.Message { +func resolveMediaRefs( + messages []providers.Message, + store media.MediaStore, + maxSize int, + currentTurnStart int, +) []providers.Message { if store == nil { return messages } + currentTurnStart = normalizeCurrentTurnStart(messages, currentTurnStart) result := make([]providers.Message, 0, len(messages)) var pendingToolImages []string @@ -96,7 +119,7 @@ func resolveMediaRefs(messages []providers.Message, store media.MediaStore, maxS mime := detectMIME(localPath, meta) pathTags = append(pathTags, buildPathTag(mime, localPath)) - if m.Role == "tool" && isTurnPromptMessage(m) && strings.HasPrefix(mime, "image/") { + if m.Role == "tool" && idx >= currentTurnStart && strings.HasPrefix(mime, "image/") { dataURL := encodeImageToDataURL(localPath, mime, info, maxSize) if dataURL != "" { pendingToolImages = append(pendingToolImages, dataURL) diff --git a/pkg/agent/agent_test.go b/pkg/agent/agent_test.go index 40c75ade..e12982b2 100644 --- a/pkg/agent/agent_test.go +++ b/pkg/agent/agent_test.go @@ -4558,6 +4558,27 @@ func (p *unexpectedTextAttachmentProvider) GetDefaultModel() string { return "unexpected-text-attachment-model" } +type unexpectedVisionProvider struct { + calls int + models []string +} + +func (p *unexpectedVisionProvider) 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) + return nil, fmt.Errorf("vision provider should not be called for this turn") +} + +func (p *unexpectedVisionProvider) GetDefaultModel() string { + return "unexpected-vision-model" +} + type loadImageThenTextFollowUpProvider struct { path string calls int @@ -4776,6 +4797,164 @@ func TestAgentLoop_UserAttachmentRoutesToImageModelAfterMediaResolution(t *testi } } +func TestAgentLoop_TextFollowUpAfterUserAttachmentStaysOnTextModel(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, + }, + }, + ModelList: []*config.ModelConfig{ + {ModelName: "text-model", Model: "openai/text-model"}, + {ModelName: "vision-model", Model: "openai/vision-model"}, + }, + } + + msgBus := bus.NewMessageBus() + textProvider := &unexpectedTextAttachmentProvider{} + al := NewAgentLoop(cfg, msgBus, textProvider) + + store := media.NewFileMediaStore() + al.SetMediaStore(store) + ref, err := store.Store(pngPath, media.MediaMeta{ContentType: "image/png"}, "test:user-attachment-followup") + if err != nil { + t.Fatalf("Store() error = %v", err) + } + + agent := al.registry.GetDefaultAgent() + if agent == nil { + t.Fatal("expected default agent") + } + + visionProvider := &resolvedImagePathVisionProvider{expectedPath: pngPath} + 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 this image", + Media: []string{ref}, + SessionKey: sessionKey, + })) + if err != nil { + t.Fatalf("first processMessage() error = %v", err) + } + if resp1 != "vision direct answer" { + t.Fatalf("first response = %q, want %q", resp1, "vision direct 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 model response" { + t.Fatalf("second response = %q, want %q", resp2, "text model response") + } + if textProvider.calls != 1 { + t.Fatalf("textProvider calls = %d, want 1", textProvider.calls) + } + if visionProvider.calls != 1 { + t.Fatalf("visionProvider calls = %d, want 1", visionProvider.calls) + } +} + +func TestAgentLoop_GenericImagePlaceholderDoesNotRouteToImageModel(t *testing.T) { + workspace := t.TempDir() + + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{ + Workspace: workspace, + ModelName: "text-model", + ImageModel: "vision-model", + MaxTokens: 4096, + MaxToolIterations: 3, + }, + }, + ModelList: []*config.ModelConfig{ + {ModelName: "text-model", Model: "openai/text-model"}, + {ModelName: "vision-model", Model: "openai/vision-model"}, + }, + } + + msgBus := bus.NewMessageBus() + textProvider := &unexpectedTextAttachmentProvider{} + al := NewAgentLoop(cfg, msgBus, textProvider) + + agent := al.registry.GetDefaultAgent() + if agent == nil { + t.Fatal("expected default agent") + } + + visionProvider := &unexpectedVisionProvider{} + agent.CandidateProviders[providers.ModelKey("openai", "vision-model")] = visionProvider + + timeoutCtx, cancel := context.WithTimeout(context.Background(), responseTimeout) + defer cancel() + + resp, err := al.processMessage(timeoutCtx, testInboundMessage(bus.InboundMessage{ + Context: bus.InboundContext{ + Channel: "telegram", + ChatID: "chat1", + ChatType: "direct", + SenderID: "user1", + MessageID: "m1", + }, + Content: "this is only a placeholder [image: photo], answer in text", + SessionKey: "agent:main:telegram:direct:user1", + })) + if err != nil { + t.Fatalf("processMessage() error = %v", err) + } + if resp != "text model response" { + t.Fatalf("response = %q, want %q", resp, "text model response") + } + if textProvider.calls != 1 { + t.Fatalf("textProvider calls = %d, want 1", textProvider.calls) + } + if visionProvider.calls != 0 { + t.Fatalf("visionProvider calls = %d, want 0", visionProvider.calls) + } +} + func TestAgentLoop_TextFollowUpAfterLoadImageStaysOnTextModel(t *testing.T) { workspace := t.TempDir() pngPath := filepath.Join(workspace, "sample.png") @@ -6264,7 +6443,7 @@ func TestResolveMediaRefs_ImageInjectsPathTag(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: "describe this", Media: []string{ref}}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) if len(result[0].Media) != 0 { t.Fatalf("expected 0 media (images use path tags), got %d", len(result[0].Media)) @@ -6297,7 +6476,7 @@ func TestResolveMediaRefs_ToolRoleImageAppendedAsUserMessage(t *testing.T) { messages := []providers.Message{ toolResultPromptMessage("Image loaded", "call_tool_result_image", []string{ref}), } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) // Tool message should have path tag but no base64 if len(result[0].Media) != 0 { @@ -6342,12 +6521,13 @@ func TestResolveMediaRefs_HistoricalToolRoleImageDoesNotAppendAsUserMessage(t *t ref, _ := store.Store(pngPath, media.MediaMeta{}, "test") messages := []providers.Message{ - {Role: "tool", Content: "Image loaded", Media: []string{ref}}, + toolResultPromptMessage("Image loaded", "call_historical_tool_result_image", []string{ref}), + {Role: "user", Content: "now summarize it in one sentence"}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 1) - if len(result) != 1 { - t.Fatalf("expected only historical tool message to remain, got %d messages", len(result)) + if len(result) != 2 { + t.Fatalf("expected historical tool message plus current follow-up, 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)) @@ -6356,6 +6536,65 @@ func TestResolveMediaRefs_HistoricalToolRoleImageDoesNotAppendAsUserMessage(t *t if !strings.Contains(result[0].Content, "[image:"+localPath+"]") { t.Fatalf("expected image path tag in historical tool content, got %q", result[0].Content) } + if result[1].Role != "user" || result[1].Content != "now summarize it in one sentence" { + t.Fatalf("expected current follow-up user message to remain untouched, got %#v", result[1]) + } +} + +func TestResolveMediaRefs_HistoricalAndCurrentToolImagesOnlyRehydrateCurrentTurn(t *testing.T) { + store := media.NewFileMediaStore() + dir := t.TempDir() + + historicalPath := filepath.Join(dir, "historical.png") + currentPath := filepath.Join(dir, "current.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(historicalPath, pngHeader, 0o644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(currentPath, pngHeader, 0o644); err != nil { + t.Fatal(err) + } + historicalRef, _ := store.Store(historicalPath, media.MediaMeta{}, "test") + currentRef, _ := store.Store(currentPath, media.MediaMeta{}, "test") + + messages := []providers.Message{ + toolResultPromptMessage("Historical image loaded", "call_hist_image", []string{historicalRef}), + {Role: "assistant", Content: "Now I will inspect a new image."}, + toolResultPromptMessage("Current image loaded", "call_current_image", []string{currentRef}), + } + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 1) + + if len(result) != 4 { + t.Fatalf("expected 4 messages (historical tool + assistant + current tool + synthetic user), got %d", len(result)) + } + if result[0].Role != "tool" || len(result[0].Media) != 0 { + t.Fatalf("historical tool message = %#v, want path-tagged tool without media", result[0]) + } + if !strings.Contains(result[0].Content, "[image:"+historicalPath+"]") { + t.Fatalf("expected historical tool content to contain %q, got %q", "[image:"+historicalPath+"]", result[0].Content) + } + if result[1].Role != "assistant" || result[1].Content != "Now I will inspect a new image." { + t.Fatalf("assistant boundary message = %#v, want untouched assistant", result[1]) + } + if result[2].Role != "tool" || len(result[2].Media) != 0 { + t.Fatalf("current tool message = %#v, want path-tagged tool without media", result[2]) + } + if !strings.Contains(result[2].Content, "[image:"+currentPath+"]") { + t.Fatalf("expected current tool content to contain %q, got %q", "[image:"+currentPath+"]", result[2].Content) + } + if result[3].Role != "user" || len(result[3].Media) != 1 { + t.Fatalf("synthetic follow-up = %#v, want one-media user message", result[3]) + } + if !strings.HasPrefix(result[3].Media[0], "data:image/png;base64,") { + t.Fatalf("expected base64 image in synthetic follow-up, got %q", result[3].Media[0][:40]) + } } func TestResolveMediaRefs_MultiToolCallPreservesOrdering(t *testing.T) { @@ -6379,7 +6618,7 @@ func TestResolveMediaRefs_MultiToolCallPreservesOrdering(t *testing.T) { 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) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) // assistant, tool#1, tool#2 must remain contiguous — no user in between if result[0].Role != "assistant" { @@ -6404,6 +6643,52 @@ func TestResolveMediaRefs_MultiToolCallPreservesOrdering(t *testing.T) { } } +func TestResolveMediaRefs_MultipleCurrentToolImagesShareSingleSyntheticFollowUp(t *testing.T) { + store := media.NewFileMediaStore() + dir := t.TempDir() + + firstPath := filepath.Join(dir, "first.png") + secondPath := filepath.Join(dir, "second.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(firstPath, pngHeader, 0o644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(secondPath, pngHeader, 0o644); err != nil { + t.Fatal(err) + } + firstRef, _ := store.Store(firstPath, media.MediaMeta{}, "test") + secondRef, _ := store.Store(secondPath, media.MediaMeta{}, "test") + + messages := []providers.Message{ + {Role: "assistant", Content: "I loaded two images for comparison."}, + toolResultPromptMessage("First image loaded", "call_first_image", []string{firstRef}), + toolResultPromptMessage("Second image loaded", "call_second_image", []string{secondRef}), + } + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) + + if len(result) != 4 { + t.Fatalf("expected assistant + 2 tool results + 1 synthetic follow-up, got %d messages", len(result)) + } + if result[3].Role != "user" { + t.Fatalf("synthetic follow-up role = %q, want user", result[3].Role) + } + if len(result[3].Media) != 2 { + t.Fatalf("synthetic follow-up media count = %d, want 2", len(result[3].Media)) + } + for i, ref := range result[3].Media { + if !strings.HasPrefix(ref, "data:image/png;base64,") { + t.Fatalf("synthetic follow-up media[%d] missing data URL prefix: %q", i, ref[:40]) + } + } +} + func TestResolveMediaRefs_OversizedImageSkipsBase64KeepsPathTag(t *testing.T) { store := media.NewFileMediaStore() dir := t.TempDir() @@ -6421,7 +6706,7 @@ func TestResolveMediaRefs_OversizedImageSkipsBase64KeepsPathTag(t *testing.T) { {Role: "user", Content: "hi", Media: []string{ref}}, } // Use a tiny limit (1KB) so the file is oversized - result := resolveMediaRefs(messages, store, 1024) + result := resolveMediaRefs(messages, store, 1024, 0) if len(result[0].Media) != 0 { t.Fatalf("expected 0 media (oversized), got %d", len(result[0].Media)) @@ -6446,7 +6731,7 @@ func TestResolveMediaRefs_UnknownTypeInjectsPath(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: "hi", Media: []string{ref}}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) if len(result[0].Media) != 0 { t.Fatalf("expected 0 media entries, got %d", len(result[0].Media)) @@ -6461,7 +6746,7 @@ func TestResolveMediaRefs_PassesThroughNonMediaRefs(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: "hi", Media: []string{"https://example.com/img.png"}}, } - result := resolveMediaRefs(messages, nil, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, nil, config.DefaultMaxMediaSize, 0) if len(result[0].Media) != 1 || result[0].Media[0] != "https://example.com/img.png" { t.Fatalf("expected passthrough of non-media:// URL, got %v", result[0].Media) @@ -6486,7 +6771,7 @@ func TestResolveMediaRefs_DoesNotMutateOriginal(t *testing.T) { } originalRef := original[0].Media[0] - resolveMediaRefs(original, store, config.DefaultMaxMediaSize) + resolveMediaRefs(original, store, config.DefaultMaxMediaSize, 0) if original[0].Media[0] != originalRef { t.Fatal("resolveMediaRefs mutated original message slice") @@ -6506,7 +6791,7 @@ func TestResolveMediaRefs_UsesMetaContentType(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: "hi", Media: []string{ref}}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) if len(result[0].Media) != 0 { t.Fatalf("expected 0 media (images use path tags), got %d", len(result[0].Media)) @@ -6530,7 +6815,7 @@ func TestResolveMediaRefs_PDFInjectsFilePath(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: "report.pdf [file]", Media: []string{ref}}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) if len(result[0].Media) != 0 { t.Fatalf("expected 0 media (non-image), got %d", len(result[0].Media)) @@ -6552,7 +6837,7 @@ func TestResolveMediaRefs_AudioInjectsAudioPath(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: "voice.ogg [audio]", Media: []string{ref}}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) if len(result[0].Media) != 0 { t.Fatalf("expected 0 media, got %d", len(result[0].Media)) @@ -6574,7 +6859,7 @@ func TestResolveMediaRefs_VideoInjectsVideoPath(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: "clip.mp4 [video]", Media: []string{ref}}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) if len(result[0].Media) != 0 { t.Fatalf("expected 0 media, got %d", len(result[0].Media)) @@ -6596,7 +6881,7 @@ func TestResolveMediaRefs_NoGenericTagAppendsPath(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: "here is my data", Media: []string{ref}}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) expected := "here is my data [file:" + csvPath + "]" if result[0].Content != expected { @@ -6688,7 +6973,7 @@ func TestResolveMediaRefs_JSONContentPrependsPathTag(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: jsonContent, Media: []string{ref}}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) want := "[image:" + pngPath + "]\n" + jsonContent if result[0].Content != want { @@ -6708,7 +6993,7 @@ func TestResolveMediaRefs_EmptyContentGetsPathTag(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: "", Media: []string{ref}}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) expected := "[file:" + docPath + "]" if result[0].Content != expected { @@ -6737,7 +7022,7 @@ func TestResolveMediaRefs_MixedImageAndFile(t *testing.T) { messages := []providers.Message{ {Role: "user", Content: "check these [file]", Media: []string{imgRef, fileRef}}, } - result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize) + result := resolveMediaRefs(messages, store, config.DefaultMaxMediaSize, 0) if len(result[0].Media) != 0 { t.Fatalf("expected 0 media (all types use path tags), got %d", len(result[0].Media)) diff --git a/pkg/agent/llm_media.go b/pkg/agent/llm_media.go index 92410677..6e27a39e 100644 --- a/pkg/agent/llm_media.go +++ b/pkg/agent/llm_media.go @@ -109,9 +109,6 @@ func sameCandidateSet(a, b []providers.FallbackCandidate) bool { func messagesContainCurrentTurnMediaTurn(messages []providers.Message) bool { for _, msg := range messages { - if !isTurnPromptMessage(msg) { - continue - } if len(msg.Media) > 0 { return true } @@ -124,7 +121,7 @@ func messagesContainCurrentTurnMediaTurn(messages []providers.Message) bool { func (p *Pipeline) routeMediaTurn(ts *turnState, exec *turnExecution) error { if p == nil || ts == nil || ts.agent == nil || exec == nil || - !messagesContainCurrentTurnMediaTurn(exec.callMessages) { + !messagesContainCurrentTurnMediaTurn(currentTurnMessages(exec.callMessages, exec.currentTurnStart)) { return nil } diff --git a/pkg/agent/pipeline_llm.go b/pkg/agent/pipeline_llm.go index 8eebe45d..0961a3d6 100644 --- a/pkg/agent/pipeline_llm.go +++ b/pkg/agent/pipeline_llm.go @@ -31,7 +31,7 @@ func (p *Pipeline) CallLLM( // PreLLM: resolve media refs (except on iteration 1 where user media is already resolved) if iteration > 1 { - exec.messages = resolveMediaRefs(exec.messages, p.MediaStore, maxMediaSize) + exec.messages = resolveMediaRefs(exec.messages, p.MediaStore, maxMediaSize, exec.currentTurnStart) } // PreLLM: graceful terminal handling @@ -388,7 +388,13 @@ func (p *Pipeline) CallLLM( fullHistory := append(append([]providers.Message(nil), trimmedHistory...), protectedTurnTail...) rebuildPromptReq := promptBuildRequestForTurn(ts, fullHistory, exec.summary, "", nil, p.Cfg) rebuildPromptReq.ActiveSkills = append([]string(nil), contextualSkills...) - return ts.agent.ContextBuilder.BuildMessagesFromPrompt(rebuildPromptReq) + rebuilt := ts.agent.ContextBuilder.BuildMessagesFromPrompt(rebuildPromptReq) + return resolveMediaRefs( + rebuilt, + p.MediaStore, + maxMediaSize, + len(rebuilt)-len(protectedTurnTail), + ) } originalHistoryCount := len(exec.history) var fit bool @@ -408,6 +414,7 @@ func (p *Pipeline) CallLLM( ) exec.history = append(trimmedStableHistory, protectedTurnTail...) exec.messages = buildMessages(trimmedStableHistory) + exec.currentTurnStart = len(exec.messages) - len(protectedTurnTail) if exec.gracefulTerminal { msgs := append([]providers.Message(nil), exec.messages...) exec.callMessages = append(msgs, ts.interruptHintMessage()) diff --git a/pkg/agent/pipeline_setup.go b/pkg/agent/pipeline_setup.go index ea05968c..0911a3b6 100644 --- a/pkg/agent/pipeline_setup.go +++ b/pkg/agent/pipeline_setup.go @@ -39,8 +39,12 @@ func (p *Pipeline) SetupTurn(ctx context.Context, ts *turnState) (*turnExecution initialPromptReq := promptBuildRequestForTurn(ts, history, summary, ts.userMessage, ts.media, cfg) initialPromptReq.ActiveSkills = append([]string(nil), contextualSkills...) messages := ts.agent.ContextBuilder.BuildMessagesFromPrompt(initialPromptReq) + currentTurnStart := len(messages) + if strings.TrimSpace(ts.userMessage) != "" || len(ts.media) > 0 { + currentTurnStart = len(messages) - 1 + } - messages = resolveMediaRefs(messages, p.MediaStore, maxMediaSize) + messages = resolveMediaRefs(messages, p.MediaStore, maxMediaSize, currentTurnStart) if !ts.opts.NoHistory { toolDefs := filterToolsByTurnProfile(ts.agent.Tools.ToProviderDefs(), ts.profile) @@ -81,7 +85,11 @@ func (p *Pipeline) SetupTurn(ctx context.Context, ts *turnState) (*turnExecution ) rebuildPromptReq.ActiveSkills = append([]string(nil), contextualSkills...) rebuilt := ts.agent.ContextBuilder.BuildMessagesFromPrompt(rebuildPromptReq) - return resolveMediaRefs(rebuilt, p.MediaStore, maxMediaSize) + rebuiltCurrentTurnStart := len(rebuilt) + if strings.TrimSpace(ts.userMessage) != "" || len(ts.media) > 0 { + rebuiltCurrentTurnStart = len(rebuilt) - 1 + } + return resolveMediaRefs(rebuilt, p.MediaStore, maxMediaSize, rebuiltCurrentTurnStart) }, ts.agent.ContextWindow, toolDefs, @@ -137,6 +145,7 @@ func (p *Pipeline) SetupTurn(ctx context.Context, ts *turnState) (*turnExecution summary, messages, ) + exec.currentTurnStart = currentTurnStart exec.activeCandidates = activeCandidates exec.activeModel = activeModel exec.activeModelConfig = resolveActiveModelConfig( diff --git a/pkg/agent/prompt_turn.go b/pkg/agent/prompt_turn.go index abf58d63..d149b491 100644 --- a/pkg/agent/prompt_turn.go +++ b/pkg/agent/prompt_turn.go @@ -217,10 +217,6 @@ func toolImageFollowUpPromptMessage(media []string) providers.Message { 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) } diff --git a/pkg/agent/turn_coord.go b/pkg/agent/turn_coord.go index 4e14c759..166ca955 100644 --- a/pkg/agent/turn_coord.go +++ b/pkg/agent/turn_coord.go @@ -150,7 +150,7 @@ func (al *AgentLoop) runTurn(ctx context.Context, ts *turnState, pipeline *Pipel // Inject pending steering messages if len(pendingMessages) > 0 { - resolvedPending := resolveMediaRefs(pendingMessages, al.mediaStore, maxMediaSize) + resolvedPending := resolveMediaRefs(pendingMessages, al.mediaStore, maxMediaSize, 0) totalContentLen := 0 for i, pm := range pendingMessages { messages = append(messages, resolvedPending[i]) @@ -431,7 +431,11 @@ func (al *AgentLoop) askSideQuestion( messages := agent.ContextBuilder.BuildMessagesFromPrompt(promptReq) maxMediaSize := al.GetConfig().Agents.Defaults.GetMaxMediaSize() - messages = resolveMediaRefs(messages, al.mediaStore, maxMediaSize) + currentTurnStart := len(messages) + if strings.TrimSpace(question) != "" || len(media) > 0 { + currentTurnStart = len(messages) - 1 + } + messages = resolveMediaRefs(messages, al.mediaStore, maxMediaSize, currentTurnStart) activeCandidates, activeModel, usedLight := al.selectCandidates(agent, question, messages) selectedModelName := sideQuestionModelName(agent, usedLight) diff --git a/pkg/agent/turn_state.go b/pkg/agent/turn_state.go index 5fc4dd40..deb653dc 100644 --- a/pkg/agent/turn_state.go +++ b/pkg/agent/turn_state.go @@ -114,10 +114,11 @@ type ActiveTurnInfo struct { type turnExecution struct { // Core message state (accumulates throughout the turn) - messages []providers.Message // built from ContextBuilder, grows per-iteration - pendingMessages []providers.Message // steering/SubTurn messages awaiting injection - history []providers.Message // from ContextManager.Assemble - summary string + messages []providers.Message // built from ContextBuilder, grows per-iteration + pendingMessages []providers.Message // steering/SubTurn messages awaiting injection + history []providers.Message // from ContextManager.Assemble + summary string + currentTurnStart int // Turn output finalContent string @@ -164,12 +165,13 @@ func newTurnExecution( messages []providers.Message, ) *turnExecution { return &turnExecution{ - history: history, - summary: summary, - messages: messages, - pendingMessages: append([]providers.Message(nil), opts.InitialSteeringMessages...), - iteration: 0, - phase: LLMPhaseSetup, + history: history, + summary: summary, + messages: messages, + pendingMessages: append([]providers.Message(nil), opts.InitialSteeringMessages...), + currentTurnStart: len(messages), + iteration: 0, + phase: LLMPhaseSetup, } }