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

312 lines
9.5 KiB
Go

package billing
import (
"errors"
"fmt"
"math"
"sort"
"strconv"
"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 `json:"value"`
Label string `json:"label,omitempty"`
Match any `json:"match,omitempty"`
StandardFactor float64 `json:"standardFactor"`
MarkupMultiplier float64 `json:"markupMultiplier"`
Enabled bool `json:"enabled"`
Note string `json:"note,omitempty"`
}
type ParameterDimension struct {
Key string `json:"key"`
Label string `json:"label,omitempty"`
DefaultValue any `json:"defaultValue,omitempty"`
BaselineValue any `json:"baselineValue"`
Tiers []ParameterTier `json:"tiers"`
}
type PriceRule struct {
ID string `json:"id"`
Provider string `json:"provider"`
Capability string `json:"capability"`
ReqKey string `json:"reqKey,omitempty"`
VariantKey string `json:"variantKey,omitempty"`
Unit Unit `json:"unit"`
QuantitySource QuantitySource `json:"quantitySource,omitempty"`
StandardUnitPriceFen int64 `json:"standardUnitPriceFen"`
MarkupMultiplier float64 `json:"markupMultiplier"`
Enabled bool `json:"enabled"`
Conditions Conditions `json:"conditions,omitempty"`
Priority int `json:"priority,omitempty"`
Note string `json:"note,omitempty"`
Source map[string]any `json:"source,omitempty"`
Dimensions []ParameterDimension `json:"parameterDimensions,omitempty"`
CreatedAt string `json:"createdAt,omitempty"`
UpdatedAt string `json:"updatedAt,omitempty"`
}
type QuoteInput struct {
Provider, Capability, ReqKey, Source, Role string
Parameters Parameters
BillingDisabled bool
}
type Quote struct {
PriceRuleID string `json:"priceRuleId"`
Provider string `json:"provider,omitempty"`
Capability string `json:"capability,omitempty"`
ReqKey string `json:"reqKey,omitempty"`
VariantKey string `json:"variantKey,omitempty"`
Unit Unit `json:"unit"`
Quantity float64 `json:"quantity"`
StandardUnitPriceFen int64 `json:"standardUnitPriceFen"`
AmountFen int64 `json:"amountFen"`
MarkupMultiplier float64 `json:"markupMultiplier"`
Currency string `json:"currency"`
Conditions Conditions `json:"conditions,omitempty"`
QuantitySource QuantitySource `json:"quantitySource,omitempty"`
Parameters Parameters `json:"parameters,omitempty"`
QuotaExempt bool `json:"quotaExempt,omitempty"`
ReservedAmountFen int64 `json:"reservedAmountFen,omitempty"`
SettlementStatus string `json:"settlementStatus,omitempty"`
}
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 {
conditions := effectiveConditions(rule)
if !rule.Enabled || rule.Provider != input.Provider || rule.Capability != input.Capability || rule.ReqKey != "" && rule.ReqKey != input.ReqKey || !conditionsMatch(conditions, input.Parameters) {
continue
}
req := 0
if rule.ReqKey != "" {
req = 1
}
rule.Conditions = conditions
matches = append(matches, candidate{rule, req, len(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, Provider: input.Provider, Capability: input.Capability, ReqKey: input.ReqKey, VariantKey: w.rule.VariantKey, Unit: w.rule.Unit, Quantity: quantity, StandardUnitPriceFen: price, MarkupMultiplier: markup, AmountFen: int64(math.Ceil(float64(price) * quantity * markup)), Currency: CurrencyCNY, Conditions: w.rule.Conditions, QuantitySource: w.rule.QuantitySource, Parameters: input.Parameters, QuotaExempt: input.Source == "platform" && input.Role == "super_admin"}, nil
}
func effectiveConditions(rule PriceRule) Conditions {
conditions := parseVariantKey(rule.VariantKey)
for key, value := range rule.Conditions {
conditions[key] = value
}
return conditions
}
func parseVariantKey(value string) Conditions {
conditions := Conditions{}
for _, part := range strings.FieldsFunc(value, func(r rune) bool { return r == ';' || r == ',' }) {
key, raw, ok := strings.Cut(part, "=")
if !ok {
continue
}
key, raw = strings.TrimSpace(key), strings.TrimSpace(raw)
if key == "ratio" {
key = "aspectRatio"
}
switch key {
case "model", "resolution", "size", "aspectRatio", "quality":
if raw != "" {
conditions[key] = strings.ToLower(raw)
}
case "duration", "imageCount", "referenceImageCount", "scale":
if number, err := strconv.ParseFloat(raw, 64); err == nil && !math.IsNaN(number) && !math.IsInf(number, 0) {
conditions[key] = number
}
}
}
return conditions
}
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
}