summaryrefslogtreecommitdiff
path: root/internal/config
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/config/config.go74
-rw-r--r--internal/config/config_test.go83
2 files changed, 146 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)
diff --git a/internal/config/config_test.go b/internal/config/config_test.go
index bb3e814..ff824ca 100644
--- a/internal/config/config_test.go
+++ b/internal/config/config_test.go
@@ -38,6 +38,73 @@ func TestLoadAppliesDefaultsAndResolvesSecrets(t *testing.T) {
}
}
+func TestLoadResolvesActiveProviderProbeConfiguration(t *testing.T) {
+ t.Setenv("TEST_UPSTREAM_KEY", "secret")
+ t.Setenv("TEST_UPSTREAM_URL", "https://example.com/v1")
+ t.Setenv("TEST_ACTIVE_PROBES", "true")
+ path := writeConfig(t, `{
+ "provider_health":{"active_probes_enabled_env":"TEST_ACTIVE_PROBES","probe_interval_seconds":15,"probe_timeout_seconds":2},
+ "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.ProviderHealth.ActiveProbesEnabled || cfg.ProviderHealth.ProbeIntervalSeconds != 15 || cfg.ProviderHealth.ProbeTimeoutSeconds != 2 {
+ t.Fatalf("unexpected provider health config: %+v", cfg.ProviderHealth)
+ }
+}
+
+func TestLoadResolvesSharedProviderHealthConfiguration(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("TEST_SHARED_HEALTH", "true")
+ path := writeConfig(t, `{
+ "control_plane":{"enabled":true},
+ "provider_health":{"shared_history_enabled_env":"TEST_SHARED_HEALTH","shared_history_stream":"test:health","shared_history_ttl_seconds":120,"shared_history_max_events":500}
+}`)
+ cfg, err := Load(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if !cfg.ProviderHealth.SharedHistoryEnabled || cfg.ProviderHealth.SharedHistoryStream != "test:health" ||
+ cfg.ProviderHealth.SharedHistoryTTLSeconds != 120 || cfg.ProviderHealth.SharedHistoryMaxEvents != 500 {
+ t.Fatalf("unexpected shared provider health config: %+v", cfg.ProviderHealth)
+ }
+}
+
+func TestLoadAllowsSharedProviderHealthWithoutRedis(t *testing.T) {
+ t.Setenv("AIGW_DATABASE_URL", "postgres://aigw:aigw@postgres/aigw")
+ t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=")
+ t.Setenv("TEST_SHARED_HEALTH", "true")
+ path := writeConfig(t, `{
+ "control_plane":{"enabled":true},
+ "provider_health":{"shared_history_enabled_env":"TEST_SHARED_HEALTH"}
+}`)
+ cfg, err := Load(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if !cfg.ProviderHealth.SharedHistoryEnabled || cfg.ControlPlane.RedisURL != "" {
+ t.Fatalf("unexpected degraded shared provider health config: %+v", cfg.ProviderHealth)
+ }
+}
+
+func TestLoadRejectsInvalidActiveProviderProbeConfiguration(t *testing.T) {
+ t.Setenv("TEST_UPSTREAM_KEY", "secret")
+ t.Setenv("TEST_UPSTREAM_URL", "https://example.com/v1")
+ path := writeConfig(t, `{
+ "provider_health":{"probe_interval_seconds":5,"probe_timeout_seconds":5},
+ "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 invalid provider probe timeout to be rejected")
+ }
+}
+
func TestLoadValidatesPublicProviderSlugs(t *testing.T) {
t.Setenv("TEST_UPSTREAM_KEY", "secret")
t.Setenv("TEST_UPSTREAM_URL", "https://example.com")
@@ -68,6 +135,22 @@ func TestLoadAcceptsOpenAIResponsesWireAPI(t *testing.T) {
}
}
+func TestLoadAcceptsOpenAIEmbeddingsWireAPI(t *testing.T) {
+ t.Setenv("TEST_UPSTREAM_KEY", "secret")
+ t.Setenv("TEST_UPSTREAM_URL", "https://example.com/v1")
+ path := writeConfig(t, `{
+ "providers": [{"id":"embeddings","protocol":"openai","wire_api":"embeddings","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}],
+ "models": [{"id":"example/embedding","capabilities":["embeddings"],"routes":[{"provider":"embeddings","upstream_model":"embedding-model"}]}]
+}`)
+ cfg, err := Load(path)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if cfg.Providers[0].WireAPI != "embeddings" || len(cfg.Models[0].Capabilities) != 1 || cfg.Models[0].Capabilities[0] != "embeddings" {
+ t.Fatalf("unexpected embeddings config: provider=%+v model=%+v", cfg.Providers[0], cfg.Models[0])
+ }
+}
+
func TestLoadRejectsIncompatibleWireAPI(t *testing.T) {
t.Setenv("TEST_UPSTREAM_KEY", "secret")
t.Setenv("TEST_UPSTREAM_URL", "https://example.com")