summaryrefslogtreecommitdiff
path: root/internal/config
diff options
context:
space:
mode:
Diffstat (limited to 'internal/config')
-rw-r--r--internal/config/config.go64
-rw-r--r--internal/config/config_test.go94
2 files changed, 158 insertions, 0 deletions
diff --git a/internal/config/config.go b/internal/config/config.go
index a67f221..047f372 100644
--- a/internal/config/config.go
+++ b/internal/config/config.go
@@ -81,6 +81,8 @@ type AdminConfig struct {
SecurityRetentionDays int `json:"security_retention_days"`
PublicURL string `json:"-"`
PublicURLEnv string `json:"public_url_env"`
+ InferencePublicURL string `json:"-"`
+ InferencePublicURLEnv string `json:"inference_public_url_env"`
Mail MailConfig `json:"mail"`
WebAuthn WebAuthnConfig `json:"webauthn"`
Token string `json:"-"`
@@ -125,7 +127,9 @@ type UpstreamHTTPConfig struct {
type ProviderConfig struct {
ID string `json:"id"`
+ Slug string `json:"slug"`
Protocol domain.Protocol `json:"protocol"`
+ WireAPI string `json:"wire_api"`
BaseURL string `json:"-"`
BaseURLEnv string `json:"base_url_env"`
APIKeyEnv string `json:"api_key_env"`
@@ -167,6 +171,7 @@ type BillingConfig struct {
type StripeConfig struct {
Enabled bool `json:"enabled"`
+ EnabledEnv string `json:"enabled_env"`
APIKeyEnv string `json:"api_key_env"`
WebhookSecretEnv string `json:"webhook_secret_env"`
SuccessURLEnv string `json:"success_url_env"`
@@ -297,6 +302,9 @@ func applyDefaults(cfg *Config) {
if cfg.Admin.PublicURLEnv == "" {
cfg.Admin.PublicURLEnv = "AIGW_PUBLIC_URL"
}
+ if cfg.Admin.InferencePublicURLEnv == "" {
+ cfg.Admin.InferencePublicURLEnv = "AIGW_INFERENCE_PUBLIC_URL"
+ }
if cfg.Admin.Mail.FromName == "" {
cfg.Admin.Mail.FromName = "AIGW"
}
@@ -400,6 +408,18 @@ func applyDefaults(cfg *Config) {
}
}
}
+ for i := range cfg.Providers {
+ if cfg.Providers[i].Slug == "" {
+ cfg.Providers[i].Slug = cfg.Providers[i].ID
+ }
+ if cfg.Providers[i].WireAPI == "" {
+ if cfg.Providers[i].Protocol == domain.ProtocolAnthropic {
+ cfg.Providers[i].WireAPI = "messages"
+ } else {
+ cfg.Providers[i].WireAPI = "chat_completions"
+ }
+ }
+ }
}
func resolveSecrets(cfg *Config) error {
@@ -446,6 +466,13 @@ func resolveSecrets(cfg *Config) error {
if err := resolveRequiredEnv(&cfg.Admin.PublicURL, cfg.Admin.PublicURLEnv, "admin.public_url"); err != nil {
return err
}
+ cfg.Admin.InferencePublicURL = strings.TrimRight(strings.TrimSpace(os.Getenv(cfg.Admin.InferencePublicURLEnv)), "/")
+ if cfg.Admin.InferencePublicURL == "" {
+ publicURL, err := url.Parse(cfg.Admin.PublicURL)
+ if err == nil && publicURL.Scheme != "" && publicURL.Host != "" {
+ cfg.Admin.InferencePublicURL = publicURL.Scheme + "://" + publicURL.Host
+ }
+ }
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))
@@ -462,6 +489,13 @@ func resolveSecrets(cfg *Config) error {
}
}
}
+ if cfg.Billing.Enabled && cfg.Billing.Stripe.EnabledEnv != "" {
+ var err error
+ cfg.Billing.Stripe.Enabled, err = envBool(cfg.Billing.Stripe.EnabledEnv)
+ if err != nil {
+ return err
+ }
+ }
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)
@@ -586,6 +620,10 @@ func Validate(cfg Config) error {
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")
}
+ inferenceURL, err := url.Parse(cfg.Admin.InferencePublicURL)
+ if err != nil || inferenceURL.Host == "" || (inferenceURL.Scheme != "http" && inferenceURL.Scheme != "https") {
+ return errors.New("admin.inference_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")
@@ -660,16 +698,30 @@ func Validate(cfg Config) error {
}
providers := make(map[string]ProviderConfig, len(cfg.Providers))
+ providerSlugs := make(map[string]string, len(cfg.Providers))
for _, provider := range cfg.Providers {
if provider.ID == "" {
return errors.New("provider id is required")
}
+ if !validProviderSlug(provider.Slug) {
+ return fmt.Errorf("provider %q: slug must be 3-64 lowercase letters, numbers, or hyphens", provider.ID)
+ }
+ if existingID, exists := providerSlugs[provider.Slug]; exists {
+ return fmt.Errorf("provider %q: duplicate public slug already used by provider %q", provider.ID, existingID)
+ }
+ providerSlugs[provider.Slug] = provider.ID
if _, exists := providers[provider.ID]; exists {
return fmt.Errorf("duplicate provider id %q", provider.ID)
}
if provider.Protocol != domain.ProtocolOpenAI && provider.Protocol != domain.ProtocolAnthropic {
return fmt.Errorf("provider %q: unsupported protocol %q", provider.ID, provider.Protocol)
}
+ if provider.Protocol == domain.ProtocolOpenAI && provider.WireAPI != "chat_completions" && provider.WireAPI != "responses" {
+ return fmt.Errorf("provider %q: wire_api must be chat_completions or responses", provider.ID)
+ }
+ if provider.Protocol == domain.ProtocolAnthropic && provider.WireAPI != "messages" {
+ return fmt.Errorf("provider %q: wire_api must be messages", provider.ID)
+ }
parsed, err := url.Parse(provider.BaseURL)
if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") {
return fmt.Errorf("provider %q: base_url must be an absolute http(s) URL", provider.ID)
@@ -719,6 +771,18 @@ func Validate(cfg Config) error {
return nil
}
+func validProviderSlug(value string) bool {
+ if len(value) < 3 || len(value) > 64 || value[0] == '-' || value[len(value)-1] == '-' {
+ return false
+ }
+ for _, character := range value {
+ if (character < 'a' || character > 'z') && (character < '0' || character > '9') && character != '-' {
+ return false
+ }
+ }
+ return true
+}
+
func (c ServerConfig) ReadHeaderTimeout() time.Duration {
return time.Duration(c.ReadHeaderTimeoutSecs) * time.Second
}
diff --git a/internal/config/config_test.go b/internal/config/config_test.go
index 8e2e9bb..bb3e814 100644
--- a/internal/config/config_test.go
+++ b/internal/config/config_test.go
@@ -27,11 +27,59 @@ func TestLoadAppliesDefaultsAndResolvesSecrets(t *testing.T) {
if cfg.Providers[0].BaseURL != "https://example.com/v1" {
t.Fatal("provider URL was not resolved")
}
+ if cfg.Providers[0].WireAPI != "chat_completions" {
+ t.Fatalf("default wire API = %q", cfg.Providers[0].WireAPI)
+ }
+ if cfg.Providers[0].Slug != "primary" {
+ t.Fatalf("default provider slug = %q", cfg.Providers[0].Slug)
+ }
if cfg.Models[0].Routes[0].Weight != 1 {
t.Fatalf("expected default route weight 1, got %d", cfg.Models[0].Routes[0].Weight)
}
}
+func TestLoadValidatesPublicProviderSlugs(t *testing.T) {
+ t.Setenv("TEST_UPSTREAM_KEY", "secret")
+ t.Setenv("TEST_UPSTREAM_URL", "https://example.com")
+ for _, providers := range []string{
+ `[{"id":"primary","slug":"Not Valid","protocol":"openai","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}]`,
+ `[{"id":"one","slug":"shared-provider","protocol":"openai","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"},{"id":"two","slug":"shared-provider","protocol":"openai","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}]`,
+ } {
+ path := writeConfig(t, `{"providers":`+providers+`,"models":[{"id":"example/model","routes":[{"provider":"primary","upstream_model":"model"}]}]}`)
+ if _, err := Load(path); err == nil {
+ t.Fatalf("expected invalid provider slugs to be rejected: %s", providers)
+ }
+ }
+}
+
+func TestLoadAcceptsOpenAIResponsesWireAPI(t *testing.T) {
+ t.Setenv("TEST_UPSTREAM_KEY", "secret")
+ t.Setenv("TEST_UPSTREAM_URL", "https://example.com")
+ path := writeConfig(t, `{
+ "providers": [{"id":"responses","protocol":"openai","wire_api":"responses","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}],
+ "models": [{"id":"example/model","routes":[{"provider":"responses","upstream_model":"gpt-example"}]}]
+}`)
+ cfg, err := Load(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if cfg.Providers[0].WireAPI != "responses" {
+ t.Fatalf("wire API = %q", cfg.Providers[0].WireAPI)
+ }
+}
+
+func TestLoadRejectsIncompatibleWireAPI(t *testing.T) {
+ t.Setenv("TEST_UPSTREAM_KEY", "secret")
+ t.Setenv("TEST_UPSTREAM_URL", "https://example.com")
+ path := writeConfig(t, `{
+ "providers": [{"id":"bad","protocol":"anthropic","wire_api":"responses","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}],
+ "models": [{"id":"example/model","routes":[{"provider":"bad","upstream_model":"model"}]}]
+}`)
+ if _, err := Load(path); err == nil {
+ t.Fatal("expected incompatible wire API to be rejected")
+ }
+}
+
func TestLoadRejectsUnknownFieldsAndTrailingData(t *testing.T) {
t.Setenv("TEST_UPSTREAM_KEY", "secret")
unknown := writeConfig(t, `{"unknown":true}`)
@@ -63,6 +111,7 @@ func TestLoadControlPlaneModeWithoutStaticRoutes(t *testing.T) {
t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=")
t.Setenv("AIGW_ADMIN_TOKEN", "admin-secret")
t.Setenv("AIGW_PUBLIC_URL", "http://localhost:8080/admin/")
+ t.Setenv("AIGW_INFERENCE_PUBLIC_URL", "https://api.example.test")
path := writeConfig(t, `{
"control_plane": {"enabled":true},
"admin": {"enabled":true}
@@ -78,6 +127,24 @@ func TestLoadControlPlaneModeWithoutStaticRoutes(t *testing.T) {
if len(cfg.Providers) != 0 || len(cfg.Models) != 0 {
t.Fatal("control-plane mode unexpectedly requires static providers or models")
}
+ if cfg.Admin.InferencePublicURL != "https://api.example.test" {
+ t.Fatalf("inference public URL = %q", cfg.Admin.InferencePublicURL)
+ }
+}
+
+func TestLoadDerivesInferenceURLFromConsoleOrigin(t *testing.T) {
+ t.Setenv("AIGW_DATABASE_URL", "postgres://aigw:aigw@postgres/aigw")
+ t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=")
+ t.Setenv("AIGW_ADMIN_TOKEN", "admin-secret")
+ t.Setenv("AIGW_PUBLIC_URL", "https://console.example.test/admin/")
+ path := writeConfig(t, `{"control_plane":{"enabled":true},"admin":{"enabled":true}}`)
+ cfg, err := Load(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if cfg.Admin.InferencePublicURL != "https://console.example.test" {
+ t.Fatalf("derived inference URL = %q", cfg.Admin.InferencePublicURL)
+ }
}
func TestLoadControlPlaneModeWithoutRedis(t *testing.T) {
@@ -128,6 +195,25 @@ func TestLoadResolvesStripeSecrets(t *testing.T) {
}
}
+func TestLoadCanDisableStripeFromEnvironmentWithoutCredentials(t *testing.T) {
+ t.Setenv("AIGW_DATABASE_URL", "postgres://aigw:aigw@postgres/aigw")
+ t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=")
+ t.Setenv("AIGW_SETTLEMENT_SPOOL_PATH", filepath.Join(t.TempDir(), "settlements.jsonl"))
+ t.Setenv("TEST_STRIPE_ENABLED", "false")
+ path := writeConfig(t, `{
+ "control_plane":{"enabled":true},
+ "billing":{"enabled":true,"currency":"usd","default_max_output_tokens":1024,"min_top_up_minor":500,"max_top_up_minor":1000000,
+ "stripe":{"enabled_env":"TEST_STRIPE_ENABLED"}}
+}`)
+ cfg, err := Load(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if cfg.Billing.Stripe.Enabled || cfg.Billing.Stripe.APIKey != "" {
+ t.Fatalf("Stripe should be disabled without credentials: %+v", cfg.Billing.Stripe)
+ }
+}
+
func TestLoadRejectsEmptyExternalServiceEnvironment(t *testing.T) {
t.Setenv("TEST_UPSTREAM_KEY", "secret")
t.Setenv("TEST_UPSTREAM_URL", "")
@@ -165,6 +251,13 @@ func TestVersionedExamplesResolveExternalServicesFromEnvironment(t *testing.T) {
if staticConfig.Providers[0].BaseURL != "https://openai.example.test/v1" {
t.Fatalf("static provider URL = %q", staticConfig.Providers[0].BaseURL)
}
+ responsesConfig, err := Load(filepath.Join("..", "..", "config.responses.example.json"))
+ if err != nil {
+ t.Fatalf("load Responses example: %v", err)
+ }
+ if len(responsesConfig.Providers) != 1 || responsesConfig.Providers[0].WireAPI != "responses" {
+ t.Fatalf("Responses example provider = %+v", responsesConfig.Providers)
+ }
t.Setenv("AIGW_DATABASE_URL", "postgres://example:secret@postgres.example.test/aigw")
t.Setenv("AIGW_REDIS_URL", "redis://redis.example.test:6379/0")
@@ -178,6 +271,7 @@ func TestVersionedExamplesResolveExternalServicesFromEnvironment(t *testing.T) {
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_ENABLED", "true")
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")