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:
Luis Pater
2026-08-27 04:19:22 +08:00
parent 4fa1de2f9b
commit 6f6856e784
2 changed files with 105 additions and 2 deletions

View File

@@ -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{}

View File

@@ -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, &param), 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))
}
}