diff options
Diffstat (limited to 'internal/config')
| -rw-r--r-- | internal/config/config.go | 171 | ||||
| -rw-r--r-- | internal/config/config_test.go | 80 |
2 files changed, 234 insertions, 17 deletions
diff --git a/internal/config/config.go b/internal/config/config.go index 97ecb21..376cf77 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -26,7 +26,8 @@ type Config struct { } type ServerConfig struct { - Address string `json:"address"` + Address string `json:"-"` + AddressEnv string `json:"address_env"` MaxBodyBytes int64 `json:"max_body_bytes"` ReadHeaderTimeoutSecs int `json:"read_header_timeout_seconds"` IdleTimeoutSecs int `json:"idle_timeout_seconds"` @@ -53,12 +54,40 @@ type ControlPlaneConfig struct { } type AdminConfig struct { - Enabled bool `json:"enabled"` - TokenEnv string `json:"token_env"` - BasePath string `json:"base_path"` - RegistrationEnabled bool `json:"registration_enabled"` - SessionTTLHours int `json:"session_ttl_hours"` - Token string `json:"-"` + Enabled bool `json:"enabled"` + TokenEnv string `json:"token_env"` + BasePath string `json:"base_path"` + RegistrationEnabled bool `json:"registration_enabled"` + SessionTTLHours int `json:"session_ttl_hours"` + PublicURL string `json:"-"` + PublicURLEnv string `json:"public_url_env"` + Mail MailConfig `json:"mail"` + WebAuthn WebAuthnConfig `json:"webauthn"` + Token string `json:"-"` +} + +type MailConfig struct { + Enabled bool `json:"enabled"` + FromName string `json:"from_name"` + TLSMode string `json:"tls_mode"` + FromAddressEnv string `json:"from_address_env"` + SMTPAddressEnv string `json:"smtp_address_env"` + SMTPUsernameEnv string `json:"smtp_username_env"` + SMTPPasswordEnv string `json:"smtp_password_env"` + SMTPImplicitTLS bool `json:"smtp_implicit_tls"` + FromAddress string `json:"-"` + SMTPAddress string `json:"-"` + SMTPUsername string `json:"-"` + SMTPPassword string `json:"-"` +} + +type WebAuthnConfig struct { + Enabled bool `json:"enabled"` + RPDisplayName string `json:"rp_display_name"` + RPIDEnv string `json:"rp_id_env"` + OriginsEnv string `json:"origins_env"` + RPID string `json:"-"` + Origins []string `json:"-"` } type UpstreamHTTPConfig struct { @@ -69,11 +98,12 @@ type UpstreamHTTPConfig struct { } type ProviderConfig struct { - ID string `json:"id"` - Protocol domain.Protocol `json:"protocol"` - BaseURL string `json:"base_url"` - APIKeyEnv string `json:"api_key_env"` - APIKey string `json:"-"` + ID string `json:"id"` + Protocol domain.Protocol `json:"protocol"` + BaseURL string `json:"-"` + BaseURLEnv string `json:"base_url_env"` + APIKeyEnv string `json:"api_key_env"` + APIKey string `json:"-"` } type ModelConfig struct { @@ -111,8 +141,10 @@ type StripeConfig struct { Enabled bool `json:"enabled"` APIKeyEnv string `json:"api_key_env"` WebhookSecretEnv string `json:"webhook_secret_env"` - SuccessURL string `json:"success_url"` - CancelURL string `json:"cancel_url"` + SuccessURL string `json:"-"` + CancelURL string `json:"-"` + SuccessURLEnv string `json:"success_url_env"` + CancelURLEnv string `json:"cancel_url_env"` APIKey string `json:"-"` WebhookSecret string `json:"-"` } @@ -148,6 +180,9 @@ func Load(path string) (Config, error) { } func applyDefaults(cfg *Config) { + if cfg.Server.AddressEnv == "" { + cfg.Server.AddressEnv = "AIGW_SERVER_ADDRESS" + } if cfg.Server.Address == "" { cfg.Server.Address = ":8080" } @@ -193,6 +228,36 @@ func applyDefaults(cfg *Config) { if cfg.Admin.SessionTTLHours == 0 { cfg.Admin.SessionTTLHours = 12 } + if cfg.Admin.PublicURLEnv == "" { + cfg.Admin.PublicURLEnv = "AIGW_PUBLIC_URL" + } + if cfg.Admin.Mail.FromName == "" { + cfg.Admin.Mail.FromName = "AIGW" + } + if cfg.Admin.Mail.TLSMode == "" { + cfg.Admin.Mail.TLSMode = "starttls" + } + if cfg.Admin.Mail.FromAddressEnv == "" { + cfg.Admin.Mail.FromAddressEnv = "AIGW_SMTP_FROM_ADDRESS" + } + if cfg.Admin.Mail.SMTPAddressEnv == "" { + cfg.Admin.Mail.SMTPAddressEnv = "AIGW_SMTP_ADDRESS" + } + if cfg.Admin.Mail.SMTPUsernameEnv == "" { + cfg.Admin.Mail.SMTPUsernameEnv = "AIGW_SMTP_USERNAME" + } + if cfg.Admin.Mail.SMTPPasswordEnv == "" { + cfg.Admin.Mail.SMTPPasswordEnv = "AIGW_SMTP_PASSWORD" + } + if cfg.Admin.WebAuthn.RPDisplayName == "" { + cfg.Admin.WebAuthn.RPDisplayName = "AIGW Console" + } + if cfg.Admin.WebAuthn.RPIDEnv == "" { + cfg.Admin.WebAuthn.RPIDEnv = "AIGW_WEBAUTHN_RP_ID" + } + if cfg.Admin.WebAuthn.OriginsEnv == "" { + cfg.Admin.WebAuthn.OriginsEnv = "AIGW_WEBAUTHN_ORIGINS" + } if cfg.UpstreamHTTP.MaxIdleConnections == 0 { cfg.UpstreamHTTP.MaxIdleConnections = 4096 } @@ -226,6 +291,12 @@ func applyDefaults(cfg *Config) { if cfg.Billing.Stripe.WebhookSecretEnv == "" { cfg.Billing.Stripe.WebhookSecretEnv = "AIGW_STRIPE_WEBHOOK_SECRET" } + if cfg.Billing.Stripe.SuccessURLEnv == "" { + cfg.Billing.Stripe.SuccessURLEnv = "AIGW_STRIPE_SUCCESS_URL" + } + if cfg.Billing.Stripe.CancelURLEnv == "" { + cfg.Billing.Stripe.CancelURLEnv = "AIGW_STRIPE_CANCEL_URL" + } for i := range cfg.Models { for j := range cfg.Models[i].Routes { if cfg.Models[i].Routes[j].Weight == 0 { @@ -236,6 +307,9 @@ func applyDefaults(cfg *Config) { } func resolveSecrets(cfg *Config) error { + if value := strings.TrimSpace(os.Getenv(cfg.Server.AddressEnv)); value != "" { + cfg.Server.Address = value + } if cfg.ControlPlane.Enabled { cfg.ControlPlane.DatabaseURL = os.Getenv(cfg.ControlPlane.DatabaseURLEnv) cfg.ControlPlane.RedisURL = os.Getenv(cfg.ControlPlane.RedisURLEnv) @@ -243,13 +317,42 @@ func resolveSecrets(cfg *Config) error { } if cfg.Admin.Enabled { cfg.Admin.Token = os.Getenv(cfg.Admin.TokenEnv) + if err := resolveRequiredEnv(&cfg.Admin.PublicURL, cfg.Admin.PublicURLEnv, "admin.public_url"); err != nil { + return err + } + if cfg.Admin.Mail.Enabled { + cfg.Admin.Mail.FromAddress = strings.TrimSpace(os.Getenv(cfg.Admin.Mail.FromAddressEnv)) + cfg.Admin.Mail.SMTPAddress = strings.TrimSpace(os.Getenv(cfg.Admin.Mail.SMTPAddressEnv)) + cfg.Admin.Mail.SMTPUsername = os.Getenv(cfg.Admin.Mail.SMTPUsernameEnv) + cfg.Admin.Mail.SMTPPassword = os.Getenv(cfg.Admin.Mail.SMTPPasswordEnv) + } + if cfg.Admin.WebAuthn.Enabled { + cfg.Admin.WebAuthn.RPID = strings.TrimSpace(os.Getenv(cfg.Admin.WebAuthn.RPIDEnv)) + for _, origin := range strings.Split(os.Getenv(cfg.Admin.WebAuthn.OriginsEnv), ",") { + if origin = strings.TrimSpace(origin); origin != "" { + cfg.Admin.WebAuthn.Origins = append(cfg.Admin.WebAuthn.Origins, origin) + } + } + } } if cfg.Billing.Enabled && cfg.Billing.Stripe.Enabled { cfg.Billing.Stripe.APIKey = os.Getenv(cfg.Billing.Stripe.APIKeyEnv) cfg.Billing.Stripe.WebhookSecret = os.Getenv(cfg.Billing.Stripe.WebhookSecretEnv) + if err := resolveRequiredEnv(&cfg.Billing.Stripe.SuccessURL, cfg.Billing.Stripe.SuccessURLEnv, "billing.stripe.success_url"); err != nil { + return err + } + if err := resolveRequiredEnv(&cfg.Billing.Stripe.CancelURL, cfg.Billing.Stripe.CancelURLEnv, "billing.stripe.cancel_url"); err != nil { + return err + } } for i := range cfg.Providers { provider := &cfg.Providers[i] + if provider.BaseURLEnv == "" { + return fmt.Errorf("provider %q: base_url_env is required", provider.ID) + } + if err := resolveRequiredEnv(&provider.BaseURL, provider.BaseURLEnv, fmt.Sprintf("provider %q base_url", provider.ID)); err != nil { + return err + } if provider.APIKeyEnv == "" { continue } @@ -261,6 +364,18 @@ func resolveSecrets(cfg *Config) error { return nil } +func resolveRequiredEnv(target *string, environment, field string) error { + if environment == "" { + return nil + } + value := strings.TrimSpace(os.Getenv(environment)) + if value == "" { + return fmt.Errorf("%s: environment variable %s is empty", field, environment) + } + *target = value + return nil +} + func Validate(cfg Config) error { if cfg.Server.MaxBodyBytes < 1024 { return errors.New("server.max_body_bytes must be at least 1024") @@ -293,6 +408,34 @@ func Validate(cfg Config) error { if cfg.Admin.SessionTTLHours < 1 || cfg.Admin.SessionTTLHours > 720 { return errors.New("admin.session_ttl_hours must be between 1 and 720") } + publicURL, err := url.Parse(cfg.Admin.PublicURL) + if err != nil || publicURL.Host == "" || (publicURL.Scheme != "http" && publicURL.Scheme != "https") { + return errors.New("admin.public_url must resolve from an environment variable to an absolute http(s) URL") + } + if cfg.Admin.Mail.Enabled { + if cfg.Admin.Mail.FromAddress == "" || cfg.Admin.Mail.SMTPAddress == "" { + return errors.New("admin.mail requires SMTP address and from address environment variables") + } + if (cfg.Admin.Mail.SMTPUsername == "") != (cfg.Admin.Mail.SMTPPassword == "") { + return errors.New("admin.mail SMTP username and password must both be set or both be empty") + } + switch cfg.Admin.Mail.TLSMode { + case "starttls", "tls", "none": + default: + return errors.New("admin.mail.tls_mode must be starttls, tls, or none") + } + } + if cfg.Admin.WebAuthn.Enabled { + if cfg.Admin.WebAuthn.RPID == "" || len(cfg.Admin.WebAuthn.Origins) == 0 { + return errors.New("admin.webauthn requires RP ID and origins environment variables") + } + for _, origin := range cfg.Admin.WebAuthn.Origins { + parsed, err := url.Parse(origin) + if err != nil || parsed.Host == "" || (parsed.Scheme != "https" && !(parsed.Scheme == "http" && parsed.Hostname() == "localhost")) { + return fmt.Errorf("admin.webauthn origin %q must use https (http is allowed only for localhost)", origin) + } + } + } } if cfg.Billing.Enabled { if !cfg.ControlPlane.Enabled { diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 2331b0f..1ecf70f 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -8,20 +8,25 @@ import ( func TestLoadAppliesDefaultsAndResolvesSecrets(t *testing.T) { t.Setenv("TEST_UPSTREAM_KEY", "secret") + t.Setenv("TEST_UPSTREAM_URL", "https://example.com/v1") + t.Setenv("AIGW_SERVER_ADDRESS", "127.0.0.1:9090") path := writeConfig(t, `{ - "providers": [{"id":"primary","protocol":"openai","base_url":"https://example.com/v1","api_key_env":"TEST_UPSTREAM_KEY"}], + "providers": [{"id":"primary","protocol":"openai","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}], "models": [{"id":"example/model","routes":[{"provider":"primary","upstream_model":"model"}]}] }`) cfg, err := Load(path) if err != nil { t.Fatal(err) } - if cfg.Server.Address != ":8080" || cfg.Server.MaxBodyBytes == 0 { + if cfg.Server.Address != "127.0.0.1:9090" || cfg.Server.MaxBodyBytes == 0 { t.Fatalf("defaults not applied: %+v", cfg.Server) } if cfg.Providers[0].APIKey != "secret" { t.Fatal("provider secret was not resolved") } + if cfg.Providers[0].BaseURL != "https://example.com/v1" { + t.Fatal("provider URL was not resolved") + } if cfg.Models[0].Routes[0].Weight != 1 { t.Fatalf("expected default route weight 1, got %d", cfg.Models[0].Routes[0].Weight) } @@ -40,11 +45,24 @@ func TestLoadRejectsUnknownFieldsAndTrailingData(t *testing.T) { } } +func TestLoadRejectsLiteralExternalServiceValues(t *testing.T) { + for _, content := range []string{ + `{"server":{"address":":9090"}}`, + `{"providers":[{"id":"primary","protocol":"openai","base_url":"https://example.test","api_key_env":"TEST_KEY"}]}`, + `{"billing":{"stripe":{"success_url":"https://console.example.test"}}}`, + } { + if _, err := Load(writeConfig(t, content)); err == nil { + t.Fatalf("expected literal external service value to be rejected: %s", content) + } + } +} + func TestLoadControlPlaneModeWithoutStaticRoutes(t *testing.T) { t.Setenv("AIGW_DATABASE_URL", "postgres://aigw:aigw@postgres/aigw") t.Setenv("AIGW_REDIS_URL", "redis://redis:6379/0") t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=") t.Setenv("AIGW_ADMIN_TOKEN", "admin-secret") + t.Setenv("AIGW_PUBLIC_URL", "http://localhost:8080/admin/") path := writeConfig(t, `{ "control_plane": {"enabled":true}, "admin": {"enabled":true} @@ -90,9 +108,11 @@ func TestLoadResolvesStripeSecrets(t *testing.T) { t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=") t.Setenv("TEST_STRIPE_KEY", "rk_test_example") t.Setenv("TEST_STRIPE_WEBHOOK", "whsec_example") + t.Setenv("TEST_STRIPE_SUCCESS", "https://console.example.test/admin/?topup=success") + t.Setenv("TEST_STRIPE_CANCEL", "https://console.example.test/admin/?topup=cancel") path := writeConfig(t, `{ "control_plane":{"enabled":true}, - "billing":{"enabled":true,"stripe":{"enabled":true,"api_key_env":"TEST_STRIPE_KEY","webhook_secret_env":"TEST_STRIPE_WEBHOOK","success_url":"http://localhost/admin/?topup=success","cancel_url":"http://localhost/admin/?topup=cancel"}} + "billing":{"enabled":true,"stripe":{"enabled":true,"api_key_env":"TEST_STRIPE_KEY","webhook_secret_env":"TEST_STRIPE_WEBHOOK","success_url_env":"TEST_STRIPE_SUCCESS","cancel_url_env":"TEST_STRIPE_CANCEL"}} }`) cfg, err := Load(path) if err != nil { @@ -101,6 +121,21 @@ func TestLoadResolvesStripeSecrets(t *testing.T) { if cfg.Billing.Stripe.APIKey != "rk_test_example" || cfg.Billing.Stripe.WebhookSecret != "whsec_example" { t.Fatal("Stripe secrets were not resolved") } + if cfg.Billing.Stripe.SuccessURL != "https://console.example.test/admin/?topup=success" || cfg.Billing.Stripe.CancelURL != "https://console.example.test/admin/?topup=cancel" { + t.Fatal("Stripe callback URLs were not resolved") + } +} + +func TestLoadRejectsEmptyExternalServiceEnvironment(t *testing.T) { + t.Setenv("TEST_UPSTREAM_KEY", "secret") + t.Setenv("TEST_UPSTREAM_URL", "") + path := writeConfig(t, `{ + "providers": [{"id":"primary","protocol":"openai","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}], + "models": [{"id":"example/model","routes":[{"provider":"primary","upstream_model":"model"}]}] +}`) + if _, err := Load(path); err == nil { + t.Fatal("expected empty provider URL environment variable to be rejected") + } } func TestLoadRejectsBillingWithoutControlPlane(t *testing.T) { @@ -110,6 +145,45 @@ func TestLoadRejectsBillingWithoutControlPlane(t *testing.T) { } } +func TestVersionedExamplesResolveExternalServicesFromEnvironment(t *testing.T) { + t.Setenv("AIGW_SERVER_ADDRESS", "127.0.0.1:18081") + t.Setenv("AIGW_API_KEYS", `[{"key":"test","key_id":"key","tenant_id":"tenant","project_id":"project","scopes":["inference"]}]`) + t.Setenv("OPENAI_BASE_URL", "https://openai.example.test/v1") + t.Setenv("OPENAI_API_KEY", "openai-secret") + t.Setenv("ANTHROPIC_BASE_URL", "https://anthropic.example.test/v1") + t.Setenv("ANTHROPIC_API_KEY", "anthropic-secret") + staticConfig, err := Load(filepath.Join("..", "..", "config.example.json")) + if err != nil { + t.Fatalf("load static example: %v", err) + } + if staticConfig.Providers[0].BaseURL != "https://openai.example.test/v1" { + t.Fatalf("static provider URL = %q", staticConfig.Providers[0].BaseURL) + } + + t.Setenv("AIGW_DATABASE_URL", "postgres://example:secret@postgres.example.test/aigw") + t.Setenv("AIGW_REDIS_URL", "redis://redis.example.test:6379/0") + t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=") + t.Setenv("AIGW_ADMIN_TOKEN", "bootstrap-secret") + t.Setenv("AIGW_PUBLIC_URL", "https://console.example.test/admin/") + t.Setenv("AIGW_SMTP_FROM_ADDRESS", "no-reply@example.test") + t.Setenv("AIGW_SMTP_ADDRESS", "smtp.example.test:587") + t.Setenv("AIGW_SMTP_USERNAME", "") + t.Setenv("AIGW_SMTP_PASSWORD", "") + t.Setenv("AIGW_WEBAUTHN_RP_ID", "console.example.test") + t.Setenv("AIGW_WEBAUTHN_ORIGINS", "https://console.example.test") + t.Setenv("AIGW_STRIPE_API_KEY", "rk_test_example") + t.Setenv("AIGW_STRIPE_WEBHOOK_SECRET", "whsec_example") + t.Setenv("AIGW_STRIPE_SUCCESS_URL", "https://console.example.test/admin/?topup=success") + t.Setenv("AIGW_STRIPE_CANCEL_URL", "https://console.example.test/admin/?topup=cancel") + controlConfig, err := Load(filepath.Join("..", "..", "config.control.example.json")) + if err != nil { + t.Fatalf("load control-plane example: %v", err) + } + if controlConfig.Billing.Stripe.SuccessURL != "https://console.example.test/admin/?topup=success" { + t.Fatalf("Stripe success URL = %q", controlConfig.Billing.Stripe.SuccessURL) + } +} + func writeConfig(t *testing.T, content string) string { t.Helper() path := filepath.Join(t.TempDir(), "config.json") |
