Files
NianAIGC/backend/internal/localstore/billing_defaults_test.go
2026-09-23 13:33:47 +08:00

189 lines
7.2 KiB
Go

package localstore_test
import (
"context"
"encoding/json"
"fmt"
"reflect"
"testing"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/localstore"
)
func TestPriceReseedingPreservesTierMarkupsAndRefreshesDefinitions(t *testing.T) {
ctx := context.Background()
store := localstore.New()
old := billing.PriceRule{
ID: "price", Provider: "test", Capability: "image.generate", Unit: billing.UnitImage,
StandardUnitPriceFen: 30, MarkupMultiplier: 1.8, Enabled: true,
Dimensions: []billing.ParameterDimension{
{Key: "size", Label: "Old size", BaselineValue: "1K", Tiers: []billing.ParameterTier{
{Value: "1K", Label: "Old tier", StandardFactor: 1, MarkupMultiplier: 1.3, Enabled: true},
{Value: "removed", StandardFactor: 1, MarkupMultiplier: 1.7, Enabled: true},
}},
{Key: "referenceImageCount", BaselineValue: 0, Tiers: []billing.ParameterTier{
{Value: 0, StandardFactor: 1, MarkupMultiplier: 1.9, Enabled: true},
}},
{Key: "removed", BaselineValue: "1K", Tiers: []billing.ParameterTier{
{Value: "1K", StandardFactor: 1, MarkupMultiplier: 5, Enabled: true},
}},
},
}
if err := store.SeedBillingPriceRules(ctx, []billing.PriceRule{old}); err != nil {
t.Fatal(err)
}
next := billing.PriceRule{
ID: old.ID, Provider: old.Provider, Capability: old.Capability, Unit: old.Unit,
StandardUnitPriceFen: 40, MarkupMultiplier: 1.2, Enabled: true, Note: "New official cost",
Dimensions: []billing.ParameterDimension{
{Key: "referenceImageCount", BaselineValue: "0", Tiers: []billing.ParameterTier{
{Value: "0", StandardFactor: 1, MarkupMultiplier: 1.2, Enabled: true},
}},
{Key: "size", Label: "New size", BaselineValue: "2K", DefaultValue: "2K", Tiers: []billing.ParameterTier{
{Value: "2K", Label: "New tier", StandardFactor: 4, MarkupMultiplier: 1.2, Enabled: true},
{Value: "1K", Label: "New label", StandardFactor: 2, MarkupMultiplier: 1.2, Enabled: false},
}},
{Key: "new", BaselineValue: "1K", Tiers: []billing.ParameterTier{
{Value: "1K", StandardFactor: 1, MarkupMultiplier: 1.2, Enabled: true},
}},
},
}
before, _ := json.Marshal(next)
for i := 0; i < 3; i++ {
if err := store.SeedBillingPriceRules(ctx, []billing.PriceRule{next}); err != nil {
t.Fatal(err)
}
got, err := store.GetBillingPriceRule(ctx, old.ID)
if err != nil || got == nil {
t.Fatalf("rule = %#v, %v", got, err)
}
want := next
want.MarkupMultiplier = old.MarkupMultiplier
want.CreatedAt, want.UpdatedAt = got.CreatedAt, got.UpdatedAt
// Decode into a fresh value so editing the expectation cannot mutate defaults.
raw, _ := json.Marshal(want)
want = billing.PriceRule{}
if err := json.Unmarshal(raw, &want); err != nil {
t.Fatal(err)
}
want.Dimensions[0].Tiers[0].MarkupMultiplier = 1.9
want.Dimensions[1].Tiers[1].MarkupMultiplier = 1.3
if !reflect.DeepEqual(*got, want) {
t.Fatalf("reseed %d: got %#v, want %#v", i, got, want)
}
}
after, _ := json.Marshal(next)
if string(before) != string(after) {
t.Fatal("seeding mutated the supplied default catalog")
}
}
func TestSeedreamTierUpdateSurvivesQuotesAndCharges(t *testing.T) {
ctx := context.Background()
store := localstore.New()
service := billing.NewService(store, nil)
if err := store.SeedBillingPriceRules(ctx, billing.DefaultBillingPriceRules()); err != nil {
t.Fatal(err)
}
for _, update := range []struct {
id string
markup float64
}{
{"base-seedream-5-0-pro", 1.3},
{"base-seedream-5-0-pro-layers", 2},
} {
rule, err := service.GetPrice(ctx, update.id)
if err != nil || rule == nil {
t.Fatalf("get rule = %#v, %v", rule, err)
}
patch := billing.PricePatch{DimensionKey: "size", TierValue: "1.5K", MarkupMultiplier: update.markup}
if err := billing.ValidatePricePatch(rule, patch); err != nil {
t.Fatal(err)
}
if _, err := service.UpdatePrice(ctx, rule.ID, patch); err != nil {
t.Fatal(err)
}
}
if _, err := store.PostWalletEntry(ctx, billing.WalletPostParams{
LedgerID: "top-up", OrganizationID: "org", Kind: "recharge", DeltaFen: 10000,
Currency: billing.CurrencyCNY, IdempotencyKey: "top-up",
}); err != nil {
t.Fatal(err)
}
ledger := billing.Ledger{Poster: store, NewID: func() string { return "charge" }}
for i := 0; i < 3; i++ {
basic, err := service.Quote(ctx, billing.QuoteCommand{
Provider: "seedream", Capability: "image.generate", ReqKey: billing.Seedream50ProModel,
Parameters: billing.Parameters{"size": "1.5K"},
})
if err != nil || basic == nil || basic.MarkupMultiplier != 1.3 || basic.AmountFen != 39 {
t.Fatalf("basic quote %d = %#v, %v", i, basic, err)
}
layers, err := service.Quote(ctx, billing.QuoteCommand{
Provider: "seedream", Capability: "image.generate", ReqKey: billing.Seedream50ProModel,
Parameters: billing.Parameters{"size": "1.5K", "layerDecomposition": true},
})
if err != nil || layers == nil || layers.MarkupMultiplier != 2 || layers.ReservedAmountFen != 510 {
t.Fatalf("layers quote %d = %#v, %v", i, layers, err)
}
actual, err := billing.CalculateSeedreamLayerAmountFen([]billing.SeedreamLayerImage{
{Width: 1024, Height: 1024}, {Width: 2048, Height: 2048},
}, layers.MarkupMultiplier)
if err != nil || actual != 90 {
t.Fatalf("layer settlement amount = %d, %v", actual, err)
}
other, err := service.Quote(ctx, billing.QuoteCommand{
Provider: "seedream", Capability: "image.generate", ReqKey: billing.Seedream50ProModel,
Parameters: billing.Parameters{"size": "2K"},
})
if err != nil || other == nil || other.MarkupMultiplier != 1.2 || other.AmountFen != 72 {
t.Fatalf("unmodified tier = %#v, %v", other, err)
}
posting, err := ledger.Charge(ctx, billing.ChargeRequest{
OrganizationID: "org", AccountID: "member", JobID: fmt.Sprintf("job-%d", i), AmountFen: basic.AmountFen,
})
if err != nil || posting.DeltaFen != -39 || posting.BalanceFen != 10000-int64(i+1)*39 {
t.Fatalf("charge = %#v, %v", posting, err)
}
// A catalog list refresh must still expose the same saved multiplier.
rules, err := service.ListPrices(ctx)
if err != nil {
t.Fatal(err)
}
for _, rule := range rules {
if rule.ID == "base-seedream-5-0-pro" && rule.Dimensions[0].Tiers[2].MarkupMultiplier != 1.3 {
t.Fatal("catalog refresh lost the saved markup")
}
}
}
}
func TestPriceReseedingUsesDefaultsForInvalidSavedTierMarkup(t *testing.T) {
for _, markup := range []float64{0, -1, 1001} {
t.Run(fmt.Sprint(markup), func(t *testing.T) {
ctx := context.Background()
store := localstore.New()
rules := billing.DefaultBillingPriceRules()
if err := store.SeedBillingPriceRules(ctx, rules); err != nil {
t.Fatal(err)
}
if _, err := store.UpdateBillingPriceRule(ctx, "base-seedream-5-0-pro", billing.PricePatch{
DimensionKey: "size", TierValue: "1.5K", MarkupMultiplier: markup,
}); err != nil {
t.Fatal(err)
}
if err := store.SeedBillingPriceRules(ctx, rules); err != nil {
t.Fatal(err)
}
rule, err := store.GetBillingPriceRule(ctx, "base-seedream-5-0-pro")
if err != nil || rule == nil {
t.Fatalf("get price = %#v, %v", rule, err)
}
if got := rule.Dimensions[0].Tiers[2].MarkupMultiplier; got != 1.2 {
t.Fatalf("invalid saved markup %v survived reseed as %v", markup, got)
}
})
}
}