package config import ( "fmt" "log/slog" "net/netip" "os" "strconv" "strings" "time" ) const defaultTrustedProxyCIDRs = "127.0.0.1/32,::1/128" type Config struct { Addr string TrustProxyHeaders bool TrustedProxyCIDRs []netip.Prefix ReadHeaderTimeout time.Duration ReadTimeout time.Duration WriteTimeout time.Duration IdleTimeout time.Duration ShutdownTimeout time.Duration MaxHeaderBytes int LogLevel slog.Level LogFormat string } func Load() (Config, error) { trustProxyHeaders, err := envBool("TRUST_PROXY_HEADERS", false) if err != nil { return Config{}, err } trustedProxyCIDRs, err := envPrefixes("TRUSTED_PROXY_CIDRS", defaultTrustedProxyCIDRs) if err != nil { return Config{}, err } readHeaderTimeout, err := envDuration("READ_HEADER_TIMEOUT", 5*time.Second) if err != nil { return Config{}, err } readTimeout, err := envDuration("READ_TIMEOUT", 10*time.Second) if err != nil { return Config{}, err } writeTimeout, err := envDuration("WRITE_TIMEOUT", 10*time.Second) if err != nil { return Config{}, err } idleTimeout, err := envDuration("IDLE_TIMEOUT", 60*time.Second) if err != nil { return Config{}, err } shutdownTimeout, err := envDuration("SHUTDOWN_TIMEOUT", 5*time.Second) if err != nil { return Config{}, err } maxHeaderBytes, err := envInt("MAX_HEADER_BYTES", 1<<20) if err != nil { return Config{}, err } logLevel, err := envLogLevel("LOG_LEVEL", slog.LevelInfo) if err != nil { return Config{}, err } logFormat, err := envLogFormat("LOG_FORMAT", "text") if err != nil { return Config{}, err } return Config{ Addr: envString("ADDR", ":8080"), TrustProxyHeaders: trustProxyHeaders, TrustedProxyCIDRs: trustedProxyCIDRs, ReadHeaderTimeout: readHeaderTimeout, ReadTimeout: readTimeout, WriteTimeout: writeTimeout, IdleTimeout: idleTimeout, ShutdownTimeout: shutdownTimeout, MaxHeaderBytes: maxHeaderBytes, LogLevel: logLevel, LogFormat: logFormat, }, nil } func envString(name, fallback string) string { value := strings.TrimSpace(os.Getenv(name)) if value == "" { return fallback } return value } func envBool(name string, fallback bool) (bool, error) { value := strings.TrimSpace(os.Getenv(name)) if value == "" { return fallback, nil } parsed, err := strconv.ParseBool(value) if err != nil { return false, fmt.Errorf("%s must be a boolean value: %w", name, err) } return parsed, nil } func envDuration(name string, fallback time.Duration) (time.Duration, error) { value := strings.TrimSpace(os.Getenv(name)) if value == "" { return fallback, nil } parsed, err := time.ParseDuration(value) if err != nil { return 0, fmt.Errorf("%s must be a duration like 5s or 1m: %w", name, err) } if parsed <= 0 { return 0, fmt.Errorf("%s must be greater than zero", name) } return parsed, nil } func envInt(name string, fallback int) (int, error) { value := strings.TrimSpace(os.Getenv(name)) if value == "" { return fallback, nil } parsed, err := strconv.Atoi(value) if err != nil { return 0, fmt.Errorf("%s must be an integer: %w", name, err) } if parsed <= 0 { return 0, fmt.Errorf("%s must be greater than zero", name) } return parsed, nil } func envPrefixes(name, fallback string) ([]netip.Prefix, error) { value := strings.TrimSpace(os.Getenv(name)) if value == "" { value = fallback } parts := strings.Split(value, ",") prefixes := make([]netip.Prefix, 0, len(parts)) for _, part := range parts { part = strings.TrimSpace(part) if part == "" { continue } if prefix, err := netip.ParsePrefix(part); err == nil { prefixes = append(prefixes, prefix.Masked()) continue } addr, err := netip.ParseAddr(part) if err != nil { return nil, fmt.Errorf("%s contains invalid CIDR or IP %q: %w", name, part, err) } prefixes = append(prefixes, netip.PrefixFrom(addr, addr.BitLen())) } if len(prefixes) == 0 { return nil, fmt.Errorf("%s must contain at least one CIDR or IP", name) } return prefixes, nil } func envLogLevel(name string, fallback slog.Level) (slog.Level, error) { value := strings.TrimSpace(os.Getenv(name)) if value == "" { return fallback, nil } var level slog.Level if err := level.UnmarshalText([]byte(strings.ToUpper(value))); err != nil { return 0, fmt.Errorf("%s must be debug, info, warn, or error: %w", name, err) } return level, nil } func envLogFormat(name, fallback string) (string, error) { value := strings.ToLower(strings.TrimSpace(os.Getenv(name))) if value == "" { return fallback, nil } switch value { case "text", "json": return value, nil default: return "", fmt.Errorf("%s must be text or json", name) } }