package main import ( "context" "errors" "flag" "fmt" "log/slog" "net/http" "os" "os/signal" "syscall" "time" "aigw/internal/adminapi" "aigw/internal/auth" "aigw/internal/billing" "aigw/internal/catalog" "aigw/internal/config" "aigw/internal/controlplane" "aigw/internal/httpapi" "aigw/internal/limits" "aigw/internal/mailer" "aigw/internal/operations" "aigw/internal/provider" "aigw/internal/providerhealth" "aigw/internal/routing" "aigw/internal/telemetry" "github.com/go-webauthn/webauthn/webauthn" ) func main() { configPath := flag.String("config", "config.json", "path to the gateway configuration") flag.Parse() logger := slog.New(slog.NewJSONHandler(os.Stdout, nil)) cfg, err := config.Load(*configPath) if err != nil { logger.Error("configuration_failed", "error", err) os.Exit(1) } ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) defer stop() if err := run(ctx, cfg, logger); err != nil { logger.Error("gateway_stopped", "error", err) os.Exit(1) } } func run(ctx context.Context, cfg config.Config, logger *slog.Logger) error { metrics := &telemetry.Metrics{} usageSink := telemetry.NewAsyncUsageLogger(logger, metrics, cfg.Observability.UsageBuffer) defer func() { closeContext, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() if err := usageSink.Close(closeContext); err != nil { logger.Warn("usage_sink_close_failed", "error", err) } }() var authenticator *auth.StaticAuthenticator var modelCatalog *catalog.Catalog var store *controlplane.Store var manager *controlplane.Manager var managerCancel context.CancelFunc var managerDone chan struct{} var requestLimiter *limits.Limiter var webAuthn *webauthn.WebAuthn if cfg.ControlPlane.Enabled { var err error store, err = controlplane.NewStore(ctx, controlplane.Options{ DatabaseURL: cfg.ControlPlane.DatabaseURL, RedisURL: cfg.ControlPlane.RedisURL, CredentialKey: cfg.ControlPlane.CredentialKey, RedisChannel: cfg.ControlPlane.RedisChannel, PreviousCredentialKeys: cfg.ControlPlane.PreviousCredentialKeys, VersionCacheKey: cfg.ControlPlane.SnapshotCacheKey, }) if err != nil { return err } defer store.Close() if cfg.ControlPlane.AutoMigrate { if err := store.Migrate(ctx); err != nil { return err } } authenticator = auth.NewDynamic(nil, cfg.Auth.AllowAnonymous) modelCatalog = catalog.NewModels(nil) requestLimiter = limits.New(cfg.ControlPlane.RedisURL, "aigw:limits", cfg.Billing.DefaultMaxOutputTokens, logger) defer requestLimiter.Close() manager = controlplane.NewManager(store, modelCatalog, authenticator, logger, time.Duration(cfg.ControlPlane.ReloadIntervalSeconds)*time.Second, requestLimiter) if _, err := manager.Reload(ctx); err != nil { return fmt.Errorf("load initial control-plane snapshot: %w", err) } managerContext, cancel := context.WithCancel(ctx) managerCancel = cancel managerDone = make(chan struct{}) go func() { defer close(managerDone) manager.Run(managerContext) }() defer func() { managerCancel() select { case <-managerDone: case <-time.After(3 * time.Second): logger.Warn("control_plane_shutdown_timed_out") } }() } else { var err error authenticator, err = auth.NewStatic(os.Getenv(cfg.Auth.KeysEnv), cfg.Auth.AllowAnonymous) if err != nil { return err } modelCatalog = catalog.New(cfg) } if cfg.Admin.Enabled && cfg.Admin.WebAuthn.Enabled { configured, err := webauthn.New(&webauthn.Config{ RPDisplayName: cfg.Admin.WebAuthn.RPDisplayName, RPID: cfg.Admin.WebAuthn.RPID, RPOrigins: cfg.Admin.WebAuthn.Origins, EncodeUserIDAsString: true, Timeouts: webauthn.TimeoutsConfig{ Login: webauthn.TimeoutConfig{Enforce: true}, Registration: webauthn.TimeoutConfig{Enforce: true}, }, }) if err != nil { return fmt.Errorf("configure WebAuthn: %w", err) } webAuthn = configured } if cfg.Admin.Enabled && cfg.Admin.Mail.Enabled { sender, err := mailer.NewSender(mailer.SMTPConfig{ Address: cfg.Admin.Mail.SMTPAddress, Username: cfg.Admin.Mail.SMTPUsername, Password: cfg.Admin.Mail.SMTPPassword, FromName: cfg.Admin.Mail.FromName, FromAddress: cfg.Admin.Mail.FromAddress, TLSMode: cfg.Admin.Mail.TLSMode, }) if err != nil { 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 if cfg.Billing.Enabled { var err error billingService, err = billing.New(ctx, billing.Options{ DatabaseURL: cfg.ControlPlane.DatabaseURL, Currency: cfg.Billing.Currency, DefaultMaxOutputTokens: cfg.Billing.DefaultMaxOutputTokens, MinTopUpMinor: cfg.Billing.MinTopUpMinor, MaxTopUpMinor: cfg.Billing.MaxTopUpMinor, 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 { billingMeter = billingService } routeHealth := providerhealth.New(providerhealth.Options{}) inferenceAPI := httpapi.New(httpapi.Options{ Authenticator: authenticator, Catalog: modelCatalog, Router: routing.New(modelCatalog, routeHealth), Forwarder: provider.New(cfg.UpstreamHTTP, metrics, routeHealth), 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, BrowserOrigin: cfg.Admin.PublicURL, }) adminHandler := http.Handler(nil) if cfg.Admin.Enabled { 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, InferencePublicURL: cfg.Admin.InferencePublicURL, DefaultLowBalanceMicros: cfg.Admin.Mail.LowBalanceMicros, Catalog: modelCatalog, ProviderHealth: routeHealth, }).Handler() } 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 } 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(): case err := <-errCh: 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 shutdownErr }