diff options
Diffstat (limited to '')
| -rw-r--r-- | internal/config/config.go | 64 |
1 files changed, 64 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 } |
