Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 9 additions & 2 deletions providers/openai/language_model_hooks.go
Original file line number Diff line number Diff line change
Expand Up @@ -297,7 +297,13 @@ func DefaultStreamProviderMetadataFunc(choice openai.ChatCompletionChoice, metad
func DefaultToPrompt(prompt fantasy.Prompt, _, _ string) ([]openai.ChatCompletionMessageParamUnion, []fantasy.CallWarning) {
var messages []openai.ChatCompletionMessageParamUnion
var warnings []fantasy.CallWarning
// All tool replies in a batch must precede synthetic user media.
var pendingToolMedia []openai.ChatCompletionMessageParamUnion
for _, msg := range prompt {
if msg.Role != fantasy.MessageRoleTool {
messages = append(messages, pendingToolMedia...)
pendingToolMedia = nil
}
switch msg.Role {
case fantasy.MessageRoleSystem:
var systemPromptParts []string
Expand Down Expand Up @@ -585,7 +591,8 @@ func DefaultToPrompt(prompt fantasy.Prompt, _, _ string) ([]openai.ChatCompletio
// OpenAI Chat Completions tool messages cannot carry image
// or audio content directly; see ToolResultMediaMessages.
mediaMessages, mediaWarnings := ToolResultMediaMessages(output, toolResultPart.ToolCallID)
messages = append(messages, mediaMessages...)
messages = append(messages, mediaMessages[0])
pendingToolMedia = append(pendingToolMedia, mediaMessages[1:]...)
warnings = append(warnings, mediaWarnings...)
default:
warnings = append(warnings, fantasy.CallWarning{
Expand All @@ -596,7 +603,7 @@ func DefaultToPrompt(prompt fantasy.Prompt, _, _ string) ([]openai.ChatCompletio
}
}
}
return messages, warnings
return append(messages, pendingToolMedia...), warnings
}

// ToolResultMediaMessages maps a tool-result media output to the chat
Expand Down
104 changes: 104 additions & 0 deletions providers/openai/tool_result_media_order_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
package openai_test

import (
"testing"

"charm.land/fantasy"
"charm.land/fantasy/providers/openai"
"charm.land/fantasy/providers/openaicompat"
openaisdk "github.com/openai/openai-go/v3"
"github.com/stretchr/testify/require"
)

func TestToolMediaFollowsToolResultBatch(t *testing.T) {
t.Parallel()

for _, provider := range []struct {
name string
toPrompt func(fantasy.Prompt, string, string) ([]openaisdk.ChatCompletionMessageParamUnion, []fantasy.CallWarning)
}{
{name: "OpenAI", toPrompt: openai.DefaultToPrompt},
{name: "OpenAICompatible", toPrompt: openaicompat.ToPromptFunc},
} {
t.Run(provider.name, func(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
separateMessages bool
secondMedia bool
followup bool
}{
{name: "SeparateMessagesMixed", separateMessages: true},
{name: "SeparateMessagesTwoMedia", separateMessages: true, secondMedia: true},
{name: "SameMessageMixed"},
{name: "SameMessageTwoMedia", secondMedia: true},
{name: "BeforeFollowingMessages", separateMessages: true, secondMedia: true, followup: true},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
media := fantasy.ToolResultOutputContentMedia{Data: "AAEC", MediaType: "image/png", Text: "first result"}
results := []fantasy.MessagePart{
fantasy.ToolResultPart{ToolCallID: "call-a", Output: media},
fantasy.ToolResultPart{ToolCallID: "call-b", Output: fantasy.ToolResultOutputContentText{Text: "second result"}},
}
imageURLs := []string{"data:image/png;base64,AAEC"}
if tc.secondMedia {
media.Data, media.Text = "AwQF", "second result"
results[1] = fantasy.ToolResultPart{ToolCallID: "call-b", Output: media}
imageURLs = append(imageURLs, "data:image/png;base64,AwQF")
}
prompt := fantasy.Prompt{{
Role: fantasy.MessageRoleAssistant,
Content: []fantasy.MessagePart{
fantasy.ToolCallPart{ToolCallID: "call-a", ToolName: "first", Input: "{}"},
fantasy.ToolCallPart{ToolCallID: "call-b", ToolName: "second", Input: "{}"},
},
}}
if tc.separateMessages {
for _, result := range results {
prompt = append(prompt, fantasy.Message{Role: fantasy.MessageRoleTool, Content: []fantasy.MessagePart{result}})
}
} else {
prompt = append(prompt, fantasy.Message{Role: fantasy.MessageRoleTool, Content: results})
}
if tc.followup {
prompt = append(prompt,
fantasy.Message{Role: fantasy.MessageRoleAssistant, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "done"}}},
fantasy.Message{Role: fantasy.MessageRoleUser, Content: []fantasy.MessagePart{fantasy.TextPart{Text: "thanks"}}},
)
}

messages, warnings := provider.toPrompt(prompt, "", "")
require.Empty(t, warnings)
wantLen := 3 + len(imageURLs)
if tc.followup {
wantLen += 2
}
require.Len(t, messages, wantLen)
require.NotNil(t, messages[0].OfAssistant)
require.Len(t, messages[0].OfAssistant.ToolCalls, 2)
for i, want := range []struct{ id, text string }{{"call-a", "first result"}, {"call-b", "second result"}} {
tool := messages[i+1].OfTool
require.NotNil(t, tool, "all tool results must precede synthetic user media")
require.Equal(t, want.id, tool.ToolCallID)
require.Equal(t, want.text, tool.Content.OfString.Value)
}
for i, url := range imageURLs {
user := messages[3+i].OfUser
require.NotNil(t, user)
require.Len(t, user.Content.OfArrayOfContentParts, 1)
image := user.Content.OfArrayOfContentParts[0].OfImageURL
require.NotNil(t, image)
require.Equal(t, url, image.ImageURL.URL)
}
if tc.followup {
require.NotNil(t, messages[wantLen-2].OfAssistant)
require.Equal(t, "done", messages[wantLen-2].OfAssistant.Content.OfString.Value)
require.NotNil(t, messages[wantLen-1].OfUser)
require.Equal(t, "thanks", messages[wantLen-1].OfUser.Content.OfString.Value)
}
})
}
})
}
}
11 changes: 9 additions & 2 deletions providers/openaicompat/language_model_hooks.go
Original file line number Diff line number Diff line change
Expand Up @@ -160,8 +160,14 @@ func ToPromptFunc(prompt fantasy.Prompt, _, _ string) ([]openaisdk.ChatCompletio
var messages []openaisdk.ChatCompletionMessageParamUnion
var warnings []fantasy.CallWarning
hasReasoning := false
// All tool replies in a batch must precede synthetic user media.
var pendingToolMedia []openaisdk.ChatCompletionMessageParamUnion

for _, msg := range prompt {
if msg.Role != fantasy.MessageRoleTool {
messages = append(messages, pendingToolMedia...)
pendingToolMedia = nil
}
switch msg.Role {
case fantasy.MessageRoleSystem:
var blocks []openaisdk.ChatCompletionContentPartTextParam
Expand Down Expand Up @@ -502,7 +508,8 @@ func ToPromptFunc(prompt fantasy.Prompt, _, _ string) ([]openaisdk.ChatCompletio
// helper, which emits a text tool message plus a synthetic
// user message holding the media.
mediaMessages, mediaWarnings := openai.ToolResultMediaMessages(output, toolResultPart.ToolCallID)
messages = append(messages, mediaMessages...)
messages = append(messages, mediaMessages[0])
pendingToolMedia = append(pendingToolMedia, mediaMessages[1:]...)
warnings = append(warnings, mediaWarnings...)
default:
warnings = append(warnings, fantasy.CallWarning{
Expand All @@ -513,7 +520,7 @@ func ToPromptFunc(prompt fantasy.Prompt, _, _ string) ([]openaisdk.ChatCompletio
}
}
}
return messages, warnings
return append(messages, pendingToolMedia...), warnings
}

// toolResultMediaUserPart maps a tool-result media output to an OpenAI chat
Expand Down
Loading