Files

355 lines
13 KiB
Go

package repository
import (
"context"
"database/sql"
"errors"
"fmt"
"strings"
"time"
"billing/internal/pricing"
)
const (
PriceSourceManual = "manual"
PriceSourceModelsDev = "models.dev"
)
// PriceSource records where the current effective local price was imported
// from. It is bookkeeping only; billing always reads the embedded Policy.
type PriceSource struct {
Kind string
CatalogID string
Revision string
FetchedAt time.Time
}
type PriceRecord struct {
Policy pricing.Policy
Source PriceSource
}
// ListPrices returns every configured exact-model policy.
func (r *SQLiteUsageRepository) ListPrices(ctx context.Context) ([]pricing.Policy, error) {
records, err := r.ListPriceRecords(ctx)
if err != nil {
return nil, err
}
policies := make([]pricing.Policy, 0, len(records))
for _, record := range records {
policies = append(policies, record.Policy)
}
return policies, nil
}
// ListPriceRecords returns effective prices together with their optional
// models.dev import link.
func (r *SQLiteUsageRepository) ListPriceRecords(ctx context.Context) ([]PriceRecord, error) {
rows, err := r.readDB.QueryContext(ctx, `
SELECT model, input_rate_micros, cache_read_rate_micros, cache_write_rate_micros, output_rate_micros,
long_context_enabled, long_context_threshold, long_context_comparison,
long_input_rate_micros, long_cache_read_rate_micros, long_cache_write_rate_micros, long_output_rate_micros,
fast_pricing_enabled, fast_multiplier_numerator, fast_multiplier_denominator,
source_kind, source_catalog_id, source_revision, source_fetched_at
FROM model_prices ORDER BY model`)
if err != nil {
return nil, fmt.Errorf("查询模型价格: %w", err)
}
defer rows.Close()
var records []PriceRecord
for rows.Next() {
var record PriceRecord
var longEnabled bool
var long pricing.LongContext
var fetchedAt string
if err := rows.Scan(
&record.Policy.Model,
&record.Policy.Base.InputMicrosPer1M, &record.Policy.Base.CacheReadMicrosPer1M,
&record.Policy.Base.CacheWriteMicrosPer1M, &record.Policy.Base.OutputMicrosPer1M,
&longEnabled, &long.ThresholdInputTokens, &long.Comparison,
&long.Rates.InputMicrosPer1M, &long.Rates.CacheReadMicrosPer1M,
&long.Rates.CacheWriteMicrosPer1M, &long.Rates.OutputMicrosPer1M,
&record.Policy.FastPricingEnabled, &record.Policy.FastMultiplier.Numerator, &record.Policy.FastMultiplier.Denominator,
&record.Source.Kind, &record.Source.CatalogID, &record.Source.Revision, &fetchedAt,
); err != nil {
return nil, fmt.Errorf("读取模型价格: %w", err)
}
if longEnabled {
record.Policy.LongContext = &long
}
if strings.TrimSpace(record.Source.Kind) == "" {
record.Source.Kind = PriceSourceManual
}
if fetchedAt != "" {
parsed, parseErr := time.Parse(time.RFC3339Nano, fetchedAt)
if parseErr != nil {
return nil, fmt.Errorf("解析模型价格来源时间: %w", parseErr)
}
record.Source.FetchedAt = parsed
}
records = append(records, record)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("遍历模型价格: %w", err)
}
return records, nil
}
// UpsertPrice atomically replaces one model policy.
func (r *SQLiteUsageRepository) UpsertPrice(ctx context.Context, policy pricing.Policy) error {
return r.UpsertPriceRecord(ctx, PriceRecord{Policy: policy, Source: PriceSource{Kind: PriceSourceManual}})
}
// UpsertPriceRecord atomically replaces one effective price and its source.
func (r *SQLiteUsageRepository) UpsertPriceRecord(ctx context.Context, record PriceRecord) error {
return r.UpsertPriceRecords(ctx, []PriceRecord{record})
}
// UpsertPriceRecords updates a confirmed catalog diff in one transaction.
func (r *SQLiteUsageRepository) UpsertPriceRecords(ctx context.Context, records []PriceRecord) error {
if len(records) == 0 {
return nil
}
for index := range records {
if strings.TrimSpace(records[index].Source.Kind) == "" {
records[index].Source.Kind = PriceSourceManual
}
if err := validatePriceRecord(records[index]); err != nil {
return err
}
}
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return fmt.Errorf("开始保存模型价格: %w", err)
}
for _, record := range records {
var exists int
if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM model_prices WHERE model=?`, record.Policy.Model).Scan(&exists); err != nil {
_ = tx.Rollback()
return fmt.Errorf("检查模型价格: %w", err)
}
if err := upsertPriceRecord(ctx, tx, record); err != nil {
_ = tx.Rollback()
return err
}
action := "新增"
if exists > 0 {
action = "修改"
}
if err := insertBusinessEvent(ctx, tx, "管理员"+action+"模型 `"+strings.ReplaceAll(record.Policy.Model, "`", "")+"` 的价格", BusinessEventSucceeded, time.Now().UTC()); err != nil {
_ = tx.Rollback()
return err
}
}
if err := tx.Commit(); err != nil {
return fmt.Errorf("提交模型价格: %w", err)
}
return nil
}
type priceExecer interface {
ExecContext(context.Context, string, ...any) (sql.Result, error)
}
func upsertPriceRecord(ctx context.Context, executor priceExecer, record PriceRecord) error {
policy := record.Policy
longEnabled := policy.LongContext != nil
long := pricing.LongContext{Comparison: "gt"}
if policy.LongContext != nil {
long = *policy.LongContext
}
fetchedAt := ""
if !record.Source.FetchedAt.IsZero() {
fetchedAt = record.Source.FetchedAt.UTC().Format(time.RFC3339Nano)
}
_, err := executor.ExecContext(ctx, `
INSERT INTO model_prices (
model, input_rate_micros, cache_read_rate_micros, cache_write_rate_micros, output_rate_micros,
long_context_enabled, long_context_threshold, long_context_comparison,
long_input_rate_micros, long_cache_read_rate_micros, long_cache_write_rate_micros, long_output_rate_micros,
fast_pricing_enabled, fast_multiplier_numerator, fast_multiplier_denominator,
source_kind, source_catalog_id, source_revision, source_fetched_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(model) DO UPDATE SET
input_rate_micros=excluded.input_rate_micros,
cache_read_rate_micros=excluded.cache_read_rate_micros,
cache_write_rate_micros=excluded.cache_write_rate_micros,
output_rate_micros=excluded.output_rate_micros,
long_context_enabled=excluded.long_context_enabled,
long_context_threshold=excluded.long_context_threshold,
long_context_comparison=excluded.long_context_comparison,
long_input_rate_micros=excluded.long_input_rate_micros,
long_cache_read_rate_micros=excluded.long_cache_read_rate_micros,
long_cache_write_rate_micros=excluded.long_cache_write_rate_micros,
long_output_rate_micros=excluded.long_output_rate_micros,
fast_pricing_enabled=excluded.fast_pricing_enabled,
fast_multiplier_numerator=excluded.fast_multiplier_numerator,
fast_multiplier_denominator=excluded.fast_multiplier_denominator,
source_kind=excluded.source_kind,
source_catalog_id=excluded.source_catalog_id,
source_revision=excluded.source_revision,
source_fetched_at=excluded.source_fetched_at,
updated_at=excluded.updated_at`,
policy.Model,
policy.Base.InputMicrosPer1M, policy.Base.CacheReadMicrosPer1M,
policy.Base.CacheWriteMicrosPer1M, policy.Base.OutputMicrosPer1M,
longEnabled, long.ThresholdInputTokens, long.Comparison,
long.Rates.InputMicrosPer1M, long.Rates.CacheReadMicrosPer1M,
long.Rates.CacheWriteMicrosPer1M, long.Rates.OutputMicrosPer1M,
policy.FastPricingEnabled, policy.FastMultiplier.Numerator, policy.FastMultiplier.Denominator,
record.Source.Kind, record.Source.CatalogID, record.Source.Revision, fetchedAt,
time.Now().UTC().Format(time.RFC3339Nano),
)
if err != nil {
return fmt.Errorf("保存模型价格: %w", err)
}
return nil
}
func validatePriceRecord(record PriceRecord) error {
if err := record.Policy.Validate(); err != nil {
return err
}
record.Source.Kind = strings.TrimSpace(record.Source.Kind)
switch record.Source.Kind {
case "", PriceSourceManual:
if record.Source.CatalogID != "" || record.Source.Revision != "" || !record.Source.FetchedAt.IsZero() {
return errors.New("手动价格不能包含目录来源")
}
case PriceSourceModelsDev:
if strings.TrimSpace(record.Source.CatalogID) == "" || strings.TrimSpace(record.Source.Revision) == "" || record.Source.FetchedAt.IsZero() {
return errors.New("models.dev 价格来源不完整")
}
default:
return fmt.Errorf("不支持的价格来源 %q", record.Source.Kind)
}
return nil
}
const costBackfillBatchSize = 1000
// BackfillMissingCosts 只补算尚未定价的历史记录,已经保存的账单金额不会随价格修改而变化。
// 每批记录独立提交,避免百万级历史记录一次性占满内存或长期占用 writer。
func (r *SQLiteUsageRepository) BackfillMissingCosts(ctx context.Context, policy pricing.Policy) (int64, error) {
if err := policy.Validate(); err != nil {
return 0, err
}
type pendingCost struct {
id int64
inputTokens int64
cacheReadTokens int64
cacheWriteTokens int64
outputTokens int64
serviceTier, speed string
}
var maxID int64
if err := r.readDB.QueryRowContext(ctx, `
SELECT COALESCE(MAX(id), 0) FROM usage_records
WHERE model = ? COLLATE NOCASE AND cost_micros IS NULL`, policy.Model).Scan(&maxID); err != nil {
return 0, fmt.Errorf("定位待补算用量范围: %w", err)
}
var updated int64
var lastID int64
for lastID < maxID {
rows, err := r.readDB.QueryContext(ctx, `
SELECT id, input_tokens, cache_read_tokens, cache_write_tokens, output_tokens, service_tier, speed
FROM usage_records
WHERE model = ? COLLATE NOCASE AND cost_micros IS NULL AND id > ? AND id <= ?
ORDER BY id LIMIT ?`, policy.Model, lastID, maxID, costBackfillBatchSize)
if err != nil {
return updated, fmt.Errorf("查询待补算用量: %w", err)
}
pending := make([]pendingCost, 0, costBackfillBatchSize)
for rows.Next() {
var item pendingCost
if err := rows.Scan(&item.id, &item.inputTokens, &item.cacheReadTokens, &item.cacheWriteTokens, &item.outputTokens, &item.serviceTier, &item.speed); err != nil {
_ = rows.Close()
return updated, fmt.Errorf("读取待补算用量: %w", err)
}
pending = append(pending, item)
}
if err := rows.Err(); err != nil {
_ = rows.Close()
return updated, fmt.Errorf("遍历待补算用量: %w", err)
}
if err := rows.Close(); err != nil {
return updated, fmt.Errorf("关闭待补算查询: %w", err)
}
if len(pending) == 0 {
break
}
lastID = pending[len(pending)-1].id
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return updated, fmt.Errorf("开始补算事务: %w", err)
}
statement, err := tx.PrepareContext(ctx, `
UPDATE usage_records
SET cost_micros = ?, price_tier = ?, fast_requested = ?, fast_pricing_applied = ?,
price_multiplier_numerator = ?, price_multiplier_denominator = ?
WHERE id = ? AND cost_micros IS NULL`)
if err != nil {
_ = tx.Rollback()
return updated, fmt.Errorf("准备补算用量价格: %w", err)
}
for _, item := range pending {
result, calculateErr := pricing.Calculate(policy, pricing.Usage{
InputTokens: item.inputTokens, CacheReadTokens: item.cacheReadTokens,
CacheWriteTokens: item.cacheWriteTokens, OutputTokens: item.outputTokens,
ServiceTier: item.serviceTier, Speed: item.speed,
})
if calculateErr != nil {
continue
}
change, updateErr := statement.ExecContext(ctx,
result.CostMicros, result.PriceTier, result.FastRequested, result.FastApplied,
result.MultiplierNumerator, result.MultiplierDenominator, item.id)
if updateErr != nil {
_ = statement.Close()
_ = tx.Rollback()
return updated, fmt.Errorf("补算用量价格: %w", updateErr)
}
count, countErr := change.RowsAffected()
if countErr != nil {
_ = statement.Close()
_ = tx.Rollback()
return updated, fmt.Errorf("读取补算数量: %w", countErr)
}
updated += count
}
if err := statement.Close(); err != nil {
_ = tx.Rollback()
return updated, fmt.Errorf("关闭补算语句: %w", err)
}
if err := tx.Commit(); err != nil {
return updated, fmt.Errorf("提交补算事务: %w", err)
}
}
return updated, nil
}
func (r *SQLiteUsageRepository) DeletePrice(ctx context.Context, model string) error {
tx, err := r.db.BeginTx(ctx, nil)
if err != nil {
return fmt.Errorf("开始删除模型价格: %w", err)
}
defer func() { _ = tx.Rollback() }()
result, err := tx.ExecContext(ctx, `DELETE FROM model_prices WHERE model = ?`, model)
if err != nil {
return fmt.Errorf("删除模型价格: %w", err)
}
if affected, _ := result.RowsAffected(); affected > 0 {
event := "管理员删除模型 `" + strings.ReplaceAll(strings.TrimSpace(model), "`", "") + "` 的价格"
if err := insertBusinessEvent(ctx, tx, event, BusinessEventSucceeded, time.Now().UTC()); err != nil {
return err
}
}
if err := tx.Commit(); err != nil {
return fmt.Errorf("提交删除模型价格: %w", err)
}
return nil
}