增加统一倍率调整按钮

This commit is contained in:
andy committed 2026-09-23 15:46:21 +08:00
1 parent be655b6569
commit 4195525dc7
23 files changed
+750 -16

No files matched your search

@@ -0,0 +1,51 @@
package billing
import (
"context"
"errors"
"math"
"testing"
)
type bulkStoreStub struct {
*quoteStoreStub
patch BulkPricePatch
err error
calls int
}
func (s *bulkStoreStub) UpdateBillingPriceRules(_ context.Context, patch BulkPricePatch) ([]PriceRule, error) {
s.calls++
s.patch = patch
return []PriceRule{{ID: "a"}}, s.err
}
func TestUpdatePricesValidatesExplicitIDsAndNormalizesMultiplier(t *testing.T) {
store := &bulkStoreStub{quoteStoreStub: &quoteStoreStub{}}
service := NewService(store, nil)
for _, patch := range []BulkPricePatch{
{MarkupMultiplier: 1.2},
{RuleIDs: []string{""}, MarkupMultiplier: 1.2},
{RuleIDs: []string{" a"}, MarkupMultiplier: 1.2},
{RuleIDs: []string{"a", "a"}, MarkupMultiplier: 1.2},
{RuleIDs: []string{"a"}, MarkupMultiplier: 0.9},
{RuleIDs: []string{"a"}, MarkupMultiplier: 1000.1},
{RuleIDs: []string{"a"}, MarkupMultiplier: math.NaN()},
{RuleIDs: []string{"a"}, MarkupMultiplier: math.Inf(1)},
} {
if _, err := service.UpdatePrices(context.Background(), patch); HTTPStatus(err) != 400 {
t.Fatalf("patch %+v: status = %d, err = %v", patch, HTTPStatus(err), err)
}
}
if store.calls != 0 {
t.Fatalf("invalid patch reached store %d times", store.calls)
}
rules, err := service.UpdatePrices(context.Background(), BulkPricePatch{RuleIDs: []string{"a"}, MarkupMultiplier: 1.234567})
if err != nil || len(rules) != 1 || store.patch.MarkupMultiplier != 1.2346 {
t.Fatalf("rules = %+v, patch = %+v, err = %v", rules, store.patch, err)
}
store.err = ErrPriceNotFound
if _, err := service.UpdatePrices(context.Background(), BulkPricePatch{RuleIDs: []string{"a"}, MarkupMultiplier: 2}); HTTPStatus(err) != 404 || !errors.Is(err, ErrPriceNotFound) {
t.Fatalf("missing rule status = %d, err = %v", HTTPStatus(err), err)
}
}
+1 -1
View File
@@ -14,7 +14,7 @@ const (
// PriceRuleSeeder is an optional store capability. Existing Store
// implementations remain source-compatible and seed-capable stores can make
// the official base catalog durable before quotes are read.
// the official base catalog durable before quotes or administrator prices are read.
type PriceRuleSeeder interface {
SeedBillingPriceRules(context.Context, []PriceRule) error
}
+8
View File
@@ -87,6 +87,13 @@ type PricePatch struct {
MarkupMultiplier float64
DimensionKey, TierValue string
}
// BulkPricePatch replaces the markup on every tier of the explicitly named rules.
// An empty ID list never means the entire catalog.
type BulkPricePatch struct {
RuleIDs []string `json:"ruleIds"`
MarkupMultiplier float64 `json:"markupMultiplier"`
}
type AdjustmentCommand struct {
OrganizationID, OperatorID, Direction, Note string
AmountFen, DeltaFen int64
@@ -113,6 +120,7 @@ type HTTPService interface {
ReadService
Quote(context.Context, QuoteCommand) (*Quote, error)
UpdatePrice(context.Context, string, PricePatch) (*PriceRule, error)
UpdatePrices(context.Context, BulkPricePatch) ([]PriceRule, error)
Adjust(context.Context, AdjustmentCommand) (AdjustmentResult, error)
}
+52 -6
View File
@@ -24,6 +24,11 @@ type Store interface {
PostBillingWalletEntry(context.Context, WalletPostParams) (WalletPosting, error)
}
// BulkPriceRuleUpdater is optional so existing Store implementations remain valid.
type BulkPriceRuleUpdater interface {
UpdateBillingPriceRules(context.Context, BulkPricePatch) ([]PriceRule, error)
}
type Service struct {
store Store
newID func() string
@@ -39,7 +44,8 @@ func NewService(store Store, newID func() string) *Service {
// SetEnabled configures whether generation quotes are required. The default is
// enabled; disabling preserves the historical optional-billing behavior and
// does not consult or seed the price catalog.
// does not consult or seed the price catalog when quoting. Administrator catalog
// reads still refresh defaults so pricing can be configured before billing starts.
func (s *Service) SetEnabled(enabled bool) *Service {
s.enabled = enabled
return s
@@ -67,10 +73,8 @@ func (s *Service) Quote(ctx context.Context, command QuoteCommand) (*Quote, erro
if !s.enabled {
return nil, nil
}
if seeder, ok := s.store.(PriceRuleSeeder); ok {
if err := seeder.SeedBillingPriceRules(ctx, DefaultBillingPriceRules()); err != nil {
return nil, err
}
if err := s.ensurePriceRules(ctx); err != nil {
return nil, err
}
rules, err := s.store.ListBillingPriceRules(ctx, false)
if err != nil {
@@ -155,7 +159,7 @@ func (s *Service) AdminOverview(ctx context.Context) (AdminOverview, error) {
if err != nil {
return AdminOverview{}, err
}
rules, err := s.store.ListBillingPriceRules(ctx, true)
rules, err := s.ListPrices(ctx)
if err != nil {
return AdminOverview{}, err
}
@@ -185,14 +189,56 @@ func (s *Service) AdminOverview(ctx context.Context) (AdminOverview, error) {
return AdminOverview{Organizations: organizations, Members: members, Ledger: ledger, PriceRules: rules}, nil
}
func (s *Service) ListPrices(ctx context.Context) ([]PriceRule, error) {
if err := s.ensurePriceRules(ctx); err != nil {
return nil, err
}
return s.store.ListBillingPriceRules(ctx, true)
}
func (s *Service) ensurePriceRules(ctx context.Context) error {
if seeder, ok := s.store.(PriceRuleSeeder); ok {
return seeder.SeedBillingPriceRules(ctx, DefaultBillingPriceRules())
}
return nil
}
func (s *Service) GetPrice(ctx context.Context, id string) (*PriceRule, error) {
return s.store.GetBillingPriceRule(ctx, id)
}
func (s *Service) UpdatePrice(ctx context.Context, id string, patch PricePatch) (*PriceRule, error) {
return s.store.UpdateBillingPriceRule(ctx, id, patch)
}
func (s *Service) UpdatePrices(ctx context.Context, patch BulkPricePatch) ([]PriceRule, error) {
if len(patch.RuleIDs) == 0 || len(patch.RuleIDs) > 1000 {
return nil, &StatusError{400, errors.New("请选择 1 至 1000 条计费规则。")}
}
if err := ValidatePricePatch(&PriceRule{}, PricePatch{MarkupMultiplier: patch.MarkupMultiplier}); err != nil {
return nil, &StatusError{400, err}
}
seen := make(map[string]struct{}, len(patch.RuleIDs))
for _, id := range patch.RuleIDs {
if strings.TrimSpace(id) == "" || id != strings.TrimSpace(id) {
return nil, &StatusError{400, errors.New("计费规则 ID 不能为空或包含首尾空格。")}
}
if _, duplicate := seen[id]; duplicate {
return nil, &StatusError{400, errors.New("计费规则 ID 不能重复。")}
}
seen[id] = struct{}{}
}
patch.MarkupMultiplier = math.Round(patch.MarkupMultiplier*10000) / 10000
updater, ok := s.store.(BulkPriceRuleUpdater)
if !ok {
return nil, errors.New("批量调整计费倍率不可用")
}
if err := s.ensurePriceRules(ctx); err != nil {
return nil, err
}
rules, err := updater.UpdateBillingPriceRules(ctx, patch)
if errors.Is(err, ErrPriceNotFound) {
return nil, &StatusError{404, ErrPriceNotFound}
}
return rules, err
}
func (s *Service) Adjust(ctx context.Context, command AdjustmentCommand) (AdjustmentResult, error) {
exists, err := s.store.BillingOrganizationExists(ctx, command.OrganizationID)
if err != nil {
+27 -1
View File
@@ -2,6 +2,7 @@ package billing
import (
"context"
"errors"
"reflect"
"testing"
)
@@ -74,6 +75,30 @@ func TestServiceQuoteDisabledReturnsNilWithoutAccessingRules(t *testing.T) {
}
}
func TestServiceCatalogReadsReturnSeedErrors(t *testing.T) {
seedErr := errors.New("catalog refresh failed")
for name, read := range map[string]func(*Service) error{
"admin overview": func(service *Service) error {
_, err := service.AdminOverview(context.Background())
return err
},
"price list": func(service *Service) error {
_, err := service.ListPrices(context.Background())
return err
},
} {
t.Run(name, func(t *testing.T) {
store := &quoteStoreStub{seedErr: seedErr}
if err := read(NewService(store, nil)); !errors.Is(err, seedErr) {
t.Fatalf("error = %v, want %v", err, seedErr)
}
if store.listCalls != 0 {
t.Fatal("returned a stale price list after catalog refresh failed")
}
})
}
}
func TestServiceQuoteFreezesConservativeSeedanceReserve(t *testing.T) {
store := &quoteStoreStub{rules: []PriceRule{{
ID: "seedance-720", Provider: "seedance", Capability: "video.generate",
@@ -230,6 +255,7 @@ type quoteStoreStub struct {
rules []PriceRule
seeded []PriceRule
seedCalls int
seedErr error
listCalls int
failOnAccess bool
}
@@ -237,7 +263,7 @@ type quoteStoreStub struct {
func (s *quoteStoreStub) SeedBillingPriceRules(_ context.Context, rules []PriceRule) error {
s.seedCalls++
s.seeded = rules
return nil
return s.seedErr
}
func (s *quoteStoreStub) ListBillingPriceRules(context.Context, bool) ([]PriceRule, error) {
s.listCalls++
+34 -1
View File
@@ -3,6 +3,7 @@ package httpapi
import (
"encoding/json"
"errors"
"io"
"math"
"net/http"
"strconv"
@@ -232,12 +233,18 @@ func (h *billingHandler) adjust(w http.ResponseWriter, r *http.Request) {
writeJSON(w, 200, result)
}
func (h *billingHandler) prices(w http.ResponseWriter, r *http.Request) {
if !allow(w, r, http.MethodGet) {
if r.Method != http.MethodGet && r.Method != http.MethodPatch {
w.Header().Set("Allow", "GET, PATCH")
writeAPIError(w, http.StatusMethodNotAllowed, "方法不允许。")
return
}
if _, ok := h.authorize(w, r, PlatformSuperAdmin); !ok {
return
}
if r.Method == http.MethodPatch {
h.updatePrices(w, r)
return
}
rules, err := h.service.ListPrices(r.Context())
if err != nil {
writeDomainError(w, err)
@@ -245,6 +252,32 @@ func (h *billingHandler) prices(w http.ResponseWriter, r *http.Request) {
}
writeJSON(w, 200, map[string]any{"priceRules": rules})
}
// Called only after the collection route has authorized a platform super admin.
func (h *billingHandler) updatePrices(w http.ResponseWriter, r *http.Request) {
var body struct {
RuleIDs []string `json:"ruleIds"`
MarkupMultiplier any `json:"markupMultiplier"`
}
decoder := json.NewDecoder(r.Body)
decoder.DisallowUnknownFields()
if err := decoder.Decode(&body); err != nil {
writeAPIError(w, 400, "请提供规则列表和上浮倍率,不能修改标准价格或参数。")
return
}
if err := decoder.Decode(&struct{}{}); err != io.EOF {
writeAPIError(w, 400, "请求必须包含一个有效的 JSON 对象。")
return
}
rules, err := h.service.UpdatePrices(r.Context(), billing.BulkPricePatch{
RuleIDs: body.RuleIDs, MarkupMultiplier: numberValue(body.MarkupMultiplier),
})
if err != nil {
writeDomainError(w, err)
return
}
writeJSON(w, 200, map[string]any{"priceRules": rules})
}
func (h *billingHandler) price(w http.ResponseWriter, r *http.Request, id string) {
if !allow(w, r, http.MethodPatch) {
return
@@ -0,0 +1,109 @@
package httpapi
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/identity"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/localstore"
)
func TestBillingBulkPriceUpdateAndLaterIndividualEdit(t *testing.T) {
service := billing.NewService(localstore.New(), nil)
h := billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeAdmin, User: identity.User{ID: "root", ClientID: "platform", Role: "super_admin"}}, service, nil)
before, err := service.ListPrices(context.Background())
if err != nil {
t.Fatal(err)
}
ids := make([]string, len(before))
for i, rule := range before {
ids[i] = rule.ID
}
response := serveJSON(t, h, http.MethodPatch, "/api/admin/billing/prices", map[string]any{"ruleIds": ids, "markupMultiplier": "1.35001"})
if response.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", response.Code, response.Body.String())
}
var result struct {
Rules []billing.PriceRule `json:"priceRules"`
}
if err := json.Unmarshal(response.Body.Bytes(), &result); err != nil {
t.Fatal(err)
}
if len(result.Rules) != len(ids) {
t.Fatalf("updated %d rules, want %d", len(result.Rules), len(ids))
}
for _, rule := range result.Rules {
if rule.MarkupMultiplier != 1.35 {
t.Fatalf("%s markup=%v", rule.ID, rule.MarkupMultiplier)
}
for _, dimension := range rule.Dimensions {
for _, tier := range dimension.Tiers {
if tier.MarkupMultiplier != 1.35 {
t.Fatalf("%s/%s/%v markup=%v", rule.ID, dimension.Key, tier.Value, tier.MarkupMultiplier)
}
}
}
}
id := "base-evolink-gpt-image-2.5-flare"
response = serveJSON(t, h, http.MethodPatch, "/api/admin/billing/prices/"+id, map[string]any{"markupMultiplier": 2, "dimensionKey": "quality", "tierValue": "high"})
if response.Code != http.StatusOK {
t.Fatalf("single tier status=%d body=%s", response.Code, response.Body.String())
}
// A new catalog read refreshes defaults and must preserve the single override.
response = serveJSON(t, h, http.MethodGet, "/api/admin/billing/prices", nil)
if response.Code != http.StatusOK {
t.Fatalf("catalog reload status=%d body=%s", response.Code, response.Body.String())
}
rule, err := service.GetPrice(context.Background(), id)
if err != nil || rule == nil {
t.Fatalf("rule=%v err=%v", rule, err)
}
for _, dimension := range rule.Dimensions {
for _, tier := range dimension.Tiers {
want := 1.35
if dimension.Key == "quality" && tier.Value == "high" {
want = 2
}
if tier.MarkupMultiplier != want {
t.Fatalf("%s/%v markup=%v want=%v", dimension.Key, tier.Value, tier.MarkupMultiplier, want)
}
}
}
}
func TestBillingBulkPriceUpdateRequiresSuperAdminAndValidBody(t *testing.T) {
for _, role := range []string{"user", "organization_admin"} {
t.Run(role, func(t *testing.T) {
service := &billingHTTPServiceStub{}
h := billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeUser, User: identity.User{ID: "user", ClientID: "platform", OrganizationID: "org", Role: role}}, service, nil)
response := serveJSON(t, h, http.MethodPatch, "/api/admin/billing/prices", map[string]any{"ruleIds": []string{"price"}, "markupMultiplier": 2})
if response.Code != http.StatusForbidden || len(service.bulkPricePatch.RuleIDs) != 0 {
t.Fatalf("status=%d mutation=%+v", response.Code, service.bulkPricePatch)
}
})
}
for _, body := range []string{
`{"ruleIds":["price"],"markupMultiplier":2,"standardUnitPriceFen":1}`,
`{"ruleIds":[1],"markupMultiplier":2}`,
`{"ruleIds":"price","markupMultiplier":2}`,
`{"ruleIds":["price"],"markupMultiplier":2} {}`,
`{"ruleIds":`,
} {
t.Run(body, func(t *testing.T) {
service := &billingHTTPServiceStub{}
h := billingTestHandler(t, identity.Session{AuthMode: identity.AuthModeAdmin, User: identity.User{ID: "root", ClientID: "platform", Role: "super_admin"}}, service, nil)
request := httptest.NewRequest(http.MethodPatch, "/api/admin/billing/prices", strings.NewReader(body))
request.AddCookie(&http.Cookie{Name: identity.SessionCookieName, Value: "signed"})
response := httptest.NewRecorder()
h.ServeHTTP(response, request)
if response.Code != http.StatusBadRequest || len(service.bulkPricePatch.RuleIDs) != 0 {
t.Fatalf("status=%d mutation=%+v body=%s", response.Code, service.bulkPricePatch, response.Body.String())
}
})
}
}
+5
View File
@@ -138,6 +138,7 @@ type billingHTTPServiceStub struct {
quote billing.QuoteCommand
adjustment billing.AdjustmentCommand
pricePatch billing.PricePatch
bulkPricePatch billing.BulkPricePatch
price *billing.PriceRule
err error
}
@@ -163,6 +164,10 @@ func (s *billingHTTPServiceStub) UpdatePrice(_ context.Context, _ string, patch
s.pricePatch = patch
return s.price, s.err
}
func (s *billingHTTPServiceStub) UpdatePrices(_ context.Context, patch billing.BulkPricePatch) ([]billing.PriceRule, error) {
s.bulkPricePatch = patch
return []billing.PriceRule{}, s.err
}
func (s *billingHTTPServiceStub) Adjust(_ context.Context, command billing.AdjustmentCommand) (billing.AdjustmentResult, error) {
s.adjustment = command
return billing.AdjustmentResult{Wallet: billing.Wallet{OrganizationID: command.OrganizationID, Currency: billing.CurrencyCNY}}, s.err
@@ -41,6 +41,7 @@ func TestRouteMethodCompatibilityHandlesOptionsBeforeApplicationAuth(t *testing.
}{
{path: "/api/health", allow: "GET, HEAD, OPTIONS"},
{path: "/api/admin/accounts", allow: "DELETE, GET, HEAD, OPTIONS, PATCH, POST, PUT"},
{path: "/api/admin/billing/prices", allow: "GET, HEAD, OPTIONS, PATCH"},
{path: "/api/v1/jobs", allow: "GET, HEAD, OPTIONS, POST"},
{path: "/api/v1/jobs/job-1/cancel", allow: "OPTIONS, POST"},
{path: "/uploads/2026/08/file.png", allow: "GET, HEAD, OPTIONS"},
+1 -1
View File
@@ -13,7 +13,7 @@ type RouteSurface struct {
}
var goRouteSurface = []RouteSurface{
{"DELETE", "/api/admin/accounts"}, {"GET", "/api/admin/accounts"}, {"PATCH", "/api/admin/accounts"}, {"POST", "/api/admin/accounts"}, {"PUT", "/api/admin/accounts"}, {"POST", "/api/admin/accounts/groups"}, {"POST", "/api/admin/accounts/password"}, {"GET", "/api/admin/billing"}, {"PATCH", "/api/admin/billing/account"}, {"POST", "/api/admin/billing/adjustments"}, {"GET", "/api/admin/billing/prices"}, {"PATCH", "/api/admin/billing/prices/{id}"}, {"DELETE", "/api/admin/organizations"}, {"GET", "/api/admin/organizations"}, {"PATCH", "/api/admin/organizations"}, {"POST", "/api/admin/organizations"}, {"GET", "/api/admin/usage"},
{"DELETE", "/api/admin/accounts"}, {"GET", "/api/admin/accounts"}, {"PATCH", "/api/admin/accounts"}, {"POST", "/api/admin/accounts"}, {"PUT", "/api/admin/accounts"}, {"POST", "/api/admin/accounts/groups"}, {"POST", "/api/admin/accounts/password"}, {"GET", "/api/admin/billing"}, {"PATCH", "/api/admin/billing/account"}, {"POST", "/api/admin/billing/adjustments"}, {"GET", "/api/admin/billing/prices"}, {"PATCH", "/api/admin/billing/prices"}, {"PATCH", "/api/admin/billing/prices/{id}"}, {"DELETE", "/api/admin/organizations"}, {"GET", "/api/admin/organizations"}, {"PATCH", "/api/admin/organizations"}, {"POST", "/api/admin/organizations"}, {"GET", "/api/admin/usage"},
{"GET", "/api/assets"}, {"POST", "/api/assets"}, {"POST", "/api/assets/upload"}, {"DELETE", "/api/assets/{id}"}, {"GET", "/api/assets/{id}/download"},
{"GET", "/api/auth/callback"}, {"GET", "/api/auth/captcha"}, {"GET", "/api/auth/login"}, {"GET", "/api/auth/logout"}, {"POST", "/api/auth/logout"}, {"GET", "/api/auth/me"}, {"POST", "/api/auth/password"}, {"POST", "/api/auth/password/change"},
{"GET", "/api/billing"}, {"POST", "/api/billing/quote"}, {"GET", "/api/generations/image"}, {"POST", "/api/generations/image"}, {"DELETE", "/api/generations/image/{id}"}, {"GET", "/api/generations/image/{id}"}, {"POST", "/api/generations/image/{id}/retry"}, {"GET", "/api/generations/video"}, {"POST", "/api/generations/video"}, {"DELETE", "/api/generations/video/{id}"}, {"GET", "/api/generations/video/{id}"},
+27
View File
@@ -100,6 +100,33 @@ func (s *Store) UpdateBillingPriceRule(_ context.Context, id string, p billing.P
r = cloneRule(r)
return &r, nil
}
func (s *Store) UpdateBillingPriceRules(_ context.Context, patch billing.BulkPricePatch) ([]billing.PriceRule, error) {
s.mu.Lock()
defer s.mu.Unlock()
if len(patch.RuleIDs) == 0 {
return nil, billing.ErrPriceNotFound
}
for _, id := range patch.RuleIDs {
if _, ok := s.priceRules[id]; !ok {
return nil, billing.ErrPriceNotFound
}
}
now := s.now().UTC().Format(time.RFC3339Nano)
out := make([]billing.PriceRule, 0, len(patch.RuleIDs))
for _, id := range patch.RuleIDs {
rule := cloneRule(s.priceRules[id])
rule.MarkupMultiplier = patch.MarkupMultiplier
for di := range rule.Dimensions {
for ti := range rule.Dimensions[di].Tiers {
rule.Dimensions[di].Tiers[ti].MarkupMultiplier = patch.MarkupMultiplier
}
}
rule.UpdatedAt = now
s.priceRules[id] = cloneRule(rule)
out = append(out, cloneRule(rule))
}
return out, nil
}
func (s *Store) SeedBillingPriceRules(_ context.Context, rules []billing.PriceRule) error {
s.mu.Lock()
defer s.mu.Unlock()
@@ -0,0 +1,69 @@
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)
}
}
@@ -0,0 +1,105 @@
package localstore_test
import (
"context"
"testing"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/localstore"
)
func TestAdminPriceReadsRefreshOldCatalogWithoutLosingMarkups(t *testing.T) {
for _, test := range []struct {
name string
disabled bool
read func(context.Context, *billing.Service) ([]billing.PriceRule, error)
}{
{"overview", false, func(ctx context.Context, service *billing.Service) ([]billing.PriceRule, error) {
overview, err := service.AdminOverview(ctx)
return overview.PriceRules, err
}},
{"prices", false, func(ctx context.Context, service *billing.Service) ([]billing.PriceRule, error) {
return service.ListPrices(ctx)
}},
{"overview billing disabled", true, func(ctx context.Context, service *billing.Service) ([]billing.PriceRule, error) {
overview, err := service.AdminOverview(ctx)
return overview.PriceRules, err
}},
{"prices billing disabled", true, func(ctx context.Context, service *billing.Service) ([]billing.PriceRule, error) {
return service.ListPrices(ctx)
}},
} {
t.Run(test.name, func(t *testing.T) {
ctx := context.Background()
store := localstore.New()
defaults := billing.DefaultBillingPriceRules()
var old billing.PriceRule
for _, rule := range defaults {
if rule.ReqKey == "gpt-image-2" && rule.Provider == "evolink" {
old = rule
break
}
}
if old.ID == "" {
t.Fatal("Image 2 default rule missing")
}
if err := store.SeedBillingPriceRules(ctx, []billing.PriceRule{old}); err != nil {
t.Fatal(err)
}
for _, patch := range []billing.PricePatch{
{DimensionKey: "quality", TierValue: "medium", MarkupMultiplier: 1.3},
{DimensionKey: "referenceImageCount", TierValue: "1–4", MarkupMultiplier: 2},
} {
if _, err := store.UpdateBillingPriceRule(ctx, old.ID, patch); err != nil {
t.Fatal(err)
}
}
service := billing.NewService(store, nil).SetEnabled(!test.disabled)
for read := 0; read < 2; read++ {
rules, err := test.read(ctx, service)
if err != nil {
t.Fatal(err)
}
if len(rules) != len(defaults) {
t.Fatalf("read %d: got %d rules, want %d", read, len(rules), len(defaults))
}
counts := map[string]int{}
for _, rule := range rules {
if rule.Provider != "evolink" || rule.Capability != "image.generate" {
continue
}
counts[rule.ReqKey]++
if rule.ReqKey == "gpt-image-2" {
if got := priceTierMarkup(t, rule, "quality", "medium"); got != 1.3 {
t.Fatalf("read %d: saved quality markup = %v, want 1.3", read, got)
}
if got := priceTierMarkup(t, rule, "referenceImageCount", "1–4"); got != 2 {
t.Fatalf("read %d: saved reference markup = %v, want 2", read, got)
}
}
}
for _, model := range []string{"gpt-image-2", "gpt-image-2.5-flare", "gpt-image-2.5-sunburst"} {
if counts[model] != 1 {
t.Fatalf("read %d: %s rule count = %d, want 1", read, model, counts[model])
}
}
}
})
}
}
func priceTierMarkup(t *testing.T, rule billing.PriceRule, dimensionKey, tierValue string) float64 {
t.Helper()
for _, dimension := range rule.Dimensions {
if dimension.Key != dimensionKey {
continue
}
for _, tier := range dimension.Tiers {
if value, ok := tier.Value.(string); ok && value == tierValue {
return tier.MarkupMultiplier
}
}
}
t.Fatalf("%s tier %s missing from %s", dimensionKey, tierValue, rule.ID)
return 0
}
+56
View File
@@ -16,6 +16,37 @@ ORDER BY provider, capability, id`
const GetBillingPriceRuleSQL = `SELECT id, provider, capability, req_key, variant_key, unit, standard_unit_price_fen, markup_multiplier, enabled, conditions, quantity_source, priority, note, source, parameter_dimensions, created_at::text, updated_at::text FROM public.billing_price_rules WHERE id = $1::text LIMIT 1`
const UpdateBillingPriceRuleSQL = `UPDATE public.billing_price_rules SET markup_multiplier = CASE WHEN $2::text = '' THEN $3::numeric ELSE markup_multiplier END, parameter_dimensions = CASE WHEN $2::text = '' THEN parameter_dimensions ELSE (SELECT jsonb_agg(CASE WHEN dimension->>'key' = $2::text THEN jsonb_set(dimension, '{tiers}', (SELECT jsonb_agg(CASE WHEN tier->>'value' = $4::text THEN jsonb_set(tier, '{markupMultiplier}', to_jsonb($3::numeric), true) ELSE tier END) FROM jsonb_array_elements(dimension->'tiers') tier), true) ELSE dimension END) FROM jsonb_array_elements(parameter_dimensions) dimension) END, updated_at = now() WHERE id = $1::text RETURNING id, provider, capability, req_key, variant_key, unit, standard_unit_price_fen, markup_multiplier, enabled, conditions, quantity_source, priority, note, source, parameter_dimensions, created_at::text, updated_at::text`
// Lock every requested row first. If even one ID is absent, eligible is false
// and this single statement updates none of them.
const UpdateBillingPriceRulesSQL = `WITH locked AS MATERIALIZED (
SELECT id FROM public.billing_price_rules WHERE id = ANY($1::text[]) FOR UPDATE
), eligible AS MATERIALIZED (
SELECT count(*) = cardinality($1::text[]) AS all_present FROM locked
), updated AS (
UPDATE public.billing_price_rules AS rule SET
markup_multiplier = $2::numeric,
parameter_dimensions = (
SELECT COALESCE(jsonb_agg(
jsonb_set(dimension.value, '{tiers}', (
SELECT COALESCE(jsonb_agg(
jsonb_set(tier.value, '{markupMultiplier}', to_jsonb($2::numeric), true)
ORDER BY tier.position
), '[]'::jsonb)
FROM jsonb_array_elements(COALESCE(dimension.value->'tiers', '[]'::jsonb))
WITH ORDINALITY AS tier(value, position)
), true) ORDER BY dimension.position
), '[]'::jsonb)
FROM jsonb_array_elements(COALESCE(rule.parameter_dimensions, '[]'::jsonb))
WITH ORDINALITY AS dimension(value, position)
),
updated_at = now()
WHERE rule.id IN (SELECT id FROM locked) AND (SELECT all_present FROM eligible)
RETURNING id, provider, capability, req_key, variant_key, unit, standard_unit_price_fen,
markup_multiplier, enabled, conditions, quantity_source, priority, note, source,
parameter_dimensions, created_at::text, updated_at::text
)
SELECT * FROM updated ORDER BY id`
const GetBillingWalletSQL = `SELECT $1::text, COALESCE(balance_fen, 0), COALESCE(total_recharged_fen, 0), COALESCE(total_charged_fen, 0), COALESCE(updated_at::text, '') FROM public.billing_wallets WHERE organization_id = $1::text UNION ALL SELECT $1::text, 0, 0, 0, '' WHERE NOT EXISTS (SELECT 1 FROM public.billing_wallets WHERE organization_id = $1::text) LIMIT 1`
const ListBillingWalletsSQL = `SELECT organization_id, balance_fen, total_recharged_fen, total_charged_fen, updated_at::text FROM public.billing_wallets ORDER BY updated_at DESC`
const ListBillingLedgerSQL = `SELECT id, organization_id, COALESCE(account_id, ''), COALESCE(job_id, ''), kind, delta_fen, balance_after_fen, currency, idempotency_key, description, metadata, created_at::text FROM public.billing_ledger WHERE ($1::text = '' OR organization_id = $1::text) AND ($2::text = '' OR account_id = $2::text) ORDER BY created_at DESC LIMIT $3::integer`
@@ -87,6 +118,31 @@ func (db *Database) UpdateBillingPriceRule(ctx context.Context, id string, patch
rule, err := scanFullPriceRule(rows)
return &rule, err
}
func (db *Database) UpdateBillingPriceRules(ctx context.Context, patch billing.BulkPricePatch) ([]billing.PriceRule, error) {
if len(patch.RuleIDs) == 0 {
return nil, billing.ErrPriceNotFound
}
rows, err := db.billingQuery(ctx, UpdateBillingPriceRulesSQL, patch.RuleIDs, patch.MarkupMultiplier)
if err != nil {
return nil, err
}
defer rows.Close()
rules := make([]billing.PriceRule, 0, len(patch.RuleIDs))
for rows.Next() {
rule, err := scanFullPriceRule(rows)
if err != nil {
return nil, err
}
rules = append(rules, rule)
}
if err := rows.Err(); err != nil {
return nil, err
}
if len(rules) != len(patch.RuleIDs) {
return nil, billing.ErrPriceNotFound
}
return rules, nil
}
func scanFullPriceRule(rows Rows) (billing.PriceRule, error) {
var rule billing.PriceRule
var req, variant, quantity, note sql.NullString
@@ -0,0 +1,32 @@
package postgres
import (
"context"
"errors"
"reflect"
"strings"
"testing"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
)
func TestUpdateBillingPriceRulesUsesSingleGuardedStatement(t *testing.T) {
q := &fakeQuerier{}
db := NewDatabase(Config{Backend: BackendPostgres}, q)
patch := billing.BulkPricePatch{RuleIDs: []string{"first", "second"}, MarkupMultiplier: 1.75}
if _, err := db.UpdateBillingPriceRules(context.Background(), patch); !errors.Is(err, billing.ErrPriceNotFound) {
t.Fatalf("empty result error = %v", err)
}
if q.sql != UpdateBillingPriceRulesSQL || !reflect.DeepEqual(q.args, []any{patch.RuleIDs, patch.MarkupMultiplier}) {
t.Fatalf("query = %q, args = %#v", q.sql, q.args)
}
for _, clause := range []string{"FOR UPDATE", "count(*) = cardinality($1::text[])", "AND (SELECT all_present FROM eligible)", "jsonb_set(tier.value", "UPDATE public.billing_price_rules"} {
if !strings.Contains(q.sql, clause) {
t.Fatalf("missing atomic bulk update clause %q", clause)
}
}
q.sql = ""
if _, err := db.UpdateBillingPriceRules(context.Background(), billing.BulkPricePatch{MarkupMultiplier: 2}); !errors.Is(err, billing.ErrPriceNotFound) || q.sql != "" {
t.Fatalf("empty ID list issued a query: %q, %v", q.sql, err)
}
}