Files
chuan 6a461bf1f3 feat: 添加持久化用量计费与管理面板
- 使用 SQLite 保存关联后的请求生命周期和用量记录
- 支持长上下文阶梯价格和可配置的 Fast 计费倍率
- 添加管理接口和可配置列的用量面板
- 保留失败、取消、重试和 compact 请求,便于计费核对
2026-08-14 21:08:02 +08:00

164 lines
5.1 KiB
Go

// Package pricing implements deterministic model-price resolution and cost calculation.
package pricing
import (
"errors"
"fmt"
"math/big"
"strings"
)
const perMillion = int64(1_000_000)
// Rates stores micro-USD rates per one million tokens.
type Rates struct {
InputMicrosPer1M int64
CacheReadMicrosPer1M int64
CacheWriteMicrosPer1M int64
OutputMicrosPer1M int64
}
// Ratio stores a multiplier exactly. Fast pricing is represented as 5/2.
type Ratio struct {
Numerator int64
Denominator int64
}
// LongContext replaces all four rates after the input threshold is reached.
type LongContext struct {
ThresholdInputTokens int64
Comparison string
Rates Rates
}
// Policy is the complete pricing configuration for one exact model name.
type Policy struct {
Model string
Base Rates
LongContext *LongContext
FastPricingEnabled bool
FastMultiplier Ratio
}
// Usage contains token buckets required by pricing. InputTokens includes cache buckets.
type Usage struct {
InputTokens int64
CacheReadTokens int64
CacheWriteTokens int64
OutputTokens int64
ServiceTier string
Speed string
}
// Result records the applied policy facts together with the final amount.
type Result struct {
CostMicros int64
PriceTier string
FastRequested bool
FastApplied bool
MultiplierNumerator int64
MultiplierDenominator int64
}
// Validate rejects incomplete or ambiguous policies before they enter the runtime catalog.
func (p Policy) Validate() error {
if strings.TrimSpace(p.Model) == "" {
return errors.New("model is required")
}
if err := p.Base.validate(); err != nil {
return fmt.Errorf("base rates: %w", err)
}
if p.FastMultiplier.Numerator <= 0 || p.FastMultiplier.Denominator <= 0 {
return errors.New("fast multiplier must be positive")
}
if p.LongContext != nil {
if p.LongContext.ThresholdInputTokens < 0 {
return errors.New("long-context threshold cannot be negative")
}
if p.LongContext.Comparison != "gt" && p.LongContext.Comparison != "gte" {
return errors.New("long-context comparison must be gt or gte")
}
if err := p.LongContext.Rates.validate(); err != nil {
return fmt.Errorf("long-context rates: %w", err)
}
}
return nil
}
func (r Rates) validate() error {
values := [...]int64{r.InputMicrosPer1M, r.CacheReadMicrosPer1M, r.CacheWriteMicrosPer1M, r.OutputMicrosPer1M}
for _, value := range values {
if value < 0 {
return errors.New("rates cannot be negative")
}
}
return nil
}
// Calculate selects the long-context tier first, then applies Fast once to the total.
func Calculate(policy Policy, usage Usage) (Result, error) {
if err := policy.Validate(); err != nil {
return Result{}, err
}
if usage.InputTokens < 0 || usage.CacheReadTokens < 0 || usage.CacheWriteTokens < 0 || usage.OutputTokens < 0 {
return Result{}, errors.New("token counts cannot be negative")
}
if usage.CacheReadTokens > usage.InputTokens-usage.CacheWriteTokens {
return Result{}, errors.New("cache token counts exceed input tokens")
}
rates := policy.Base
tier := "base"
if long := policy.LongContext; long != nil && thresholdMatched(usage.InputTokens, long.ThresholdInputTokens, long.Comparison) {
rates = long.Rates
tier = "long_context"
}
multiplier := Ratio{Numerator: 1, Denominator: 1}
fastRequested := isFast(usage.ServiceTier, usage.Speed)
fastApplied := policy.FastPricingEnabled && fastRequested
if fastApplied {
multiplier = policy.FastMultiplier
}
uncachedInput := usage.InputTokens - usage.CacheReadTokens - usage.CacheWriteTokens
numerator := new(big.Int)
addCostTerm(numerator, uncachedInput, rates.InputMicrosPer1M)
addCostTerm(numerator, usage.CacheReadTokens, rates.CacheReadMicrosPer1M)
addCostTerm(numerator, usage.CacheWriteTokens, rates.CacheWriteMicrosPer1M)
addCostTerm(numerator, usage.OutputTokens, rates.OutputMicrosPer1M)
numerator.Mul(numerator, big.NewInt(multiplier.Numerator))
denominator := big.NewInt(perMillion * multiplier.Denominator)
// The complete request is rounded once using round-half-up.
numerator.Add(numerator, new(big.Int).Quo(new(big.Int).Set(denominator), big.NewInt(2)))
numerator.Quo(numerator, denominator)
if !numerator.IsInt64() {
return Result{}, errors.New("calculated cost exceeds int64")
}
return Result{
CostMicros: numerator.Int64(),
PriceTier: tier,
FastRequested: fastRequested,
FastApplied: fastApplied,
MultiplierNumerator: multiplier.Numerator,
MultiplierDenominator: multiplier.Denominator,
}, nil
}
func isFast(serviceTier, speed string) bool {
tier := strings.ToLower(strings.TrimSpace(serviceTier))
return tier == "priority" || tier == "fast" || strings.EqualFold(strings.TrimSpace(speed), "fast")
}
func thresholdMatched(input, threshold int64, comparison string) bool {
if comparison == "gte" {
return input >= threshold
}
return input > threshold
}
func addCostTerm(total *big.Int, tokens, rate int64) {
term := new(big.Int).Mul(big.NewInt(tokens), big.NewInt(rate))
total.Add(total, term)
}