增加统一倍率调整按钮
This commit is contained in:
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: "eStoreStub{}}
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 := "eStoreStub{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 := "eStoreStub{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++
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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"},
|
||||
|
||||
@@ -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}"},
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user