Files
smallwebwaf/internal/config/config_test.go
T
clawbot 808e69f442
check / check (push) Waiting to run
Log the rest of the request log's fields (closes #79)
Each request log line now has the fields "Request log" in SPEC.md lists
whose features are built: instance (SWWAF_INSTANCE_NAME), scheme,
request_id (a trusted proxy's X-Request-ID or a new one, sent on to the
app), forwarded_for, client_group, content_type, content_length, the
headers SWWAF_LOG_REQUEST_HEADERS names, has_authorization, has_cookie,
websocket, response_content_type, cache_control, location, counts and
the timings. Authorization, Cookie and Set-Cookie values are never
logged. An entry of SWWAF_LOG_REQUEST_HEADERS that is not a header name,
or is Host or Transfer-Encoding, stops the start.

Deviation: counts has request totals only.
Deviation: SWWAF_INSTANCE_NAME is on request lines only.

Model: opus-5-5
2026-10-06 17:26:21 +02:00

600 lines
19 KiB
Go

package config_test
import (
"bytes"
"encoding/json"
"log/slog"
"maps"
"net/netip"
"os"
"slices"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/config"
)
// The settings, by name.
const (
listenAddr = "SWWAF_LISTEN_ADDR"
upstreamURL = "SWWAF_UPSTREAM_URL"
mode = "SWWAF_MODE"
trustedProxies = "SWWAF_TRUSTED_PROXIES"
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
clientHeaderMaxBytes = "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES"
clientIdleTimeout = "SWWAF_CLIENT_IDLE_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"
allowNets = "SWWAF_ALLOW_NETS"
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
denyNets = "SWWAF_DENY_NETS"
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
rateLimitPerHour = "SWWAF_RATE_LIMIT_PER_HOUR"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
deniedCountries = "SWWAF_DENIED_COUNTRIES"
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
banResponse = "SWWAF_BAN_RESPONSE"
limitBanDuration = "SWWAF_LIMIT_BAN_DURATION"
limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW"
maxBanDuration = "SWWAF_MAX_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS"
banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX"
stateDir = "SWWAF_STATE_DIR"
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
metricsTopN = "SWWAF_METRICS_TOP_N"
instanceName = "SWWAF_INSTANCE_NAME"
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
)
// defaultLogRequestHeaders is the default of SWWAF_LOG_REQUEST_HEADERS.
const defaultLogRequestHeaders = "accept,accept-language,accept-encoding," +
"content-type,origin,range"
// token is a token of 32 characters, the shortest allowed.
const token = "0123456789abcdef0123456789abcdef"
// off switches a timeout, a size limit or a rate 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",
Observe: false,
ClientRequestTimeout: time.Minute,
ClientRequestHeaderMaxBytes: 32 << 10,
ClientIdleTimeout: 2 * time.Minute,
ClientResponseTimeout: 30 * time.Minute,
UpstreamRequestTimeout: time.Minute,
UpstreamResponseTimeout: 30 * time.Minute,
RequestMaxBytes: 100 << 20,
ResponseMaxBytes: 5 << 30,
RateLimitPerMinute: 1000,
RateLimitPerHour: 10000,
RateLimitPerDay: 50000,
BanResponse: 403,
LimitBanDuration: time.Hour,
LimitBanRepeatWindow: 24 * time.Hour,
MaxBanDuration: 7 * 24 * time.Hour,
MaxBans: 5000,
BanScopeV4Prefix: 32,
StateDir: "/var/lib/smallwebwaf",
StateWriteDelay: 10 * time.Second,
StateCounterInterval: 15 * time.Minute,
MetricsToken: "",
MetricsTopN: 50,
})
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")
wantNetblocks(t, cfg.AllowNets)
wantNetblocks(t, cfg.RateLimitExemptNets)
wantNetblocks(t, cfg.DenyNets)
wantCountries(t, deniedCountries, cfg.DeniedCountries)
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries)
hostname, err := os.Hostname()
if err != nil || hostname == "" || cfg.InstanceName != hostname {
t.Errorf("%s is %q, want the host's name %q (%v)", instanceName,
cfg.InstanceName, hostname, err)
}
wantHeaders := strings.Split(defaultLogRequestHeaders, ",")
if !slices.Equal(cfg.LogRequestHeaders, wantHeaders) {
t.Errorf("%s gave %v, want %v", logRequestHeaders, cfg.LogRequestHeaders,
wantHeaders)
}
}
func TestValuesAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
listenAddr: "127.0.0.1:9000",
upstreamURL: "https://app.internal:8443/",
mode: "observe",
trustedProxies: " 192.0.2.1, 10.1.2.3/8 ,2001:db8::/32",
clientRequestTimeout: "90s",
clientHeaderMaxBytes: "8K",
clientIdleTimeout: "5m",
clientResponseTimeout: "7d",
upstreamRequestTimeout: "1h30m",
upstreamResponseTimeout: off,
requestMaxBytes: "512K",
responseMaxBytes: "1234",
allowNets: "192.0.2.7",
rateLimitExemptNets: "2001:db8::/48, 10.9.8.7",
denyNets: "198.51.100.0/24",
rateLimitPerMinute: "60",
rateLimitPerHour: "600",
rateLimitPerDay: "6000",
deniedCountries: "cn, RU,kp,Xk",
allowedCountries: "de",
banResponse: "429",
limitBanDuration: "15m",
limitBanRepeatWindow: "2d",
maxBanDuration: "30d",
maxBans: "100",
banScopeV4Prefix: "24",
stateDir: "/srv/waf-state",
stateWriteDelay: "500ms",
stateCounterInterval: "1h",
metricsToken: token,
metricsTopN: "10",
})
wantSettings(t, cfg, config.Config{
ListenAddr: "127.0.0.1:9000",
Observe: true,
ClientRequestTimeout: 90 * time.Second,
ClientRequestHeaderMaxBytes: 8 << 10,
ClientIdleTimeout: 5 * time.Minute,
ClientResponseTimeout: 7 * 24 * time.Hour,
UpstreamRequestTimeout: 90 * time.Minute,
UpstreamResponseTimeout: 0,
RequestMaxBytes: 512 << 10,
ResponseMaxBytes: 1234,
RateLimitPerMinute: 60,
RateLimitPerHour: 600,
RateLimitPerDay: 6000,
BanResponse: 429,
LimitBanDuration: 15 * time.Minute,
LimitBanRepeatWindow: 48 * time.Hour,
MaxBanDuration: 30 * 24 * time.Hour,
MaxBans: 100,
BanScopeV4Prefix: 24,
StateDir: "/srv/waf-state",
StateWriteDelay: 500 * time.Millisecond,
StateCounterInterval: time.Hour,
MetricsToken: token,
MetricsTopN: 10,
})
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")
wantNetblocks(t, cfg.AllowNets, "192.0.2.7/32")
wantNetblocks(t, cfg.RateLimitExemptNets, "2001:db8::/48", "10.9.8.7/32")
wantNetblocks(t, cfg.DenyNets, "198.51.100.0/24")
wantCountries(t, deniedCountries, cfg.DeniedCountries, "CN", "RU", "KP", "XK")
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE")
}
func TestInstanceNameAndLoggedHeadersAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
instanceName: "fsn1app1/gitea",
logRequestHeaders: " Accept , X-Custom",
})
if cfg.InstanceName != "fsn1app1/gitea" ||
!slices.Equal(cfg.LogRequestHeaders, []string{"accept", "x-custom"}) {
t.Errorf("%s is %q and %s gives %v", instanceName, cfg.InstanceName,
logRequestHeaders, cfg.LogRequestHeaders)
}
}
func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{
deniedCountries: "cn,ru",
allowedCountries: "de,RU",
}.lookupEnv)
if err == nil {
t.Fatal("ru on both country lists was accepted")
}
if !strings.HasPrefix(err.Error(), allowedCountries+": ") ||
!strings.Contains(err.Error(), `"RU"`) {
t.Errorf("error %q does not name %s and RU", err, allowedCountries)
}
}
func TestSizesAndOff(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
requestMaxBytes: "3G",
responseMaxBytes: off,
clientRequestTimeout: off,
clientIdleTimeout: off,
})
if cfg.RequestMaxBytes != 3<<30 || cfg.ResponseMaxBytes != 0 ||
cfg.ClientRequestTimeout != 0 || cfg.ClientIdleTimeout != 0 {
t.Errorf("3G, off, off and off read as %d, %d, %s and %s",
cfg.RequestMaxBytes, cfg.ResponseMaxBytes, cfg.ClientRequestTimeout,
cfg.ClientIdleTimeout)
}
}
func TestRequestHeaderMaxBytesJustOver4K(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{clientHeaderMaxBytes: "4097"})
if cfg.ClientRequestHeaderMaxBytes != 4097 {
t.Errorf("4097 read as %d", cfg.ClientRequestHeaderMaxBytes)
}
}
func TestRequestHeaderMaxBytesRefusalNeverOffersOff(t *testing.T) {
t.Parallel()
for _, value := range []string{"32KB", "0", "4K", off} {
t.Run(value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(
environment{clientHeaderMaxBytes: value}.lookupEnv)
want := clientHeaderMaxBytes + `: "` + value +
`" is not a size of more than 4K, such as 32K`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
})
}
}
func TestRateLimitsOff(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
rateLimitPerMinute: off,
rateLimitPerHour: off,
rateLimitPerDay: off,
})
if cfg.RateLimitPerMinute != 0 || cfg.RateLimitPerHour != 0 ||
cfg.RateLimitPerDay != 0 {
t.Errorf("off read as %d, %d and %d",
cfg.RateLimitPerMinute, cfg.RateLimitPerHour, cfg.RateLimitPerDay)
}
}
func TestBanResponseCloseIsZero(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{banResponse: "close"})
if cfg.BanResponse != 0 {
t.Errorf("close read as %d, want 0", cfg.BanResponse)
}
}
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://:8081"},
{upstreamURL, "http://127.0.0.1:0"},
{upstreamURL, "http://127.0.0.1:99999"},
{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"},
{mode, "Observe"}, {mode, "block"}, {mode, ""},
{trustedProxies, "10.0.0.0/33"},
{trustedProxies, "traefik"},
{trustedProxies, "10.0.0.0/8,,192.168.0.0/16"},
{trustedProxies, "fe80::1%eth0"},
{allowNets, "192.0.2.0/24,monitoring"},
{rateLimitExemptNets, "2001:db8::/129"},
{denyNets, "198.51.100.0/24,"},
{clientRequestTimeout, "60"}, {clientRequestTimeout, ""},
{clientIdleTimeout, "0s"},
{clientIdleTimeout, "2 minutes"},
{clientResponseTimeout, "1y"},
{upstreamRequestTimeout, "-1s"},
{upstreamResponseTimeout, "0s"},
{upstreamResponseTimeout, "1.5d"},
{requestMaxBytes, "100MB"},
{requestMaxBytes, "100m"},
{requestMaxBytes, "1.5M"},
{responseMaxBytes, "0"},
{responseMaxBytes, "-5"},
{responseMaxBytes, "99999999999G"},
{rateLimitPerMinute, ""},
{rateLimitPerMinute, "1K"},
{rateLimitPerHour, "0"},
{rateLimitPerHour, "1.5"},
{rateLimitPerDay, "-1"},
{rateLimitPerDay, "lots"},
{deniedCountries, "nk"},
{deniedCountries, "kp,,ir"},
{deniedCountries, "prk"},
{deniedCountries, "408"},
{deniedCountries, "k"},
{deniedCountries, "eu"},
{deniedCountries, "un"},
{deniedCountries, "su"},
{allowedCountries, "ac"},
{allowedCountries, "uk"},
{allowedCountries, "zz"},
{allowedCountries, "de,germany"},
{banResponse, "404"}, {banResponse, "drop"}, {banResponse, ""},
{limitBanDuration, off}, {limitBanDuration, "0s"}, {limitBanDuration, "1"},
{limitBanRepeatWindow, off}, {limitBanRepeatWindow, "-1h"},
{maxBanDuration, off}, {maxBanDuration, "1w"},
{maxBans, off}, {maxBans, "0"}, {maxBans, "5K"},
{banScopeV4Prefix, "33"}, {banScopeV4Prefix, "-1"}, {banScopeV4Prefix, "/24"},
{stateDir, ""}, {stateDir, "state"}, {stateDir, "./var/lib/smallwebwaf"},
{stateWriteDelay, off}, {stateWriteDelay, "0s"},
{stateCounterInterval, off}, {stateCounterInterval, "15"},
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
{logRequestHeaders, "accept,,origin"}, {logRequestHeaders, "accept;origin"},
{logRequestHeaders, "accept language"}, {logRequestHeaders, "x-foo:"},
{logRequestHeaders, "host"}, {logRequestHeaders, "accept,Host"},
{logRequestHeaders, "transfer-encoding"}, {logRequestHeaders, "TRANSFER-ENCODING"},
} {
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 TestHostOrTransferEncodingStopsTheStart(t *testing.T) {
t.Parallel()
// Only Host's message points to the field host.
for value, want := range map[string]string{
"Host": `"Host" is taken out of every request by Go's HTTP server, ` +
"so it can never be logged; the request's host is the field host",
"transfer-encoding": `"transfer-encoding" is taken out of every ` +
"request by Go's HTTP server, so it can never be logged",
} {
t.Run(value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{logRequestHeaders: value}.lookupEnv)
if err == nil || err.Error() != logRequestHeaders+": "+want {
t.Errorf("error %v, want %s: %s", err, logRequestHeaders, want)
}
})
}
}
func TestShortTokenStopsTheStartWithoutShowingIt(t *testing.T) {
t.Parallel()
// Characters are counted, not bytes: each é takes two.
for _, value := range []string{"", token[1:], strings.Repeat("é", 31)} {
t.Run(value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{metricsToken: value}.lookupEnv)
want := metricsToken + ": is shorter than 32 characters"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
})
}
}
func TestTokenIsLoggedMasked(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{metricsToken: token})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
if strings.Contains(out.String(), token) ||
!strings.Contains(out.String(), `"`+metricsToken+`":"********"`) {
t.Errorf("the token is not logged masked: %s", out.String())
}
}
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)
}
hostname, _ := os.Hostname()
want := map[string]string{
listenAddr: ":8080",
upstreamURL: "http://127.0.0.1:8081",
mode: "enforce",
trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
clientRequestTimeout: "45s",
clientHeaderMaxBytes: "32K",
clientIdleTimeout: "120s",
clientResponseTimeout: "30m",
upstreamRequestTimeout: "60s",
upstreamResponseTimeout: "30m",
requestMaxBytes: "100M",
responseMaxBytes: "5G",
allowNets: "",
rateLimitExemptNets: "",
denyNets: "",
rateLimitPerMinute: "1000",
rateLimitPerHour: "10000",
rateLimitPerDay: "50000",
deniedCountries: "",
allowedCountries: "",
banResponse: "403",
limitBanDuration: "1h",
limitBanRepeatWindow: "24h",
maxBanDuration: "7d",
maxBans: "5000",
banScopeV4Prefix: "32",
stateDir: "/var/lib/smallwebwaf",
stateWriteDelay: "10s",
stateCounterInterval: "15m",
metricsToken: "",
metricsTopN: "50",
instanceName: hostname,
logRequestHeaders: defaultLogRequestHeaders,
}
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.Observe != want.Observe ||
got.ClientRequestTimeout != want.ClientRequestTimeout ||
got.ClientRequestHeaderMaxBytes != want.ClientRequestHeaderMaxBytes ||
got.ClientIdleTimeout != want.ClientIdleTimeout ||
got.ClientResponseTimeout != want.ClientResponseTimeout ||
got.UpstreamRequestTimeout != want.UpstreamRequestTimeout ||
got.UpstreamResponseTimeout != want.UpstreamResponseTimeout ||
got.RequestMaxBytes != want.RequestMaxBytes ||
got.ResponseMaxBytes != want.ResponseMaxBytes ||
got.RateLimitPerMinute != want.RateLimitPerMinute ||
got.RateLimitPerHour != want.RateLimitPerHour ||
got.RateLimitPerDay != want.RateLimitPerDay {
t.Errorf("settings\n%+v\nwant\n%+v", got, want)
}
wantBanSettings(t, got, want)
}
// wantBanSettings checks the settings for bans, the state files and the
// metrics.
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper()
if got.BanResponse != want.BanResponse ||
got.LimitBanDuration != want.LimitBanDuration ||
got.LimitBanRepeatWindow != want.LimitBanRepeatWindow ||
got.MaxBanDuration != want.MaxBanDuration ||
got.MaxBans != want.MaxBans ||
got.BanScopeV4Prefix != want.BanScopeV4Prefix {
t.Errorf("ban settings\n%+v\nwant\n%+v", got, want)
}
if got.StateDir != want.StateDir ||
got.StateWriteDelay != want.StateWriteDelay ||
got.StateCounterInterval != want.StateCounterInterval {
t.Errorf("state settings\n%+v\nwant\n%+v", got, want)
}
if got.MetricsToken != want.MetricsToken || got.MetricsTopN != want.MetricsTopN {
t.Errorf("metrics token %q and top %d, want %q and %d",
got.MetricsToken, got.MetricsTopN, want.MetricsToken, want.MetricsTopN)
}
}
// 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)
}
}
// wantCountries checks the list of countries the setting name gave.
func wantCountries(t *testing.T, name string, got []string, want ...string) {
t.Helper()
if !slices.Equal(got, want) {
t.Errorf("%s gave %v, want %v", name, got, want)
}
}