From eadb2ffe85c43cf6fc741c9823cd28eedb4a844c Mon Sep 17 00:00:00 2001 From: Chia Date: Wed, 5 Aug 2026 22:01:29 +1200 Subject: feat: harden prepaid billing and commercial operations --- cmd/aigw/main.go | 149 ++++++++++++++++++++++++++++++----------- cmd/aigw/main_test.go | 14 ++++ cmd/migrate/main.go | 10 +++ cmd/reconcile-billing/main.go | 79 ++++++++++++++++++++++ cmd/rotate-credentials/main.go | 41 ++++++++++++ 5 files changed, 253 insertions(+), 40 deletions(-) create mode 100644 cmd/aigw/main_test.go create mode 100644 cmd/reconcile-billing/main.go create mode 100644 cmd/rotate-credentials/main.go (limited to 'cmd') diff --git a/cmd/aigw/main.go b/cmd/aigw/main.go index c7d9ed3..108a849 100644 --- a/cmd/aigw/main.go +++ b/cmd/aigw/main.go @@ -21,6 +21,7 @@ import ( "aigw/internal/httpapi" "aigw/internal/limits" "aigw/internal/mailer" + "aigw/internal/operations" "aigw/internal/provider" "aigw/internal/routing" "aigw/internal/telemetry" @@ -69,7 +70,8 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error { store, err = controlplane.NewStore(ctx, controlplane.Options{ DatabaseURL: cfg.ControlPlane.DatabaseURL, RedisURL: cfg.ControlPlane.RedisURL, CredentialKey: cfg.ControlPlane.CredentialKey, RedisChannel: cfg.ControlPlane.RedisChannel, - VersionCacheKey: cfg.ControlPlane.SnapshotCacheKey, + PreviousCredentialKeys: cfg.ControlPlane.PreviousCredentialKeys, + VersionCacheKey: cfg.ControlPlane.SnapshotCacheKey, }) if err != nil { return err @@ -138,6 +140,15 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error { return fmt.Errorf("configure SMTP: %w", err) } go mailer.NewWorker(store, sender, logger).Run(ctx) + go store.RunMailNotificationWorker(ctx, controlplane.MailNotificationConfig{ + LowBalanceMicros: cfg.Admin.Mail.LowBalanceMicros, + SpendAnomalyMultiplier: cfg.Admin.Mail.SpendAnomalyMultiplier, + SpendAnomalyMinMicros: cfg.Admin.Mail.SpendAnomalyMinMicros, + Interval: time.Duration(cfg.Admin.Mail.NotificationIntervalSeconds) * time.Second, + }, logger) + } + if cfg.Admin.Enabled { + go store.RunRetentionWorker(ctx, cfg.Admin.AuditRetentionDays, cfg.Admin.SecurityRetentionDays, logger) } var billingService *billing.Service @@ -150,11 +161,18 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error { StripeEnabled: cfg.Billing.Stripe.Enabled, StripeAPIKey: cfg.Billing.Stripe.APIKey, StripeWebhookSecret: cfg.Billing.Stripe.WebhookSecret, StripeSuccessURL: cfg.Billing.Stripe.SuccessURL, StripeCancelURL: cfg.Billing.Stripe.CancelURL, + StripePortalReturnURL: cfg.Billing.Stripe.PortalReturnURL, + StripeAutomaticTax: cfg.Billing.Stripe.AutomaticTaxEnabled, + StripeProductTaxCode: cfg.Billing.Stripe.ProductTaxCode, + SettlementSpoolPath: cfg.Billing.SettlementSpoolPath, + Metrics: metrics, }) if err != nil { return err } defer billingService.Close() + go billingService.RunSettlementWorker(ctx) + go billingService.RunStripeOperations(ctx) } var billingMeter billing.Meter if billingService != nil { @@ -162,58 +180,109 @@ func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error { } inferenceAPI := httpapi.New(httpapi.Options{ - Authenticator: authenticator, - Catalog: modelCatalog, - Router: routing.New(modelCatalog), - Forwarder: provider.New(cfg.UpstreamHTTP, metrics), - UsageSink: usageSink, - BillingMeter: billingMeter, - Limiter: requestLimiter, - UsageRecorder: store, - Metrics: metrics, - Logger: logger, - MaxBodyBytes: cfg.Server.MaxBodyBytes, - ExposeMetrics: cfg.Observability.ExposeMetrics, + Authenticator: authenticator, + Catalog: modelCatalog, + Router: routing.New(modelCatalog), + Forwarder: provider.New(cfg.UpstreamHTTP, metrics), + UsageSink: usageSink, + BillingMeter: billingMeter, + Limiter: requestLimiter, + UsageRecorder: optionalUsageRecorder(store), + Metrics: metrics, + Logger: logger, + MaxBodyBytes: cfg.Server.MaxBodyBytes, + ExposeMetrics: cfg.Observability.ExposeMetrics, + DeploymentRegion: cfg.Server.DeploymentRegion, }) - root := http.NewServeMux() - root.Handle("/", inferenceAPI.Handler()) + adminHandler := http.Handler(nil) if cfg.Admin.Enabled { - adminHandler := adminapi.New(adminapi.Options{ + adminHandler = adminapi.New(adminapi.Options{ Store: store, Manager: manager, Billing: billingService, Token: cfg.Admin.Token, Logger: logger, Prefix: cfg.Admin.BasePath, RegistrationEnabled: cfg.Admin.RegistrationEnabled, SessionTTL: time.Duration(cfg.Admin.SessionTTLHours) * time.Hour, Currency: cfg.Billing.Currency, PublicURL: cfg.Admin.PublicURL, WebAuthn: webAuthn, MailEnabled: cfg.Admin.Mail.Enabled, }).Handler() - root.Handle(cfg.Admin.BasePath, adminHandler) - root.Handle(cfg.Admin.BasePath+"/", adminHandler) } - if billingService != nil && billingService.StripeEnabled() { - root.Handle("/billing/stripe/webhook", billingService.WebhookHandler()) + operationHandler := operations.Handler{Store: store, Manager: manager, Billing: billingService, Metrics: metrics, + MaxSnapshotAge: 3 * time.Duration(cfg.ControlPlane.ReloadIntervalSeconds) * time.Second} + return serve(ctx, cfg, logger, store, inferenceAPI, adminHandler, billingService, operationHandler) +} + +func optionalUsageRecorder(store *controlplane.Store) httpapi.UsageRecorder { + if store == nil { + return nil } + return store +} - server := &http.Server{ - Addr: cfg.Server.Address, Handler: root, - ReadHeaderTimeout: cfg.Server.ReadHeaderTimeout(), - IdleTimeout: cfg.Server.IdleTimeout(), - } - errCh := make(chan error, 1) - go func() { - logger.Info("gateway_listening", "address", cfg.Server.Address, "control_plane", cfg.ControlPlane.Enabled, - "billing", cfg.Billing.Enabled, "stripe", cfg.Billing.Stripe.Enabled) - errCh <- server.ListenAndServe() - }() +func serve(ctx context.Context, cfg config.Config, logger *slog.Logger, store *controlplane.Store, inferenceAPI *httpapi.API, adminHandler http.Handler, billingService *billing.Service, operationHandler http.Handler) error { + type listener struct { + name, address string + handler http.Handler + } + var listeners []listener + webhookMux := http.NewServeMux() + hasWebhooks := false + if billingService != nil && billingService.StripeEnabled() { + webhookMux.Handle("/billing/stripe/webhook", billingService.WebhookHandler()) + hasWebhooks = true + } + if cfg.Admin.Enabled && cfg.Admin.Mail.Enabled && cfg.Admin.Mail.FeedbackSecret != "" { + webhookMux.Handle("/mail/feedback", store.MailFeedbackHandler(cfg.Admin.Mail.FeedbackSecret)) + hasWebhooks = true + } + if cfg.Server.SplitListeners { + listeners = append(listeners, listener{"inference", cfg.Server.PublicAddress, inferenceAPI.InferenceHandler()}, listener{"operations", cfg.Server.OperationsAddress, operationHandler}) + if adminHandler != nil { + listeners = append(listeners, listener{"admin", cfg.Server.AdminAddress, adminHandler}) + } + if hasWebhooks { + listeners = append(listeners, listener{"webhooks", cfg.Server.WebhookAddress, webhookMux}) + } + } else { + root := http.NewServeMux() + root.Handle("/healthz", operationHandler) + root.Handle("/readyz", operationHandler) + root.Handle("/metrics", operationHandler) + root.Handle("/", inferenceAPI.Handler()) + if adminHandler != nil { + root.Handle(cfg.Admin.BasePath, adminHandler) + root.Handle(cfg.Admin.BasePath+"/", adminHandler) + } + if hasWebhooks { + root.Handle("/billing/stripe/webhook", webhookMux) + root.Handle("/mail/feedback", webhookMux) + } + listeners = []listener{{"combined", cfg.Server.Address, root}} + } + servers := make([]*http.Server, 0, len(listeners)) + errCh := make(chan error, len(listeners)) + for _, item := range listeners { + handler, err := httpapi.TrustProxyHeaders(item.handler, cfg.Server.TrustedProxyCIDRs, cfg.Server.RequireHTTPS && item.name != "operations") + if err != nil { + return err + } + server := &http.Server{Addr: item.address, Handler: handler, ReadHeaderTimeout: cfg.Server.ReadHeaderTimeout(), IdleTimeout: cfg.Server.IdleTimeout()} + servers = append(servers, server) + go func(item listener, server *http.Server) { + logger.Info("http_listener_started", "name", item.name, "address", item.address) + errCh <- server.ListenAndServe() + }(item, server) + } select { case <-ctx.Done(): - shutdownContext, cancel := context.WithTimeout(context.Background(), cfg.Server.ShutdownTimeout()) - defer cancel() - if err := server.Shutdown(shutdownContext); err != nil { - return fmt.Errorf("shutdown HTTP server: %w", err) - } - return nil case err := <-errCh: - if errors.Is(err, http.ErrServerClosed) { - return nil + if !errors.Is(err, http.ErrServerClosed) { + return fmt.Errorf("serve HTTP: %w", err) + } + } + shutdownCtx, cancel := context.WithTimeout(context.Background(), cfg.Server.ShutdownTimeout()) + defer cancel() + var shutdownErr error + for _, server := range servers { + if err := server.Shutdown(shutdownCtx); err != nil && shutdownErr == nil { + shutdownErr = err } - return fmt.Errorf("serve HTTP: %w", err) } + return shutdownErr } diff --git a/cmd/aigw/main_test.go b/cmd/aigw/main_test.go new file mode 100644 index 0000000..449c35d --- /dev/null +++ b/cmd/aigw/main_test.go @@ -0,0 +1,14 @@ +package main + +import ( + "testing" + + "aigw/internal/controlplane" +) + +func TestOptionalUsageRecorderRejectsTypedNilStore(t *testing.T) { + var store *controlplane.Store + if recorder := optionalUsageRecorder(store); recorder != nil { + t.Fatal("nil control-plane store became a non-nil usage recorder") + } +} diff --git a/cmd/migrate/main.go b/cmd/migrate/main.go index 20452f9..a24b7bc 100644 --- a/cmd/migrate/main.go +++ b/cmd/migrate/main.go @@ -12,6 +12,7 @@ import ( func main() { environment := flag.String("database-url-env", "AIGW_DATABASE_URL", "environment variable containing the PostgreSQL URL") + status := flag.Bool("status", false, "print the latest applied migration instead of changing the schema") flag.Parse() databaseURL := os.Getenv(*environment) if databaseURL == "" { @@ -20,6 +21,15 @@ func main() { } ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() + if *status { + migration, err := controlplane.MigrationStatusDatabase(ctx, databaseURL) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + fmt.Printf("version=%d name=%s checksum=%s applied_at=%s\n", migration.Version, migration.Name, migration.Checksum, migration.AppliedAt.Format(time.RFC3339)) + return + } if err := controlplane.MigrateDatabase(ctx, databaseURL); err != nil { fmt.Fprintln(os.Stderr, err) os.Exit(1) diff --git a/cmd/reconcile-billing/main.go b/cmd/reconcile-billing/main.go new file mode 100644 index 0000000..79329ad --- /dev/null +++ b/cmd/reconcile-billing/main.go @@ -0,0 +1,79 @@ +package main + +import ( + "context" + "encoding/json" + "flag" + "fmt" + "os" + "time" + + "aigw/internal/billing" +) + +func main() { + databaseURLEnv := flag.String("database-url-env", "AIGW_DATABASE_URL", "environment variable containing the PostgreSQL URL") + stripeKeyEnv := flag.String("stripe-key-env", "AIGW_STRIPE_API_KEY", "environment variable containing the Stripe restricted key") + currency := flag.String("currency", "usd", "wallet currency") + limit := flag.Int("limit", 200, "maximum recent top-up orders to reconcile") + resolveMissing := flag.Bool("resolve-confirmed-missing", false, "close uncredited pending orders whose Stripe Sessions are confirmed missing") + resolutionReason := flag.String("resolution-reason", "maintenance reconciliation confirmed the uncredited Checkout Session is absent from the configured Stripe account", "auditable reason for resolving missing orders") + flag.Parse() + databaseURL, stripeKey := os.Getenv(*databaseURLEnv), os.Getenv(*stripeKeyEnv) + if databaseURL == "" || stripeKey == "" { + fmt.Fprintln(os.Stderr, "database and Stripe key environment variables are required") + os.Exit(1) + } + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute) + defer cancel() + service, err := billing.New(ctx, billing.Options{DatabaseURL: databaseURL, Currency: *currency, StripeEnabled: true, StripeAPIKey: stripeKey}) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + defer service.Close() + result, err := service.Reconcile(ctx, *limit) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + if *resolveMissing { + for _, mismatch := range result.Mismatches { + if mismatch["type"] != "stripe_session_missing" { + continue + } + orderID, _ := mismatch["order_id"].(string) + order, getErr := service.GetTopUpOrder(ctx, "", orderID) + if getErr != nil { + fmt.Fprintln(os.Stderr, getErr) + os.Exit(1) + } + var resolveErr error + if order.Status == "paid" { + _, resolveErr = service.ReverseMissingTopUpCredit(ctx, order.TenantID, order.ID, + billing.ResolveMissingTopUpInput{Reason: *resolutionReason}, billing.ResolutionActor{ID: "reconcile-billing", Type: "maintenance"}) + } else { + _, resolveErr = service.ResolveMissingTopUp(ctx, order.TenantID, order.ID, + billing.ResolveMissingTopUpInput{Reason: *resolutionReason}, billing.ResolutionActor{ID: "reconcile-billing", Type: "maintenance"}) + } + if resolveErr != nil { + fmt.Fprintln(os.Stderr, resolveErr) + os.Exit(1) + } + } + result, err = service.Reconcile(ctx, *limit) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + } + encoder := json.NewEncoder(os.Stdout) + encoder.SetIndent("", " ") + if err := encoder.Encode(result); err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + if result.MismatchCount > 0 { + os.Exit(2) + } +} diff --git a/cmd/rotate-credentials/main.go b/cmd/rotate-credentials/main.go new file mode 100644 index 0000000..598035d --- /dev/null +++ b/cmd/rotate-credentials/main.go @@ -0,0 +1,41 @@ +package main + +import ( + "context" + "fmt" + "os" + "strings" + "time" + + "aigw/internal/controlplane" +) + +func main() { + databaseURL := strings.TrimSpace(os.Getenv("AIGW_DATABASE_URL")) + current := strings.TrimSpace(os.Getenv("AIGW_CREDENTIAL_KEY")) + previousRaw := os.Getenv("AIGW_CREDENTIAL_PREVIOUS_KEYS") + if databaseURL == "" || current == "" || strings.TrimSpace(previousRaw) == "" { + fmt.Fprintln(os.Stderr, "AIGW_DATABASE_URL, AIGW_CREDENTIAL_KEY, and AIGW_CREDENTIAL_PREVIOUS_KEYS are required") + os.Exit(2) + } + previous := []string{} + for _, value := range strings.Split(previousRaw, ",") { + if value = strings.TrimSpace(value); value != "" { + previous = append(previous, value) + } + } + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Minute) + defer cancel() + store, err := controlplane.NewStore(ctx, controlplane.Options{DatabaseURL: databaseURL, CredentialKey: current, PreviousCredentialKeys: previous}) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + defer store.Close() + count, err := store.RotateCredentials(ctx) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + fmt.Printf("re-encrypted %d credential records\n", count) +} -- cgit v1.2.3