summaryrefslogtreecommitdiff
path: root/internal/billing/service.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/billing/service.go66
1 files changed, 48 insertions, 18 deletions
diff --git a/internal/billing/service.go b/internal/billing/service.go
index 30f0e32..8a91839 100644
--- a/internal/billing/service.go
+++ b/internal/billing/service.go
@@ -25,24 +25,29 @@ import (
const microsPerUnit = int64(1_000_000)
type Service struct {
- db *pgxpool.Pool
- currency string
- defaultMaxOutputTokens int64
- minTopUpMinor int64
- maxTopUpMinor int64
- stripeEnabled bool
- 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
+ db *pgxpool.Pool
+ currency string
+ defaultMaxOutputTokens int64
+ minTopUpMinor int64
+ maxTopUpMinor int64
+ stripeEnabled bool
+ stripeWebhookSecret string
+ stripeSuccessURL string
+ stripeCancelURL string
+ stripePortalReturnURL string
+ stripeAutomaticTax bool
+ stripeProductTaxCode string
+ integrationIdentifier string
+ createStripeCheckout stripeCheckoutCreator
+ createStripeCustomer stripeCustomerCreator
+ updateStripeCustomer stripeCustomerUpdater
+ retrieveStripeSetupIntent stripeSetupIntentRetriever
+ createStripePaymentIntent stripePaymentIntentCreator
+ retrieveStripePaymentIntent stripePaymentIntentRetriever
+ stripeClient *stripe.Client
+ settlementSpoolPath string
+ metrics OperationalMetrics
+ spoolMu sync.Mutex
}
func New(ctx context.Context, options Options) (*Service, error) {
@@ -68,6 +73,11 @@ 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.createStripeCustomer = service.stripeClient.V1Customers.Create
+ service.updateStripeCustomer = service.stripeClient.V1Customers.Update
+ service.retrieveStripeSetupIntent = service.stripeClient.V1SetupIntents.Retrieve
+ service.createStripePaymentIntent = service.stripeClient.V1PaymentIntents.Create
+ service.retrieveStripePaymentIntent = service.stripeClient.V1PaymentIntents.Retrieve
}
return service, nil
}
@@ -131,6 +141,21 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error {
return ErrQuotaExceeded
}
}
+ if input.Principal.MonthlySpendMicros > 0 {
+ period := time.Date(time.Now().UTC().Year(), time.Now().UTC().Month(), 1, 0, 0, 0, 0, time.UTC)
+ nextPeriod := period.AddDate(0, 1, 0)
+ 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 monthly spend quota: %w", err)
+ }
+ limit := input.Principal.MonthlySpendMicros
+ if reserved > limit || used > limit-reserved || pending > limit-used-reserved {
+ return ErrQuotaExceeded
+ }
+ }
if balance-held < reserved {
return ErrInsufficientBalance
}
@@ -200,6 +225,11 @@ func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error {
if err != nil {
return fmt.Errorf("load billing reservation: %w", err)
}
+ if _, err := tx.Exec(ctx, `UPDATE api_keys
+ SET last_used_at = GREATEST(COALESCE(last_used_at, $2), $2)
+ WHERE id = $1`, keyID, event.StartedAt); err != nil {
+ return fmt.Errorf("update API key last used time: %w", err)
+ }
if status != "pending" {
return tx.Commit(ctx)
}