feat(pricing): pickModelPrice 按输入媒体取媒体价/默认价

This commit is contained in:
2026-09-01 14:55:22 +08:00
parent 4260125040
commit 1825ea5581
2 changed files with 238 additions and 2 deletions
+13 -2
View File
@@ -278,6 +278,16 @@ func matchModelRule(r *modelRules, u *ChargeUsage) *modelRule {
return nil
}
// pickModelPrice 命中规则后按输入媒体取价:mediaPrices 命中使用媒体价,否则默认 price。
func pickModelPrice(rule *modelRule, usage *ChargeUsage) *modelPrice {
if len(rule.MediaPrices) > 0 {
if p, ok := rule.MediaPrices[usage.MediaType]; ok {
return p
}
}
return rule.Price
}
// ====================== model token 计算器(per_1K / per_1M ======================
// modelTokenCalculator 模型 token 计价:首条命中规则后按基准换算。
@@ -353,7 +363,7 @@ func (c modelTokenCalculator) charge(rules interface{}, usage *ChargeUsage) (flo
if rule == nil {
return 0, errors.New("无匹配计费规则")
}
p := rule.Price
p := pickModelPrice(rule, usage)
// 防负成本:异常/边界用量报 cached > prompt 时钳为 0,避免结算出现负数(M3)。
promptInput := usage.PromptTokens - usage.CachedTokens
if promptInput < 0 {
@@ -402,7 +412,8 @@ func (c modelUnitCalculator) charge(rules interface{}, usage *ChargeUsage) (floa
if rule == nil {
return 0, errors.New("无匹配计费规则")
}
return ceilFen(roundCost(c.usage(usage) / c.base * rule.Price.UnitPrice)), nil
p := pickModelPrice(rule, usage)
return ceilFen(roundCost(c.usage(usage) / c.base * p.UnitPrice)), nil
}
// 用量提取函数
+225
View File
@@ -108,3 +108,228 @@ func TestUnitValidateMediaPricesValid(t *testing.T) {
t.Fatalf("合法 mediaPrices 应通过: %v", err)
}
}
func TestTokenChargeMediaPricesAudio(t *testing.T) {
c := modelTokenCalculator{base: 1e6}
rules, err := c.validate(`{"unit":"per_1M","rules":[{"name":"r","price":{"input":0.3,"output":1.8,"cacheHit":0.12},"mediaPrices":{"audio":{"input":4.5,"output":1.8,"cacheHit":1.8}}}]}`)
if err != nil {
t.Fatalf("validate: %v", err)
}
got, err := c.charge(rules, &ChargeUsage{PromptTokens: 1e6, MediaType: "audio"})
if err != nil {
t.Fatalf("charge: %v", err)
}
if got != 4.5 {
t.Fatalf("含音频应收媒体价 4.5,实收 %v", got)
}
}
func TestTokenChargeDefaultWithoutMedia(t *testing.T) {
c := modelTokenCalculator{base: 1e6}
rules, _ := c.validate(`{"unit":"per_1M","rules":[{"name":"r","price":{"input":0.3,"output":1.8,"cacheHit":0.12},"mediaPrices":{"audio":{"input":4.5,"output":1.8,"cacheHit":1.8}}}]}`)
got, err := c.charge(rules, &ChargeUsage{PromptTokens: 1e6, MediaType: "text"})
if err != nil {
t.Fatalf("charge: %v", err)
}
if got != 0.3 {
t.Fatalf("不含媒体应收默认价 0.3,实收 %v", got)
}
}
func TestTokenChargeMediaPricesMultiType(t *testing.T) {
c := modelTokenCalculator{base: 1000}
rules, _ := c.validate(`{"unit":"per_1K","rules":[{"name":"r","price":{"input":0.3},"mediaPrices":{"audio":{"input":4.5},"video":{"input":6}}}]}`)
got, err := c.charge(rules, &ChargeUsage{PromptTokens: 1000, MediaType: "video"})
if err != nil {
t.Fatalf("charge: %v", err)
}
if got != 6 {
t.Fatalf("video 应收媒体价 6,实收 %v", got)
}
}
func TestTokenChargeMediaPricesUnmatchedTypeUsesDefault(t *testing.T) {
// mediaPrices 只有 audiovideo 输入落默认价
c := modelTokenCalculator{base: 1000}
rules, _ := c.validate(`{"unit":"per_1K","rules":[{"name":"r","price":{"input":0.3},"mediaPrices":{"audio":{"input":4.5}}}]}`)
got, err := c.charge(rules, &ChargeUsage{PromptTokens: 1000, MediaType: "video"})
if err != nil {
t.Fatalf("charge: %v", err)
}
if got != 0.3 {
t.Fatalf("未配 mediaPrices 的媒体应收默认价 0.3,实收 %v", got)
}
}
func TestUnitChargeMediaPrices(t *testing.T) {
c := modelUnitCalculator{base: 1, usage: usageDurationSec}
rules, err := c.validate(`{"unit":"per_second","rules":[{"name":"r","price":{"unitPrice":1},"mediaPrices":{"audio":{"unitPrice":2}}}]}`)
if err != nil {
t.Fatalf("validate: %v", err)
}
got, err := c.charge(rules, &ChargeUsage{DurationSec: 3, MediaType: "audio"})
if err != nil {
t.Fatalf("charge: %v", err)
}
if got != 6 {
t.Fatalf("audio 单位价应收 6,实收 %v", got)
}
}
func TestUnitChargeMediaPricesDefault(t *testing.T) {
c := modelUnitCalculator{base: 1, usage: usageDurationSec}
rules, _ := c.validate(`{"unit":"per_second","rules":[{"name":"r","price":{"unitPrice":1},"mediaPrices":{"audio":{"unitPrice":2}}}]}`)
got, err := c.charge(rules, &ChargeUsage{DurationSec: 3, MediaType: "video"})
if err != nil {
t.Fatalf("charge: %v", err)
}
if got != 3 {
t.Fatalf("未配媒体价应收默认 3,实收 %v", got)
}
}
// 阶梯算费 fixture(原 fixture 的 mediaType:"text" 已随 match 移除)
const modelTokenTieredJSON = `{
"unit":"per_1M","tiered":true,
"rules":[
{"name":"思考≤32000","match":{"thinking":true,"inputLengthMax":32000},"price":{"input":0.6,"output":3.6,"cacheHit":0.12}},
{"name":"思考>32000","match":{"thinking":true,"inputLengthMin":32001},"price":{"input":0.9,"output":5.4,"cacheHit":0.18}}
]}`
// 阶梯算费:prompt 20000 → (0.6*20000+3.6*1000)/1e6=0.0156→0.0250000 → 0.0504→0.06
func TestModelTokenCalculator_TieredMatch(t *testing.T) {
c := modelTokenCalculator{base: 1e6}
rules, err := c.validate(modelTokenTieredJSON)
if err != nil {
t.Fatalf("validate err: %v", err)
}
got, err := c.charge(rules, &ChargeUsage{PromptTokens: 20000, CompletionTokens: 1000, Thinking: boolPtr(true)})
if err != nil {
t.Fatalf("charge err: %v", err)
}
if got != 0.02 {
t.Fatalf("tier1 got %v want 0.02", got)
}
got, err = c.charge(rules, &ChargeUsage{PromptTokens: 50000, CompletionTokens: 1000, Thinking: boolPtr(true)})
if err != nil {
t.Fatalf("charge err: %v", err)
}
if got != 0.06 {
t.Fatalf("tier2 got %v want 0.06", got)
}
}
// 无匹配报错:thinking=false 不在任一档
func TestModelTokenCalculator_NoMatch(t *testing.T) {
c := modelTokenCalculator{base: 1e6}
rules, err := c.validate(modelTokenTieredJSON)
if err != nil {
t.Fatalf("validate err: %v", err)
}
_, err = c.charge(rules, &ChargeUsage{PromptTokens: 1000, Thinking: boolPtr(false)})
if err == nil {
t.Fatal("expected no-match error")
}
}
// cache-hit 算费:(2000-1000)/1000*1 + 1000/1000*0.1 + 500/1000*2 = 2.1
func TestModelTokenCalculator_CacheHit(t *testing.T) {
c := modelTokenCalculator{base: 1000}
rules, err := c.validate(`{"unit":"per_1K","tiered":false,"rules":[{"name":"统一","price":{"input":1,"output":2,"cacheHit":0.1}}]}`)
if err != nil {
t.Fatalf("validate err: %v", err)
}
got, err := c.charge(rules, &ChargeUsage{PromptTokens: 2000, CompletionTokens: 500, CachedTokens: 1000})
if err != nil {
t.Fatalf("charge err: %v", err)
}
if got != 2.1 {
t.Fatalf("cache got %v want 2.1", got)
}
}
// M3cached > prompt 时金额不为负(promptInput 钳为 0):0 + 3000/1000*0.1 = 0.3
func TestModelTokenCalculator_CachedGreaterThanPrompt(t *testing.T) {
c := modelTokenCalculator{base: 1000}
rules, err := c.validate(`{"unit":"per_1K","tiered":false,"rules":[{"name":"统一","price":{"input":1,"output":2,"cacheHit":0.1}}]}`)
if err != nil {
t.Fatalf("validate err: %v", err)
}
got, err := c.charge(rules, &ChargeUsage{PromptTokens: 1000, CompletionTokens: 0, CachedTokens: 3000})
if err != nil {
t.Fatalf("charge err: %v", err)
}
if got < 0 {
t.Fatalf("got negative amount %v", got)
}
if got != 0.3 {
t.Fatalf("got %v want 0.3", got)
}
}
// P2validate 拒绝 unit 与实例 base 不匹配的配置
func TestModelTokenCalculator_ValidateUnitBaseMismatch(t *testing.T) {
c := modelTokenCalculator{base: 1e6}
if _, err := c.validate(`{"unit":"per_1K","tiered":false,"rules":[{"name":"统一","price":{"input":1}}]}`); err == nil {
t.Fatal("expected unit/base mismatch error for per_1M instance with per_1K unit")
}
c2 := modelTokenCalculator{base: 1000}
if _, err := c2.validate(`{"unit":"per_1M","tiered":false,"rules":[{"name":"统一","price":{"input":1}}]}`); err == nil {
t.Fatal("expected unit/base mismatch error for per_1K instance with per_1M unit")
}
if _, err := c.validate(`{"unit":"per_1M","tiered":false,"rules":[{"name":"统一","price":{"input":1}}]}`); err != nil {
t.Fatalf("unexpected err: %v", err)
}
if _, err := c2.validate(`{"unit":"per_1K","tiered":false,"rules":[{"name":"统一","price":{"input":1}}]}`); err != nil {
t.Fatalf("unexpected err: %v", err)
}
}
// P1validate 拒绝缺 price 的规则(token 计算器)
func TestModelTokenCalculator_ValidateNilPrice(t *testing.T) {
c := modelTokenCalculator{base: 1e6}
if _, err := c.validate(`{"unit":"per_1M","tiered":false,"rules":[{"name":"无价"}]}`); err == nil {
t.Fatal("expected nil price validation error")
}
}
// P1validate 拒绝缺 price 的规则(单位计算器)
func TestModelUnitCalculator_ValidateNilPrice(t *testing.T) {
c := modelUnitCalculator{base: 60, usage: usageDurationSec}
if _, err := c.validate(`{"unit":"per_minute","tiered":false,"rules":[{"name":"无价"}]}`); err == nil {
t.Fatal("expected nil price validation error")
}
}
// 单位计算器多路径:按分钟(90s/60*2=3)、按字(10000*0.0005=5)、按张×分辨率(3*0.5=1.5
func TestModelUnitCalculator_Paths(t *testing.T) {
c := modelUnitCalculator{base: 60, usage: usageDurationSec}
rules, err := c.validate(`{"unit":"per_minute","tiered":false,"rules":[{"name":"统一","price":{"unitPrice":2}}]}`)
if err != nil {
t.Fatalf("validate err: %v", err)
}
got, err := c.charge(rules, &ChargeUsage{DurationSec: 90})
if err != nil || got != 3 {
t.Fatalf("per_minute got %v want 3, err %v", got, err)
}
c2 := modelUnitCalculator{base: 1, usage: usageCharCount}
rules2, err := c2.validate(`{"unit":"per_char","tiered":false,"rules":[{"name":"统一","price":{"unitPrice":0.0005}}]}`)
if err != nil {
t.Fatalf("validate err: %v", err)
}
got, err = c2.charge(rules2, &ChargeUsage{CharCount: 10000})
if err != nil || got != 5 {
t.Fatalf("per_char got %v want 5, err %v", got, err)
}
c3 := modelUnitCalculator{base: 1, usage: usageImageCount}
rules3, err := c3.validate(`{"unit":"per_1","tiered":false,"rules":[
{"name":"512x512","match":{"outputResolution":"512x512"},"price":{"unitPrice":0.2}},
{"name":"1024x1024","match":{"outputResolution":"1024x1024"},"price":{"unitPrice":0.5}}]}`)
if err != nil {
t.Fatalf("validate err: %v", err)
}
got, err = c3.charge(rules3, &ChargeUsage{ImageCount: 3, OutputResolution: "1024x1024"})
if err != nil || got != 1.5 {
t.Fatalf("per_1 got %v want 1.5, err %v", got, err)
}
}