Skip to content
Open
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
58 changes: 56 additions & 2 deletions service/relayconvert/internal/oai_chat/to_claude_messages_req.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,56 @@ type openRouterRequestReasoning struct {
Exclude bool `json:"exclude,omitempty"`
}

func convertOpenAIToolResultContentToClaude(c *gin.Context, message dto.Message) (any, error) {
contentItems, ok := message.Content.([]any)
if !ok {
return message.Content, nil
}

convertedContent := make([]any, 0, len(contentItems))
convertedAnyImage := false
for _, contentItem := range contentItems {
contentMap, ok := contentItem.(map[string]any)
if !ok || contentMap["type"] != dto.ContentTypeImageURL {
convertedContent = append(convertedContent, contentItem)
continue
}

mediaMessages := (&dto.Message{Content: []any{contentItem}}).ParseContent()
if len(mediaMessages) != 1 {
return nil, fmt.Errorf("invalid image_url tool result content")
}
source := mediaMessages[0].ToFileSource()
if source == nil {
return nil, fmt.Errorf("invalid image_url tool result source")
}

base64Data, mimeType, err := relaymedia.ResolveBase64Data(c, source, "formatting tool result image for Claude")
if err != nil {
return nil, fmt.Errorf("get tool result file data failed: %s", err.Error())
}

contentType := "image"
if strings.HasPrefix(mimeType, "application/pdf") {
contentType = "document"
}
convertedContent = append(convertedContent, dto.ClaudeMediaMessage{
Type: contentType,
Source: &dto.ClaudeMessageSource{
Type: "base64",
MediaType: mimeType,
Data: base64Data,
},
})
convertedAnyImage = true
}

if !convertedAnyImage {
return message.Content, nil
}
return convertedContent, nil
}

func OpenAIChatRequestToClaudeMessages(c *gin.Context, textRequest dto.GeneralOpenAIRequest) (*dto.ClaudeRequest, error) {
claudeTools := make([]any, 0, len(textRequest.Tools))

Expand Down Expand Up @@ -299,6 +349,10 @@ func OpenAIChatRequestToClaudeMessages(c *gin.Context, textRequest dto.GeneralOp
Role: message.Role,
}
if message.Role == "tool" {
toolResultContent, err := convertOpenAIToolResultContentToClaude(c, message)
if err != nil {
return nil, err
}
if len(claudeMessages) > 0 && claudeMessages[len(claudeMessages)-1].Role == "user" {
lastClaudeMessage := claudeMessages[len(claudeMessages)-1]
if content, ok := lastClaudeMessage.Content.(string); ok {
Expand All @@ -312,7 +366,7 @@ func OpenAIChatRequestToClaudeMessages(c *gin.Context, textRequest dto.GeneralOp
lastClaudeMessage.Content = append(lastClaudeMessage.Content.([]dto.ClaudeMediaMessage), dto.ClaudeMediaMessage{
Type: "tool_result",
ToolUseId: message.ToolCallId,
Content: message.Content,
Content: toolResultContent,
})
claudeMessages[len(claudeMessages)-1] = lastClaudeMessage
continue
Expand All @@ -323,7 +377,7 @@ func OpenAIChatRequestToClaudeMessages(c *gin.Context, textRequest dto.GeneralOp
{
Type: "tool_result",
ToolUseId: message.ToolCallId,
Content: message.Content,
Content: toolResultContent,
},
}
} else if message.IsStringContent() && message.ToolCalls == nil {
Expand Down
106 changes: 106 additions & 0 deletions service/relayconvert/internal/oai_chat/to_claude_messages_req_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
package oaichat

import (
"encoding/json"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/dto"
relaymedia "github.com/QuantumNous/new-api/service/relayconvert/internal/media"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

func TestOpenAIChatRequestToClaudeMessagesConvertsToolResultImageURL(t *testing.T) {
relaymedia.SetMediaResolver(relaymedia.MediaResolver{
GetBase64Data: func(_ *gin.Context, source types.FileSource, _ ...string) (string, string, error) {
assert.Equal(t, "https://example.com/tool-result.png", source.GetRawData())
return "resolved-image-data", "image/png", nil
},
})
t.Cleanup(func() {
relaymedia.SetMediaResolver(relaymedia.MediaResolver{})
})

request := toolResultConversionRequest([]any{
map[string]any{"type": "text", "text": "Image result:"},
map[string]any{
"type": dto.ContentTypeImageURL,
"image_url": map[string]any{
"url": "https://example.com/tool-result.png",
},
},
})

claudeRequest, err := OpenAIChatRequestToClaudeMessages(nil, request)
require.NoError(t, err)
require.Len(t, claudeRequest.Messages, 3)

toolResultBlocks, ok := claudeRequest.Messages[2].Content.([]dto.ClaudeMediaMessage)
require.True(t, ok)
require.Len(t, toolResultBlocks, 1)
assert.Equal(t, "tool_result", toolResultBlocks[0].Type)
assert.Equal(t, "call_1", toolResultBlocks[0].ToolUseId)

toolContent, ok := toolResultBlocks[0].Content.([]any)
require.True(t, ok)
require.Len(t, toolContent, 2)
assert.Equal(t, map[string]any{"type": "text", "text": "Image result:"}, toolContent[0])

imageBlock, ok := toolContent[1].(dto.ClaudeMediaMessage)
require.True(t, ok)
assert.Equal(t, "image", imageBlock.Type)
require.NotNil(t, imageBlock.Source)
assert.Equal(t, "base64", imageBlock.Source.Type)
assert.Equal(t, "image/png", imageBlock.Source.MediaType)
assert.Equal(t, "resolved-image-data", imageBlock.Source.Data)

payload, err := common.Marshal(claudeRequest)
require.NoError(t, err)
assert.NotContains(t, string(payload), `"image_url"`)
assert.Contains(t, string(payload), `"type":"image"`)
}

func TestOpenAIChatRequestToClaudeMessagesPreservesNonImageURLToolContent(t *testing.T) {
original := []any{
map[string]any{"type": "text", "text": "plain result"},
map[string]any{
"type": "image",
"source": map[string]any{
"type": "base64",
"media_type": "image/png",
"data": "already-converted",
},
},
map[string]any{"type": "custom", "value": "preserve me"},
}

claudeRequest, err := OpenAIChatRequestToClaudeMessages(nil, toolResultConversionRequest(original))
require.NoError(t, err)
require.Len(t, claudeRequest.Messages, 3)

toolResultBlocks, ok := claudeRequest.Messages[2].Content.([]dto.ClaudeMediaMessage)
require.True(t, ok)
require.Len(t, toolResultBlocks, 1)
assert.Equal(t, original, toolResultBlocks[0].Content)
}

func toolResultConversionRequest(toolContent any) dto.GeneralOpenAIRequest {
return dto.GeneralOpenAIRequest{
Model: "test-model",
Messages: []dto.Message{
{Role: "user", Content: "Use read_image."},
{
Role: "assistant",
ToolCalls: json.RawMessage(`[{"id":"call_1","type":"function","function":{"name":"read_image","arguments":"{}"}}]`),
},
{
Role: "tool",
ToolCallId: "call_1",
Content: toolContent,
},
},
}
}