mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 11:33:18 +08:00
fix(gemini): preserve image output in chat completions
This commit is contained in:
@@ -473,7 +473,7 @@ func geminiResponseToChatCompletions(
|
||||
rawData []byte,
|
||||
usageOverride *ClaudeUsage,
|
||||
) (*apicompat.ChatCompletionsResponse, *ClaudeUsage, error) {
|
||||
claudeRespMap, usage := convertGeminiToClaudeMessage(geminiResp, originalModel, rawData)
|
||||
claudeRespMap, usage := convertGeminiToClaudeMessage(geminiResp, originalModel, rawData, true)
|
||||
if usageOverride != nil && (usageOverride.InputTokens > 0 || usageOverride.OutputTokens > 0 || usageOverride.CacheReadInputTokens > 0) {
|
||||
usage = usageOverride
|
||||
if usageMap, ok := claudeRespMap["usage"].(map[string]any); ok {
|
||||
|
||||
@@ -0,0 +1,157 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestGeminiResponseToChatCompletionsPreservesInlineData(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
parts []any
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "image only",
|
||||
parts: []any{
|
||||
map[string]any{"inlineData": map[string]any{"mimeType": "image/png", "data": "aW1hZ2U="}},
|
||||
},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "text and image",
|
||||
parts: []any{
|
||||
map[string]any{"text": "rendered image:\n"},
|
||||
map[string]any{"inlineData": map[string]any{"mimeType": "image/webp", "data": "d2VicA=="}},
|
||||
},
|
||||
want: "rendered image:\n",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
geminiResp := map[string]any{
|
||||
"candidates": []any{map[string]any{
|
||||
"content": map[string]any{"parts": tt.parts},
|
||||
"finishReason": "STOP",
|
||||
}},
|
||||
}
|
||||
rawData, err := json.Marshal(geminiResp)
|
||||
require.NoError(t, err)
|
||||
|
||||
got, _, err := geminiResponseToChatCompletions(geminiResp, "gemini-test", rawData, nil)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, got.Choices, 1)
|
||||
|
||||
var content string
|
||||
require.NoError(t, json.Unmarshal(got.Choices[0].Message.Content, &content))
|
||||
require.Equal(t, tt.want, content)
|
||||
require.Equal(t, "stop", got.Choices[0].FinishReason)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGeminiResponseToChatCompletionsOmitsInvalidInlineData(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
inlineData map[string]any
|
||||
}{
|
||||
{
|
||||
name: "unsupported MIME type",
|
||||
inlineData: map[string]any{"mimeType": "image/svg+xml", "data": "PHN2Zz48L3N2Zz4="},
|
||||
},
|
||||
{
|
||||
name: "malformed MIME type",
|
||||
inlineData: map[string]any{"mimeType": "image/png; charset=utf-8", "data": "aW1hZ2U="},
|
||||
},
|
||||
{
|
||||
name: "malformed base64",
|
||||
inlineData: map[string]any{"mimeType": "image/png", "data": "not-valid-base64!!!"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
geminiResp := map[string]any{
|
||||
"candidates": []any{map[string]any{
|
||||
"content": map[string]any{"parts": []any{map[string]any{"text": "before"}, map[string]any{"inlineData": tt.inlineData}, map[string]any{"text": "after"}}},
|
||||
"finishReason": "STOP",
|
||||
}},
|
||||
}
|
||||
rawData, err := json.Marshal(geminiResp)
|
||||
require.NoError(t, err)
|
||||
|
||||
got, _, err := geminiResponseToChatCompletions(geminiResp, "gemini-test", rawData, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
var content string
|
||||
require.NoError(t, json.Unmarshal(got.Choices[0].Message.Content, &content))
|
||||
require.Equal(t, "beforeafter", content)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConvertGeminiToClaudeMessageOmitsInlineDataForAnthropicMessages(t *testing.T) {
|
||||
geminiResp := map[string]any{
|
||||
"candidates": []any{map[string]any{
|
||||
"content": map[string]any{"parts": []any{
|
||||
map[string]any{"text": "before"},
|
||||
map[string]any{"inlineData": map[string]any{"mimeType": "image/png", "data": "aW1hZ2U="}},
|
||||
map[string]any{"functionCall": map[string]any{"name": "get_weather", "args": map[string]any{"city": "Paris"}}},
|
||||
map[string]any{"text": "after"},
|
||||
}},
|
||||
"finishReason": "STOP",
|
||||
}},
|
||||
}
|
||||
rawData, err := json.Marshal(geminiResp)
|
||||
require.NoError(t, err)
|
||||
|
||||
withInlineData, _ := convertGeminiToClaudeMessage(geminiResp, "gemini-test", rawData, true)
|
||||
contentWithInlineData := withInlineData["content"].([]any)
|
||||
require.Len(t, contentWithInlineData, 4)
|
||||
require.Equal(t, map[string]any{"type": "text", "text": "before"}, contentWithInlineData[0])
|
||||
require.Equal(t, map[string]any{"type": "text", "text": ""}, contentWithInlineData[1])
|
||||
require.Equal(t, "tool_use", contentWithInlineData[2].(map[string]any)["type"])
|
||||
require.Equal(t, "get_weather", contentWithInlineData[2].(map[string]any)["name"])
|
||||
require.Equal(t, map[string]any{"type": "text", "text": "after"}, contentWithInlineData[3])
|
||||
|
||||
withoutInlineData, _ := convertGeminiToClaudeMessage(geminiResp, "gemini-test", rawData, false)
|
||||
contentWithoutInlineData := withoutInlineData["content"].([]any)
|
||||
require.Len(t, contentWithoutInlineData, 3)
|
||||
require.Equal(t, map[string]any{"type": "text", "text": "before"}, contentWithoutInlineData[0])
|
||||
require.Equal(t, "tool_use", contentWithoutInlineData[1].(map[string]any)["type"])
|
||||
require.Equal(t, "get_weather", contentWithoutInlineData[1].(map[string]any)["name"])
|
||||
require.Equal(t, map[string]any{"type": "text", "text": "after"}, contentWithoutInlineData[2])
|
||||
}
|
||||
|
||||
func TestGeminiResponseToChatCompletionsRetainsTextAndToolBehavior(t *testing.T) {
|
||||
geminiResp := map[string]any{
|
||||
"candidates": []any{map[string]any{
|
||||
"content": map[string]any{"parts": []any{
|
||||
map[string]any{"text": "checking"},
|
||||
map[string]any{"functionCall": map[string]any{
|
||||
"name": "get_weather",
|
||||
"args": map[string]any{"city": "Paris"},
|
||||
}},
|
||||
}},
|
||||
"finishReason": "STOP",
|
||||
}},
|
||||
}
|
||||
rawData, err := json.Marshal(geminiResp)
|
||||
require.NoError(t, err)
|
||||
|
||||
got, _, err := geminiResponseToChatCompletions(geminiResp, "gemini-test", rawData, nil)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, got.Choices, 1)
|
||||
|
||||
choice := got.Choices[0]
|
||||
var content string
|
||||
require.NoError(t, json.Unmarshal(choice.Message.Content, &content))
|
||||
require.Equal(t, "checking", content)
|
||||
require.Equal(t, "tool_calls", choice.FinishReason)
|
||||
require.Len(t, choice.Message.ToolCalls, 1)
|
||||
require.Equal(t, "get_weather", choice.Message.ToolCalls[0].Function.Name)
|
||||
require.JSONEq(t, `{"city":"Paris"}`, choice.Message.ToolCalls[0].Function.Arguments)
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
@@ -1068,7 +1069,7 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
|
||||
return nil, s.writeClaudeError(c, http.StatusBadGateway, "upstream_error", "Failed to read upstream stream")
|
||||
}
|
||||
collectedBytes, _ := json.Marshal(collected)
|
||||
claudeResp, usageObj2 := convertGeminiToClaudeMessage(collected, originalModel, collectedBytes)
|
||||
claudeResp, usageObj2 := convertGeminiToClaudeMessage(collected, originalModel, collectedBytes, false)
|
||||
c.JSON(http.StatusOK, claudeResp)
|
||||
usage = usageObj2
|
||||
if usageObj != nil && (usageObj.InputTokens > 0 || usageObj.OutputTokens > 0) {
|
||||
@@ -1965,7 +1966,7 @@ func (s *GeminiMessagesCompatService) handleNonStreamingResponse(c *gin.Context,
|
||||
return nil, s.writeClaudeError(c, http.StatusBadGateway, "upstream_error", "Failed to parse upstream response")
|
||||
}
|
||||
|
||||
claudeResp, usage := convertGeminiToClaudeMessage(geminiResp, originalModel, unwrappedBody)
|
||||
claudeResp, usage := convertGeminiToClaudeMessage(geminiResp, originalModel, unwrappedBody, false)
|
||||
c.JSON(http.StatusOK, claudeResp)
|
||||
|
||||
return usage, nil
|
||||
@@ -2717,7 +2718,7 @@ func unwrapGeminiResponse(raw []byte) ([]byte, error) {
|
||||
return raw, nil
|
||||
}
|
||||
|
||||
func convertGeminiToClaudeMessage(geminiResp map[string]any, originalModel string, rawData []byte) (map[string]any, *ClaudeUsage) {
|
||||
func convertGeminiToClaudeMessage(geminiResp map[string]any, originalModel string, rawData []byte, includeInlineData bool) (map[string]any, *ClaudeUsage) {
|
||||
usage := extractGeminiUsage(rawData)
|
||||
if usage == nil {
|
||||
usage = &ClaudeUsage{}
|
||||
@@ -2740,6 +2741,16 @@ func convertGeminiToClaudeMessage(geminiResp map[string]any, originalModel strin
|
||||
"text": text,
|
||||
})
|
||||
}
|
||||
if inlineData, ok := pm["inlineData"].(map[string]any); includeInlineData && ok {
|
||||
mimeType, _ := inlineData["mimeType"].(string)
|
||||
data, _ := inlineData["data"].(string)
|
||||
if isGeminiInlineImageMIMEType(mimeType) && isValidBase64(data) {
|
||||
contentBlocks = append(contentBlocks, map[string]any{
|
||||
"type": "text",
|
||||
"text": fmt.Sprintf("", mimeType, data),
|
||||
})
|
||||
}
|
||||
}
|
||||
if fc, ok := pm["functionCall"].(map[string]any); ok {
|
||||
name, _ := fc["name"].(string)
|
||||
if strings.TrimSpace(name) == "" {
|
||||
@@ -2782,6 +2793,23 @@ func convertGeminiToClaudeMessage(geminiResp map[string]any, originalModel strin
|
||||
return resp, usage
|
||||
}
|
||||
|
||||
func isGeminiInlineImageMIMEType(mimeType string) bool {
|
||||
switch mimeType {
|
||||
case "image/gif", "image/jpeg", "image/png", "image/webp":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func isValidBase64(data string) bool {
|
||||
if data == "" {
|
||||
return false
|
||||
}
|
||||
_, err := base64.StdEncoding.DecodeString(data)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
func extractGeminiUsage(data []byte) *ClaudeUsage {
|
||||
usage := gjson.GetBytes(data, "usageMetadata")
|
||||
if !usage.Exists() {
|
||||
|
||||
Reference in New Issue
Block a user