diff options
Diffstat (limited to 'cmd')
| -rw-r--r-- | cmd/aigw/main.go | 186 |
1 files changed, 186 insertions, 0 deletions
diff --git a/cmd/aigw/main.go b/cmd/aigw/main.go new file mode 100644 index 0000000..de929de --- /dev/null +++ b/cmd/aigw/main.go @@ -0,0 +1,186 @@ +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) + } +} |
