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/provider" "aigw/internal/routing" "aigw/internal/telemetry" ) 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 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, 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) } 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, }) if err != nil { return err } defer billingService.Close() } var billingMeter billing.Meter if billingService != nil { billingMeter = billingService } 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, }) root := http.NewServeMux() root.Handle("/", inferenceAPI.Handler()) 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, }).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()) } 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() }() 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 } return fmt.Errorf("serve HTTP: %w", err) } }