feat(pricing): pickModelPrice 按输入媒体取媒体价/默认价
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
// 用量提取函数
|
||||
|
||||
@@ -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 只有 audio,video 输入落默认价
|
||||
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.02;50000 → 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)
|
||||
}
|
||||
}
|
||||
|
||||
// M3:cached > 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)
|
||||
}
|
||||
}
|
||||
|
||||
// P2:validate 拒绝 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)
|
||||
}
|
||||
}
|
||||
|
||||
// P1:validate 拒绝缺 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")
|
||||
}
|
||||
}
|
||||
|
||||
// P1:validate 拒绝缺 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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user