openmeter / app /config /server.go
Leon4gr45's picture
Upload folder using huggingface_hub (part 4)
1f10f31 verified
Raw
History Blame Contribute Delete
5.72 kB
package config
import (
"errors"
"fmt"
"net/http"
"net/netip"
"time"
"github.com/samber/lo"
"github.com/spf13/viper"
"github.com/openmeterio/openmeter/pkg/models"
)
// ServerConfig holds HTTP server timeout configuration.
type ServerConfig struct {
ReadHeaderTimeout time.Duration
ReadTimeout time.Duration
WriteTimeout time.Duration
IdleTimeout time.Duration
ResponseValidation ResponseValidationConfig
ClientIPMiddleware ClientIPMiddlewareConfig
}
// ResponseValidationConfig controls optional post-response OpenAPI validation on the v3 API.
type ResponseValidationConfig struct {
Mode ResponseValidationMode
}
type ResponseValidationMode string
const (
// ResponseValidationModeOff disables response validation. This is the default.
ResponseValidationModeOff ResponseValidationMode = "off"
// ResponseValidationModeUnstable validates only routes marked x-unstable: true in the spec.
ResponseValidationModeUnstable ResponseValidationMode = "unstable"
// ResponseValidationModeAll validates every route in the v3 spec.
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)
// ClientIPMiddlewareConfig configures the middleware that extracts the client IP address from the HTTP request.
// See: https://adam-p.ca/blog/2022/03/x-forwarded-for/
type ClientIPMiddlewareConfig struct {
Source ClientIPSource
// Header defines the header name in the HTTP request containing the real client IP address.
// Set this only if ClientIPSourceHeader is used as Source.
// Only use headers your proxy unconditionally overwrites on every request,
// e.g. "X-Real-IP" (Nginx ngx_http_realip_module), "CF-Connecting-IP" (Cloudflare), or "X-Client-IP" (Apache mod_remoteip).
// Pass-through headers like "True-Client-IP", "X-Azure-ClientIP", or "Fastly-Client-IP" are client-spoofable
// unless your edge strips the inbound value.
Header string
// TrustedIPPrefixes lists IP prefixes for trusted proxies.
// Set this only if the ClientIPSourceXFF is used as Source.
TrustedIPPrefixes []string
// TrustedProxies defines the number of trusted proxies.
// Set this only if the ClientIPSourceXFF is used as Source.
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")
}
// chi's ClientIPFromHeader takes the LAST header value, which for the append-style
// X-Forwarded-For header is the nearest proxy hop, not the client.
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 {
// Use the same parser as chi's ClientIPFromXFF (netip.MustParsePrefix), which is
// stricter than net.ParseCIDR; a mismatch would panic at middleware construction.
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
}
// chi's ClientIPFromXFFTrustedProxies panics if the count is < 1.
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...)
}
// ConfigureServer sets defaults for HTTP server timeouts.
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)
}