package localstore_test import ( "context" "errors" "testing" "git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing" "git.nianxx.cn/wangxuming/NianAIGC/backend/internal/localstore" ) func TestBulkPriceUpdateIsScopedAtomicAndSurvivesSingleEditAndReseed(t *testing.T) { ctx := context.Background() store := localstore.New() defaults := []billing.PriceRule{ {ID: "a", Provider: "a", Capability: "image.generate", Unit: billing.UnitImage, StandardUnitPriceFen: 20, MarkupMultiplier: 1.2, Enabled: true, Dimensions: []billing.ParameterDimension{{Key: "quality", Tiers: []billing.ParameterTier{{Value: "low", StandardFactor: 1, MarkupMultiplier: 1.2, Enabled: true}, {Value: "high", StandardFactor: 2, MarkupMultiplier: 1.2, Enabled: false}}}}}, {ID: "b", Provider: "b", Capability: "image.generate", Unit: billing.UnitImage, StandardUnitPriceFen: 30, MarkupMultiplier: 1.2, Enabled: false, Dimensions: []billing.ParameterDimension{{Key: "quality", Tiers: []billing.ParameterTier{{Value: "low", StandardFactor: 1, MarkupMultiplier: 1.2, Enabled: true}}}}}, {ID: "c", Provider: "c", Capability: "image.generate", Unit: billing.UnitImage, StandardUnitPriceFen: 40, MarkupMultiplier: 1.2, Enabled: true}, } if err := store.SeedBillingPriceRules(ctx, defaults); err != nil { t.Fatal(err) } if _, err := store.UpdateBillingPriceRules(ctx, billing.BulkPricePatch{RuleIDs: []string{"a", "missing"}, MarkupMultiplier: 2}); !errors.Is(err, billing.ErrPriceNotFound) { t.Fatalf("missing ID error = %v", err) } if _, err := store.UpdateBillingPriceRules(ctx, billing.BulkPricePatch{MarkupMultiplier: 2}); !errors.Is(err, billing.ErrPriceNotFound) { t.Fatalf("empty ID error = %v", err) } before, _ := store.GetBillingPriceRule(ctx, "a") if before.MarkupMultiplier != 1.2 || before.Dimensions[0].Tiers[0].MarkupMultiplier != 1.2 { t.Fatalf("partial write on missing ID: %+v", before) } updated, err := store.UpdateBillingPriceRules(ctx, billing.BulkPricePatch{RuleIDs: []string{"a", "b"}, MarkupMultiplier: 1.8}) if err != nil || len(updated) != 2 { t.Fatalf("bulk update = %+v, %v", updated, err) } for _, id := range []string{"a", "b"} { rule, _ := store.GetBillingPriceRule(ctx, id) if rule.MarkupMultiplier != 1.8 { t.Fatalf("%s rule markup = %g", id, rule.MarkupMultiplier) } for _, dimension := range rule.Dimensions { for _, tier := range dimension.Tiers { if tier.MarkupMultiplier != 1.8 { t.Fatalf("%s tier markup = %g", id, tier.MarkupMultiplier) } } } } a, _ := store.GetBillingPriceRule(ctx, "a") if a.StandardUnitPriceFen != 20 || a.Dimensions[0].Tiers[1].Enabled || a.Dimensions[0].Tiers[1].StandardFactor != 2 { t.Fatalf("bulk update changed non-markup fields: %+v", a) } c, _ := store.GetBillingPriceRule(ctx, "c") if c.MarkupMultiplier != 1.2 { t.Fatalf("unselected rule changed: %+v", c) } if _, err := store.UpdateBillingPriceRule(ctx, "a", billing.PricePatch{DimensionKey: "quality", TierValue: "low", MarkupMultiplier: 2.3}); err != nil { t.Fatal(err) } if err := store.SeedBillingPriceRules(ctx, defaults); err != nil { t.Fatal(err) } a, _ = store.GetBillingPriceRule(ctx, "a") b, _ := store.GetBillingPriceRule(ctx, "b") if a.MarkupMultiplier != 1.8 || a.Dimensions[0].Tiers[0].MarkupMultiplier != 2.3 || a.Dimensions[0].Tiers[1].MarkupMultiplier != 1.8 || b.MarkupMultiplier != 1.8 || b.Dimensions[0].Tiers[0].MarkupMultiplier != 1.8 { t.Fatalf("single edit or bulk update lost after reseed: a=%+v b=%+v", a, b) } }