Files
NianAIGC/backend/internal/localstore/billing.go

168 lines
5.3 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) 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
}
if r.CreatedAt == "" {
r.CreatedAt = now
}
r.UpdatedAt = now
s.priceRules[r.ID] = cloneRule(r)
}
return nil
}
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
}