mirror of
https://github.com/alibaba/higress.git
synced 2026-07-24 05:10:34 +08:00
fix(ai-proxy): fix cache token double-counting in Claude-to-OpenAI protocol conversion (#4149)
Signed-off-by: Xijun Dai <daixijun1990@gmail.com> Co-authored-by: 澄潭 <zty98751@alibaba-inc.com> Co-authored-by: Kent Dong <ch3cho@qq.com>
This commit is contained in:
@@ -353,6 +353,26 @@ func (c *ClaudeToOpenAIConverter) ConvertClaudeRequestToOpenAIWithOptions(body [
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// computeClaudeInputTokens computes Claude-compatible input_tokens from OpenAI usage.
|
||||
//
|
||||
// In OpenAI's API, prompt_tokens includes cached_tokens (subset relationship).
|
||||
// In Claude's API, input_tokens should NOT include cache tokens.
|
||||
// We detect the OpenAI-standard semantics by checking if total_tokens == prompt_tokens + completion_tokens.
|
||||
// For providers like Bedrock where prompt_tokens does NOT include cache tokens
|
||||
// (total_tokens != prompt_tokens + completion_tokens), we return prompt_tokens as-is.
|
||||
func computeClaudeInputTokens(u *usage) int {
|
||||
if u == nil {
|
||||
return 0
|
||||
}
|
||||
promptTokens := u.PromptTokens
|
||||
if u.PromptTokensDetails != nil && u.PromptTokensDetails.CachedTokens > 0 {
|
||||
if u.TotalTokens > 0 && u.TotalTokens == promptTokens+u.CompletionTokens {
|
||||
return promptTokens - u.PromptTokensDetails.CachedTokens
|
||||
}
|
||||
}
|
||||
return promptTokens
|
||||
}
|
||||
|
||||
// ConvertOpenAIResponseToClaude converts an OpenAI response back to Claude format
|
||||
func (c *ClaudeToOpenAIConverter) ConvertOpenAIResponseToClaude(ctx wrapper.HttpContext, body []byte) ([]byte, error) {
|
||||
log.Debugf("[OpenAI->Claude] Original OpenAI response body: %s", string(body))
|
||||
@@ -373,7 +393,7 @@ func (c *ClaudeToOpenAIConverter) ConvertOpenAIResponseToClaude(ctx wrapper.Http
|
||||
// Only include usage if it's available
|
||||
if openaiResponse.Usage != nil {
|
||||
claudeResponse.Usage = claudeTextGenUsage{
|
||||
InputTokens: openaiResponse.Usage.PromptTokens,
|
||||
InputTokens: computeClaudeInputTokens(openaiResponse.Usage),
|
||||
OutputTokens: openaiResponse.Usage.CompletionTokens,
|
||||
}
|
||||
if openaiResponse.Usage.PromptTokensDetails != nil {
|
||||
@@ -582,7 +602,7 @@ func (c *ClaudeToOpenAIConverter) buildClaudeStreamResponse(ctx wrapper.HttpCont
|
||||
// Only include usage if it's available
|
||||
if openaiResponse.Usage != nil {
|
||||
message.Usage = claudeTextGenUsage{
|
||||
InputTokens: openaiResponse.Usage.PromptTokens,
|
||||
InputTokens: computeClaudeInputTokens(openaiResponse.Usage),
|
||||
OutputTokens: 0,
|
||||
}
|
||||
}
|
||||
@@ -828,7 +848,7 @@ func (c *ClaudeToOpenAIConverter) buildClaudeStreamResponse(ctx wrapper.HttpCont
|
||||
}
|
||||
if openaiResponse.Usage != nil {
|
||||
message.Usage = claudeTextGenUsage{
|
||||
InputTokens: openaiResponse.Usage.PromptTokens,
|
||||
InputTokens: computeClaudeInputTokens(openaiResponse.Usage),
|
||||
OutputTokens: 0,
|
||||
}
|
||||
}
|
||||
@@ -985,7 +1005,7 @@ func (c *ClaudeToOpenAIConverter) buildClaudeStreamResponse(ctx wrapper.HttpCont
|
||||
openaiResponse.Usage.PromptTokens, openaiResponse.Usage.CompletionTokens)
|
||||
|
||||
usage := &claudeTextGenUsage{
|
||||
InputTokens: openaiResponse.Usage.PromptTokens,
|
||||
InputTokens: computeClaudeInputTokens(openaiResponse.Usage),
|
||||
OutputTokens: openaiResponse.Usage.CompletionTokens,
|
||||
}
|
||||
if openaiResponse.Usage.PromptTokensDetails != nil {
|
||||
|
||||
@@ -1099,6 +1099,72 @@ func TestClaudeToOpenAIConverter_ConvertOpenAIResponseToClaude(t *testing.T) {
|
||||
require.NotNil(t, toolContent.Input)
|
||||
assert.Equal(t, "/Users/zhangty/git/higress/README.md", (*toolContent.Input)["file_path"])
|
||||
})
|
||||
|
||||
t.Run("openai_standard_usage_with_cached_tokens", func(t *testing.T) {
|
||||
// OpenAI standard: prompt_tokens includes cached_tokens.
|
||||
// total_tokens == prompt_tokens + completion_tokens -> subtract cached_tokens.
|
||||
openaiResponse := `{
|
||||
"id": "chatcmpl-test",
|
||||
"model": "gpt-4o",
|
||||
"object": "chat.completion",
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"finish_reason": "stop",
|
||||
"message": {"role": "assistant", "content": "Hello!"}
|
||||
}],
|
||||
"usage": {
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 20,
|
||||
"total_tokens": 120,
|
||||
"prompt_tokens_details": {"cached_tokens": 60}
|
||||
}
|
||||
}`
|
||||
|
||||
result, err := converter.ConvertOpenAIResponseToClaude(nil, []byte(openaiResponse))
|
||||
require.NoError(t, err)
|
||||
|
||||
var claudeResp claudeTextGenResponse
|
||||
err = json.Unmarshal(result, &claudeResp)
|
||||
require.NoError(t, err)
|
||||
|
||||
// input_tokens should exclude cache: 100 - 60 = 40
|
||||
assert.Equal(t, 40, claudeResp.Usage.InputTokens)
|
||||
assert.Equal(t, 20, claudeResp.Usage.OutputTokens)
|
||||
assert.Equal(t, 60, claudeResp.Usage.CacheReadInputTokens)
|
||||
})
|
||||
|
||||
t.Run("bedrock_style_usage_with_cached_tokens", func(t *testing.T) {
|
||||
// Bedrock-style: prompt_tokens does NOT include cached_tokens.
|
||||
// total_tokens (180) != prompt_tokens (100) + completion_tokens (20) -> do NOT subtract.
|
||||
openaiResponse := `{
|
||||
"id": "chatcmpl-test",
|
||||
"model": "gpt-4o",
|
||||
"object": "chat.completion",
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"finish_reason": "stop",
|
||||
"message": {"role": "assistant", "content": "Hello!"}
|
||||
}],
|
||||
"usage": {
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 20,
|
||||
"total_tokens": 180,
|
||||
"prompt_tokens_details": {"cached_tokens": 60}
|
||||
}
|
||||
}`
|
||||
|
||||
result, err := converter.ConvertOpenAIResponseToClaude(nil, []byte(openaiResponse))
|
||||
require.NoError(t, err)
|
||||
|
||||
var claudeResp claudeTextGenResponse
|
||||
err = json.Unmarshal(result, &claudeResp)
|
||||
require.NoError(t, err)
|
||||
|
||||
// input_tokens should NOT be adjusted: 100 (prompt_tokens already excludes cache)
|
||||
assert.Equal(t, 100, claudeResp.Usage.InputTokens)
|
||||
assert.Equal(t, 20, claudeResp.Usage.OutputTokens)
|
||||
assert.Equal(t, 60, claudeResp.Usage.CacheReadInputTokens)
|
||||
})
|
||||
}
|
||||
|
||||
func TestProviderConfigSupportsMessageReasoningContent(t *testing.T) {
|
||||
@@ -1354,6 +1420,28 @@ func TestClaudeToOpenAIConverter_ConvertOpenAIStreamResponseToClaude_WithCachedT
|
||||
|
||||
resultStr := string(result)
|
||||
assert.Contains(t, resultStr, "\"type\":\"message_delta\"")
|
||||
// OpenAI standard: prompt_tokens (100) includes cached_tokens (60).
|
||||
// Claude semantics: input_tokens should exclude cache tokens, so 100 - 60 = 40.
|
||||
assert.Contains(t, resultStr, "\"input_tokens\":40")
|
||||
assert.Contains(t, resultStr, "\"output_tokens\":20")
|
||||
assert.Contains(t, resultStr, "\"cache_read_input_tokens\":60")
|
||||
}
|
||||
|
||||
func TestClaudeToOpenAIConverter_ConvertOpenAIStreamResponseToClaude_BedrockStyleUsage(t *testing.T) {
|
||||
converter := &ClaudeToOpenAIConverter{}
|
||||
|
||||
// Bedrock-style usage: prompt_tokens does NOT include cached_tokens.
|
||||
// total_tokens (180) != prompt_tokens (100) + completion_tokens (20),
|
||||
// so computeClaudeInputTokens should NOT subtract cached_tokens.
|
||||
streamChunk := "data: {\"id\":\"chatcmpl-test\",\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"}}]}\n\n" +
|
||||
"data: {\"id\":\"chatcmpl-test\",\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":20,\"total_tokens\":180,\"prompt_tokens_details\":{\"cached_tokens\":60}}}\n\n"
|
||||
|
||||
result, err := converter.ConvertOpenAIStreamResponseToClaude(nil, []byte(streamChunk))
|
||||
require.NoError(t, err)
|
||||
|
||||
resultStr := string(result)
|
||||
assert.Contains(t, resultStr, "\"type\":\"message_delta\"")
|
||||
// Bedrock: prompt_tokens already excludes cache, so input_tokens = prompt_tokens = 100.
|
||||
assert.Contains(t, resultStr, "\"input_tokens\":100")
|
||||
assert.Contains(t, resultStr, "\"output_tokens\":20")
|
||||
assert.Contains(t, resultStr, "\"cache_read_input_tokens\":60")
|
||||
|
||||
Reference in New Issue
Block a user