summaryrefslogtreecommitdiff
path: root/cmd/aigw
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 /cmd/aigw
parentcd0dd91ab93653631904f2ea0e574ccde6d60339 (diff)
feat: harden prepaid billing and commercial operations
Diffstat (limited to 'cmd/aigw')
-rw-r--r--cmd/aigw/main.go149
-rw-r--r--cmd/aigw/main_test.go14
2 files changed, 123 insertions, 40 deletions
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")
+ }
+}