From fe8a616aa3037043241556cdea956fd26d0620a6 Mon Sep 17 00:00:00 2001 From: Dylan <1990016@gmail.com> Date: Fri, 24 Jul 2026 00:31:32 +0800 Subject: [PATCH] fix(usage): classify partial token accounting correctly --- internal/redisqueue/plugin.go | 2 +- internal/redisqueue/plugin_test.go | 34 ++++++- .../runtime/executor/helps/usage_helpers.go | 33 ++++--- .../executor/helps/usage_helpers_test.go | 19 +++- sdk/cliproxy/usage/accounting.go | 92 ++++++++++++++++++- sdk/cliproxy/usage/accounting_test.go | 65 +++++++++++++ 6 files changed, 223 insertions(+), 22 deletions(-) diff --git a/internal/redisqueue/plugin.go b/internal/redisqueue/plugin.go index 784532eb4..f43eadd26 100644 --- a/internal/redisqueue/plugin.go +++ b/internal/redisqueue/plugin.go @@ -65,7 +65,7 @@ func (p *usageQueuePlugin) HandleUsage(ctx context.Context, record coreusage.Rec } responseServiceTier := strings.TrimSpace(record.ResponseServiceTier) - usageDetail := coreusage.EnsureTokenBreakdown(record.Detail) + usageDetail := coreusage.EnsureTokenBreakdownForProvider(record.Detail, record.Provider, record.ExecutorType) tokens := tokenStats{ InputTokens: usageDetail.InputTokens, OutputTokens: usageDetail.OutputTokens, diff --git a/internal/redisqueue/plugin_test.go b/internal/redisqueue/plugin_test.go index 682e02cd3..5f41507f2 100644 --- a/internal/redisqueue/plugin_test.go +++ b/internal/redisqueue/plugin_test.go @@ -62,7 +62,7 @@ func TestUsageQueuePluginPayloadIncludesStableFieldsAndSuccess(t *testing.T) { requireMissingField(t, payload, "request_service_tier") requireStringField(t, payload, "response_service_tier", "default") requireIntField(t, payload, "accounting_version", coreusage.TokenAccountingSchemaVersion) - requireTokenBreakdown(t, payload, coreusage.TokenAccountingQualityUnclassified, 30) + requireTokenBreakdown(t, payload, coreusage.TokenAccountingQualityComplete, 30) requireTokensBoolField(t, payload, "cache_read_tokens_present", true) requireHeaderField(t, payload, "response_headers", "X-Upstream-Request-Id", []string{"upstream-req-1"}) requireHeaderField(t, payload, "response_headers", "Retry-After", []string{"30"}) @@ -72,6 +72,38 @@ func TestUsageQueuePluginPayloadIncludesStableFieldsAndSuccess(t *testing.T) { }) } +func TestUsageQueuePluginNormalizesDirectSDKUsageByProvider(t *testing.T) { + tests := []struct { + provider string + wantTotal int + }{ + {provider: "openai", wantTotal: 130}, + {provider: "gemini", wantTotal: 142}, + } + for _, tt := range tests { + t.Run(tt.provider, func(t *testing.T) { + withEnabledQueue(t, func() { + ctx := internallogging.WithResponseStatusHolder(context.Background()) + internallogging.SetResponseStatus(ctx, http.StatusOK) + + (&usageQueuePlugin{}).HandleUsage(ctx, coreusage.Record{ + Provider: tt.provider, + Model: "direct-sdk-model", + Detail: coreusage.Detail{ + InputTokens: 100, + OutputTokens: 30, + ReasoningTokens: 12, + }, + }) + + payload := popSinglePayload(t) + requireIntField(t, requireTokensPayload(t, payload), "total_tokens", tt.wantTotal) + requireTokenBreakdown(t, payload, coreusage.TokenAccountingQualityComplete, int64(tt.wantTotal)) + }) + }) + } +} + func TestUsageQueuePluginPayloadIncludesGenerateFalse(t *testing.T) { withEnabledQueue(t, func() { ctx := internallogging.WithResponseStatusHolder(context.Background()) diff --git a/internal/runtime/executor/helps/usage_helpers.go b/internal/runtime/executor/helps/usage_helpers.go index d5f2309e2..fc8f71afa 100644 --- a/internal/runtime/executor/helps/usage_helpers.go +++ b/internal/runtime/executor/helps/usage_helpers.go @@ -177,7 +177,7 @@ func (r *UsageReporter) buildAdditionalModelRecord(model string, detail usage.De if model == "" { return usage.Record{}, false } - detail = normalizeUsageDetailTotal(detail) + detail = normalizeUsageDetailTotal(detail, r.provider, r.executorType) if !hasNonZeroTokenUsage(detail) { return usage.Record{}, false } @@ -201,14 +201,14 @@ func (r *UsageReporter) publishWithOutcome(ctx context.Context, detail usage.Det if r == nil { return } - detail = normalizeUsageDetailTotal(detail) + detail = normalizeUsageDetailTotal(detail, r.provider, r.executorType) r.once.Do(func() { r.publishRecord(ctx, r.buildRecord(detail, failed, fail)) }) } -func normalizeUsageDetailTotal(detail usage.Detail) usage.Detail { - return usage.EnsureTokenBreakdown(detail) +func normalizeUsageDetailTotal(detail usage.Detail, provider, executorType string) usage.Detail { + return usage.EnsureTokenBreakdownForProvider(detail, provider, executorType) } func hasNonZeroTokenUsage(detail usage.Detail) bool { @@ -551,11 +551,14 @@ func hasOpenAIStyleUsageTokenFields(usageNode gjson.Result) bool { if !usageNode.Exists() || !usageNode.IsObject() { return false } + return usageNode.Get("total_tokens").Exists() || hasOpenAIStyleUsageBucketFields(usageNode) +} + +func hasOpenAIStyleUsageBucketFields(usageNode gjson.Result) bool { return usageNode.Get("prompt_tokens").Exists() || usageNode.Get("input_tokens").Exists() || usageNode.Get("completion_tokens").Exists() || usageNode.Get("output_tokens").Exists() || - usageNode.Get("total_tokens").Exists() || usageNode.Get("prompt_tokens_details.cached_tokens").Exists() || usageNode.Get("input_tokens_details.cached_tokens").Exists() || usageNode.Get("prompt_tokens_details.cache_write_tokens").Exists() || @@ -605,14 +608,18 @@ func parseOpenAIStyleUsageNode(usageNode gjson.Result) usage.Detail { if reasoning.Exists() { detail.ReasoningTokens = reasoning.Int() } - detail.TokenBreakdown = usage.NewSubsetTokenBreakdown( - detail.InputTokens, - detail.CacheReadTokens, - detail.CacheCreationTokens, - detail.OutputTokens, - detail.ReasoningTokens, - detail.TotalTokens, - ) + if hasOpenAIStyleUsageBucketFields(usageNode) { + detail.TokenBreakdown = usage.NewSubsetTokenBreakdown( + detail.InputTokens, + detail.CacheReadTokens, + detail.CacheCreationTokens, + detail.OutputTokens, + detail.ReasoningTokens, + detail.TotalTokens, + ) + } else { + detail.TokenBreakdown = usage.NewUnclassifiedTokenBreakdown(detail.TotalTokens) + } if detail.TotalTokens == 0 { detail.TotalTokens = detail.TokenBreakdown.TotalTokens } diff --git a/internal/runtime/executor/helps/usage_helpers_test.go b/internal/runtime/executor/helps/usage_helpers_test.go index 0f5adba6a..4511e033b 100644 --- a/internal/runtime/executor/helps/usage_helpers_test.go +++ b/internal/runtime/executor/helps/usage_helpers_test.go @@ -69,6 +69,21 @@ func TestParseOpenAIUsageResponses(t *testing.T) { } } +func TestParseOpenAIUsageTotalOnlyIsUnclassified(t *testing.T) { + detail := ParseOpenAIUsage([]byte(`{"usage":{"total_tokens":42}}`)) + if !detail.TokenBreakdown.Valid() || detail.TokenBreakdown.Quality != usage.TokenAccountingQualityUnclassified || + detail.TotalTokens != 42 || detail.TokenBreakdown.UnclassifiedTokens != 42 { + t.Fatalf("detail = %+v", detail) + } +} + +func TestParseOpenAIUsageExplicitZeroBucketsRemainInconsistent(t *testing.T) { + detail := ParseOpenAIUsage([]byte(`{"usage":{"input_tokens":0,"output_tokens":0,"total_tokens":42}}`)) + if !detail.TokenBreakdown.Valid() || detail.TokenBreakdown.Quality != usage.TokenAccountingQualityInconsistent { + t.Fatalf("detail = %+v", detail) + } +} + func TestParseCodexUsageIncludesCacheWriteTokens(t *testing.T) { data := []byte(`{"response":{"service_tier":"priority","usage":{"input_tokens":100,"output_tokens":20,"total_tokens":120,"input_tokens_details":{"cached_tokens":30,"cache_write_tokens":40}}}}`) detail, ok := ParseCodexUsage(data) @@ -354,11 +369,11 @@ func TestNormalizeUsageDetailTotalDoesNotDoubleCountReasoning(t *testing.T) { InputTokens: 100, OutputTokens: 30, ReasoningTokens: 12, - }) + }, "openai", "") if detail.TotalTokens != 130 { t.Fatalf("total tokens = %d, want 130", detail.TotalTokens) } - if detail.TokenBreakdown.Quality != usage.TokenAccountingQualityUnclassified || detail.TokenBreakdown.UnclassifiedTokens != 130 { + if detail.TokenBreakdown.Quality != usage.TokenAccountingQualityComplete || detail.TokenBreakdown.Output.ReasoningTokens != 12 { t.Fatalf("token breakdown = %+v", detail.TokenBreakdown) } } diff --git a/sdk/cliproxy/usage/accounting.go b/sdk/cliproxy/usage/accounting.go index 307a607e5..6429cb7fc 100644 --- a/sdk/cliproxy/usage/accounting.go +++ b/sdk/cliproxy/usage/accounting.go @@ -1,5 +1,7 @@ package usage +import "strings" + // TokenAccountingSchemaVersion identifies the canonical token accounting contract. const TokenAccountingSchemaVersion = 2 @@ -12,6 +14,15 @@ const ( TokenAccountingQualityUnclassified TokenAccountingQuality = "unclassified" ) +type tokenAccountingSemantics uint8 + +const ( + tokenAccountingSemanticsUnknown tokenAccountingSemantics = iota + tokenAccountingSemanticsSubset + tokenAccountingSemanticsIndependent + tokenAccountingSemanticsSeparateReasoning +) + // TokenInputBreakdown contains mutually exclusive input token buckets. type TokenInputBreakdown struct { TotalTokens int64 `json:"total_tokens"` @@ -188,12 +199,15 @@ func NewUnclassifiedTokenBreakdown(total int64) TokenBreakdown { // EnsureTokenBreakdown attaches a valid v2 breakdown to legacy or direct SDK // usage details without guessing whether reasoning is already inside output. func EnsureTokenBreakdown(detail Detail) Detail { + return EnsureTokenBreakdownForProvider(detail, "", "") +} + +// EnsureTokenBreakdownForProvider attaches a valid v2 breakdown to legacy or +// direct SDK usage details using the known provider's token semantics. Unknown +// providers remain unclassified instead of guessing how their buckets overlap. +func EnsureTokenBreakdownForProvider(detail Detail, provider, executorType string) Detail { if !detail.TokenBreakdown.Valid() { - total := detail.TotalTokens - if total == 0 { - total = detail.InputTokens + detail.OutputTokens - } - detail.TokenBreakdown = NewUnclassifiedTokenBreakdown(total) + detail.TokenBreakdown = tokenBreakdownForProvider(detail, provider, executorType) } if detail.TotalTokens == 0 { detail.TotalTokens = detail.TokenBreakdown.TotalTokens @@ -201,6 +215,74 @@ func EnsureTokenBreakdown(detail Detail) Detail { return detail } +func tokenBreakdownForProvider(detail Detail, provider, executorType string) TokenBreakdown { + switch tokenAccountingSemanticsFor(provider, executorType) { + case tokenAccountingSemanticsSubset: + return NewSubsetTokenBreakdown( + detail.InputTokens, + detail.CacheReadTokens, + detail.CacheCreationTokens, + detail.OutputTokens, + detail.ReasoningTokens, + detail.TotalTokens, + ) + case tokenAccountingSemanticsIndependent: + return NewIndependentTokenBreakdown( + detail.InputTokens, + detail.CacheReadTokens, + detail.CacheCreationTokens, + detail.OutputTokens, + detail.ReasoningTokens, + detail.TotalTokens, + ) + case tokenAccountingSemanticsSeparateReasoning: + return NewSeparateReasoningTokenBreakdown( + detail.InputTokens, + detail.CacheReadTokens, + detail.CacheCreationTokens, + detail.OutputTokens, + detail.ReasoningTokens, + detail.TotalTokens, + ) + default: + total := detail.TotalTokens + if total == 0 { + var okTotal bool + total, okTotal = nonNegativeSum(detail.InputTokens, detail.OutputTokens) + if !okTotal { + return inconsistentTokenBreakdown(detail.TotalTokens, 0) + } + } + return NewUnclassifiedTokenBreakdown(total) + } +} + +func tokenAccountingSemanticsFor(provider, executorType string) tokenAccountingSemantics { + normalizedProvider := strings.ToLower(strings.TrimSpace(provider)) + normalizedExecutor := strings.ToLower(strings.TrimSpace(executorType)) + value := strings.TrimSpace(normalizedProvider + " " + normalizedExecutor) + if value == "" || value == "unknown" || value == "unknown unknown" { + return tokenAccountingSemanticsUnknown + } + if normalizedExecutor == "openaicompatexecutor" || normalizedProvider == "openai-compatibility" || strings.HasPrefix(normalizedProvider, "openai-compatible-") { + return tokenAccountingSemanticsSubset + } + if strings.Contains(value, "claude") || strings.Contains(value, "anthropic") { + return tokenAccountingSemanticsIndependent + } + for _, marker := range []string{"gemini", "aistudio", "antigravity", "vertex", "interaction"} { + if strings.Contains(value, marker) { + return tokenAccountingSemanticsSeparateReasoning + } + } + for _, marker := range []string{"openai", "codex", "xai", "grok", "kimi", "qwen", "deepseek", "openrouter"} { + if strings.Contains(value, marker) { + return tokenAccountingSemanticsSubset + } + } + return tokenAccountingSemanticsUnknown +} + func inconsistentTokenBreakdown(total, fallback int64) TokenBreakdown { resolved := total if resolved <= 0 { diff --git a/sdk/cliproxy/usage/accounting_test.go b/sdk/cliproxy/usage/accounting_test.go index f5e7834ea..3cca043bd 100644 --- a/sdk/cliproxy/usage/accounting_test.go +++ b/sdk/cliproxy/usage/accounting_test.go @@ -54,3 +54,68 @@ func TestNewUnclassifiedTokenBreakdownDoesNotGuessBuckets(t *testing.T) { t.Fatalf("breakdown = %+v", breakdown) } } + +func TestEnsureTokenBreakdownForProviderUsesKnownSemantics(t *testing.T) { + tests := []struct { + name string + provider string + executorType string + detail Detail + wantTotal int64 + wantInput int64 + wantOutput int64 + }{ + { + name: "OpenAI subsets cache and reasoning", + provider: "openai", + detail: Detail{InputTokens: 100, OutputTokens: 30, ReasoningTokens: 12, CacheReadTokens: 40, CacheCreationTokens: 10}, + wantTotal: 130, + wantInput: 100, + wantOutput: 30, + }, + { + name: "OpenAI compatible executor takes precedence", + provider: "anthropic", + executorType: "OpenAICompatExecutor", + detail: Detail{InputTokens: 100, OutputTokens: 30, ReasoningTokens: 12, CacheReadTokens: 40, CacheCreationTokens: 10}, + wantTotal: 130, + wantInput: 100, + wantOutput: 30, + }, + { + name: "Gemini keeps reasoning separate", + provider: "gemini", + detail: Detail{InputTokens: 100, OutputTokens: 30, ReasoningTokens: 12, CacheReadTokens: 40, CacheCreationTokens: 10}, + wantTotal: 142, + wantInput: 100, + wantOutput: 42, + }, + { + name: "Claude keeps cache and reasoning independent", + provider: "anthropic", + detail: Detail{InputTokens: 100, OutputTokens: 30, ReasoningTokens: 12, CacheReadTokens: 40, CacheCreationTokens: 10}, + wantTotal: 192, + wantInput: 150, + wantOutput: 42, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + detail := EnsureTokenBreakdownForProvider(tt.detail, tt.provider, tt.executorType) + if !detail.TokenBreakdown.Valid() || detail.TokenBreakdown.Quality != TokenAccountingQualityComplete { + t.Fatalf("token breakdown = %+v", detail.TokenBreakdown) + } + if detail.TotalTokens != tt.wantTotal || detail.TokenBreakdown.TotalTokens != tt.wantTotal || + detail.TokenBreakdown.Input.TotalTokens != tt.wantInput || detail.TokenBreakdown.Output.TotalTokens != tt.wantOutput { + t.Fatalf("detail = %+v, want total=%d input=%d output=%d", detail, tt.wantTotal, tt.wantInput, tt.wantOutput) + } + }) + } +} + +func TestEnsureTokenBreakdownForUnknownProviderDoesNotGuessReasoning(t *testing.T) { + detail := EnsureTokenBreakdownForProvider(Detail{InputTokens: 100, OutputTokens: 30, ReasoningTokens: 12}, "plugin-provider", "") + if detail.TotalTokens != 130 || detail.TokenBreakdown.Quality != TokenAccountingQualityUnclassified || detail.TokenBreakdown.UnclassifiedTokens != 130 { + t.Fatalf("detail = %+v", detail) + } +}