first commit
This commit is contained in:
215
internal/config/config.go
Normal file
215
internal/config/config.go
Normal file
@@ -0,0 +1,215 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
130
internal/config/config_test.go
Normal file
130
internal/config/config_test.go
Normal file
@@ -0,0 +1,130 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestLoadUsesProductionDefaults(t *testing.T) {
|
||||
clearConfigEnv(t)
|
||||
|
||||
cfg, err := Load()
|
||||
if err != nil {
|
||||
t.Fatalf("load config: %v", err)
|
||||
}
|
||||
|
||||
if cfg.Addr != ":8080" {
|
||||
t.Fatalf("got addr %q, want %q", cfg.Addr, ":8080")
|
||||
}
|
||||
|
||||
if cfg.TrustProxyHeaders {
|
||||
t.Fatal("proxy headers should not be trusted by default")
|
||||
}
|
||||
|
||||
wantProxy := netip.MustParsePrefix("127.0.0.1/32")
|
||||
if len(cfg.TrustedProxyCIDRs) == 0 || cfg.TrustedProxyCIDRs[0] != wantProxy {
|
||||
t.Fatalf("got trusted proxies %v, want first %v", cfg.TrustedProxyCIDRs, wantProxy)
|
||||
}
|
||||
|
||||
if cfg.ReadHeaderTimeout != 5*time.Second {
|
||||
t.Fatalf("got read header timeout %s", cfg.ReadHeaderTimeout)
|
||||
}
|
||||
|
||||
if cfg.LogLevel != slog.LevelInfo {
|
||||
t.Fatalf("got log level %s, want info", cfg.LogLevel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadParsesEnvironment(t *testing.T) {
|
||||
clearConfigEnv(t)
|
||||
t.Setenv("ADDR", "127.0.0.1:9090")
|
||||
t.Setenv("TRUST_PROXY_HEADERS", "true")
|
||||
t.Setenv("TRUSTED_PROXY_CIDRS", "10.0.0.0/8,192.0.2.10")
|
||||
t.Setenv("READ_HEADER_TIMEOUT", "3s")
|
||||
t.Setenv("READ_TIMEOUT", "4s")
|
||||
t.Setenv("WRITE_TIMEOUT", "5s")
|
||||
t.Setenv("IDLE_TIMEOUT", "30s")
|
||||
t.Setenv("SHUTDOWN_TIMEOUT", "7s")
|
||||
t.Setenv("MAX_HEADER_BYTES", "8192")
|
||||
t.Setenv("LOG_LEVEL", "debug")
|
||||
t.Setenv("LOG_FORMAT", "json")
|
||||
|
||||
cfg, err := Load()
|
||||
if err != nil {
|
||||
t.Fatalf("load config: %v", err)
|
||||
}
|
||||
|
||||
if cfg.Addr != "127.0.0.1:9090" {
|
||||
t.Fatalf("got addr %q", cfg.Addr)
|
||||
}
|
||||
|
||||
if !cfg.TrustProxyHeaders {
|
||||
t.Fatal("expected proxy headers to be trusted")
|
||||
}
|
||||
|
||||
if got, want := cfg.TrustedProxyCIDRs[1], netip.MustParsePrefix("192.0.2.10/32"); got != want {
|
||||
t.Fatalf("got proxy %v, want %v", got, want)
|
||||
}
|
||||
|
||||
if cfg.ReadTimeout != 4*time.Second || cfg.WriteTimeout != 5*time.Second {
|
||||
t.Fatalf("got read/write timeouts %s/%s", cfg.ReadTimeout, cfg.WriteTimeout)
|
||||
}
|
||||
|
||||
if cfg.MaxHeaderBytes != 8192 {
|
||||
t.Fatalf("got max header bytes %d", cfg.MaxHeaderBytes)
|
||||
}
|
||||
|
||||
if cfg.LogLevel != slog.LevelDebug || cfg.LogFormat != "json" {
|
||||
t.Fatalf("got logging %s/%s", cfg.LogLevel, cfg.LogFormat)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadRejectsInvalidEnvironment(t *testing.T) {
|
||||
tests := map[string]string{
|
||||
"TRUST_PROXY_HEADERS": "yes-please",
|
||||
"TRUSTED_PROXY_CIDRS": "not-a-cidr",
|
||||
"TRUSTED_PROXY_CIDRS_EMPTY": ",",
|
||||
"READ_TIMEOUT": "forever",
|
||||
"MAX_HEADER_BYTES": "0",
|
||||
"LOG_LEVEL": "verbose",
|
||||
"LOG_FORMAT": "xml",
|
||||
}
|
||||
|
||||
for name, value := range tests {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
clearConfigEnv(t)
|
||||
envName := name
|
||||
if name == "TRUSTED_PROXY_CIDRS_EMPTY" {
|
||||
envName = "TRUSTED_PROXY_CIDRS"
|
||||
}
|
||||
|
||||
t.Setenv(envName, value)
|
||||
|
||||
if _, err := Load(); err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func clearConfigEnv(t *testing.T) {
|
||||
t.Helper()
|
||||
|
||||
for _, name := range []string{
|
||||
"ADDR",
|
||||
"TRUST_PROXY_HEADERS",
|
||||
"TRUSTED_PROXY_CIDRS",
|
||||
"READ_HEADER_TIMEOUT",
|
||||
"READ_TIMEOUT",
|
||||
"WRITE_TIMEOUT",
|
||||
"IDLE_TIMEOUT",
|
||||
"SHUTDOWN_TIMEOUT",
|
||||
"MAX_HEADER_BYTES",
|
||||
"LOG_LEVEL",
|
||||
"LOG_FORMAT",
|
||||
} {
|
||||
t.Setenv(name, "")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user