summaryrefslogtreecommitdiff
path: root/internal/billing/service.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-05 22:01:29 +1200
committerChia <Chia@93.nz>2026-08-05 22:07:50 +1200
commiteadb2ffe85c43cf6fc741c9823cd28eedb4a844c (patch)
tree1aba2536d57360da403aa35c9ced58b615c7064e /internal/billing/service.go
parentcd0dd91ab93653631904f2ea0e574ccde6d60339 (diff)
feat: harden prepaid billing and commercial operations
Diffstat (limited to '')
-rw-r--r--internal/billing/service.go319
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"`