296 lines
11 KiB
Go
296 lines
11 KiB
Go
package plugin
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"math/big"
|
|
"net/http"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"billing/internal/modelcatalog"
|
|
"billing/internal/pricing"
|
|
"billing/internal/repository"
|
|
)
|
|
|
|
type priceRatesDTO struct {
|
|
InputPer1M string `json:"input_per_1m"`
|
|
CacheReadPer1M string `json:"cache_read_per_1m"`
|
|
CacheWritePer1M string `json:"cache_write_per_1m"`
|
|
OutputPer1M string `json:"output_per_1m"`
|
|
}
|
|
|
|
type longContextDTO struct {
|
|
ThresholdInputTokens int64 `json:"threshold_input_tokens"`
|
|
Comparison string `json:"comparison"`
|
|
priceRatesDTO
|
|
}
|
|
|
|
type priceDTO struct {
|
|
Model string `json:"model"`
|
|
Base priceRatesDTO `json:"base"`
|
|
LongContext *longContextDTO `json:"long_context,omitempty"`
|
|
FastPricingEnabled bool `json:"fast_pricing_enabled"`
|
|
FastMultiplier string `json:"fast_multiplier"`
|
|
Source priceSourceDTO `json:"source"`
|
|
}
|
|
|
|
type priceSourceDTO struct {
|
|
Kind string `json:"kind"`
|
|
CatalogID string `json:"catalog_id,omitempty"`
|
|
Provider string `json:"provider,omitempty"`
|
|
Model string `json:"model,omitempty"`
|
|
Revision string `json:"revision,omitempty"`
|
|
FetchedAt time.Time `json:"fetched_at,omitzero"`
|
|
}
|
|
|
|
type priceDeleteRequest struct {
|
|
Model string `json:"model"`
|
|
}
|
|
|
|
func (a *App) listPrices() ManagementResponse {
|
|
a.mu.RLock()
|
|
defer a.mu.RUnlock()
|
|
if a.store == nil {
|
|
return jsonManagementResponse(http.StatusServiceUnavailable, map[string]any{"error": map[string]string{"code": "database_unavailable", "message": "价格数据库尚未初始化"}})
|
|
}
|
|
records, err := a.store.ListPriceRecords(context.Background())
|
|
if err != nil {
|
|
return jsonManagementResponse(http.StatusInternalServerError, map[string]any{"error": map[string]string{"code": "database_error", "message": err.Error()}})
|
|
}
|
|
prices := make([]priceDTO, 0, len(records))
|
|
for _, record := range records {
|
|
prices = append(prices, priceRecordToDTO(record))
|
|
}
|
|
return jsonManagementResponse(http.StatusOK, map[string]any{"prices": prices})
|
|
}
|
|
|
|
func (a *App) putPrice(body []byte) ManagementResponse {
|
|
var dto priceDTO
|
|
if err := decodeJSONBody(body, &dto); err != nil {
|
|
return jsonManagementResponse(http.StatusBadRequest, map[string]any{"error": map[string]string{"code": "invalid_price", "message": err.Error()}})
|
|
}
|
|
policy, err := dtoToPolicy(dto)
|
|
if err != nil {
|
|
return jsonManagementResponse(http.StatusBadRequest, map[string]any{"error": map[string]string{"code": "invalid_price", "message": err.Error()}})
|
|
}
|
|
return a.persistPriceRecord(repository.PriceRecord{Policy: policy, Source: repository.PriceSource{Kind: repository.PriceSourceManual}})
|
|
}
|
|
|
|
func (a *App) persistPriceRecord(record repository.PriceRecord) ManagementResponse {
|
|
return a.persistPriceRecordFromCatalog(record, nil)
|
|
}
|
|
|
|
func (a *App) persistPriceRecordFromCatalog(record repository.PriceRecord, expectedCatalog *modelcatalog.Manager) ManagementResponse {
|
|
a.priceMu.Lock()
|
|
defer a.priceMu.Unlock()
|
|
return a.persistPriceRecordLocked(record, expectedCatalog)
|
|
}
|
|
|
|
func (a *App) persistPriceRecordLocked(record repository.PriceRecord, expectedCatalog *modelcatalog.Manager) ManagementResponse {
|
|
a.mu.Lock()
|
|
store := a.store
|
|
if store == nil {
|
|
a.mu.Unlock()
|
|
return jsonManagementResponse(http.StatusServiceUnavailable, map[string]any{"error": map[string]string{"code": "database_unavailable", "message": "价格数据库尚未初始化"}})
|
|
}
|
|
if expectedCatalog != nil && a.catalog != expectedCatalog {
|
|
a.mu.Unlock()
|
|
return jsonManagementResponse(http.StatusServiceUnavailable, map[string]any{"error": map[string]string{"code": "catalog_reconfigured", "message": "价格目录配置已经变化,请重新搜索"}})
|
|
}
|
|
if err := store.UpsertPriceRecord(context.Background(), record); err != nil {
|
|
a.mu.Unlock()
|
|
return jsonManagementResponse(http.StatusInternalServerError, map[string]any{"error": map[string]string{"code": "database_error", "message": err.Error()}})
|
|
}
|
|
a.prices[normalizeModelName(record.Policy.Model)] = record.Policy
|
|
a.mu.Unlock()
|
|
a.mu.RLock()
|
|
currentStore := a.store == store
|
|
a.mu.RUnlock()
|
|
if !currentStore {
|
|
return jsonManagementResponse(http.StatusServiceUnavailable, map[string]any{"error": map[string]string{"code": "database_reconfigured", "message": "价格已经保存,请在重新配置后重试历史补算"}})
|
|
}
|
|
if _, err := store.BackfillMissingCosts(context.Background(), record.Policy); err != nil {
|
|
a.mu.RLock()
|
|
reconfigured := a.store != store
|
|
a.mu.RUnlock()
|
|
if reconfigured {
|
|
return jsonManagementResponse(http.StatusServiceUnavailable, map[string]any{"error": map[string]string{"code": "database_reconfigured", "message": "价格已经保存,重新配置中止了历史补算"}})
|
|
}
|
|
return jsonManagementResponse(http.StatusInternalServerError, map[string]any{"error": map[string]string{"code": "database_error", "message": err.Error()}})
|
|
}
|
|
return jsonManagementResponse(http.StatusOK, priceRecordToDTO(record))
|
|
}
|
|
|
|
func (a *App) deletePrice(body []byte) ManagementResponse {
|
|
var request priceDeleteRequest
|
|
if err := decodeJSONBody(body, &request); err != nil {
|
|
return jsonManagementResponse(http.StatusBadRequest, map[string]any{"error": map[string]string{"code": "invalid_price", "message": err.Error()}})
|
|
}
|
|
request.Model = strings.TrimSpace(request.Model)
|
|
if request.Model == "" {
|
|
return jsonManagementResponse(http.StatusBadRequest, map[string]any{"error": map[string]string{"code": "invalid_price", "message": "model 不能为空"}})
|
|
}
|
|
a.priceMu.Lock()
|
|
defer a.priceMu.Unlock()
|
|
a.mu.Lock()
|
|
store := a.store
|
|
if store == nil {
|
|
a.mu.Unlock()
|
|
return jsonManagementResponse(http.StatusServiceUnavailable, map[string]any{"error": map[string]string{"code": "database_unavailable", "message": "价格数据库尚未初始化"}})
|
|
}
|
|
if err := store.DeletePrice(context.Background(), request.Model); err != nil {
|
|
a.mu.Unlock()
|
|
return jsonManagementResponse(http.StatusInternalServerError, map[string]any{"error": map[string]string{"code": "database_error", "message": err.Error()}})
|
|
}
|
|
defer a.mu.Unlock()
|
|
delete(a.prices, normalizeModelName(request.Model))
|
|
return jsonManagementResponse(http.StatusOK, map[string]any{"deleted": request.Model})
|
|
}
|
|
|
|
func dtoToPolicy(dto priceDTO) (pricing.Policy, error) {
|
|
base, err := dtoRates(dto.Base)
|
|
if err != nil {
|
|
return pricing.Policy{}, err
|
|
}
|
|
multiplierText := strings.TrimSpace(dto.FastMultiplier)
|
|
if multiplierText == "" {
|
|
multiplierText = "2.5"
|
|
}
|
|
multiplier, err := parseRatio(multiplierText)
|
|
if err != nil {
|
|
return pricing.Policy{}, errors.New("fast_multiplier 必须是正数,最多保留六位小数")
|
|
}
|
|
policy := pricing.Policy{Model: strings.TrimSpace(dto.Model), Base: base, FastPricingEnabled: dto.FastPricingEnabled, FastMultiplier: multiplier}
|
|
if dto.LongContext != nil {
|
|
rates, ratesErr := dtoRates(dto.LongContext.priceRatesDTO)
|
|
if ratesErr != nil {
|
|
return pricing.Policy{}, ratesErr
|
|
}
|
|
policy.LongContext = &pricing.LongContext{ThresholdInputTokens: dto.LongContext.ThresholdInputTokens, Comparison: dto.LongContext.Comparison, Rates: rates}
|
|
}
|
|
if err := policy.Validate(); err != nil {
|
|
return pricing.Policy{}, err
|
|
}
|
|
return policy, nil
|
|
}
|
|
|
|
func dtoRates(dto priceRatesDTO) (pricing.Rates, error) {
|
|
values := []*int64{new(int64), new(int64), new(int64), new(int64)}
|
|
texts := []string{dto.InputPer1M, dto.CacheReadPer1M, dto.CacheWritePer1M, dto.OutputPer1M}
|
|
for index, text := range texts {
|
|
value, err := parseDecimalMicros(text)
|
|
if err != nil {
|
|
return pricing.Rates{}, errors.New("价格必须是非负数字,最多保留六位小数")
|
|
}
|
|
*values[index] = value
|
|
}
|
|
return pricing.Rates{InputMicrosPer1M: *values[0], CacheReadMicrosPer1M: *values[1], CacheWriteMicrosPer1M: *values[2], OutputMicrosPer1M: *values[3]}, nil
|
|
}
|
|
|
|
func policyToDTO(policy pricing.Policy) priceDTO {
|
|
dto := priceDTO{Model: policy.Model, Base: ratesToDTO(policy.Base), FastPricingEnabled: policy.FastPricingEnabled, FastMultiplier: formatRatio(policy.FastMultiplier), Source: priceSourceDTO{Kind: repository.PriceSourceManual}}
|
|
if policy.LongContext != nil {
|
|
dto.LongContext = &longContextDTO{ThresholdInputTokens: policy.LongContext.ThresholdInputTokens, Comparison: policy.LongContext.Comparison, priceRatesDTO: ratesToDTO(policy.LongContext.Rates)}
|
|
}
|
|
return dto
|
|
}
|
|
|
|
func priceRecordToDTO(record repository.PriceRecord) priceDTO {
|
|
dto := policyToDTO(record.Policy)
|
|
source := record.Source
|
|
if source.Kind == "" {
|
|
source.Kind = repository.PriceSourceManual
|
|
}
|
|
dto.Source = priceSourceDTO{Kind: source.Kind, CatalogID: source.CatalogID, Revision: source.Revision, FetchedAt: source.FetchedAt}
|
|
if source.CatalogID != "" {
|
|
parts := strings.SplitN(source.CatalogID, "/", 2)
|
|
dto.Source.Provider = parts[0]
|
|
if len(parts) == 2 {
|
|
dto.Source.Model = parts[1]
|
|
}
|
|
}
|
|
return dto
|
|
}
|
|
|
|
func ratesToDTO(rates pricing.Rates) priceRatesDTO {
|
|
return priceRatesDTO{InputPer1M: formatMicros(rates.InputMicrosPer1M), CacheReadPer1M: formatMicros(rates.CacheReadMicrosPer1M), CacheWritePer1M: formatMicros(rates.CacheWriteMicrosPer1M), OutputPer1M: formatMicros(rates.OutputMicrosPer1M)}
|
|
}
|
|
|
|
func parseDecimalMicros(text string) (int64, error) {
|
|
text = strings.TrimSpace(text)
|
|
if text == "" || strings.HasPrefix(text, "-") || strings.Count(text, ".") > 1 {
|
|
return 0, errors.New("invalid decimal")
|
|
}
|
|
parts := strings.SplitN(text, ".", 2)
|
|
if parts[0] == "" {
|
|
parts[0] = "0"
|
|
}
|
|
fraction := ""
|
|
if len(parts) == 2 {
|
|
fraction = parts[1]
|
|
}
|
|
if len(fraction) > 6 {
|
|
return 0, errors.New("too many decimal places")
|
|
}
|
|
for len(fraction) < 6 {
|
|
fraction += "0"
|
|
}
|
|
whole := new(big.Int)
|
|
if _, ok := whole.SetString(parts[0]+fraction, 10); !ok || !whole.IsInt64() {
|
|
return 0, errors.New("invalid decimal")
|
|
}
|
|
return whole.Int64(), nil
|
|
}
|
|
|
|
func parseRatio(text string) (pricing.Ratio, error) {
|
|
numerator, err := parseDecimalMicros(text)
|
|
if err != nil || numerator <= 0 {
|
|
return pricing.Ratio{}, errors.New("invalid ratio")
|
|
}
|
|
denominator := int64(1_000_000)
|
|
divisor := gcd(numerator, denominator)
|
|
return pricing.Ratio{Numerator: numerator / divisor, Denominator: denominator / divisor}, nil
|
|
}
|
|
|
|
func gcd(a, b int64) int64 {
|
|
for b != 0 {
|
|
a, b = b, a%b
|
|
}
|
|
return a
|
|
}
|
|
|
|
func formatMicros(value int64) string {
|
|
prefix := ""
|
|
if value < 0 {
|
|
prefix = "-"
|
|
value = -value
|
|
}
|
|
whole := value / 1_000_000
|
|
fraction := strconv.FormatInt(value%1_000_000+1_000_000, 10)[1:]
|
|
fraction = strings.TrimRight(fraction, "0")
|
|
if fraction == "" {
|
|
return prefix + strconv.FormatInt(whole, 10)
|
|
}
|
|
return prefix + strconv.FormatInt(whole, 10) + "." + fraction
|
|
}
|
|
|
|
func formatRatio(value pricing.Ratio) string {
|
|
if value.Denominator == 0 {
|
|
return ""
|
|
}
|
|
rational := new(big.Rat).SetFrac(big.NewInt(value.Numerator), big.NewInt(value.Denominator))
|
|
return strings.TrimRight(strings.TrimRight(rational.FloatString(6), "0"), ".")
|
|
}
|
|
|
|
func decodeJSONBody(body []byte, target any) error {
|
|
decoder := json.NewDecoder(strings.NewReader(string(body)))
|
|
decoder.DisallowUnknownFields()
|
|
if err := decoder.Decode(target); err != nil {
|
|
return err
|
|
}
|
|
return nil
|
|
}
|