mirror of
https://github.com/router-for-me/CLIProxyAPI.git
synced 2026-09-03 06:35:00 +08:00
fix(usage): classify partial token accounting correctly
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user