From fff3419b2a7aee0a744ce93229c3d0860d485262 Mon Sep 17 00:00:00 2001 From: zhouyayu Date: Sun, 20 Sep 2026 20:45:55 +0800 Subject: [PATCH] fix(responses): preserve cached input token usage --- changelog/unreleased/responses-cache-usage.md | 7 ++ internal/gateway/compat_stream.go | 6 +- internal/gateway/openai_stream.go | 6 +- internal/gateway/responses.go | 26 ++++++-- internal/gateway/responses_test.go | 64 +++++++++++++++++++ 5 files changed, 102 insertions(+), 7 deletions(-) create mode 100644 changelog/unreleased/responses-cache-usage.md create mode 100644 internal/gateway/responses_test.go diff --git a/changelog/unreleased/responses-cache-usage.md b/changelog/unreleased/responses-cache-usage.md new file mode 100644 index 0000000..623c041 --- /dev/null +++ b/changelog/unreleased/responses-cache-usage.md @@ -0,0 +1,7 @@ +### English + +- Preserve cached input token usage in Responses API output for both streaming and non-streaming requests without double-counting total tokens. + +### 中文 + +- 在 Responses API 的流式与非流式输出中保留缓存输入 Token 用量,同时避免在总 Token 数中重复计数。 diff --git a/internal/gateway/compat_stream.go b/internal/gateway/compat_stream.go index 0c00562..4654c35 100644 --- a/internal/gateway/compat_stream.go +++ b/internal/gateway/compat_stream.go @@ -608,7 +608,11 @@ func RelayResponsesStream(writer io.Writer, body io.Reader, requestID, model str return stats, err } } - completed := responsesResponse(requestID, model, content, output.reasoning.String(), calls, derefInt(stats.PromptTokens), derefInt(stats.CompletionTokens)) + completed := responsesResponse( + requestID, model, content, output.reasoning.String(), calls, + derefInt(stats.PromptTokens), derefInt(stats.CompletionTokens), + stats.CacheReadTokens, stats.CachedTokens, + ) if err := eventWriter.write("response.completed", map[string]any{"type": "response.completed", "response": completed}); err != nil { return stats, err } diff --git a/internal/gateway/openai_stream.go b/internal/gateway/openai_stream.go index 7dffc2d..616ef11 100644 --- a/internal/gateway/openai_stream.go +++ b/internal/gateway/openai_stream.go @@ -428,10 +428,14 @@ func ParseStreamUsageLine(line string) (StreamRelayStats, bool) { if credits == nil { credits = parsed.Usage.Credit } + cacheRead := parsed.Usage.CacheReadTokens + if cacheRead == nil { + cacheRead = parsed.Usage.PromptDetails.CachedTokens + } return StreamRelayStats{ PromptTokens: parsed.Usage.PromptTokens, CompletionTokens: parsed.Usage.CompletionTokens, - CacheReadTokens: parsed.Usage.CacheReadTokens, + CacheReadTokens: cacheRead, CacheWriteTokens: parsed.Usage.CacheWriteTokens, CachedTokens: parsed.Usage.PromptDetails.CachedTokens, UsageSource: firstNonEmpty(parsed.Usage.Source, "estimate"), diff --git a/internal/gateway/responses.go b/internal/gateway/responses.go index e0c210c..d069186 100644 --- a/internal/gateway/responses.go +++ b/internal/gateway/responses.go @@ -49,7 +49,11 @@ func (h *Handler) HandleResponses(w http.ResponseWriter, r *http.Request) { CachedTokens: result.CachedTokens, UsageSource: result.UsageSource, Credits: result.Credits, ConsumedCredits: result.ConsumedCredits, Model: result.Model, }, nil, result.AttemptCount, result.ReasoningLevel) - response := responsesResponse(execution.RequestID, firstNonEmpty(result.Model, execution.PublicModel), result.Content, result.Reasoning, decodeOpenAIToolCalls(result.ToolCalls), result.PromptTokens, result.CompletionTokens) + response := responsesResponse( + execution.RequestID, firstNonEmpty(result.Model, execution.PublicModel), result.Content, result.Reasoning, + decodeOpenAIToolCalls(result.ToolCalls), result.PromptTokens, result.CompletionTokens, + result.CacheReadTokens, result.CachedTokens, + ) translate.RestoreResponseToolNames(response, execution.Request.ResponseToolNames) writeJSON(w, http.StatusOK, response) } @@ -100,11 +104,16 @@ func writeCompatibilityOpenAIError(w http.ResponseWriter, err error) { WriteClassifiedErr(w, err) } -func responsesResponse(requestID, model, content, reasoning string, toolCalls []proxyToolCall, promptTokens, completionTokens int) map[string]any { +func responsesResponse( + requestID, model, content, reasoning string, + toolCalls []proxyToolCall, + promptTokens, completionTokens int, + cacheReadTokens, cachedTokens *int, +) map[string]any { return map[string]any{ "id": "resp_" + requestID, "object": "response", "created_at": time.Now().Unix(), "status": "completed", "model": model, "output": responsesOutputItems(requestID, content, reasoning, toolCalls), - "usage": responsesUsage(promptTokens, completionTokens), + "usage": responsesUsage(promptTokens, completionTokens, cacheReadTokens, cachedTokens), } } @@ -132,6 +141,13 @@ func responseFunctionCallItem(requestID string, callIndex int, call proxyToolCal } } -func responsesUsage(promptTokens, completionTokens int) map[string]any { - return map[string]any{"input_tokens": promptTokens, "output_tokens": completionTokens, "total_tokens": promptTokens + completionTokens} +func responsesUsage(promptTokens, completionTokens int, cacheReadTokens, cachedTokens *int) map[string]any { + usage := map[string]any{"input_tokens": promptTokens, "output_tokens": completionTokens, "total_tokens": promptTokens + completionTokens} + if cacheReadTokens == nil { + cacheReadTokens = cachedTokens + } + if cacheReadTokens != nil { + usage["input_tokens_details"] = map[string]any{"cached_tokens": *cacheReadTokens} + } + return usage } diff --git a/internal/gateway/responses_test.go b/internal/gateway/responses_test.go new file mode 100644 index 0000000..8ce4a3a --- /dev/null +++ b/internal/gateway/responses_test.go @@ -0,0 +1,64 @@ +package gateway + +import "testing" + +func TestResponsesUsagePreservesCachedInputTokens(t *testing.T) { + zero := 0 + read := 64 + cached := 48 + + withoutCache := responsesUsage(100, 20, nil, nil) + if _, ok := withoutCache["input_tokens_details"]; ok { + t.Fatalf("missing cache usage must not be fabricated: %#v", withoutCache) + } + + withZero := responsesUsage(100, 20, &zero, nil) + zeroDetails := withZero["input_tokens_details"].(map[string]any) + if zeroDetails["cached_tokens"] != 0 { + t.Fatalf("explicit zero cache usage was lost: %#v", withZero) + } + + withFallback := responsesUsage(100, 20, nil, &cached) + fallbackDetails := withFallback["input_tokens_details"].(map[string]any) + if fallbackDetails["cached_tokens"] != 48 { + t.Fatalf("cached token fallback mismatch: %#v", withFallback) + } + + withTopLevel := responsesUsage(100, 20, &read, &cached) + topLevelDetails := withTopLevel["input_tokens_details"].(map[string]any) + if topLevelDetails["cached_tokens"] != 64 { + t.Fatalf("cache_read_tokens must win: %#v", withTopLevel) + } + if withTopLevel["total_tokens"] != 120 { + t.Fatalf("cached input must not be added twice: %#v", withTopLevel) + } +} + +func TestParseStreamUsageLineCacheReadFallback(t *testing.T) { + for _, test := range []struct { + name string + usage string + want *int + }{ + {name: "detail only", usage: `"prompt_tokens_details":{"cached_tokens":2176}`, want: ptrInt(2176)}, + {name: "explicit zero", usage: `"prompt_tokens_details":{"cached_tokens":0}`, want: ptrInt(0)}, + {name: "top-level wins", usage: `"cache_read_tokens":12,"prompt_tokens_details":{"cached_tokens":2176}`, want: ptrInt(12)}, + {name: "unknown stays absent", usage: `"prompt_tokens":16`, want: nil}, + } { + t.Run(test.name, func(t *testing.T) { + stats, ok := ParseStreamUsageLine(`data: {"usage":{` + test.usage + `}}`) + if !ok { + t.Fatal("usage not parsed") + } + if test.want == nil { + if stats.CacheReadTokens != nil { + t.Fatalf("fabricated cache read: %v", *stats.CacheReadTokens) + } + return + } + if stats.CacheReadTokens == nil || *stats.CacheReadTokens != *test.want { + t.Fatalf("cache read = %v, want %d", stats.CacheReadTokens, *test.want) + } + }) + } +}