fix: strip extra_content from tool calls for non-Google providers
This commit is contained in:
parent
86da6a7d56
commit
8ec55a50a8
2 changed files with 26 additions and 6 deletions
|
|
@ -117,7 +117,7 @@ func (p *Provider) Chat(
|
||||||
|
|
||||||
requestBody := map[string]any{
|
requestBody := map[string]any{
|
||||||
"model": model,
|
"model": model,
|
||||||
"messages": serializeMessages(messages),
|
"messages": serializeMessages(messages, p.apiBase),
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(tools) > 0 {
|
if len(tools) > 0 {
|
||||||
|
|
@ -401,13 +401,33 @@ type openaiMessage struct {
|
||||||
ToolCallID string `json:"tool_call_id,omitempty"`
|
ToolCallID string `json:"tool_call_id,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func supportsExtraContent(apiBase string) bool {
|
||||||
|
u, err := url.Parse(apiBase)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
host := u.Hostname()
|
||||||
|
return strings.HasPrefix(host, "generativelanguage") || strings.HasSuffix(host, ".googleapis.com")
|
||||||
|
}
|
||||||
|
|
||||||
// serializeMessages converts internal Message structs to the OpenAI wire format.
|
// serializeMessages converts internal Message structs to the OpenAI wire format.
|
||||||
// - Strips SystemParts (unknown to third-party endpoints)
|
// - Strips SystemParts (unknown to third-party endpoints)
|
||||||
// - Converts messages with Media to multipart content format (text + image_url parts)
|
// - Converts messages with Media to multipart content format (text + image_url parts)
|
||||||
// - Preserves ToolCallID, ToolCalls, and ReasoningContent for all messages
|
// - Preserves ToolCallID, ToolCalls, and ReasoningContent for all messages
|
||||||
func serializeMessages(messages []Message) []any {
|
func serializeMessages(messages []Message, apiBase string) []any {
|
||||||
|
supportsExtra := supportsExtraContent(apiBase)
|
||||||
|
|
||||||
out := make([]any, 0, len(messages))
|
out := make([]any, 0, len(messages))
|
||||||
for _, m := range messages {
|
for _, m := range messages {
|
||||||
|
if len(m.ToolCalls) > 0 && !supportsExtra {
|
||||||
|
cleanCalls := make([]ToolCall, len(m.ToolCalls))
|
||||||
|
for i, tc := range m.ToolCalls {
|
||||||
|
tc.ExtraContent = nil
|
||||||
|
cleanCalls[i] = tc
|
||||||
|
}
|
||||||
|
m.ToolCalls = cleanCalls
|
||||||
|
}
|
||||||
|
|
||||||
if len(m.Media) == 0 {
|
if len(m.Media) == 0 {
|
||||||
out = append(out, openaiMessage{
|
out = append(out, openaiMessage{
|
||||||
Role: m.Role,
|
Role: m.Role,
|
||||||
|
|
|
||||||
|
|
@ -648,7 +648,7 @@ func TestSerializeMessages_PlainText(t *testing.T) {
|
||||||
{Role: "user", Content: "hello"},
|
{Role: "user", Content: "hello"},
|
||||||
{Role: "assistant", Content: "hi", ReasoningContent: "thinking..."},
|
{Role: "assistant", Content: "hi", ReasoningContent: "thinking..."},
|
||||||
}
|
}
|
||||||
result := serializeMessages(messages)
|
result := serializeMessages(messages, "")
|
||||||
|
|
||||||
data, err := json.Marshal(result)
|
data, err := json.Marshal(result)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|
@ -670,7 +670,7 @@ func TestSerializeMessages_WithMedia(t *testing.T) {
|
||||||
messages := []protocoltypes.Message{
|
messages := []protocoltypes.Message{
|
||||||
{Role: "user", Content: "describe this", Media: []string{"data:image/png;base64,abc123"}},
|
{Role: "user", Content: "describe this", Media: []string{"data:image/png;base64,abc123"}},
|
||||||
}
|
}
|
||||||
result := serializeMessages(messages)
|
result := serializeMessages(messages, "")
|
||||||
|
|
||||||
data, _ := json.Marshal(result)
|
data, _ := json.Marshal(result)
|
||||||
var msgs []map[string]any
|
var msgs []map[string]any
|
||||||
|
|
@ -703,7 +703,7 @@ func TestSerializeMessages_MediaWithToolCallID(t *testing.T) {
|
||||||
messages := []protocoltypes.Message{
|
messages := []protocoltypes.Message{
|
||||||
{Role: "tool", Content: "image result", Media: []string{"data:image/png;base64,xyz"}, ToolCallID: "call_1"},
|
{Role: "tool", Content: "image result", Media: []string{"data:image/png;base64,xyz"}, ToolCallID: "call_1"},
|
||||||
}
|
}
|
||||||
result := serializeMessages(messages)
|
result := serializeMessages(messages, "")
|
||||||
|
|
||||||
data, _ := json.Marshal(result)
|
data, _ := json.Marshal(result)
|
||||||
var msgs []map[string]any
|
var msgs []map[string]any
|
||||||
|
|
@ -833,7 +833,7 @@ func TestSerializeMessages_StripsSystemParts(t *testing.T) {
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
result := serializeMessages(messages)
|
result := serializeMessages(messages, "")
|
||||||
|
|
||||||
data, _ := json.Marshal(result)
|
data, _ := json.Marshal(result)
|
||||||
raw := string(data)
|
raw := string(data)
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue