diff options
| author | Chia <Chia@93.nz> | 2026-08-05 22:01:29 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-05 22:07:50 +1200 |
| commit | eadb2ffe85c43cf6fc741c9823cd28eedb4a844c (patch) | |
| tree | 1aba2536d57360da403aa35c9ced58b615c7064e /internal/billing/service.go | |
| parent | cd0dd91ab93653631904f2ea0e574ccde6d60339 (diff) | |
feat: harden prepaid billing and commercial operations
Diffstat (limited to '')
| -rw-r--r-- | internal/billing/service.go | 319 |
1 files changed, 308 insertions, 11 deletions
diff --git a/internal/billing/service.go b/internal/billing/service.go index b87a0bd..30f0e32 100644 --- a/internal/billing/service.go +++ b/internal/billing/service.go @@ -1,6 +1,7 @@ package billing import ( + "bufio" "context" "crypto/rand" "encoding/hex" @@ -8,13 +9,17 @@ import ( "errors" "fmt" "math/big" + "os" + "path/filepath" "strings" + "sync" "time" "aigw/internal/domain" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" + "github.com/stripe/stripe-go/v86" ) const microsPerUnit = int64(1_000_000) @@ -29,8 +34,15 @@ type Service struct { stripeWebhookSecret string stripeSuccessURL string stripeCancelURL string + stripePortalReturnURL string + stripeAutomaticTax bool + stripeProductTaxCode string integrationIdentifier string createStripeCheckout stripeCheckoutCreator + stripeClient *stripe.Client + settlementSpoolPath string + metrics OperationalMetrics + spoolMu sync.Mutex } func New(ctx context.Context, options Options) (*Service, error) { @@ -47,10 +59,15 @@ func New(ctx context.Context, options Options) (*Service, error) { minTopUpMinor: options.MinTopUpMinor, maxTopUpMinor: options.MaxTopUpMinor, stripeEnabled: options.StripeEnabled, stripeWebhookSecret: options.StripeWebhookSecret, stripeSuccessURL: options.StripeSuccessURL, stripeCancelURL: options.StripeCancelURL, + stripePortalReturnURL: options.StripePortalReturnURL, stripeAutomaticTax: options.StripeAutomaticTax, + stripeProductTaxCode: options.StripeProductTaxCode, integrationIdentifier: "aigw_balance_" + randomLetters(8), + settlementSpoolPath: strings.TrimSpace(options.SettlementSpoolPath), + metrics: options.Metrics, } if options.StripeEnabled { - service.createStripeCheckout = newStripeCheckoutCreator(options.StripeAPIKey) + service.stripeClient = stripe.NewClient(options.StripeAPIKey) + service.createStripeCheckout = service.stripeClient.V1CheckoutSessions.Create } return service, nil } @@ -67,10 +84,15 @@ func (s *Service) Currency() string { return s.currency } +func (s *Service) Ping(ctx context.Context) error { return s.db.Ping(ctx) } + func (s *Service) Authorize(ctx context.Context, input Authorization) error { if input.RequestID == "" || input.Principal.TenantID == "" || input.Principal.ProjectID == "" || input.Principal.KeyID == "" { return errors.New("billing authorization identity is incomplete") } + 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) if err != nil { return err @@ -101,7 +123,7 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { var used, pending int64 if err := tx.QueryRow(ctx, `SELECT COALESCE((SELECT cost_micros FROM usage_monthly_rollups WHERE project_id=$1 AND period_start=$2),0), - COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE project_id=$1 AND status='pending' AND created_at >= $2 AND created_at < $3),0)`, + COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE project_id=$1 AND status IN ('pending','metering_failed') AND created_at >= $2 AND created_at < $3),0)`, input.Principal.ProjectID, period, nextPeriod).Scan(&used, &pending); err != nil { return fmt.Errorf("read monthly spend quota: %w", err) } @@ -115,15 +137,19 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { if _, err := tx.Exec(ctx, ` INSERT INTO billing_reservations ( request_id, tenant_id, project_id, key_id, public_model, currency, reserved_micros, + price_version_id, input_price_micros_per_million, output_price_micros_per_million, cache_read_price_micros_per_million, cache_write_price_micros_per_million) - VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)`, + VALUES ($1,$2,$3,$4,$5,$6,$7,NULLIF($8,'')::uuid,$9,$10,$11,$12)`, input.RequestID, input.Principal.TenantID, input.Principal.ProjectID, input.Principal.KeyID, - input.Model.ID, s.currency, reserved, input.Model.InputPriceMicrosPerMillion, + input.Model.ID, s.currency, reserved, input.Model.PriceVersionID, input.Model.InputPriceMicrosPerMillion, input.Model.OutputPriceMicrosPerMillion, input.Model.CacheReadPriceMicrosPerMillion, input.Model.CacheWritePriceMicrosPerMillion); err != nil { return fmt.Errorf("create billing reservation: %w", err) } + if _, err := tx.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, input.RequestID); err != nil { + return fmt.Errorf("create settlement job: %w", err) + } if _, err := tx.Exec(ctx, ` UPDATE tenant_wallets SET reserved_micros = reserved_micros + $2, updated_at = now() WHERE tenant_id = $1`, input.Principal.TenantID, reserved); err != nil { @@ -135,7 +161,27 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { return nil } -func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error { +func (s *Service) EnqueueSettlement(ctx context.Context, event domain.UsageEvent) error { + payload, err := json.Marshal(event) + if err != nil { + return fmt.Errorf("encode settlement event: %w", err) + } + _, err = s.db.Exec(ctx, ` + INSERT INTO billing_settlement_jobs (request_id, event, status, available_at, updated_at) + VALUES ($1,$2,'pending',now(),now()) + ON CONFLICT (request_id) DO UPDATE SET event=EXCLUDED.event, + status=CASE WHEN billing_settlement_jobs.status='done' THEN 'done' ELSE 'pending' END, + available_at=now(), locked_at=NULL, last_error='', updated_at=now()`, event.RequestID, payload) + if err == nil { + return nil + } + if spoolErr := s.appendSettlementSpool(payload); spoolErr != nil { + return fmt.Errorf("enqueue settlement in PostgreSQL: %v; append durable spool: %w", err, spoolErr) + } + return nil +} + +func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error { tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) if err != nil { return fmt.Errorf("begin usage settlement: %w", err) @@ -159,7 +205,51 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error { } actualCost := int64(0) - if event.StatusCode >= 200 && event.StatusCode < 300 { + billableSuccess := event.StatusCode >= 200 && event.StatusCode < 300 && event.Success + if billableSuccess && !event.UsageReported { + // Fail closed: keep the authorization hold in place and make the request + // visible to reconciliation. Releasing it would turn an unmetered success + // into a free request; guessing tokens here could overcharge the customer. + if _, err := tx.Exec(ctx, `UPDATE billing_reservations SET status='metering_failed', settled_at=now() + WHERE request_id=$1`, event.RequestID); err != nil { + return fmt.Errorf("mark unmetered reservation: %w", err) + } + var usageAlreadyRecorded bool + if err := tx.QueryRow(ctx, `SELECT true FROM usage_events WHERE request_id=$1 FOR UPDATE`, event.RequestID).Scan(&usageAlreadyRecorded); err != nil && !errors.Is(err, pgx.ErrNoRows) { + return fmt.Errorf("lock unmetered usage event: %w", err) + } + 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,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'`, + 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.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens); err != nil { + return fmt.Errorf("persist unmetered usage event: %w", err) + } + if !usageAlreadyRecorded { + period := time.Date(event.StartedAt.UTC().Year(), event.StartedAt.UTC().Month(), 1, 0, 0, 0, 0, time.UTC) + if _, err := tx.Exec(ctx, `INSERT INTO usage_monthly_rollups + (period_start,tenant_id,project_id,request_count,successful_requests,input_tokens,output_tokens,total_tokens) + VALUES ($1,$2,$3,1,1,$4,$5,$6) + ON CONFLICT (project_id,period_start) DO UPDATE SET + request_count=usage_monthly_rollups.request_count+1, + successful_requests=usage_monthly_rollups.successful_requests+1, + input_tokens=usage_monthly_rollups.input_tokens+EXCLUDED.input_tokens, + output_tokens=usage_monthly_rollups.output_tokens+EXCLUDED.output_tokens, + total_tokens=usage_monthly_rollups.total_tokens+EXCLUDED.total_tokens,updated_at=now()`, + period, tenantID, projectID, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens); err != nil { + return fmt.Errorf("roll up unmetered usage event: %w", err) + } + } + return tx.Commit(ctx) + } + if billableSuccess { actualCost, err = usageCost(event.Usage, inputPrice, outputPrice, cacheReadPrice, cacheWritePrice) if err != nil { return err @@ -195,8 +285,9 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error { return fmt.Errorf("lock existing usage event: %w", err) } if usageAlreadyRecorded { - if _, err := tx.Exec(ctx, `UPDATE usage_events SET cost_micros=$2, charged_micros=$3, uncollected_micros=$4 WHERE request_id=$1`, - event.RequestID, actualCost, charged, uncollected); err != nil { + 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 { return fmt.Errorf("apply usage charge: %w", err) } } else if _, err := tx.Exec(ctx, ` @@ -204,14 +295,14 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error { 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, - cost_micros, charged_micros, uncollected_micros) - 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) + 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) 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.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens, - actualCost, charged, uncollected); err != nil { + actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil { return fmt.Errorf("persist usage event: %w", err) } if charged > 0 { @@ -251,6 +342,199 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error { return nil } +// RunSettlementWorker processes jobs using PostgreSQL row locks so any gateway +// instance can resume work left by another instance after a crash. +func (s *Service) RunSettlementWorker(ctx context.Context) { + ticker := time.NewTicker(time.Second) + defer ticker.Stop() + for { + s.drainSettlementSpool(ctx) + s.recoverStaleSettlements(ctx) + for i := 0; i < 32; i++ { + processed, err := s.processSettlementJob(ctx) + if err != nil || !processed { + break + } + } + select { + case <-ctx.Done(): + return + case <-ticker.C: + } + } +} + +func (s *Service) recoverStaleSettlements(ctx context.Context) { + // A process may die after authorization but before it can attach a response + // event. Releasing after an hour prevents permanent holds; such synthetic + // events stay visible in Usage for reconciliation. + rows, err := s.db.Query(ctx, `SELECT request_id,tenant_id::text,project_id::text,key_id::text,public_model,created_at + FROM billing_reservations WHERE status='pending' AND created_at<now()-interval '1 hour' + AND EXISTS(SELECT 1 FROM billing_settlement_jobs j WHERE j.request_id=billing_reservations.request_id AND j.status='awaiting_event') LIMIT 100`) + if err != nil { + return + } + defer rows.Close() + for rows.Next() { + var event domain.UsageEvent + if rows.Scan(&event.RequestID, &event.TenantID, &event.ProjectID, &event.KeyID, &event.PublicModel, &event.StartedAt) != nil { + continue + } + event.StatusCode = 500 + event.Success = false + event.ErrorType = "gateway_interrupted_before_settlement" + event.DurationMS = time.Since(event.StartedAt).Milliseconds() + _ = s.EnqueueSettlement(ctx, event) + } + _, _ = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='retry',locked_at=NULL,available_at=now(),updated_at=now(),last_error='recovered stale processing lease' + WHERE status='processing' AND locked_at<now()-interval '5 minutes'`) +} + +func (s *Service) processSettlementJob(ctx context.Context) (bool, error) { + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return false, err + } + defer tx.Rollback(ctx) + var requestID string + var payload []byte + err = tx.QueryRow(ctx, ` + WITH selected AS ( + SELECT request_id FROM billing_settlement_jobs + WHERE status IN ('pending','retry') AND available_at <= now() + ORDER BY available_at, created_at FOR UPDATE SKIP LOCKED LIMIT 1 + ) + UPDATE billing_settlement_jobs j SET status='processing', attempts=attempts+1, + locked_at=now(), updated_at=now() + FROM selected WHERE j.request_id=selected.request_id + RETURNING j.request_id, j.event`, + ).Scan(&requestID, &payload) + if errors.Is(err, pgx.ErrNoRows) { + return false, tx.Commit(ctx) + } + if err != nil { + return false, err + } + if err := tx.Commit(ctx); err != nil { + return false, err + } + var event domain.UsageEvent + if err := json.Unmarshal(payload, &event); err != nil { + s.retrySettlement(ctx, requestID, fmt.Errorf("decode settlement event: %w", err)) + return true, err + } + jobCtx, cancel := context.WithTimeout(ctx, 15*time.Second) + err = s.settle(jobCtx, event) + cancel() + if err != nil { + s.retrySettlement(ctx, requestID, err) + return true, err + } + _, err = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='done', locked_at=NULL, + last_error='', updated_at=now(), completed_at=now() WHERE request_id=$1`, requestID) + return true, err +} + +func (s *Service) retrySettlement(ctx context.Context, requestID string, cause error) { + message := cause.Error() + if len(message) > 1000 { + message = message[:1000] + } + _, _ = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='retry', locked_at=NULL, + last_error=$2, available_at=now() + make_interval(secs => LEAST(300, power(2, LEAST(attempts, 8))::int)), + updated_at=now() WHERE request_id=$1`, requestID, message) +} + +func (s *Service) appendSettlementSpool(payload []byte) error { + if s.settlementSpoolPath == "" { + return errors.New("settlement spool path is not configured") + } + s.spoolMu.Lock() + defer s.spoolMu.Unlock() + if err := os.MkdirAll(filepath.Dir(s.settlementSpoolPath), 0o700); err != nil { + return err + } + file, err := os.OpenFile(s.settlementSpoolPath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o600) + if err != nil { + return err + } + defer file.Close() + if _, err := file.Write(append(payload, '\n')); err != nil { + return err + } + return file.Sync() +} + +func (s *Service) drainSettlementSpool(ctx context.Context) { + if s.settlementSpoolPath == "" { + return + } + s.spoolMu.Lock() + defer s.spoolMu.Unlock() + file, err := os.Open(s.settlementSpoolPath) + if errors.Is(err, os.ErrNotExist) { + return + } + if err != nil { + return + } + var pending [][]byte + scanner := bufio.NewScanner(file) + scanner.Buffer(make([]byte, 64*1024), 2<<20) + for scanner.Scan() { + line := append([]byte(nil), scanner.Bytes()...) + var event domain.UsageEvent + if json.Unmarshal(line, &event) != nil || event.RequestID == "" { + pending = append(pending, line) + continue + } + if _, err := s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET event=$2, status=CASE WHEN status='done' THEN 'done' ELSE 'pending' END, + available_at=now(), locked_at=NULL, updated_at=now() WHERE request_id=$1`, event.RequestID, line); err != nil { + pending = append(pending, line) + } + } + _ = file.Close() + if scanner.Err() != nil { + return + } + temporary := s.settlementSpoolPath + ".tmp" + out, err := os.OpenFile(temporary, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o600) + if err != nil { + return + } + for _, line := range pending { + _, _ = out.Write(append(line, '\n')) + } + _ = out.Sync() + _ = out.Close() + _ = os.Rename(temporary, s.settlementSpoolPath) +} + +func (s *Service) SettlementQueueStatus(ctx context.Context) (SettlementQueueStatus, error) { + var result SettlementQueueStatus + err := s.db.QueryRow(ctx, `SELECT + count(*) FILTER (WHERE status='awaiting_event'), count(*) FILTER (WHERE status='pending'), + count(*) FILTER (WHERE status='processing'), count(*) FILTER (WHERE status='retry'), + min(created_at) FILTER (WHERE status IN ('awaiting_event','pending','processing','retry')) + FROM billing_settlement_jobs`).Scan(&result.AwaitingEvent, &result.Pending, &result.Processing, &result.Retrying, &result.OldestPending) + if err != nil { + return result, err + } + if s.settlementSpoolPath != "" { + s.spoolMu.Lock() + file, openErr := os.Open(s.settlementSpoolPath) + if openErr == nil { + scanner := bufio.NewScanner(file) + for scanner.Scan() { + result.SpoolRecords++ + } + _ = file.Close() + } + s.spoolMu.Unlock() + } + return result, nil +} + func boolToInt(value bool) int { if value { return 1 @@ -258,8 +542,21 @@ func boolToInt(value bool) int { return 0 } +func meteringStatus(event domain.UsageEvent) string { + if event.StatusCode < 200 || event.StatusCode >= 300 || !event.Success { + return "upstream_failed" + } + if event.UsageReported { + return "reported" + } + 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 + } var limits struct { MaxTokens int64 `json:"max_tokens"` MaxCompletionTokens int64 `json:"max_completion_tokens"` |
