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
|
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
|
// ConvertOpenAIResponseToClaude converts an OpenAI response back to Claude format
|
||||||
func (c *ClaudeToOpenAIConverter) ConvertOpenAIResponseToClaude(ctx wrapper.HttpContext, body []byte) ([]byte, error) {
|
func (c *ClaudeToOpenAIConverter) ConvertOpenAIResponseToClaude(ctx wrapper.HttpContext, body []byte) ([]byte, error) {
|
||||||
log.Debugf("[OpenAI->Claude] Original OpenAI response body: %s", string(body))
|
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
|
// Only include usage if it's available
|
||||||
if openaiResponse.Usage != nil {
|
if openaiResponse.Usage != nil {
|
||||||
claudeResponse.Usage = claudeTextGenUsage{
|
claudeResponse.Usage = claudeTextGenUsage{
|
||||||
InputTokens: openaiResponse.Usage.PromptTokens,
|
InputTokens: computeClaudeInputTokens(openaiResponse.Usage),
|
||||||
OutputTokens: openaiResponse.Usage.CompletionTokens,
|
OutputTokens: openaiResponse.Usage.CompletionTokens,
|
||||||
}
|
}
|
||||||
if openaiResponse.Usage.PromptTokensDetails != nil {
|
if openaiResponse.Usage.PromptTokensDetails != nil {
|
||||||
@@ -582,7 +602,7 @@ func (c *ClaudeToOpenAIConverter) buildClaudeStreamResponse(ctx wrapper.HttpCont
|
|||||||
// Only include usage if it's available
|
// Only include usage if it's available
|
||||||
if openaiResponse.Usage != nil {
|
if openaiResponse.Usage != nil {
|
||||||
message.Usage = claudeTextGenUsage{
|
message.Usage = claudeTextGenUsage{
|
||||||
InputTokens: openaiResponse.Usage.PromptTokens,
|
InputTokens: computeClaudeInputTokens(openaiResponse.Usage),
|
||||||
OutputTokens: 0,
|
OutputTokens: 0,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -828,7 +848,7 @@ func (c *ClaudeToOpenAIConverter) buildClaudeStreamResponse(ctx wrapper.HttpCont
|
|||||||
}
|
}
|
||||||
if openaiResponse.Usage != nil {
|
if openaiResponse.Usage != nil {
|
||||||
message.Usage = claudeTextGenUsage{
|
message.Usage = claudeTextGenUsage{
|
||||||
InputTokens: openaiResponse.Usage.PromptTokens,
|
InputTokens: computeClaudeInputTokens(openaiResponse.Usage),
|
||||||
OutputTokens: 0,
|
OutputTokens: 0,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -985,7 +1005,7 @@ func (c *ClaudeToOpenAIConverter) buildClaudeStreamResponse(ctx wrapper.HttpCont
|
|||||||
openaiResponse.Usage.PromptTokens, openaiResponse.Usage.CompletionTokens)
|
openaiResponse.Usage.PromptTokens, openaiResponse.Usage.CompletionTokens)
|
||||||
|
|
||||||
usage := &claudeTextGenUsage{
|
usage := &claudeTextGenUsage{
|
||||||
InputTokens: openaiResponse.Usage.PromptTokens,
|
InputTokens: computeClaudeInputTokens(openaiResponse.Usage),
|
||||||
OutputTokens: openaiResponse.Usage.CompletionTokens,
|
OutputTokens: openaiResponse.Usage.CompletionTokens,
|
||||||
}
|
}
|
||||||
if openaiResponse.Usage.PromptTokensDetails != nil {
|
if openaiResponse.Usage.PromptTokensDetails != nil {
|
||||||
|
|||||||
@@ -1099,6 +1099,72 @@ func TestClaudeToOpenAIConverter_ConvertOpenAIResponseToClaude(t *testing.T) {
|
|||||||
require.NotNil(t, toolContent.Input)
|
require.NotNil(t, toolContent.Input)
|
||||||
assert.Equal(t, "/Users/zhangty/git/higress/README.md", (*toolContent.Input)["file_path"])
|
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) {
|
func TestProviderConfigSupportsMessageReasoningContent(t *testing.T) {
|
||||||
@@ -1354,6 +1420,28 @@ func TestClaudeToOpenAIConverter_ConvertOpenAIStreamResponseToClaude_WithCachedT
|
|||||||
|
|
||||||
resultStr := string(result)
|
resultStr := string(result)
|
||||||
assert.Contains(t, resultStr, "\"type\":\"message_delta\"")
|
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, "\"input_tokens\":100")
|
||||||
assert.Contains(t, resultStr, "\"output_tokens\":20")
|
assert.Contains(t, resultStr, "\"output_tokens\":20")
|
||||||
assert.Contains(t, resultStr, "\"cache_read_input_tokens\":60")
|
assert.Contains(t, resultStr, "\"cache_read_input_tokens\":60")
|
||||||
|
|||||||
Reference in New Issue
Block a user