package config_test import ( "bytes" "encoding/json" "log/slog" "maps" "net/netip" "slices" "strings" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/config" ) // The settings, by name. const ( listenAddr = "SWWAF_LISTEN_ADDR" upstreamURL = "SWWAF_UPSTREAM_URL" trustedProxies = "SWWAF_TRUSTED_PROXIES" clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT" clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT" upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT" upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT" requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES" responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES" ) // off switches a timeout or a size limit off. const off = "off" // environment is a set of environment variables, for FromEnvironment. type environment map[string]string // lookupEnv reads one of the variables, as os.LookupEnv does. func (e environment) lookupEnv(name string) (string, bool) { value, ok := e[name] return value, ok } // fromEnvironment reads the settings from env, which must be valid. func fromEnvironment(t *testing.T, env environment) *config.Config { t.Helper() cfg, err := config.FromEnvironment(env.lookupEnv) if err != nil { t.Fatalf("settings %v: %v", env, err) } return cfg } func TestDefaults(t *testing.T) { t.Parallel() cfg := fromEnvironment(t, environment{}) wantSettings(t, cfg, config.Config{ ListenAddr: ":8080", ClientRequestTimeout: time.Minute, ClientResponseTimeout: 30 * time.Minute, UpstreamRequestTimeout: time.Minute, UpstreamResponseTimeout: 30 * time.Minute, RequestMaxBytes: 100 << 20, ResponseMaxBytes: 5 << 30, }) if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" { t.Errorf("%s is %s", upstreamURL, cfg.UpstreamURL) } wantNetblocks(t, cfg.TrustedProxies, "10.0.0.0/8", "172.16.0.0/12", "192.168.0.0/16") } func TestValuesAsSet(t *testing.T) { t.Parallel() cfg := fromEnvironment(t, environment{ listenAddr: "127.0.0.1:9000", upstreamURL: "https://app.internal:8443/", trustedProxies: " 192.0.2.1, 10.1.2.3/8 ,2001:db8::/32", clientRequestTimeout: "90s", clientResponseTimeout: "7d", upstreamRequestTimeout: "1h30m", upstreamResponseTimeout: off, requestMaxBytes: "512K", responseMaxBytes: "1234", }) wantSettings(t, cfg, config.Config{ ListenAddr: "127.0.0.1:9000", ClientRequestTimeout: 90 * time.Second, ClientResponseTimeout: 7 * 24 * time.Hour, UpstreamRequestTimeout: 90 * time.Minute, UpstreamResponseTimeout: 0, RequestMaxBytes: 512 << 10, ResponseMaxBytes: 1234, }) if cfg.UpstreamURL.String() != "https://app.internal:8443/" { t.Errorf("%s is %s", upstreamURL, cfg.UpstreamURL) } wantNetblocks(t, cfg.TrustedProxies, "192.0.2.1/32", "10.0.0.0/8", "2001:db8::/32") } func TestSizesAndOff(t *testing.T) { t.Parallel() cfg := fromEnvironment(t, environment{ requestMaxBytes: "3G", responseMaxBytes: off, clientRequestTimeout: off, }) if cfg.RequestMaxBytes != 3<<30 || cfg.ResponseMaxBytes != 0 || cfg.ClientRequestTimeout != 0 { t.Errorf("3G, off and off read as %d, %d and %s", cfg.RequestMaxBytes, cfg.ResponseMaxBytes, cfg.ClientRequestTimeout) } } func TestTrustedProxiesSetButEmptyTrustNothing(t *testing.T) { t.Parallel() cfg := fromEnvironment(t, environment{trustedProxies: ""}) if len(cfg.TrustedProxies) != 0 { t.Errorf("trusted proxies %v, want none", cfg.TrustedProxies) } } func TestInvalidValueStopsTheStart(t *testing.T) { t.Parallel() for _, tc := range []struct{ name, value string }{ {listenAddr, "8080"}, {listenAddr, ":http"}, {listenAddr, ":65536"}, {upstreamURL, "127.0.0.1:8081"}, {upstreamURL, "ftp://127.0.0.1:8081"}, {upstreamURL, "http://"}, {upstreamURL, "http://127.0.0.1:8081/app"}, {upstreamURL, "http://127.0.0.1:8081/?a=1"}, {upstreamURL, "http://user:secret@127.0.0.1:8081"}, {trustedProxies, "10.0.0.0/33"}, {trustedProxies, "traefik"}, {trustedProxies, "10.0.0.0/8,,192.168.0.0/16"}, {trustedProxies, "fe80::1%eth0"}, {clientRequestTimeout, "60"}, {clientRequestTimeout, ""}, {clientResponseTimeout, "1y"}, {upstreamRequestTimeout, "-1s"}, {upstreamResponseTimeout, "0s"}, {upstreamResponseTimeout, "1.5d"}, {requestMaxBytes, "100MB"}, {requestMaxBytes, "100m"}, {requestMaxBytes, "1.5M"}, {responseMaxBytes, "0"}, {responseMaxBytes, "-5"}, {responseMaxBytes, "99999999999G"}, } { t.Run(tc.name+"="+tc.value, func(t *testing.T) { t.Parallel() _, err := config.FromEnvironment(environment{tc.name: tc.value}.lookupEnv) if err == nil { t.Fatalf("%s=%q was accepted", tc.name, tc.value) } if !strings.HasPrefix(err.Error(), tc.name+": ") { t.Errorf("error %q does not name %s", err, tc.name) } }) } } func TestLogsEachSettingWithItsValue(t *testing.T) { t.Parallel() cfg := fromEnvironment(t, environment{clientRequestTimeout: "45s"}) var out bytes.Buffer slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg) var line struct { Settings map[string]string `json:"settings"` } err := json.Unmarshal(out.Bytes(), &line) if err != nil { t.Fatalf("decode %s: %v", out.Bytes(), err) } want := map[string]string{ listenAddr: ":8080", upstreamURL: "http://127.0.0.1:8081", trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16", clientRequestTimeout: "45s", clientResponseTimeout: "30m", upstreamRequestTimeout: "60s", upstreamResponseTimeout: "30m", requestMaxBytes: "100M", responseMaxBytes: "5G", } if !maps.Equal(line.Settings, want) { t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want) } } // wantSettings checks the settings that are plain values. func wantSettings(t *testing.T, got *config.Config, want config.Config) { t.Helper() if got.ListenAddr != want.ListenAddr || got.ClientRequestTimeout != want.ClientRequestTimeout || got.ClientResponseTimeout != want.ClientResponseTimeout || got.UpstreamRequestTimeout != want.UpstreamRequestTimeout || got.UpstreamResponseTimeout != want.UpstreamResponseTimeout || got.RequestMaxBytes != want.RequestMaxBytes || got.ResponseMaxBytes != want.ResponseMaxBytes { t.Errorf("settings\n%+v\nwant\n%+v", got, want) } } // wantNetblocks checks a list of netblocks. func wantNetblocks(t *testing.T, got []netip.Prefix, want ...string) { t.Helper() gotText := make([]string, 0, len(got)) for _, netblock := range got { gotText = append(gotText, netblock.String()) } if !slices.Equal(gotText, want) { t.Errorf("netblocks %v, want %v", gotText, want) } }