diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/config/config.go | 64 | ||||
| -rw-r--r-- | internal/config/config_test.go | 94 |
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") |
