mirror of
https://github.com/router-for-me/CLIProxyAPI.git
synced 2026-09-03 06:35:00 +08:00
fix(gemini): map cached content tokens to claude cache read usage
- Deduct `cachedContentTokenCount` from `promptTokenCount` for `usage.input_tokens`. - Set `usage.cache_read_input_tokens` when cached tokens are present in streaming and non-streaming responses. Closes: #5238
This commit is contained in:
@@ -264,8 +264,16 @@ func ConvertGeminiResponseToClaude(_ context.Context, _ string, originalRequestR
|
||||
|
||||
thoughtsTokenCount := usageResult.Get("thoughtsTokenCount").Int()
|
||||
candidatesTokenCount := usageResult.Get("candidatesTokenCount").Int()
|
||||
cachedTokenCount := usageResult.Get("cachedContentTokenCount").Int()
|
||||
promptTokenCount := usageResult.Get("promptTokenCount").Int() - cachedTokenCount
|
||||
if promptTokenCount < 0 {
|
||||
promptTokenCount = 0
|
||||
}
|
||||
template, _ = sjson.SetBytes(template, "usage.output_tokens", candidatesTokenCount+thoughtsTokenCount)
|
||||
template, _ = sjson.SetBytes(template, "usage.input_tokens", usageResult.Get("promptTokenCount").Int())
|
||||
template, _ = sjson.SetBytes(template, "usage.input_tokens", promptTokenCount)
|
||||
if cachedTokenCount > 0 {
|
||||
template, _ = sjson.SetBytes(template, "usage.cache_read_input_tokens", cachedTokenCount)
|
||||
}
|
||||
|
||||
appendEvent("message_delta", string(template))
|
||||
(*param).(*Params).HasFinalEvents = true
|
||||
@@ -296,10 +304,17 @@ func ConvertGeminiResponseToClaudeNonStream(_ context.Context, _ string, origina
|
||||
out, _ = sjson.SetBytes(out, "id", root.Get("responseId").String())
|
||||
out, _ = sjson.SetBytes(out, "model", root.Get("modelVersion").String())
|
||||
|
||||
inputTokens := root.Get("usageMetadata.promptTokenCount").Int()
|
||||
cachedTokens := root.Get("usageMetadata.cachedContentTokenCount").Int()
|
||||
inputTokens := root.Get("usageMetadata.promptTokenCount").Int() - cachedTokens
|
||||
if inputTokens < 0 {
|
||||
inputTokens = 0
|
||||
}
|
||||
outputTokens := root.Get("usageMetadata.candidatesTokenCount").Int() + root.Get("usageMetadata.thoughtsTokenCount").Int()
|
||||
out, _ = sjson.SetBytes(out, "usage.input_tokens", inputTokens)
|
||||
out, _ = sjson.SetBytes(out, "usage.output_tokens", outputTokens)
|
||||
if cachedTokens > 0 {
|
||||
out, _ = sjson.SetBytes(out, "usage.cache_read_input_tokens", cachedTokens)
|
||||
}
|
||||
|
||||
parts := root.Get("candidates.0.content.parts")
|
||||
textBuilder := strings.Builder{}
|
||||
|
||||
@@ -202,3 +202,91 @@ func TestConvertGeminiResponseToClaudeNonStream_TrailingSignatureOnlyPart(t *tes
|
||||
t.Fatalf("unexpected text block: %s", textBlock.Raw)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConvertGeminiResponseToClaude_UsageWithCachedContentTokenCount(t *testing.T) {
|
||||
requestJSON := []byte(`{"model":"gemini-2.5-pro","messages":[{"role":"user","content":"hi"}]}`)
|
||||
chunk := []byte(`{
|
||||
"candidates": [{
|
||||
"content": {
|
||||
"parts": [{"text": "Hello world"}]
|
||||
},
|
||||
"finishReason": "STOP"
|
||||
}],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 100,
|
||||
"candidatesTokenCount": 7,
|
||||
"cachedContentTokenCount": 91
|
||||
},
|
||||
"modelVersion": "gemini-2.5-pro",
|
||||
"responseId": "resp-usage-cache"
|
||||
}`)
|
||||
|
||||
var param any
|
||||
ctx := context.Background()
|
||||
output := bytes.Join(ConvertGeminiResponseToClaude(ctx, "gemini-2.5-pro", requestJSON, requestJSON, chunk, ¶m), nil)
|
||||
outputText := string(output)
|
||||
|
||||
if !strings.Contains(outputText, `"type":"message_delta"`) {
|
||||
t.Fatalf("expected message_delta event in output, got: %s", outputText)
|
||||
}
|
||||
|
||||
foundMessageDelta := false
|
||||
// Find the message_delta event data
|
||||
for _, line := range strings.Split(outputText, "\n") {
|
||||
if strings.HasPrefix(line, "data: ") && strings.Contains(line, `"type":"message_delta"`) {
|
||||
foundMessageDelta = true
|
||||
deltaJSON := gjson.Parse(strings.TrimPrefix(line, "data: "))
|
||||
inputTokens := deltaJSON.Get("usage.input_tokens").Int()
|
||||
if inputTokens != 9 {
|
||||
t.Fatalf("expected usage.input_tokens = 9 (100 - 91), got %d. Payload: %s", inputTokens, line)
|
||||
}
|
||||
cacheReadTokens := deltaJSON.Get("usage.cache_read_input_tokens").Int()
|
||||
if cacheReadTokens != 91 {
|
||||
t.Fatalf("expected usage.cache_read_input_tokens = 91, got %d. Payload: %s", cacheReadTokens, line)
|
||||
}
|
||||
outputTokens := deltaJSON.Get("usage.output_tokens").Int()
|
||||
if outputTokens != 7 {
|
||||
t.Fatalf("expected usage.output_tokens = 7, got %d. Payload: %s", outputTokens, line)
|
||||
}
|
||||
}
|
||||
}
|
||||
if !foundMessageDelta {
|
||||
t.Fatalf("failed to locate parsed message_delta event in payload: %s", outputText)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConvertGeminiResponseToClaudeNonStream_UsageWithCachedContentTokenCount(t *testing.T) {
|
||||
requestJSON := []byte(`{"model":"gemini-2.5-pro","messages":[{"role":"user","content":"hi"}]}`)
|
||||
geminiResponse := []byte(`{
|
||||
"candidates": [{
|
||||
"content": {
|
||||
"parts": [{"text": "Hello world"}]
|
||||
},
|
||||
"finishReason": "STOP"
|
||||
}],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 100,
|
||||
"candidatesTokenCount": 7,
|
||||
"cachedContentTokenCount": 91
|
||||
},
|
||||
"modelVersion": "gemini-2.5-pro",
|
||||
"responseId": "resp-usage-cache-nonstream"
|
||||
}`)
|
||||
|
||||
ctx := context.Background()
|
||||
output := ConvertGeminiResponseToClaudeNonStream(ctx, "gemini-2.5-pro", requestJSON, requestJSON, geminiResponse, nil)
|
||||
outputJSON := gjson.ParseBytes(output)
|
||||
|
||||
inputTokens := outputJSON.Get("usage.input_tokens").Int()
|
||||
if inputTokens != 9 {
|
||||
t.Fatalf("expected usage.input_tokens = 9 (100 - 91), got %d. Output: %s", inputTokens, string(output))
|
||||
}
|
||||
cacheReadTokens := outputJSON.Get("usage.cache_read_input_tokens").Int()
|
||||
if cacheReadTokens != 91 {
|
||||
t.Fatalf("expected usage.cache_read_input_tokens = 91, got %d. Output: %s", cacheReadTokens, string(output))
|
||||
}
|
||||
outputTokens := outputJSON.Get("usage.output_tokens").Int()
|
||||
if outputTokens != 7 {
|
||||
t.Fatalf("expected usage.output_tokens = 7, got %d. Output: %s", outputTokens, string(output))
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user