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

254 lines
6.4 KiB
Go

package billing
import (
"errors"
"fmt"
"math"
"sort"
"strings"
)
const CurrencyCNY = "CNY"
var (
ErrAmbiguousPriceRule = errors.New("billing price rules are ambiguous")
ErrPriceRuleNotFound = errors.New("billing price rule not found")
ErrParameterTier = errors.New("billing parameter tier not found")
)
type Unit string
const (
UnitRequest Unit = "request"
UnitImage Unit = "image"
UnitVideoSecond Unit = "video_second"
)
type QuantitySource string
const (
QuantityRequest QuantitySource = "request"
QuantityImageCount QuantitySource = "image_count"
QuantityDuration QuantitySource = "duration"
)
type Parameters map[string]any
type Conditions map[string]any
type ConditionRange struct {
Min *float64
Max *float64
Values []any
}
type ParameterTier struct {
Value any
Match any
StandardFactor, MarkupMultiplier float64
Enabled bool
}
type ParameterDimension struct {
Key string
DefaultValue any
BaselineValue any
Tiers []ParameterTier
}
type PriceRule struct {
ID, Provider, Capability, ReqKey string
Unit Unit
QuantitySource QuantitySource
StandardUnitPriceFen int64
MarkupMultiplier float64
Enabled bool
Conditions Conditions
Priority int
Dimensions []ParameterDimension
}
type QuoteInput struct {
Provider, Capability, ReqKey, Source, Role string
Parameters Parameters
BillingDisabled bool
}
type Quote struct {
PriceRuleID string
Unit Unit
Quantity float64
StandardUnitPriceFen, AmountFen int64
MarkupMultiplier float64
Currency string
QuotaExempt bool
}
type Catalog struct{ Rules []PriceRule }
func (c Catalog) Quote(input QuoteInput) (*Quote, error) {
if input.BillingDisabled || input.Provider == "mock" {
return nil, nil
}
type candidate struct {
rule PriceRule
req, conditions, priority int
}
var matches []candidate
for _, rule := range c.Rules {
if !rule.Enabled || rule.Provider != input.Provider || rule.Capability != input.Capability || rule.ReqKey != "" && rule.ReqKey != input.ReqKey || !conditionsMatch(rule.Conditions, input.Parameters) {
continue
}
req := 0
if rule.ReqKey != "" {
req = 1
}
matches = append(matches, candidate{rule, req, len(rule.Conditions), rule.Priority})
}
if len(matches) == 0 {
return nil, ErrPriceRuleNotFound
}
sort.Slice(matches, func(i, j int) bool {
a, b := matches[i], matches[j]
if a.req != b.req {
return a.req > b.req
}
if a.conditions != b.conditions {
return a.conditions > b.conditions
}
if a.priority != b.priority {
return a.priority > b.priority
}
return a.rule.ID < b.rule.ID
})
w := matches[0]
if len(matches) > 1 && matches[1].req == w.req && matches[1].conditions == w.conditions && matches[1].priority == w.priority {
return nil, fmt.Errorf("%w: %s and %s", ErrAmbiguousPriceRule, w.rule.ID, matches[1].rule.ID)
}
price, markup, err := tierPrice(w.rule, input.Parameters)
if err != nil {
return nil, err
}
quantity := quantityFor(w.rule, input.Parameters)
return &Quote{PriceRuleID: w.rule.ID, Unit: w.rule.Unit, Quantity: quantity, StandardUnitPriceFen: price, MarkupMultiplier: markup, AmountFen: int64(math.Ceil(float64(price) * quantity * markup)), Currency: CurrencyCNY, QuotaExempt: input.Source == "platform" && input.Role == "super_admin"}, nil
}
func tierPrice(rule PriceRule, params Parameters) (int64, float64, error) {
if len(rule.Dimensions) == 0 {
return rule.StandardUnitPriceFen, rule.MarkupMultiplier, nil
}
factor, markup := 1.0, 1.0
for _, dimension := range rule.Dimensions {
actual, ok := params[dimension.Key]
if !ok {
actual = dimension.DefaultValue
if actual == nil {
actual = dimension.BaselineValue
}
}
found := false
for _, tier := range dimension.Tiers {
if tier.Enabled && conditionMatches(firstNonNil(tier.Match, tier.Value), actual) {
if tier.StandardFactor <= 0 || tier.MarkupMultiplier < 1 {
return 0, 0, ErrParameterTier
}
factor *= tier.StandardFactor
markup = math.Max(markup, tier.MarkupMultiplier)
found = true
break
}
}
if !found {
return 0, 0, ErrParameterTier
}
}
return int64(math.Ceil(float64(rule.StandardUnitPriceFen) * factor)), markup, nil
}
func quantityFor(rule PriceRule, params Parameters) float64 {
source := rule.QuantitySource
if source == "" {
if rule.Unit == UnitImage {
source = QuantityImageCount
} else if rule.Unit == UnitVideoSecond {
source = QuantityDuration
} else {
source = QuantityRequest
}
}
if source == QuantityRequest {
return 1
}
key := "imageCount"
if source == QuantityDuration {
key = "duration"
}
n, ok := number(params[key])
if !ok || n <= 0 {
return 1
}
return math.Ceil(n)
}
func conditionsMatch(cs Conditions, p Parameters) bool {
for k, c := range cs {
if !conditionMatches(c, p[k]) {
return false
}
}
return true
}
func conditionMatches(condition, actual any) bool {
if actual == nil {
return false
}
switch c := condition.(type) {
case ConditionRange:
n, ok := number(actual)
if !ok {
return false
}
if c.Min != nil && n < *c.Min {
return false
}
if c.Max != nil && n > *c.Max {
return false
}
for _, v := range c.Values {
if scalarEqual(v, actual) {
return true
}
}
return len(c.Values) == 0
case map[string]any:
n, ok := number(actual)
if min, yes := number(c["min"]); yes && (!ok || n < min) {
return false
}
if max, yes := number(c["max"]); yes && (!ok || n > max) {
return false
}
return true
default:
return scalarEqual(condition, actual)
}
}
func scalarEqual(a, b any) bool {
if x, ok := number(a); ok {
y, ok := number(b)
return ok && x == y
}
return strings.EqualFold(strings.TrimSpace(fmt.Sprint(a)), strings.TrimSpace(fmt.Sprint(b)))
}
func number(v any) (float64, bool) {
switch n := v.(type) {
case int:
return float64(n), true
case int64:
return float64(n), true
case float64:
return n, !math.IsNaN(n) && !math.IsInf(n, 0)
case float32:
return float64(n), true
default:
return 0, false
}
}
func firstNonNil(a, b any) any {
if a != nil {
return a
}
return b
}