diff options
Diffstat (limited to 'internal/billing/service.go')
| -rw-r--r-- | internal/billing/service.go | 115 |
1 files changed, 85 insertions, 30 deletions
diff --git a/internal/billing/service.go b/internal/billing/service.go index 8a91839..8a5f266 100644 --- a/internal/billing/service.go +++ b/internal/billing/service.go @@ -39,6 +39,7 @@ type Service struct { stripeProductTaxCode string integrationIdentifier string createStripeCheckout stripeCheckoutCreator + createStripePortalSession stripePortalSessionCreator createStripeCustomer stripeCustomerCreator updateStripeCustomer stripeCustomerUpdater retrieveStripeSetupIntent stripeSetupIntentRetriever @@ -73,6 +74,7 @@ func New(ctx context.Context, options Options) (*Service, error) { if options.StripeEnabled { service.stripeClient = stripe.NewClient(options.StripeAPIKey) service.createStripeCheckout = service.stripeClient.V1CheckoutSessions.Create + service.createStripePortalSession = service.stripeClient.V1BillingPortalSessions.Create service.createStripeCustomer = service.stripeClient.V1Customers.Create service.updateStripeCustomer = service.stripeClient.V1Customers.Update service.retrieveStripeSetupIntent = service.stripeClient.V1SetupIntents.Retrieve @@ -103,7 +105,7 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { if input.Model.PriceCurrency != "" && input.Model.PriceCurrency != s.currency { return fmt.Errorf("model price currency %s does not match wallet currency %s", input.Model.PriceCurrency, s.currency) } - reserved, err := reservationCost(input.Model, input.Body, s.defaultMaxOutputTokens) + reserved, err := reservationCost(input.Model, input.Body, s.defaultMaxOutputTokens, input.Protocol) if err != nil { return err } @@ -156,6 +158,22 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { return ErrQuotaExceeded } } + if input.Principal.DailySpendMicros > 0 { + now := time.Now().UTC() + period := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) + nextPeriod := period.AddDate(0, 0, 1) + var used, pending int64 + if err := tx.QueryRow(ctx, `SELECT + COALESCE((SELECT sum(cost_micros) FROM usage_events WHERE key_id=$1 AND started_at >= $2 AND started_at < $3),0), + COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE key_id=$1 AND status IN ('pending','metering_failed') AND created_at >= $2 AND created_at < $3),0)`, + input.Principal.KeyID, period, nextPeriod).Scan(&used, &pending); err != nil { + return fmt.Errorf("read API key daily spend quota: %w", err) + } + limit := input.Principal.DailySpendMicros + if reserved > limit || used > limit-reserved || pending > limit-used-reserved { + return ErrDailyQuotaExceeded + } + } if balance-held < reserved { return ErrInsufficientBalance } @@ -250,15 +268,15 @@ func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error { } if _, err := tx.Exec(ctx, `INSERT INTO usage_events ( request_id,tenant_id,project_id,key_id,public_model,provider_id,upstream_model,protocol,stream, - status_code,success,error_type,attempts,started_at,duration_ms,input_tokens,output_tokens,total_tokens, + status_code,success,error_type,attempts,started_at,duration_ms,ttft_ms,input_tokens,output_tokens,total_tokens, cache_creation_input_tokens,cache_read_input_tokens,cost_micros,charged_micros,uncollected_micros, usage_reported,metering_status) VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,'usage_not_reported',$12,$13,$14, - $15,$16,$17,$18,$19,0,0,0,false,'missing') - ON CONFLICT (request_id) DO UPDATE SET error_type='usage_not_reported',usage_reported=false,metering_status='missing'`, + $15,$16,$17,$18,$19,$20,0,0,0,false,'missing') + ON CONFLICT (request_id) DO UPDATE SET error_type='usage_not_reported',ttft_ms=EXCLUDED.ttft_ms,usage_reported=false,metering_status='missing'`, event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel, string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.Attempts, event.StartedAt, - event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, + event.DurationMS, event.TTFTMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens); err != nil { return fmt.Errorf("persist unmetered usage event: %w", err) } @@ -316,21 +334,21 @@ func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error { } if usageAlreadyRecorded { if _, err := tx.Exec(ctx, `UPDATE usage_events SET cost_micros=$2, charged_micros=$3, uncollected_micros=$4, - usage_reported=$5,metering_status=$6 WHERE request_id=$1`, - event.RequestID, actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil { + ttft_ms=GREATEST(ttft_ms,$5),usage_reported=$6,metering_status=$7 WHERE request_id=$1`, + event.RequestID, actualCost, charged, uncollected, event.TTFTMS, event.UsageReported, meteringStatus(event)); err != nil { return fmt.Errorf("apply usage charge: %w", err) } } else if _, err := tx.Exec(ctx, ` INSERT INTO usage_events ( request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model, protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms, - input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, + ttft_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status) - VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25) + VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26) ON CONFLICT (request_id) DO NOTHING`, event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel, string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.ErrorType, event.Attempts, - event.StartedAt, event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, + event.StartedAt, event.DurationMS, event.TTFTMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens, actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil { return fmt.Errorf("persist usage event: %w", err) @@ -582,25 +600,32 @@ func meteringStatus(event domain.UsageEvent) string { return "missing" } -func reservationCost(model domain.Model, body []byte, defaultMaxOutput int64) (int64, error) { - maxOutput := defaultMaxOutput - if model.MaxOutputTokens > 0 && (maxOutput == 0 || model.MaxOutputTokens < maxOutput) { - maxOutput = model.MaxOutputTokens +func reservationCost(model domain.Model, body []byte, defaultMaxOutput int64, protocols ...domain.Protocol) (int64, error) { + protocol := domain.ProtocolOpenAI + if len(protocols) > 0 && protocols[0] != "" { + protocol = protocols[0] } + maxOutput := int64(0) var limits struct { MaxTokens int64 `json:"max_tokens"` MaxCompletionTokens int64 `json:"max_completion_tokens"` MaxOutputTokens int64 `json:"max_output_tokens"` } - if json.Unmarshal(body, &limits) == nil { - explicitMax := int64(0) - for _, value := range []int64{limits.MaxTokens, limits.MaxCompletionTokens, limits.MaxOutputTokens} { - if value > explicitMax { - explicitMax = value - } + if protocol != domain.ProtocolOpenAIEmbeddings { + maxOutput = defaultMaxOutput + if model.MaxOutputTokens > 0 && (maxOutput == 0 || model.MaxOutputTokens < maxOutput) { + maxOutput = model.MaxOutputTokens } - if explicitMax > 0 { - maxOutput = explicitMax + if json.Unmarshal(body, &limits) == nil { + explicitMax := int64(0) + for _, value := range []int64{limits.MaxTokens, limits.MaxCompletionTokens, limits.MaxOutputTokens} { + if value > explicitMax { + explicitMax = value + } + } + if explicitMax > 0 { + maxOutput = explicitMax + } } } cacheReservePrice := model.CacheReadPriceMicrosPerMillion @@ -618,21 +643,51 @@ func usageCost(usage domain.Usage, inputPrice, outputPrice, cacheReadPrice, cach } func calculateCost(input, output, cacheRead, cacheWrite, inputPrice, outputPrice, cacheReadPrice, cacheWritePrice int64) (int64, error) { - values := []int64{input, output, cacheRead, cacheWrite, inputPrice, outputPrice, cacheReadPrice, cacheWritePrice} - for _, value := range values { - if value < 0 { - return 0, errors.New("billing values cannot be negative") + return calculateMeteredCost([]meteredCharge{ + {Unit: domain.MeteringUnitToken, Quantity: input, PriceMicros: inputPrice, PerQuantity: microsPerUnit}, + {Unit: domain.MeteringUnitToken, Quantity: output, PriceMicros: outputPrice, PerQuantity: microsPerUnit}, + {Unit: domain.MeteringUnitToken, Quantity: cacheRead, PriceMicros: cacheReadPrice, PerQuantity: microsPerUnit}, + {Unit: domain.MeteringUnitToken, Quantity: cacheWrite, PriceMicros: cacheWritePrice, PerQuantity: microsPerUnit}, + }) +} + +type meteredCharge struct { + Unit domain.MeteringUnit + Quantity int64 + PriceMicros int64 + PerQuantity int64 +} + +// calculateMeteredCost is the common fixed-point primitive for token, image, +// and duration pricing. Token rates use PerQuantity=1_000_000; image and second +// rates can use PerQuantity=1 without changing wallet or ledger arithmetic. +func calculateMeteredCost(charges []meteredCharge) (int64, error) { + byScale := make(map[int64]*big.Int) + for _, charge := range charges { + if charge.Unit != domain.MeteringUnitToken && charge.Unit != domain.MeteringUnitImage && charge.Unit != domain.MeteringUnitSecond { + return 0, fmt.Errorf("unsupported metering unit %q", charge.Unit) + } + if charge.Quantity < 0 || charge.PriceMicros < 0 || charge.PerQuantity <= 0 { + return 0, errors.New("metering quantity, price, or scale is invalid") + } + if charge.Quantity == 0 || charge.PriceMicros == 0 { + continue + } + component := new(big.Int).Mul(big.NewInt(charge.Quantity), big.NewInt(charge.PriceMicros)) + if byScale[charge.PerQuantity] == nil { + byScale[charge.PerQuantity] = new(big.Int) } + byScale[charge.PerQuantity].Add(byScale[charge.PerQuantity], component) } total := new(big.Int) - for _, pair := range [][2]int64{{input, inputPrice}, {output, outputPrice}, {cacheRead, cacheReadPrice}, {cacheWrite, cacheWritePrice}} { - total.Add(total, new(big.Int).Mul(big.NewInt(pair[0]), big.NewInt(pair[1]))) + for scale, numerator := range byScale { + numerator.Add(numerator, big.NewInt(scale-1)) + numerator.Div(numerator, big.NewInt(scale)) + total.Add(total, numerator) } if total.Sign() == 0 { return 0, nil } - total.Add(total, big.NewInt(microsPerUnit-1)) - total.Div(total, big.NewInt(microsPerUnit)) if !total.IsInt64() { return 0, errors.New("calculated charge exceeds supported range") } |
