| package config |
|
|
| import ( |
| "errors" |
| "fmt" |
| "net/http" |
| "net/netip" |
| "time" |
|
|
| "github.com/samber/lo" |
| "github.com/spf13/viper" |
|
|
| "github.com/openmeterio/openmeter/pkg/models" |
| ) |
|
|
| |
| type ServerConfig struct { |
| ReadHeaderTimeout time.Duration |
| ReadTimeout time.Duration |
| WriteTimeout time.Duration |
| IdleTimeout time.Duration |
|
|
| ResponseValidation ResponseValidationConfig |
|
|
| ClientIPMiddleware ClientIPMiddlewareConfig |
| } |
|
|
| |
| type ResponseValidationConfig struct { |
| Mode ResponseValidationMode |
| } |
|
|
| type ResponseValidationMode string |
|
|
| const ( |
| |
| ResponseValidationModeOff ResponseValidationMode = "off" |
| |
| ResponseValidationModeUnstable ResponseValidationMode = "unstable" |
| |
| ResponseValidationModeAll ResponseValidationMode = "all" |
| ) |
|
|
| func (m ResponseValidationMode) Enabled() bool { |
| return m != "" && m != ResponseValidationModeOff |
| } |
|
|
| func (m ResponseValidationMode) Validate() error { |
| switch m { |
| case "", ResponseValidationModeOff, ResponseValidationModeUnstable, ResponseValidationModeAll: |
| return nil |
| default: |
| return errors.New("invalid response validation mode (allowed: off, unstable, all)") |
| } |
| } |
|
|
| type ClientIPSource string |
|
|
| const ( |
| ClientIPSourceRemoteAddr ClientIPSource = "remote-address" |
| ClientIPSourceHeader ClientIPSource = "header" |
| ClientIPSourceXFF ClientIPSource = "x-forwarded-for" |
| ) |
|
|
| var _ models.Validator = (*ClientIPMiddlewareConfig)(nil) |
|
|
| |
| |
| type ClientIPMiddlewareConfig struct { |
| Source ClientIPSource |
|
|
| |
| |
| |
| |
| |
| |
| Header string |
|
|
| |
| |
| TrustedIPPrefixes []string |
|
|
| |
| |
| TrustedProxies int |
| } |
|
|
| func (c ClientIPMiddlewareConfig) Validate() error { |
| switch c.Source { |
| case ClientIPSourceRemoteAddr: |
| return nil |
| case ClientIPSourceHeader: |
| if c.Header == "" { |
| return errors.New("missing client IP header") |
| } |
|
|
| |
| |
| if http.CanonicalHeaderKey(c.Header) == "X-Forwarded-For" { |
| return fmt.Errorf("X-Forwarded-For cannot be used as client IP header, use the %s source instead", ClientIPSourceXFF) |
| } |
|
|
| return nil |
| case ClientIPSourceXFF: |
| if len(c.TrustedIPPrefixes) > 0 { |
| |
| |
| invalidPrefixes := lo.Filter(c.TrustedIPPrefixes, func(prefix string, _ int) bool { |
| _, err := netip.ParsePrefix(prefix) |
|
|
| return err != nil |
| }) |
|
|
| if len(invalidPrefixes) > 0 { |
| return fmt.Errorf("invalid trusted IP prefixes: %+v", invalidPrefixes) |
| } |
|
|
| return nil |
| } |
|
|
| |
| if c.TrustedProxies < 1 { |
| return fmt.Errorf("either trusted IP prefixes or a positive number of trusted proxies must be set if real client IP source is set to %s", ClientIPSourceXFF) |
| } |
|
|
| return nil |
| default: |
| return fmt.Errorf("invalid client IP source: %s", c.Source) |
| } |
| } |
|
|
| func (c ServerConfig) Validate() error { |
| var errs []error |
|
|
| if c.ReadHeaderTimeout < 0 { |
| errs = append(errs, errors.New("readHeaderTimeout must be non-negative")) |
| } |
|
|
| if c.ReadTimeout < 0 { |
| errs = append(errs, errors.New("readTimeout must be non-negative")) |
| } |
|
|
| if c.WriteTimeout < 0 { |
| errs = append(errs, errors.New("writeTimeout must be non-negative")) |
| } |
|
|
| if c.IdleTimeout < 0 { |
| errs = append(errs, errors.New("idleTimeout must be non-negative")) |
| } |
|
|
| if err := c.ResponseValidation.Mode.Validate(); err != nil { |
| errs = append(errs, err) |
| } |
|
|
| if err := c.ClientIPMiddleware.Validate(); err != nil { |
| errs = append(errs, err) |
| } |
|
|
| return errors.Join(errs...) |
| } |
|
|
| |
| func ConfigureServer(v *viper.Viper, prefixes ...string) { |
| prefixer := NewViperKeyPrefixer(prefixes...) |
|
|
| v.SetDefault(prefixer("readHeaderTimeout"), 10*time.Second) |
| v.SetDefault(prefixer("readTimeout"), 60*time.Second) |
| v.SetDefault(prefixer("writeTimeout"), 90*time.Second) |
| v.SetDefault(prefixer("idleTimeout"), 120*time.Second) |
|
|
| v.SetDefault(prefixer("responseValidation.mode"), string(ResponseValidationModeOff)) |
|
|
| v.SetDefault(prefixer("clientIPMiddleware.source"), ClientIPSourceRemoteAddr) |
| v.SetDefault(prefixer("clientIPMiddleware.header"), "") |
| v.SetDefault(prefixer("clientIPMiddleware.trustedIPPrefixes"), nil) |
| v.SetDefault(prefixer("clientIPMiddleware.trustedProxies"), 0) |
| } |
|
|