diff options
Diffstat (limited to 'cmd/aigw')
| -rw-r--r-- | cmd/aigw/main.go | 149 | ||||
| -rw-r--r-- | cmd/aigw/main_test.go | 14 |
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") + } +} |
