diff --git a/pkg/providers/openai_compat/provider.go b/pkg/providers/openai_compat/provider.go index 9ab6c800b..4bd3b99f2 100644 --- a/pkg/providers/openai_compat/provider.go +++ b/pkg/providers/openai_compat/provider.go @@ -126,7 +126,7 @@ func (p *Provider) Chat( requestBody := map[string]any{ "model": model, - "messages": serializeMessages(messages), + "messages": serializeMessages(messages, p.apiBase), } if len(tools) > 0 { @@ -415,13 +415,33 @@ type openaiMessage struct { 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. // - Strips SystemParts (unknown to third-party endpoints) // - Converts messages with Media to multipart content format (text + image_url parts) // - 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)) 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 { out = append(out, openaiMessage{ Role: m.Role, diff --git a/pkg/providers/openai_compat/provider_test.go b/pkg/providers/openai_compat/provider_test.go index 41f278a1b..358832aad 100644 --- a/pkg/providers/openai_compat/provider_test.go +++ b/pkg/providers/openai_compat/provider_test.go @@ -648,7 +648,7 @@ func TestSerializeMessages_PlainText(t *testing.T) { {Role: "user", Content: "hello"}, {Role: "assistant", Content: "hi", ReasoningContent: "thinking..."}, } - result := serializeMessages(messages) + result := serializeMessages(messages, "") data, err := json.Marshal(result) if err != nil { @@ -670,7 +670,7 @@ func TestSerializeMessages_WithMedia(t *testing.T) { messages := []protocoltypes.Message{ {Role: "user", Content: "describe this", Media: []string{"data:image/png;base64,abc123"}}, } - result := serializeMessages(messages) + result := serializeMessages(messages, "") data, _ := json.Marshal(result) var msgs []map[string]any @@ -703,7 +703,7 @@ func TestSerializeMessages_MediaWithToolCallID(t *testing.T) { messages := []protocoltypes.Message{ {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) 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) raw := string(data)