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) } }) } }