fix(usage): classify partial token accounting correctly

This commit is contained in:
Dylan
2026-07-24 00:31:32 +08:00
parent 416a080174
commit fe8a616aa3
6 changed files with 223 additions and 22 deletions

View File

@@ -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,

View File

@@ -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())

View File

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

View File

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

View File

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

View File

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