fix(agent): limit media routing to the active turn
This commit is contained in:
parent
58926f76b9
commit
adb89a16b9
6 changed files with 255 additions and 35 deletions
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue