summaryrefslogtreecommitdiff
path: root/cmd
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--cmd/aigw/main.go186
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)
+ }
+}