diff options
Diffstat (limited to 'internal/config/config.go')
| -rw-r--r-- | internal/config/config.go | 74 |
1 files changed, 63 insertions, 11 deletions
diff --git a/internal/config/config.go b/internal/config/config.go index 047f372..2df74cf 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -15,15 +15,16 @@ import ( ) type Config struct { - Server ServerConfig `json:"server"` - Auth AuthConfig `json:"auth"` - ControlPlane ControlPlaneConfig `json:"control_plane"` - Admin AdminConfig `json:"admin"` - UpstreamHTTP UpstreamHTTPConfig `json:"upstream_http"` - Providers []ProviderConfig `json:"providers"` - Models []ModelConfig `json:"models"` - Billing BillingConfig `json:"billing"` - Observability ObservabilityConfig `json:"observability"` + Server ServerConfig `json:"server"` + Auth AuthConfig `json:"auth"` + ControlPlane ControlPlaneConfig `json:"control_plane"` + Admin AdminConfig `json:"admin"` + UpstreamHTTP UpstreamHTTPConfig `json:"upstream_http"` + ProviderHealth ProviderHealthConfig `json:"provider_health"` + Providers []ProviderConfig `json:"providers"` + Models []ModelConfig `json:"models"` + Billing BillingConfig `json:"billing"` + Observability ObservabilityConfig `json:"observability"` } type ServerConfig struct { @@ -125,6 +126,18 @@ type UpstreamHTTPConfig struct { ResponseHeaderTimeoutSecs int `json:"response_header_timeout_seconds"` } +type ProviderHealthConfig struct { + ActiveProbesEnabledEnv string `json:"active_probes_enabled_env"` + ActiveProbesEnabled bool `json:"-"` + SharedHistoryEnabledEnv string `json:"shared_history_enabled_env"` + SharedHistoryEnabled bool `json:"-"` + SharedHistoryStream string `json:"shared_history_stream"` + SharedHistoryTTLSeconds int `json:"shared_history_ttl_seconds"` + SharedHistoryMaxEvents int64 `json:"shared_history_max_events"` + ProbeIntervalSeconds int `json:"probe_interval_seconds"` + ProbeTimeoutSeconds int `json:"probe_timeout_seconds"` +} + type ProviderConfig struct { ID string `json:"id"` Slug string `json:"slug"` @@ -139,6 +152,7 @@ type ProviderConfig struct { type ModelConfig struct { ID string `json:"id"` OwnedBy string `json:"owned_by"` + Capabilities []string `json:"capabilities"` InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"` OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"` CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"` @@ -359,6 +373,27 @@ func applyDefaults(cfg *Config) { if cfg.UpstreamHTTP.ResponseHeaderTimeoutSecs == 0 { cfg.UpstreamHTTP.ResponseHeaderTimeoutSecs = 60 } + if cfg.ProviderHealth.ActiveProbesEnabledEnv == "" { + cfg.ProviderHealth.ActiveProbesEnabledEnv = "AIGW_PROVIDER_ACTIVE_PROBES_ENABLED" + } + if cfg.ProviderHealth.SharedHistoryEnabledEnv == "" { + cfg.ProviderHealth.SharedHistoryEnabledEnv = "AIGW_PROVIDER_SHARED_HISTORY_ENABLED" + } + if cfg.ProviderHealth.SharedHistoryStream == "" { + cfg.ProviderHealth.SharedHistoryStream = "aigw:provider-health:events" + } + if cfg.ProviderHealth.SharedHistoryTTLSeconds == 0 { + cfg.ProviderHealth.SharedHistoryTTLSeconds = 900 + } + if cfg.ProviderHealth.SharedHistoryMaxEvents == 0 { + cfg.ProviderHealth.SharedHistoryMaxEvents = 20000 + } + if cfg.ProviderHealth.ProbeIntervalSeconds == 0 { + cfg.ProviderHealth.ProbeIntervalSeconds = 30 + } + if cfg.ProviderHealth.ProbeTimeoutSeconds == 0 { + cfg.ProviderHealth.ProbeTimeoutSeconds = 5 + } if cfg.Observability.UsageBuffer == 0 { cfg.Observability.UsageBuffer = 8192 } @@ -451,6 +486,14 @@ func resolveSecrets(cfg *Config) error { return err } cfg.Server.DeploymentRegion = strings.ToLower(strings.TrimSpace(os.Getenv(cfg.Server.DeploymentRegionEnv))) + cfg.ProviderHealth.ActiveProbesEnabled, err = envBool(cfg.ProviderHealth.ActiveProbesEnabledEnv) + if err != nil { + return err + } + cfg.ProviderHealth.SharedHistoryEnabled, err = envBool(cfg.ProviderHealth.SharedHistoryEnabledEnv) + if err != nil { + return err + } if cfg.ControlPlane.Enabled { cfg.ControlPlane.DatabaseURL = os.Getenv(cfg.ControlPlane.DatabaseURLEnv) cfg.ControlPlane.RedisURL = os.Getenv(cfg.ControlPlane.RedisURLEnv) @@ -571,6 +614,15 @@ func Validate(cfg Config) error { if cfg.Observability.UsageBuffer < 1 { return errors.New("observability.usage_buffer must be positive") } + if cfg.ProviderHealth.ProbeIntervalSeconds < 5 || cfg.ProviderHealth.ProbeIntervalSeconds > 3600 || + cfg.ProviderHealth.ProbeTimeoutSeconds < 1 || cfg.ProviderHealth.ProbeTimeoutSeconds >= cfg.ProviderHealth.ProbeIntervalSeconds { + return errors.New("provider_health probe interval must be 5-3600 seconds and timeout must be shorter than the interval") + } + if cfg.ProviderHealth.SharedHistoryTTLSeconds < 60 || cfg.ProviderHealth.SharedHistoryTTLSeconds > 86400 || + cfg.ProviderHealth.SharedHistoryMaxEvents < 100 || cfg.ProviderHealth.SharedHistoryMaxEvents > 1_000_000 || + strings.TrimSpace(cfg.ProviderHealth.SharedHistoryStream) == "" { + return errors.New("provider_health shared history requires a stream, TTL of 60-86400 seconds, and 100-1000000 events") + } if cfg.Server.SplitListeners { seen := map[string]string{} for name, address := range map[string]string{"public": cfg.Server.PublicAddress, "admin": cfg.Server.AdminAddress, "webhook": cfg.Server.WebhookAddress, "operations": cfg.Server.OperationsAddress} { @@ -716,8 +768,8 @@ func Validate(cfg Config) error { 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.ProtocolOpenAI && provider.WireAPI != "chat_completions" && provider.WireAPI != "responses" && provider.WireAPI != "embeddings" { + return fmt.Errorf("provider %q: wire_api must be chat_completions, responses, or embeddings", provider.ID) } if provider.Protocol == domain.ProtocolAnthropic && provider.WireAPI != "messages" { return fmt.Errorf("provider %q: wire_api must be messages", provider.ID) |
