diff --git a/internal/redisqueue/plugin_test.go b/internal/redisqueue/plugin_test.go index 5f41507f2..a55d83af4 100644 --- a/internal/redisqueue/plugin_test.go +++ b/internal/redisqueue/plugin_test.go @@ -144,7 +144,7 @@ func TestUsageQueuePluginPayloadDefaultsGenerateTrueWhenOmitted(t *testing.T) { }) } -func TestUsageQueuePluginMarksCanonicalZeroCacheRead(t *testing.T) { +func TestUsageQueuePluginPreservesLegacyCachedOnlyUsage(t *testing.T) { withEnabledQueue(t, func() { ctx := internallogging.WithResponseStatusHolder(context.Background()) internallogging.SetResponseStatus(ctx, http.StatusOK) @@ -153,21 +153,16 @@ func TestUsageQueuePluginMarksCanonicalZeroCacheRead(t *testing.T) { Provider: "openai", Model: "gpt-5.4", Detail: coreusage.Detail{ - CachedTokens: 13, - CacheReadTokens: 0, + CachedTokens: 13, }, }) payload := popSinglePayload(t) requireTokensBoolField(t, payload, "cache_read_tokens_present", true) tokens := requireTokensPayload(t, payload) - var cacheReadTokens int64 - if errUnmarshal := json.Unmarshal(tokens["cache_read_tokens"], &cacheReadTokens); errUnmarshal != nil { - t.Fatalf("unmarshal cache_read_tokens: %v", errUnmarshal) - } - if cacheReadTokens != 0 { - t.Fatalf("cache_read_tokens = %d, want 0", cacheReadTokens) - } + requireIntField(t, tokens, "cache_read_tokens", 13) + requireIntField(t, tokens, "total_tokens", 13) + requireTokenBreakdown(t, payload, coreusage.TokenAccountingQualityUnclassified, 13) }) } diff --git a/internal/runtime/executor/helps/usage_helpers.go b/internal/runtime/executor/helps/usage_helpers.go index fc8f71afa..52e1687f3 100644 --- a/internal/runtime/executor/helps/usage_helpers.go +++ b/internal/runtime/executor/helps/usage_helpers.go @@ -609,14 +609,35 @@ func parseOpenAIStyleUsageNode(usageNode gjson.Result) usage.Detail { detail.ReasoningTokens = reasoning.Int() } if hasOpenAIStyleUsageBucketFields(usageNode) { - detail.TokenBreakdown = usage.NewSubsetTokenBreakdown( - detail.InputTokens, - detail.CacheReadTokens, - detail.CacheCreationTokens, - detail.OutputTokens, - detail.ReasoningTokens, - detail.TotalTokens, - ) + if inputNode.Exists() && outputNode.Exists() { + detail.TokenBreakdown = usage.NewSubsetTokenBreakdown( + detail.InputTokens, + detail.CacheReadTokens, + detail.CacheCreationTokens, + detail.OutputTokens, + detail.ReasoningTokens, + detail.TotalTokens, + ) + } else { + cacheReadTokens := detail.CacheReadTokens + cacheCreationTokens := detail.CacheCreationTokens + if !inputNode.Exists() { + cacheReadTokens = 0 + cacheCreationTokens = 0 + } + reasoningTokens := detail.ReasoningTokens + if !outputNode.Exists() { + reasoningTokens = 0 + } + detail.TokenBreakdown = usage.NewPartialSubsetTokenBreakdown( + detail.InputTokens, + cacheReadTokens, + cacheCreationTokens, + detail.OutputTokens, + reasoningTokens, + detail.TotalTokens, + ) + } } else { detail.TokenBreakdown = usage.NewUnclassifiedTokenBreakdown(detail.TotalTokens) } @@ -691,16 +712,28 @@ func parseClaudeUsageNode(usageNode gjson.Result) usage.Detail { func parseGeminiFamilyUsageDetail(node gjson.Result) usage.Detail { cachedTokens := node.Get("cachedContentTokenCount").Int() + toolUseTokens := firstExistingUsageNode(node, "toolUsePromptTokenCount", "tool_use_prompt_token_count").Int() + inputTokens, okInput := safeUsageTokenSum(node.Get("promptTokenCount").Int(), toolUseTokens) detail := usage.Detail{ - InputTokens: node.Get("promptTokenCount").Int(), + InputTokens: inputTokens, OutputTokens: node.Get("candidatesTokenCount").Int(), ReasoningTokens: node.Get("thoughtsTokenCount").Int(), TotalTokens: node.Get("totalTokenCount").Int(), CachedTokens: cachedTokens, CacheReadTokens: cachedTokens, } + if !okInput { + detail.TokenBreakdown = invalidUsageTokenBreakdown(detail.TotalTokens) + return detail + } if detail.TotalTokens == 0 { - detail.TotalTokens = detail.InputTokens + detail.OutputTokens + detail.ReasoningTokens + var okTotal bool + detail.TotalTokens, okTotal = safeUsageTokenSum(detail.InputTokens, detail.OutputTokens, detail.ReasoningTokens) + if !okTotal { + detail.TotalTokens = 0 + detail.TokenBreakdown = invalidUsageTokenBreakdown(0) + return detail + } } detail.TokenBreakdown = usage.NewSeparateReasoningTokenBreakdown( detail.InputTokens, @@ -715,8 +748,13 @@ func parseGeminiFamilyUsageDetail(node gjson.Result) usage.Detail { func parseInteractionsUsageDetail(node gjson.Result) usage.Detail { cacheRead := firstExistingUsageNode(node, "cache_read_tokens", "cacheReadTokens") + toolUseTokens := firstExistingUsageNode(node, "tool_use_tokens", "total_tool_use_tokens", "toolUseTokens", "totalToolUseTokens").Int() + inputTokens, okInput := safeUsageTokenSum( + firstExistingUsageNode(node, "input_tokens", "prompt_tokens", "total_input_tokens").Int(), + toolUseTokens, + ) detail := usage.Detail{ - InputTokens: firstExistingUsageNode(node, "input_tokens", "prompt_tokens", "total_input_tokens").Int(), + InputTokens: inputTokens, OutputTokens: firstExistingUsageNode(node, "output_tokens", "completion_tokens", "total_output_tokens").Int(), ReasoningTokens: firstExistingUsageNode(node, "reasoning_tokens", "thoughtsTokenCount", "total_thought_tokens").Int(), TotalTokens: firstExistingUsageNode(node, "total_tokens", "totalTokenCount").Int(), @@ -724,11 +762,21 @@ func parseInteractionsUsageDetail(node gjson.Result) usage.Detail { CacheReadTokens: cacheRead.Int(), CacheCreationTokens: firstExistingUsageNode(node, "cache_creation_tokens", "cacheCreationTokens", "cache_write_tokens", "cacheWriteTokens").Int(), } + if !okInput { + detail.TokenBreakdown = invalidUsageTokenBreakdown(detail.TotalTokens) + return detail + } if !cacheRead.Exists() && detail.CachedTokens > 0 { detail.CacheReadTokens = detail.CachedTokens } if detail.TotalTokens == 0 { - detail.TotalTokens = detail.InputTokens + detail.OutputTokens + detail.ReasoningTokens + var okTotal bool + detail.TotalTokens, okTotal = safeUsageTokenSum(detail.InputTokens, detail.OutputTokens, detail.ReasoningTokens) + if !okTotal { + detail.TotalTokens = 0 + detail.TokenBreakdown = invalidUsageTokenBreakdown(0) + return detail + } } detail.TokenBreakdown = usage.NewSeparateReasoningTokenBreakdown( detail.InputTokens, @@ -829,6 +877,29 @@ func firstExistingUsageNode(root gjson.Result, paths ...string) gjson.Result { return gjson.Result{} } +func safeUsageTokenSum(values ...int64) (int64, bool) { + var total int64 + for _, value := range values { + if value < 0 || total > int64(^uint64(0)>>1)-value { + return 0, false + } + total += value + } + return total, true +} + +func invalidUsageTokenBreakdown(total int64) usage.TokenBreakdown { + if total < 0 { + total = 0 + } + return usage.TokenBreakdown{ + SchemaVersion: usage.TokenAccountingSchemaVersion, + Quality: usage.TokenAccountingQualityInconsistent, + TotalTokens: total, + UnclassifiedTokens: total, + } +} + func ParseAntigravityUsage(data []byte) usage.Detail { usageNode := gjson.ParseBytes(data) node := usageNode.Get("response.usageMetadata") diff --git a/internal/runtime/executor/helps/usage_helpers_test.go b/internal/runtime/executor/helps/usage_helpers_test.go index 4511e033b..0ce00217b 100644 --- a/internal/runtime/executor/helps/usage_helpers_test.go +++ b/internal/runtime/executor/helps/usage_helpers_test.go @@ -77,6 +77,14 @@ func TestParseOpenAIUsageTotalOnlyIsUnclassified(t *testing.T) { } } +func TestParseOpenAIUsagePartialBucketsPreserveKnownTokens(t *testing.T) { + detail := ParseOpenAIUsage([]byte(`{"usage":{"input_tokens":10,"total_tokens":15}}`)) + if !detail.TokenBreakdown.Valid() || detail.TokenBreakdown.Quality != usage.TokenAccountingQualityUnclassified || + detail.TokenBreakdown.Input.TotalTokens != 10 || detail.TokenBreakdown.UnclassifiedTokens != 5 { + 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 { @@ -339,6 +347,33 @@ func TestParseGeminiUsageNormalizesCachedContent(t *testing.T) { } } +func TestParseGeminiUsageIncludesToolUsePromptTokens(t *testing.T) { + detail := ParseGeminiUsage([]byte(`{"usageMetadata":{"promptTokenCount":10,"candidatesTokenCount":2,"thoughtsTokenCount":3,"toolUsePromptTokenCount":5,"totalTokenCount":20}}`)) + if detail.InputTokens != 15 || detail.TotalTokens != 20 { + t.Fatalf("detail = %+v", detail) + } + if !detail.TokenBreakdown.Valid() || detail.TokenBreakdown.Quality != usage.TokenAccountingQualityComplete || + detail.TokenBreakdown.Input.UncachedTokens != 15 || detail.TokenBreakdown.Output.ReasoningTokens != 3 { + t.Fatalf("token breakdown = %+v", detail.TokenBreakdown) + } +} + +func TestParseGeminiUsageRejectsInvalidToolUseSums(t *testing.T) { + tests := map[string]string{ + "negative": `{"usageMetadata":{"promptTokenCount":10,"toolUsePromptTokenCount":-1,"totalTokenCount":10}}`, + "overflow": `{"usageMetadata":{"promptTokenCount":9223372036854775807,"toolUsePromptTokenCount":1,"totalTokenCount":9223372036854775807}}`, + } + for name, payload := range tests { + t.Run(name, func(t *testing.T) { + detail := ParseGeminiUsage([]byte(payload)) + if detail.InputTokens < 0 || !detail.TokenBreakdown.Valid() || + detail.TokenBreakdown.Quality != usage.TokenAccountingQualityInconsistent { + t.Fatalf("detail = %+v", detail) + } + }) + } +} + func TestParseInteractionsUsage(t *testing.T) { detail := ParseInteractionsUsage([]byte(`{"usage":{"input_tokens":3,"output_tokens":4,"reasoning_tokens":5,"cached_tokens":2}}`)) if detail.InputTokens != 3 { @@ -385,6 +420,17 @@ func TestParseInteractionsUsageNormalizesCacheWriteAlias(t *testing.T) { } } +func TestParseInteractionsUsageIncludesToolUseTokens(t *testing.T) { + detail := ParseInteractionsUsage([]byte(`{"usage":{"total_input_tokens":2,"total_output_tokens":6,"total_thought_tokens":3,"total_tool_use_tokens":4,"total_tokens":15}}`)) + if detail.InputTokens != 6 || detail.OutputTokens != 6 || detail.ReasoningTokens != 3 || detail.TotalTokens != 15 { + t.Fatalf("detail = %+v", detail) + } + if !detail.TokenBreakdown.Valid() || detail.TokenBreakdown.Quality != usage.TokenAccountingQualityComplete || + detail.TokenBreakdown.Input.UncachedTokens != 6 || detail.TokenBreakdown.Output.TotalTokens != 9 { + t.Fatalf("token breakdown = %+v", detail.TokenBreakdown) + } +} + func TestParseInteractionsStreamUsage(t *testing.T) { detail, ok := ParseInteractionsStreamUsage([]byte(`{"type":"interaction.completed","interaction":{"usage":{"input_tokens":2,"output_tokens":6,"total_tokens":8}}}`)) if !ok { diff --git a/sdk/cliproxy/usage/accounting.go b/sdk/cliproxy/usage/accounting.go index 6429cb7fc..85e89ea94 100644 --- a/sdk/cliproxy/usage/accounting.go +++ b/sdk/cliproxy/usage/accounting.go @@ -114,6 +114,46 @@ func NewSubsetTokenBreakdown(inputTotal, cacheRead, cacheWrite, outputTotal, rea } } +// NewPartialSubsetTokenBreakdown preserves known subset buckets while assigning +// an authoritative remainder to the unclassified bucket. +func NewPartialSubsetTokenBreakdown(inputTotal, cacheRead, cacheWrite, outputTotal, reasoning, total int64) TokenBreakdown { + cacheTotal, okCache := nonNegativeSum(cacheRead, cacheWrite) + expectedTotal, okExpected := nonNegativeSum(inputTotal, outputTotal) + if !okCache || !okExpected || inputTotal < 0 || outputTotal < 0 || reasoning < 0 || + cacheTotal > inputTotal || reasoning > outputTotal || total < 0 { + return inconsistentTokenBreakdown(total, expectedTotal) + } + resolvedTotal := total + if resolvedTotal == 0 { + resolvedTotal = expectedTotal + } + if resolvedTotal < expectedTotal { + return inconsistentTokenBreakdown(total, expectedTotal) + } + unclassified := resolvedTotal - expectedTotal + quality := TokenAccountingQualityComplete + if unclassified > 0 { + quality = TokenAccountingQualityUnclassified + } + return TokenBreakdown{ + SchemaVersion: TokenAccountingSchemaVersion, + Quality: quality, + TotalTokens: resolvedTotal, + Input: TokenInputBreakdown{ + TotalTokens: inputTotal, + UncachedTokens: inputTotal - cacheTotal, + CacheReadTokens: cacheRead, + CacheWriteTokens: cacheWrite, + }, + Output: TokenOutputBreakdown{ + TotalTokens: outputTotal, + NonReasoningTokens: outputTotal - reasoning, + ReasoningTokens: reasoning, + }, + UnclassifiedTokens: unclassified, + } +} + // NewIndependentTokenBreakdown normalizes protocols where uncached input, // cache reads, cache writes, non-reasoning output, and reasoning are separate. func NewIndependentTokenBreakdown(uncachedInput, cacheRead, cacheWrite, nonReasoningOutput, reasoning, total int64) TokenBreakdown { @@ -207,7 +247,13 @@ func EnsureTokenBreakdown(detail Detail) Detail { // providers remain unclassified instead of guessing how their buckets overlap. func EnsureTokenBreakdownForProvider(detail Detail, provider, executorType string) Detail { if !detail.TokenBreakdown.Valid() { - detail.TokenBreakdown = tokenBreakdownForProvider(detail, provider, executorType) + semantics := tokenAccountingSemanticsFor(provider, executorType) + if detail.CacheReadTokens == 0 && detail.CachedTokens > 0 && detail.InputTokens == 0 && + detail.OutputTokens == 0 && detail.ReasoningTokens == 0 && detail.CacheCreationTokens == 0 && detail.TotalTokens == 0 && + (semantics == tokenAccountingSemanticsSubset || semantics == tokenAccountingSemanticsSeparateReasoning) { + detail.CacheReadTokens = detail.CachedTokens + } + detail.TokenBreakdown = tokenBreakdownForSemantics(detail, semantics) } if detail.TotalTokens == 0 { detail.TotalTokens = detail.TokenBreakdown.TotalTokens @@ -215,8 +261,18 @@ func EnsureTokenBreakdownForProvider(detail Detail, provider, executorType strin return detail } -func tokenBreakdownForProvider(detail Detail, provider, executorType string) TokenBreakdown { - switch tokenAccountingSemanticsFor(provider, executorType) { +func tokenBreakdownForSemantics(detail Detail, semantics tokenAccountingSemantics) TokenBreakdown { + if detail.TotalTokens == 0 && detail.InputTokens == 0 && detail.OutputTokens == 0 { + if total, okTotal := unclassifiedTokenLowerBound(detail); !okTotal { + return inconsistentTokenBreakdown(detail.TotalTokens, 0) + } else if total > 0 && (semantics == tokenAccountingSemanticsUnknown || + semantics == tokenAccountingSemanticsSubset || + (semantics == tokenAccountingSemanticsSeparateReasoning && + (detail.CacheReadTokens > 0 || detail.CacheCreationTokens > 0 || detail.CachedTokens > 0))) { + return NewUnclassifiedTokenBreakdown(total) + } + } + switch semantics { case tokenAccountingSemanticsSubset: return NewSubsetTokenBreakdown( detail.InputTokens, @@ -248,7 +304,7 @@ func tokenBreakdownForProvider(detail Detail, provider, executorType string) Tok total := detail.TotalTokens if total == 0 { var okTotal bool - total, okTotal = nonNegativeSum(detail.InputTokens, detail.OutputTokens) + total, okTotal = unclassifiedTokenLowerBound(detail) if !okTotal { return inconsistentTokenBreakdown(detail.TotalTokens, 0) } @@ -257,6 +313,25 @@ func tokenBreakdownForProvider(detail Detail, provider, executorType string) Tok } } +func unclassifiedTokenLowerBound(detail Detail) (int64, bool) { + cacheTokens, okCache := nonNegativeSum(detail.CacheReadTokens, detail.CacheCreationTokens) + if !okCache || detail.InputTokens < 0 || detail.OutputTokens < 0 || detail.ReasoningTokens < 0 || detail.CachedTokens < 0 { + return 0, false + } + inputTotal := detail.InputTokens + if cacheTokens > inputTotal { + inputTotal = cacheTokens + } + if detail.CachedTokens > inputTotal { + inputTotal = detail.CachedTokens + } + outputTotal := detail.OutputTokens + if detail.ReasoningTokens > outputTotal { + outputTotal = detail.ReasoningTokens + } + return nonNegativeSum(inputTotal, outputTotal) +} + func tokenAccountingSemanticsFor(provider, executorType string) tokenAccountingSemantics { normalizedProvider := strings.ToLower(strings.TrimSpace(provider)) normalizedExecutor := strings.ToLower(strings.TrimSpace(executorType)) diff --git a/sdk/cliproxy/usage/accounting_test.go b/sdk/cliproxy/usage/accounting_test.go index 3cca043bd..4c1e13434 100644 --- a/sdk/cliproxy/usage/accounting_test.go +++ b/sdk/cliproxy/usage/accounting_test.go @@ -15,6 +15,17 @@ func TestNewSubsetTokenBreakdownAvoidsCacheAndReasoningDoubleCount(t *testing.T) } } +func TestNewPartialSubsetTokenBreakdownPreservesKnownBuckets(t *testing.T) { + breakdown := NewPartialSubsetTokenBreakdown(10, 4, 0, 0, 0, 15) + if !breakdown.Valid() { + t.Fatalf("breakdown is invalid: %+v", breakdown) + } + if breakdown.Quality != TokenAccountingQualityUnclassified || breakdown.Input.TotalTokens != 10 || + breakdown.UnclassifiedTokens != 5 { + t.Fatalf("breakdown = %+v", breakdown) + } +} + func TestNewIndependentTokenBreakdownKeepsClaudeCacheBucketsIndependent(t *testing.T) { breakdown := NewIndependentTokenBreakdown(30, 7, 13, 5, 0, 55) if !breakdown.Valid() { @@ -119,3 +130,33 @@ func TestEnsureTokenBreakdownForUnknownProviderDoesNotGuessReasoning(t *testing. t.Fatalf("detail = %+v", detail) } } + +func TestEnsureTokenBreakdownForUnknownProviderPreservesAuxiliaryOnlyUsage(t *testing.T) { + detail := EnsureTokenBreakdownForProvider(Detail{ReasoningTokens: 12, CacheReadTokens: 7}, "plugin-provider", "") + if detail.TotalTokens != 19 || detail.TokenBreakdown.Quality != TokenAccountingQualityUnclassified || detail.TokenBreakdown.UnclassifiedTokens != 19 { + t.Fatalf("detail = %+v", detail) + } +} + +func TestEnsureTokenBreakdownForGeminiClassifiesReasoningOnlyUsage(t *testing.T) { + detail := EnsureTokenBreakdownForProvider(Detail{ReasoningTokens: 12}, "gemini", "") + if detail.TotalTokens != 12 || detail.TokenBreakdown.Quality != TokenAccountingQualityComplete || + detail.TokenBreakdown.Output.ReasoningTokens != 12 { + t.Fatalf("detail = %+v", detail) + } +} + +func TestEnsureTokenBreakdownPreservesLegacyCachedOnlyUsage(t *testing.T) { + detail := EnsureTokenBreakdownForProvider(Detail{CachedTokens: 13}, "openai", "") + if detail.TotalTokens != 13 || detail.CacheReadTokens != 13 || detail.TokenBreakdown.Quality != TokenAccountingQualityUnclassified || + detail.TokenBreakdown.UnclassifiedTokens != 13 { + t.Fatalf("detail = %+v", detail) + } +} + +func TestEnsureTokenBreakdownDoesNotOverrideCanonicalZeroCacheRead(t *testing.T) { + detail := EnsureTokenBreakdownForProvider(Detail{CachedTokens: 13, CacheCreationTokens: 13}, "openai", "") + if detail.CacheReadTokens != 0 { + t.Fatalf("detail = %+v", detail) + } +}