220 lines
7.1 KiB
Go
220 lines
7.1 KiB
Go
package localstore
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"sort"
|
|
"time"
|
|
|
|
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/billing"
|
|
)
|
|
|
|
func (s *Store) BillingWallet(_ context.Context, organizationID string) (billing.Wallet, error) {
|
|
s.mu.RLock()
|
|
defer s.mu.RUnlock()
|
|
if w, ok := s.wallets[organizationID]; ok {
|
|
return w, nil
|
|
}
|
|
return billing.Wallet{OrganizationID: organizationID, Currency: billing.CurrencyCNY}, nil
|
|
}
|
|
func (s *Store) BillingWallets(_ context.Context) ([]billing.Wallet, error) {
|
|
s.mu.RLock()
|
|
defer s.mu.RUnlock()
|
|
out := make([]billing.Wallet, 0, len(s.wallets))
|
|
for _, w := range s.wallets {
|
|
out = append(out, w)
|
|
}
|
|
sort.Slice(out, func(i, j int) bool { return out[i].UpdatedAt > out[j].UpdatedAt })
|
|
return out, nil
|
|
}
|
|
func (s *Store) BillingLedger(_ context.Context, org, account string, limit int) ([]billing.LedgerEntry, error) {
|
|
s.mu.RLock()
|
|
defer s.mu.RUnlock()
|
|
if limit <= 0 || limit > 500 {
|
|
limit = 500
|
|
}
|
|
out := []billing.LedgerEntry{}
|
|
for i := len(s.ledger) - 1; i >= 0 && len(out) < limit; i-- {
|
|
e := s.ledger[i]
|
|
if org != "" && e.OrganizationID != org || account != "" && e.AccountID != account {
|
|
continue
|
|
}
|
|
e.Metadata = cloneMap(e.Metadata)
|
|
out = append(out, e)
|
|
}
|
|
return out, nil
|
|
}
|
|
func (s *Store) ListBillingPriceRules(_ context.Context, include bool) ([]billing.PriceRule, error) {
|
|
s.mu.RLock()
|
|
defer s.mu.RUnlock()
|
|
out := []billing.PriceRule{}
|
|
for _, r := range s.priceRules {
|
|
if include || r.Enabled {
|
|
out = append(out, cloneRule(r))
|
|
}
|
|
}
|
|
sort.Slice(out, func(i, j int) bool {
|
|
if out[i].Provider != out[j].Provider {
|
|
return out[i].Provider < out[j].Provider
|
|
}
|
|
if out[i].Capability != out[j].Capability {
|
|
return out[i].Capability < out[j].Capability
|
|
}
|
|
return out[i].ID < out[j].ID
|
|
})
|
|
return out, nil
|
|
}
|
|
func (s *Store) GetBillingPriceRule(_ context.Context, id string) (*billing.PriceRule, error) {
|
|
s.mu.RLock()
|
|
defer s.mu.RUnlock()
|
|
r, ok := s.priceRules[id]
|
|
if !ok {
|
|
return nil, nil
|
|
}
|
|
r = cloneRule(r)
|
|
return &r, nil
|
|
}
|
|
func (s *Store) UpdateBillingPriceRule(_ context.Context, id string, p billing.PricePatch) (*billing.PriceRule, error) {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
r, ok := s.priceRules[id]
|
|
if !ok {
|
|
return nil, nil
|
|
}
|
|
if p.DimensionKey == "" {
|
|
r.MarkupMultiplier = p.MarkupMultiplier
|
|
} else {
|
|
for di := range r.Dimensions {
|
|
if r.Dimensions[di].Key != p.DimensionKey {
|
|
continue
|
|
}
|
|
for ti := range r.Dimensions[di].Tiers {
|
|
if fmt.Sprint(r.Dimensions[di].Tiers[ti].Value) == p.TierValue {
|
|
r.Dimensions[di].Tiers[ti].MarkupMultiplier = p.MarkupMultiplier
|
|
}
|
|
}
|
|
}
|
|
}
|
|
r.UpdatedAt = s.now().UTC().Format(time.RFC3339Nano)
|
|
s.priceRules[id] = cloneRule(r)
|
|
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()
|
|
now := s.now().UTC().Format(time.RFC3339Nano)
|
|
for _, r := range rules {
|
|
existing, ok := s.priceRules[r.ID]
|
|
if ok {
|
|
r.MarkupMultiplier = existing.MarkupMultiplier
|
|
r.CreatedAt = existing.CreatedAt
|
|
r.Dimensions = preserveTierMarkups(r.Dimensions, existing.Dimensions)
|
|
}
|
|
if r.CreatedAt == "" {
|
|
r.CreatedAt = now
|
|
}
|
|
r.UpdatedAt = now
|
|
s.priceRules[r.ID] = cloneRule(r)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// Defaults own the catalog shape and official costs; administrators own the
|
|
// markup for each matching dimension/tier. Copy before editing so a reused
|
|
// default catalog is not changed by one store's customizations.
|
|
func preserveTierMarkups(defaults, existing []billing.ParameterDimension) []billing.ParameterDimension {
|
|
markups := make(map[[2]string]float64)
|
|
for _, dimension := range existing {
|
|
for _, tier := range dimension.Tiers {
|
|
if tier.MarkupMultiplier >= 1 && tier.MarkupMultiplier <= 1000 {
|
|
markups[[2]string{dimension.Key, fmt.Sprint(tier.Value)}] = tier.MarkupMultiplier
|
|
}
|
|
}
|
|
}
|
|
merged := append([]billing.ParameterDimension(nil), defaults...)
|
|
for di := range merged {
|
|
merged[di].Tiers = append([]billing.ParameterTier(nil), merged[di].Tiers...)
|
|
for ti := range merged[di].Tiers {
|
|
if markup, ok := markups[[2]string{merged[di].Key, fmt.Sprint(merged[di].Tiers[ti].Value)}]; ok {
|
|
merged[di].Tiers[ti].MarkupMultiplier = markup
|
|
}
|
|
}
|
|
}
|
|
return merged
|
|
}
|
|
|
|
func (s *Store) PostBillingWalletEntry(ctx context.Context, p billing.WalletPostParams) (billing.WalletPosting, error) {
|
|
return s.PostWalletEntry(ctx, p)
|
|
}
|
|
func (s *Store) PostWalletEntry(_ context.Context, p billing.WalletPostParams) (billing.WalletPosting, error) {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
if old, ok := s.postings[p.IdempotencyKey]; ok {
|
|
if !walletParamsEqual(old.params, p) {
|
|
return billing.WalletPosting{}, fmt.Errorf("BILLING_IDEMPOTENCY_PAYLOAD_MISMATCH")
|
|
}
|
|
return old.posting, nil
|
|
}
|
|
w := s.wallets[p.OrganizationID]
|
|
w.OrganizationID = p.OrganizationID
|
|
if w.Currency == "" {
|
|
w.Currency = billing.CurrencyCNY
|
|
}
|
|
next := w.BalanceFen + p.DeltaFen
|
|
if next < 0 {
|
|
return billing.WalletPosting{}, fmt.Errorf("BILLING_INSUFFICIENT_BALANCE")
|
|
}
|
|
now := s.now().UTC()
|
|
w.BalanceFen = next
|
|
if p.DeltaFen > 0 && (p.Kind == "recharge" || p.Kind == "adjustment") {
|
|
w.TotalRechargedFen += p.DeltaFen
|
|
}
|
|
if p.Kind == "charge" && p.DeltaFen < 0 {
|
|
w.TotalChargedFen += -p.DeltaFen
|
|
}
|
|
w.UpdatedAt = now.Format(time.RFC3339Nano)
|
|
s.wallets[p.OrganizationID] = w
|
|
posting := billing.WalletPosting{LedgerID: p.LedgerID, BalanceAfterFen: next, BalanceFen: next, TotalRechargedFen: w.TotalRechargedFen, TotalChargedFen: w.TotalChargedFen, DeltaFen: p.DeltaFen, CreatedAt: now, UpdatedAt: now}
|
|
s.ledger = append(s.ledger, billing.LedgerEntry{ID: p.LedgerID, OrganizationID: p.OrganizationID, AccountID: p.AccountID, JobID: p.JobID, Kind: p.Kind, DeltaFen: p.DeltaFen, BalanceAfterFen: next, Currency: defaultCurrency(p.Currency), IdempotencyKey: p.IdempotencyKey, Description: p.Description, Metadata: cloneMap(p.Metadata), CreatedAt: now.Format(time.RFC3339Nano)})
|
|
p.Metadata = cloneMap(p.Metadata)
|
|
s.postings[p.IdempotencyKey] = walletRecord{params: p, posting: posting}
|
|
return posting, nil
|
|
}
|
|
func walletParamsEqual(a, b billing.WalletPostParams) bool {
|
|
return a.OrganizationID == b.OrganizationID && a.AccountID == b.AccountID && a.JobID == b.JobID && a.Kind == b.Kind && a.DeltaFen == b.DeltaFen && defaultCurrency(a.Currency) == defaultCurrency(b.Currency)
|
|
}
|
|
func defaultCurrency(v string) string {
|
|
if v == "" {
|
|
return billing.CurrencyCNY
|
|
}
|
|
return v
|
|
}
|