summaryrefslogtreecommitdiff
path: root/internal/billing/service.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/billing/service.go')
-rw-r--r--internal/billing/service.go115
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")
}