Merge pull request #6132 from wucm667/fix/issue-6125-anthropic-cache-ttl

fix: prevent duplicate Anthropic cache TTL billing
This commit is contained in:
Wesley Liddick
2026-08-25 11:02:29 +08:00
committed by GitHub
6 changed files with 168 additions and 26 deletions
+31 -3
View File
@@ -5,6 +5,7 @@ import (
"errors"
"fmt"
"log"
"math"
"strings"
"sync"
"time"
@@ -1370,16 +1371,43 @@ func (s *BillingService) computeTokenBreakdown(
// multiplier 用于长上下文等场景下的整体价格缩放(普通调用传 1.0 即可)。
func (s *BillingService) computeCacheCreationCost(pricing *ModelPricing, tokens UsageTokens, price, multiplier float64) float64 {
if pricing.SupportsCacheBreakdown && (pricing.CacheCreation5mPrice > 0 || pricing.CacheCreation1hPrice > 0) {
if tokens.CacheCreation5mTokens == 0 && tokens.CacheCreation1hTokens == 0 && tokens.CacheCreationTokens > 0 {
cacheCreation5mTokens, cacheCreation1hTokens := normalizeCacheCreationBreakdown(tokens)
if cacheCreation5mTokens == 0 && cacheCreation1hTokens == 0 && tokens.CacheCreationTokens > 0 {
// API 未返回 ephemeral 明细,回退到全部按 5m 单价计费
return float64(tokens.CacheCreationTokens) * pricing.CacheCreation5mPrice * multiplier
}
return float64(tokens.CacheCreation5mTokens)*pricing.CacheCreation5mPrice*multiplier +
float64(tokens.CacheCreation1hTokens)*pricing.CacheCreation1hPrice*multiplier
return float64(cacheCreation5mTokens)*pricing.CacheCreation5mPrice*multiplier +
float64(cacheCreation1hTokens)*pricing.CacheCreation1hPrice*multiplier
}
return float64(tokens.CacheCreationTokens) * price * multiplier
}
// normalizeCacheCreationBreakdown caps contradictory 5m/1h details at an explicitly
// positive aggregate while retaining their reported ratio as closely as integer tokens allow.
func normalizeCacheCreationBreakdown(tokens UsageTokens) (int, int) {
cacheCreation5mTokens := tokens.CacheCreation5mTokens
cacheCreation1hTokens := tokens.CacheCreation1hTokens
aggregate := tokens.CacheCreationTokens
if cacheCreation5mTokens < 0 {
cacheCreation5mTokens = 0
}
if cacheCreation1hTokens < 0 {
cacheCreation1hTokens = 0
}
if aggregate <= 0 || (cacheCreation5mTokens <= aggregate && cacheCreation1hTokens <= aggregate-cacheCreation5mTokens) {
return cacheCreation5mTokens, cacheCreation1hTokens
}
detailTotal := float64(cacheCreation5mTokens) + float64(cacheCreation1hTokens)
normalized5mTokens := math.Round(float64(aggregate) * float64(cacheCreation5mTokens) / detailTotal)
if normalized5mTokens >= float64(aggregate) {
cacheCreation5mTokens = aggregate
} else {
cacheCreation5mTokens = int(normalized5mTokens)
}
return cacheCreation5mTokens, aggregate - cacheCreation5mTokens
}
// calculatePerRequestCost 按次/图片计费
func (s *BillingService) calculatePerRequestCost(resolved *ResolvedPricing, input CostInput) (*CostBreakdown, error) {
units := input.UsageUnits
@@ -1384,6 +1384,122 @@ func TestCalculateCost_SupportsCacheBreakdown(t *testing.T) {
require.InDelta(t, expected5m+expected1h, cost.CacheCreationCost, 1e-10)
}
func TestComputeCacheCreationCost_CapsContradictoryBreakdownAtAggregate(t *testing.T) {
svc := &BillingService{}
pricing := &ModelPricing{
SupportsCacheBreakdown: true,
CacheCreation5mPrice: 1,
CacheCreation1hPrice: 1,
}
tokens := UsageTokens{
CacheCreationTokens: 463184,
CacheCreation5mTokens: 463184,
CacheCreation1hTokens: 463184,
}
cost := svc.computeCacheCreationCost(pricing, tokens, 0, 1)
require.Equal(t, float64(tokens.CacheCreationTokens), cost,
"billed cache-creation token equivalent must not exceed the positive aggregate")
}
func TestNormalizeCacheCreationBreakdown_BillingSafetyInvariant(t *testing.T) {
tests := []struct {
name string
tokens UsageTokens
want5m int
want1h int
}{
{
name: "preserves ratio when capping",
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: 90, CacheCreation1hTokens: 60},
want5m: 60,
want1h: 40,
},
{
name: "details below aggregate unchanged",
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: 30, CacheCreation1hTokens: 60},
want5m: 30,
want1h: 60,
},
{
name: "absent 5m detail unchanged",
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation1hTokens: 60},
want5m: 0,
want1h: 60,
},
{
name: "absent 1h detail unchanged",
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: 30},
want5m: 30,
want1h: 0,
},
{
name: "negative detail clamped",
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: -50, CacheCreation1hTokens: 60},
want5m: 0,
want1h: 60,
},
{
name: "negative detail cannot hide oversized positive detail",
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: -50, CacheCreation1hTokens: 150},
want5m: 0,
want1h: 100,
},
{
name: "integer boundary details capped without overflow",
tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: int(^uint(0) >> 1), CacheCreation1hTokens: int(^uint(0) >> 1)},
want5m: 50,
want1h: 50,
},
{
name: "integer boundary aggregate avoids float conversion overflow",
tokens: UsageTokens{CacheCreationTokens: int(^uint(0) >> 1), CacheCreation5mTokens: int(^uint(0) >> 1), CacheCreation1hTokens: 1},
want5m: int(^uint(0) >> 1),
want1h: 0,
},
{
name: "zero aggregate unchanged",
tokens: UsageTokens{CacheCreation5mTokens: 90, CacheCreation1hTokens: 60},
want5m: 90,
want1h: 60,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got5m, got1h := normalizeCacheCreationBreakdown(tt.tokens)
require.Equal(t, tt.want5m, got5m)
require.Equal(t, tt.want1h, got1h)
})
}
}
func TestComputeCacheCreationCost_PreservesZeroDetailFallback(t *testing.T) {
svc := &BillingService{}
pricing := &ModelPricing{
SupportsCacheBreakdown: true,
CacheCreation5mPrice: 4e-6,
CacheCreation1hPrice: 5e-6,
}
tests := []struct {
name string
tokens UsageTokens
}{
{name: "zero details", tokens: UsageTokens{CacheCreationTokens: 100}},
{name: "one negative detail", tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: -25}},
{name: "both negative details", tokens: UsageTokens{CacheCreationTokens: 100, CacheCreation5mTokens: -25, CacheCreation1hTokens: -75}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cost := svc.computeCacheCreationCost(pricing, tt.tokens, 0, 1)
require.InDelta(t, 100*4e-6, cost, 1e-12)
})
}
}
func TestCalculateCost_LargeTokenCount(t *testing.T) {
svc := newTestBillingService()
@@ -1295,21 +1295,20 @@ func TestGatewayService_ParseSSEUsagePassthrough_MessageStartFallbacks(t *testin
}
func TestGatewayService_ParseSSEUsagePassthrough_MessageDeltaSelectiveOverwrite(t *testing.T) {
usage := &ClaudeUsage{
InputTokens: 10,
CacheCreation5mTokens: 2,
CacheCreation1hTokens: 6,
}
data := `{"type":"message_delta","usage":{"input_tokens":0,"output_tokens":5,"cache_creation_input_tokens":8,"cache_read_input_tokens":0,"cached_tokens":11,"cache_creation":{"ephemeral_5m_input_tokens":1,"ephemeral_1h_input_tokens":0}}}`
usage := &ClaudeUsage{}
start := `{"type":"message_start","message":{"usage":{"input_tokens":10,"cache_creation_input_tokens":463184,"cache_creation":{"ephemeral_5m_input_tokens":0,"ephemeral_1h_input_tokens":463184}}}}`
parseSSEUsagePassthrough(start, usage)
data := `{"type":"message_delta","usage":{"input_tokens":0,"output_tokens":5,"cache_creation_input_tokens":463184,"cache_read_input_tokens":0,"cached_tokens":11,"cache_creation":{"ephemeral_5m_input_tokens":463184,"ephemeral_1h_input_tokens":0}}}`
parseSSEUsagePassthrough(data, usage)
require.Equal(t, 10, usage.InputTokens, "message_delta 中 0 值不应覆盖已有 input_tokens")
require.Equal(t, 5, usage.OutputTokens)
require.Equal(t, 8, usage.CacheCreationInputTokens)
require.Equal(t, 463184, usage.CacheCreationInputTokens)
require.Equal(t, 11, usage.CacheReadInputTokens, "cache_read_input_tokens 为空时应回退到 cached_tokens")
require.Equal(t, 1, usage.CacheCreation5mTokens)
require.Equal(t, 6, usage.CacheCreation1hTokens, "message_delta 中 0 值不应覆盖已有 1h 明细")
require.Equal(t, 463184, usage.CacheCreation5mTokens)
require.Equal(t, 0, usage.CacheCreation1hTokens)
}
func TestGatewayService_ParseSSEUsagePassthrough_NoopCases(t *testing.T) {
@@ -669,10 +669,10 @@ func parseSSEUsagePassthrough(data string, usage *ClaudeUsage) {
cc5m := deltaUsage.Get("cache_creation.ephemeral_5m_input_tokens")
cc1h := deltaUsage.Get("cache_creation.ephemeral_1h_input_tokens")
if cc5m.Exists() && cc5m.Int() > 0 {
if cc5m.Exists() {
usage.CacheCreation5mTokens = int(cc5m.Int())
}
if cc1h.Exists() && cc1h.Int() > 0 {
if cc1h.Exists() {
usage.CacheCreation1hTokens = int(cc1h.Int())
}
}
@@ -82,20 +82,19 @@ func TestParseSSEUsage_DeltaOverwritesWithNonZero(t *testing.T) {
require.Equal(t, 60, usage.CacheReadInputTokens)
}
func TestParseSSEUsage_DeltaDoesNotResetCacheCreationBreakdown(t *testing.T) {
func TestParseSSEUsage_DeltaAuthoritativelyUpdatesCacheCreationBreakdown(t *testing.T) {
svc := newMinimalGatewayService()
usage := &ClaudeUsage{}
// 先在 message_start 中写入非零 5m/1h 明细
svc.parseSSEUsage(`{"type":"message_start","message":{"usage":{"input_tokens":100,"cache_creation":{"ephemeral_5m_input_tokens":30,"ephemeral_1h_input_tokens":70}}}}`, usage)
require.Equal(t, 30, usage.CacheCreation5mTokens)
require.Equal(t, 70, usage.CacheCreation1hTokens)
svc.parseSSEUsage(`{"type":"message_start","message":{"usage":{"cache_creation_input_tokens":463184,"cache_creation":{"ephemeral_5m_input_tokens":0,"ephemeral_1h_input_tokens":463184}}}}`, usage)
require.Equal(t, 463184, usage.CacheCreationInputTokens)
require.Equal(t, 0, usage.CacheCreation5mTokens)
require.Equal(t, 463184, usage.CacheCreation1hTokens)
// 后续 delta 带默认 0,不应覆盖已有非零值
svc.parseSSEUsage(`{"type":"message_delta","usage":{"output_tokens":12,"cache_creation":{"ephemeral_5m_input_tokens":0,"ephemeral_1h_input_tokens":0}}}`, usage)
require.Equal(t, 30, usage.CacheCreation5mTokens, "delta 的 0 值不应重置 5m 明细")
require.Equal(t, 70, usage.CacheCreation1hTokens, "delta 的 0 值不应重置 1h 明细")
require.Equal(t, 12, usage.OutputTokens)
svc.parseSSEUsage(`{"type":"message_delta","usage":{"cache_creation_input_tokens":463184,"cache_creation":{"ephemeral_5m_input_tokens":463184,"ephemeral_1h_input_tokens":0}}}`, usage)
require.Equal(t, 463184, usage.CacheCreationInputTokens)
require.Equal(t, 463184, usage.CacheCreation5mTokens)
require.Equal(t, 0, usage.CacheCreation1hTokens)
}
func TestParseSSEUsage_InvalidJSON(t *testing.T) {
@@ -1237,11 +1237,11 @@ func (s *GatewayService) extractSSEUsagePatch(event map[string]any) *sseUsagePat
patch.hasCacheReadInput = true
}
if cc, ok := usageObj["cache_creation"].(map[string]any); ok {
if v, exists := parseSSEUsageInt(cc["ephemeral_5m_input_tokens"]); exists && v > 0 {
if v, exists := parseSSEUsageInt(cc["ephemeral_5m_input_tokens"]); exists {
patch.cacheCreation5mTokens = v
patch.hasCacheCreation5m = true
}
if v, exists := parseSSEUsageInt(cc["ephemeral_1h_input_tokens"]); exists && v > 0 {
if v, exists := parseSSEUsageInt(cc["ephemeral_1h_input_tokens"]); exists {
patch.cacheCreation1hTokens = v
patch.hasCacheCreation1h = true
}