Files
smallwebwaf/internal/config/config_test.go
T
clawbot f2fcf11aed
check / check (push) Waiting to run
Byte limits per client over a minute, an hour and a day (closes #20)
SWWAF_BYTES_LIMIT_PER_MINUTE, _PER_HOUR and _PER_DAY (10G, 20G, 50G)
and SWWAF_BYTES_COUNT (both). A request's bytes are counted once its
answer has ended, for a request passed to the app that the rate limits
count; what a WebSocket carries each way, once it closes. Bytes over a
limit ban the client as a broken rate limit does, and cut nothing
short. clients.json keeps the byte buckets, the log line's counts carry
the byte totals, ban notes say what the limit is on, and the limit hits
metric is labelled by kind.

Judgement call: limit_hit names a byte window minute_bytes, hour_bytes
or day_bytes, as counts names the byte totals.
Judgement call: in observe mode, the bytes of a request enforce mode
would have refused are not counted.

Model: opus-5-5
2026-10-07 10:18:13 +00:00

1488 lines
48 KiB
Go

package config_test
import (
"bytes"
"crypto/x509"
"encoding/json"
"log/slog"
"maps"
"net/http"
"net/netip"
"os"
"path/filepath"
"reflect"
"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"
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
bytesLimitPerMinute = "SWWAF_BYTES_LIMIT_PER_MINUTE"
bytesLimitPerHour = "SWWAF_BYTES_LIMIT_PER_HOUR"
bytesLimitPerDay = "SWWAF_BYTES_LIMIT_PER_DAY"
bytesCount = "SWWAF_BYTES_COUNT"
lookupSource = "SWWAF_LOOKUP_SOURCE"
lookupDBPath = "SWWAF_LOOKUP_DB_PATH"
lookupTimeout = "SWWAF_LOOKUP_TIMEOUT"
addLookupHeaders = "SWWAF_ADD_LOOKUP_HEADERS"
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"
attackBanDuration = "SWWAF_ATTACK_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"
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
metricsTopN = "SWWAF_METRICS_TOP_N"
instanceName = "SWWAF_INSTANCE_NAME"
logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS"
rulesDir = "SWWAF_RULES_DIR"
rulesEnabled = "SWWAF_RULES_ENABLED"
logRemoteURL = "SWWAF_LOG_REMOTE_URL"
logRemoteTLSCAFile = "SWWAF_LOG_REMOTE_TLS_CA_FILE"
logRemoteBuffer = "SWWAF_LOG_REMOTE_BUFFER"
logRemoteFacility = "SWWAF_LOG_REMOTE_FACILITY"
logRemoteAppName = "SWWAF_LOG_REMOTE_APP_NAME"
alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL"
alertWebhookHeaders = "SWWAF_ALERT_WEBHOOK_HEADERS"
alertSlackWebhookURL = "SWWAF_ALERT_SLACK_WEBHOOK_URL"
alertNtfyURL = "SWWAF_ALERT_NTFY_URL"
alertNtfyToken = "SWWAF_ALERT_NTFY_TOKEN" //nolint:gosec // the setting's name
alertEvents = "SWWAF_ALERT_EVENTS"
alertCooldown = "SWWAF_ALERT_COOLDOWN"
alertMaxPerHour = "SWWAF_ALERT_MAX_PER_HOUR"
)
// defaultAlertEvents is the default of SWWAF_ALERT_EVENTS, and
// defaultAlertCooldown that of SWWAF_ALERT_COOLDOWN.
const (
defaultAlertEvents = "ban,permanent_ban,waf_block,anomaly,reputation_hit," +
"source_failure,file_error"
defaultAlertCooldown = "15m"
)
// defaultLogRequestHeaders is the default of SWWAF_LOG_REQUEST_HEADERS.
const defaultLogRequestHeaders = "accept,accept-language,accept-encoding," +
"content-type,origin,range"
// testCA is a CA certificate, of which only that it reads matters here.
const testCA = `-----BEGIN CERTIFICATE-----
MIIBkzCCATmgAwIBAgIUeySaE27dnr6A2HijrMB13gTLUKIwCgYIKoZIzj0EAwIw
HjEcMBoGA1UEAwwTc21hbGx3ZWJ3YWYgdGVzdCBDQTAgFw0yNjEwMDYxNDI2MTNa
GA8yMTI2MDkxMjE0MjYxM1owHjEcMBoGA1UEAwwTc21hbGx3ZWJ3YWYgdGVzdCBD
QTBZMBMGByqGSM49AgEGCCqGSM49AwEHA0IABKDEhcWKKhet2KgSdME+iEPxyEyn
2sd9IdElbt8DM2SfCdB2JsXo0C07UNZaywMPMfn/n8LNI/PKwu+N2uX7gfSjUzBR
MB0GA1UdDgQWBBRZo3BPLv0KbV4drw6JI1JIUKmoRzAfBgNVHSMEGDAWgBRZo3BP
Lv0KbV4drw6JI1JIUKmoRzAPBgNVHRMBAf8EBTADAQH/MAoGCCqGSM49BAMCA0gA
MEUCICC5k+76UpWoSwVbZA+atu5WcALEOJGqwOUWua3zemhcAiEA8Hfdxgwp0z2v
rlG9y/jrJb6ORy3kTLWo2EA0BA67vuI=
-----END CERTIFICATE-----
`
// token is a token of 32 characters, the shortest allowed, and
// otherToken another.
const (
token = "0123456789abcdef0123456789abcdef"
otherToken = "fedcba9876543210fedcba9876543210"
)
// instance is an SWWAF_INSTANCE_NAME that is a valid app name too, and
// remoteURL an SWWAF_LOG_REMOTE_URL, for the tests that send the lines.
const (
instance = "fsn1app1/gitea"
remoteURL = "syslog+udp://192.0.2.1:514"
)
// off switches a timeout, a size limit or a rate limit off.
const off = "off"
// enabled is true, as a setting's value.
const enabled = "true"
// defaultLookupSource is the default of SWWAF_LOOKUP_SOURCE, and
// fileSource the source that is the lookup database.
const (
defaultLookupSource = "geojs"
fileSource = "file"
)
// 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,
AttackBanDuration: 7 * 24 * time.Hour,
MaxBans: 5000,
BanScopeV4Prefix: 32,
StateDir: "/var/lib/smallwebwaf",
StateWriteDelay: 10 * time.Second,
StateCounterInterval: 15 * time.Minute,
MetricsToken: "",
MetricsTopN: 50,
RulesDir: "/etc/smallwebwaf/rules.d",
RulesEnabled: true,
})
wantLookupSettings(t, cfg, config.Config{
LookupSource: defaultLookupSource, LookupTimeout: time.Second,
})
wantByteLimitSettings(t, cfg, config.Config{
BytesLimitPerMinute: 10 << 30, BytesLimitPerHour: 20 << 30,
BytesLimitPerDay: 50 << 30, BytesCount: "both",
})
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)
}
if len(cfg.RateLimitExemptPaths) != 0 {
t.Errorf("%s gave %v, want none", rateLimitExemptPaths, cfg.RateLimitExemptPaths)
}
}
func TestNoSingleRequestBreaksAByteLimitAtTheDefaults(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{})
// The largest request body and the largest response, both counted.
largest := cfg.RequestMaxBytes + cfg.ResponseMaxBytes
for _, limit := range []int64{
cfg.BytesLimitPerMinute, cfg.BytesLimitPerHour, cfg.BytesLimitPerDay,
} {
if largest > limit {
t.Errorf("a request of %d bytes breaks the byte limit of %d", largest, limit)
}
}
}
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",
attackBanDuration: "1d",
maxBans: "100",
banScopeV4Prefix: "24",
stateDir: "/srv/waf-state",
stateWriteDelay: "500ms",
stateCounterInterval: "1h",
metricsToken: token,
metricsTopN: "10",
rulesDir: "/srv/waf-rules",
rulesEnabled: "false",
})
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,
AttackBanDuration: 24 * time.Hour,
MaxBans: 100,
BanScopeV4Prefix: 24,
StateDir: "/srv/waf-state",
StateWriteDelay: 500 * time.Millisecond,
StateCounterInterval: time.Hour,
MetricsToken: token,
MetricsTopN: 10,
RulesDir: "/srv/waf-rules",
RulesEnabled: false,
})
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 TestByteLimitSettingsAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
bytesLimitPerMinute: "512M",
bytesLimitPerHour: off,
bytesLimitPerDay: "100000",
bytesCount: "response",
})
wantByteLimitSettings(t, cfg, config.Config{
BytesLimitPerMinute: 512 << 20, BytesLimitPerHour: 0,
BytesLimitPerDay: 100000, BytesCount: "response",
})
}
func TestRateLimitExemptPathsAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{rateLimitExemptPaths: "/assets/, /favicon.ico"})
if !slices.Equal(cfg.RateLimitExemptPaths, []string{"/assets/", "/favicon.ico"}) {
t.Errorf("%s gave %v, want /assets/ and /favicon.ico",
rateLimitExemptPaths, cfg.RateLimitExemptPaths)
}
}
func TestPathPrefixNotStartingWithSlashStopsTheStart(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(
environment{rateLimitExemptPaths: "/favicon.ico,assets/"}.lookupEnv)
want := rateLimitExemptPaths + `: "assets/" is not a path prefix ` +
`starting with /, such as /assets/`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
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 TestRemoteLogSettingsDefaults(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{instanceName: instance})
if cfg.LogRemoteURL != nil || cfg.LogRemoteTLSCAs != nil ||
cfg.LogRemoteBuffer != 10000 || cfg.LogRemoteFacility != 16 ||
cfg.LogRemoteAppName != instance {
t.Errorf("remote log settings %v, %v, %d, %d and %q, want no URL, no "+
"certificates, 10000, 16 and %s's %s", cfg.LogRemoteURL,
cfg.LogRemoteTLSCAs, cfg.LogRemoteBuffer, cfg.LogRemoteFacility,
cfg.LogRemoteAppName, instanceName, instance)
}
}
func TestRemoteLogSettingsAsSet(t *testing.T) {
t.Parallel()
caFile := filepath.Join(t.TempDir(), "ca.pem")
err := os.WriteFile(caFile, []byte(testCA), 0o600)
if err != nil {
t.Fatalf("write %s: %v", caFile, err)
}
cfg := fromEnvironment(t, environment{
logRemoteURL: "syslog+tls://logs.example:6514",
logRemoteTLSCAFile: caFile,
logRemoteBuffer: "500",
logRemoteFacility: "daemon",
logRemoteAppName: instance,
})
roots := x509.NewCertPool()
roots.AppendCertsFromPEM([]byte(testCA))
if cfg.LogRemoteURL.String() != "syslog+tls://logs.example:6514" ||
!roots.Equal(cfg.LogRemoteTLSCAs) || cfg.LogRemoteBuffer != 500 ||
cfg.LogRemoteFacility != 3 || cfg.LogRemoteAppName != instance {
t.Errorf("remote log settings %v, %v, %d, %d and %q", cfg.LogRemoteURL,
cfg.LogRemoteTLSCAs, cfg.LogRemoteBuffer, cfg.LogRemoteFacility,
cfg.LogRemoteAppName)
}
}
func TestRemoteLogURLForms(t *testing.T) {
t.Parallel()
for _, value := range []string{
"syslog+udp://192.0.2.1:514",
"syslog+tcp://[2001:db8::1]:514",
"syslog+tls://logs.example:6514/",
} {
cfg := fromEnvironment(t, environment{logRemoteURL: value})
if cfg.LogRemoteURL.String() != value {
t.Errorf("%s read as %v", value, cfg.LogRemoteURL)
}
}
cfg := fromEnvironment(t, environment{logRemoteURL: ""})
if cfg.LogRemoteURL != nil {
t.Errorf("set but empty, %s read as %v", logRemoteURL, cfg.LogRemoteURL)
}
}
func TestRemoteLogFacilitiesByNumber(t *testing.T) {
t.Parallel()
for name, number := range map[string]int{
"kern": 0, "user": 1, "auth": 4, "authpriv": 10, "ftp": 11,
"local0": 16, "local5": 21, "local7": 23,
} {
cfg := fromEnvironment(t, environment{logRemoteFacility: name})
if cfg.LogRemoteFacility != number {
t.Errorf("%s read as %d, want %d", name, cfg.LogRemoteFacility, number)
}
}
}
func TestInvalidRemoteLogSettingStopsTheStart(t *testing.T) {
t.Parallel()
for _, tc := range []struct{ name, value string }{
{logRemoteURL, "logs.example:514"},
{logRemoteURL, "syslog://logs.example:514"},
{logRemoteURL, "http://logs.example:514"},
{logRemoteURL, "syslog+udp://logs.example"},
{logRemoteURL, "syslog+tcp://:514"},
{logRemoteURL, "syslog+tcp://logs.example:0"},
{logRemoteURL, "syslog+tls://logs.example:65536"},
{logRemoteURL, "syslog+tls://user@logs.example:6514"},
{logRemoteURL, "syslog+tcp://logs.example:514/app"},
{logRemoteURL, "syslog+tcp://logs.example:514?tls=1"},
{logRemoteTLSCAFile, "/nonexistent/ca.pem"},
{logRemoteBuffer, off}, {logRemoteBuffer, "0"}, {logRemoteBuffer, "10K"},
{logRemoteFacility, "local8"}, {logRemoteFacility, "LOCAL0"},
{logRemoteFacility, "16"}, {logRemoteFacility, ""},
{logRemoteAppName, ""}, {logRemoteAppName, "my app"},
{logRemoteAppName, "gitéa"}, {logRemoteAppName, strings.Repeat("a", 49)},
} {
_, err := config.FromEnvironment(environment{tc.name: tc.value}.lookupEnv)
if err == nil || !strings.HasPrefix(err.Error(), tc.name+": ") {
t.Errorf("%s=%q: error %v, want one naming it", tc.name, tc.value, err)
}
}
}
func TestRemoteLogCAFileWithoutCertificateStopsTheStart(t *testing.T) {
t.Parallel()
caFile := filepath.Join(t.TempDir(), "ca.pem")
err := os.WriteFile(caFile, []byte("not a certificate\n"), 0o600)
if err != nil {
t.Fatalf("write %s: %v", caFile, err)
}
_, err = config.FromEnvironment(environment{logRemoteTLSCAFile: caFile}.lookupEnv)
want := logRemoteTLSCAFile + `: "` + caFile + `" holds no PEM certificate`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestInstanceNameNotAnAppNameStopsTheStartOnlyWhileSending(t *testing.T) {
t.Parallel()
const spaced = "fsn1 app1"
sending := environment{logRemoteURL: remoteURL, instanceName: spaced}
_, err := config.FromEnvironment(sending.lookupEnv)
want := logRemoteAppName + `: is unset, and ` + instanceName +
` "fsn1 app1", its default, is not 1 to 48 printable ASCII characters ` +
`without a space, such as gitea`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
cfg := fromEnvironment(t, environment{instanceName: spaced})
if cfg.LogRemoteAppName != spaced {
t.Errorf("not sending, %s is %q", logRemoteAppName, cfg.LogRemoteAppName)
}
sending[logRemoteAppName] = instance
cfg = fromEnvironment(t, sending)
if cfg.LogRemoteAppName != instance {
t.Errorf("set to %s, %s is %q", instance, logRemoteAppName,
cfg.LogRemoteAppName)
}
}
func TestAppNameSetStopsTheStartWhileSending(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{
logRemoteURL: remoteURL,
instanceName: instance,
logRemoteAppName: "my app",
}.lookupEnv)
want := logRemoteAppName + `: "my app" is not 1 to 48 printable ASCII ` +
`characters without a space, such as gitea`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestAlertSettingsDefaults(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{})
if cfg.AlertWebhookURL != nil || len(cfg.AlertWebhookHeaders) != 0 ||
cfg.AlertSlackWebhookURL != nil || cfg.AlertNtfyURL != nil ||
cfg.AlertNtfyToken != "" ||
strings.Join(cfg.AlertEvents, ",") != defaultAlertEvents ||
cfg.AlertCooldown != 15*time.Minute || cfg.AlertMaxPerHour != 60 {
t.Errorf("alert settings %v, %v, %v, %v, %q, %v, %s and %d, want no URLs, "+
"no headers, no token, %s, 15m and 60", cfg.AlertWebhookURL,
cfg.AlertWebhookHeaders, cfg.AlertSlackWebhookURL, cfg.AlertNtfyURL,
cfg.AlertNtfyToken, cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour,
defaultAlertEvents)
}
}
func TestAlertSettingsAsSet(t *testing.T) {
t.Parallel()
const (
webhook = "https://alerts.example:8443/hooks/waf?team=ops"
slack = "https://hooks.slack.example/services/T0123/B4567/abcdef"
ntfy = "https://ntfy.example/smallwebwaf-alerts"
token = "tk_0123456789abcdefghijklmnopq"
)
cfg := fromEnvironment(t, environment{
alertWebhookURL: webhook,
alertWebhookHeaders: "Authorization: Bearer abc:def , x-team:ops",
alertSlackWebhookURL: slack,
alertNtfyURL: ntfy,
alertNtfyToken: token,
alertEvents: "ban, file_error",
alertCooldown: "1h",
alertMaxPerHour: "10",
})
headers := http.Header{"Authorization": {"Bearer abc:def"}, "X-Team": {"ops"}}
if cfg.AlertWebhookURL.String() != webhook ||
!reflect.DeepEqual(cfg.AlertWebhookHeaders, headers) ||
cfg.AlertSlackWebhookURL.String() != slack || cfg.AlertNtfyURL.String() != ntfy ||
cfg.AlertNtfyToken != token ||
!slices.Equal(cfg.AlertEvents, []string{"ban", "file_error"}) ||
cfg.AlertCooldown != time.Hour || cfg.AlertMaxPerHour != 10 {
t.Errorf("alert settings %v, %v, %v, %v, %q, %v, %s and %d", cfg.AlertWebhookURL,
cfg.AlertWebhookHeaders, cfg.AlertSlackWebhookURL, cfg.AlertNtfyURL,
cfg.AlertNtfyToken, cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour)
}
}
func TestAlertSettingsSetEmptyOrOff(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
alertWebhookURL: "", alertSlackWebhookURL: "", alertNtfyURL: "",
alertNtfyToken: "", alertEvents: "", alertCooldown: off, alertMaxPerHour: off,
})
if cfg.AlertWebhookURL != nil || cfg.AlertSlackWebhookURL != nil ||
cfg.AlertNtfyURL != nil || cfg.AlertNtfyToken != "" ||
len(cfg.AlertEvents) != 0 || cfg.AlertCooldown != 0 || cfg.AlertMaxPerHour != 0 {
t.Errorf("set empty or off, alert settings %v, %v, %v, %q, %v, %s and %d",
cfg.AlertWebhookURL, cfg.AlertSlackWebhookURL, cfg.AlertNtfyURL,
cfg.AlertNtfyToken, cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour)
}
}
func TestInvalidAlertSettingStopsTheStart(t *testing.T) {
t.Parallel()
wantStartStopped(t, []struct{ name, value string }{
{alertWebhookURL, "alerts.example/smallwebwaf"},
{alertWebhookURL, "ftp://alerts.example/"},
{alertWebhookURL, "https:///smallwebwaf"},
{alertWebhookURL, "https://user:password@alerts.example/"},
{alertWebhookURL, "https://alerts.example/#top"},
{alertWebhookURL, "https://alerts.example:0/"},
{alertWebhookURL, "https://alerts.example:65536/"},
{alertSlackWebhookURL, "hooks.slack.example/services/T0123"},
{alertSlackWebhookURL, "https://user:password@hooks.slack.example/"},
{alertNtfyURL, "ntfy://ntfy.example/smallwebwaf-alerts"},
{alertNtfyURL, "https://ntfy.example/smallwebwaf-alerts#top"},
{alertWebhookHeaders, "Authorization"},
{alertWebhookHeaders, "X Team:ops"},
{alertWebhookHeaders, ":ops"},
{alertWebhookHeaders, "X-Team:ops,"},
{alertWebhookHeaders, "X-Team:o\r\nps"},
{alertEvents, "bans"},
{alertEvents, "summary"},
{alertEvents, "ban,,file_error"},
{alertCooldown, "0"},
{alertCooldown, "soon"},
{alertMaxPerHour, "0"},
{alertMaxPerHour, "-1"},
{alertMaxPerHour, "1.5"},
})
}
func TestWebhookHeadersAreLoggedMaskedAndNeverShown(t *testing.T) {
t.Parallel()
const secret = "Bearer 0123456789abcdef"
cfg := fromEnvironment(t, environment{
alertWebhookHeaders: "Authorization:" + secret + ",X-Team:ops",
})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
logged := out.String()
if strings.Contains(logged, secret) || strings.Contains(logged, "ops") ||
!strings.Contains(logged,
`"`+alertWebhookHeaders+`":"Authorization:********,X-Team:********"`) {
t.Errorf("the headers are not logged masked: %s", logged)
}
// An item that is not a header is named by its place, not shown.
_, err := config.FromEnvironment(environment{
alertWebhookHeaders: "X-Team:ops," + secret,
}.lookupEnv)
want := alertWebhookHeaders + ": item 2 is not a header name followed by : " +
"and the header's value, such as Authorization:Bearer <token>"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestWebhookURLIsLoggedWithoutItsPathOrQueryAndNeverShown(t *testing.T) {
t.Parallel()
const secret = "T0123/B4567/abcdef"
for value, want := range map[string]string{
"https://hooks.example/services/" + secret: "https://hooks.example/********",
"https://hooks.example:8443?token=" + secret: "https://hooks.example:8443/********",
"http://[2001:db8::1]:8080": "http://[2001:db8::1]:8080",
} {
cfg := fromEnvironment(t, environment{alertWebhookURL: value})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
logged := out.String()
if strings.Contains(logged, secret) ||
!strings.Contains(logged, `"`+alertWebhookURL+`":"`+want+`"`) {
t.Errorf("%s is not logged as %s: %s", value, want, logged)
}
}
// A value that is not such a URL is not shown either.
for _, value := range []string{
"ftp://hooks.example/services/" + secret,
"https://hooks.example/services/%zz" + secret,
} {
_, err := config.FromEnvironment(environment{alertWebhookURL: value}.lookupEnv)
want := alertWebhookURL + ": is not an http or https URL without a user or " +
"a fragment, such as https://alerts.example/smallwebwaf"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
}
func TestSlackAndNtfySettingsAreLoggedWithoutTheirSecrets(t *testing.T) {
t.Parallel()
const token = "tk_0123456789abcdefghijklmnopq"
cfg := fromEnvironment(t, environment{
alertSlackWebhookURL: "https://hooks.slack.example/services/T0123/B4567/abcdef",
alertNtfyURL: "https://ntfy.example/smallwebwaf-alerts",
alertNtfyToken: token,
})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
logged := out.String()
for _, want := range []string{
`"` + alertSlackWebhookURL + `":"https://hooks.slack.example/********"`,
`"` + alertNtfyURL + `":"https://ntfy.example/********"`,
`"` + alertNtfyToken + `":"********"`,
} {
if !strings.Contains(logged, want) {
t.Errorf("no %s in the settings logged: %s", want, logged)
}
}
for _, secret := range []string{"T0123", "smallwebwaf-alerts", token} {
if strings.Contains(logged, secret) {
t.Errorf("%s in the settings logged: %s", secret, logged)
}
}
}
func TestNtfyTokenWithAControlCharacterStopsTheStartWithoutShowingIt(t *testing.T) {
t.Parallel()
// A file saved with Windows line ends keeps the carriage return.
for name, env := range map[string]environment{
"set": {alertNtfyToken: token + "\r"},
"in a file": {alertNtfyToken + "_FILE": writeFile(t, token+"\r\n")},
} {
_, err := config.FromEnvironment(env.lookupEnv)
want := alertNtfyToken + ": holds a control character, such as the " +
"carriage return of a Windows line end"
if err == nil || err.Error() != want {
t.Errorf("%s: error %v, want %s", name, err, want)
}
}
}
func TestInstanceNameWithAControlCharacterStopsTheStartOnlyWithNtfySet(t *testing.T) {
t.Parallel()
const name = "fsn1app1\r"
_, err := config.FromEnvironment(environment{
instanceName: name, alertNtfyURL: "https://ntfy.example/smallwebwaf-alerts",
}.lookupEnv)
want := instanceName + `: "fsn1app1\r" holds a control character, such as the ` +
`carriage return of a Windows line end, and is sent to ntfy in a header ` +
`while ` + alertNtfyURL + ` is set`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
cfg := fromEnvironment(t, environment{instanceName: name})
if cfg.InstanceName != name {
t.Errorf("not sending to ntfy, %s is %q", instanceName, cfg.InstanceName)
}
}
func TestInstanceNameNotUTF8StopsTheStart(t *testing.T) {
t.Parallel()
// café saved in Latin-1.
const latin1 = "caf\xe9"
for name, env := range map[string]environment{
"set": {instanceName: latin1},
"in a file": {instanceName + "_FILE": writeFile(t, latin1+"\n")},
} {
_, err := config.FromEnvironment(env.lookupEnv)
want := instanceName + `: "caf\xe9" is not valid UTF-8`
if err == nil || err.Error() != want {
t.Errorf("%s: error %v, want %s", name, err, want)
}
}
}
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 TestLookupSettingsAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
lookupTimeout: "500ms", addLookupHeaders: enabled,
})
wantLookupSettings(t, cfg, config.Config{
LookupSource: defaultLookupSource, LookupTimeout: 500 * time.Millisecond,
AddLookupHeaders: true,
})
cfg = fromEnvironment(t, environment{lookupSource: off})
wantLookupSettings(t, cfg, config.Config{
LookupSource: off, LookupTimeout: time.Second,
})
}
func TestLookupDBPathGoesWithTheFileSourceAlone(t *testing.T) {
t.Parallel()
const path = "/var/lib/ipinfo/ipinfo_lite.mmdb"
cfg := fromEnvironment(t, environment{lookupSource: fileSource, lookupDBPath: path})
if cfg.LookupSource != fileSource || cfg.LookupDBPath != path {
t.Errorf("lookups from %q in %q, want file in %q",
cfg.LookupSource, cfg.LookupDBPath, path)
}
for _, tc := range []struct {
env environment
want string
}{
{
environment{lookupSource: fileSource},
lookupSource + ": is file while " + lookupDBPath +
" is unset; it names the file to look clients up in",
},
{
environment{lookupSource: fileSource, lookupDBPath: ""},
lookupSource + ": is file while " + lookupDBPath +
" is unset; it names the file to look clients up in",
},
{
environment{lookupDBPath: path},
lookupDBPath + ": is set while " + lookupSource + " is geojs; only file reads it",
},
{
environment{lookupSource: off, lookupDBPath: path},
lookupDBPath + ": is set while " + lookupSource + " is off; only file reads it",
},
} {
_, err := config.FromEnvironment(tc.env.lookupEnv)
if err == nil || err.Error() != tc.want {
t.Errorf("settings %v: error %v, want %s", tc.env, err, tc.want)
}
}
}
func TestSettingNeedingLookupsStopsTheStartWhileTheyAreOff(t *testing.T) {
t.Parallel()
for name, value := range map[string]string{
deniedCountries: "kp",
allowedCountries: "de",
addLookupHeaders: enabled,
} {
t.Run(name, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{
lookupSource: off, name: value,
}.lookupEnv)
if err == nil || !strings.HasPrefix(err.Error(), name+": ") ||
!strings.Contains(err.Error(), lookupSource+" is off") {
t.Errorf("error %v, want one naming %s and %s", err, name, lookupSource)
}
})
}
// Set empty, the country lists need nothing looked up.
fromEnvironment(t, environment{
lookupSource: off, deniedCountries: "", 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()
wantStartStopped(t, []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"},
{rateLimitExemptPaths, "/assets/,,/static/"},
{bytesLimitPerMinute, "10GB"}, {bytesLimitPerHour, "0"},
{bytesLimitPerDay, "-1G"},
{bytesCount, "all"}, {bytesCount, "Both"}, {bytesCount, ""},
{lookupSource, "ipinfo"}, {lookupSource, "GeoJS"}, {lookupSource, ""},
{lookupTimeout, off}, {lookupTimeout, "0s"}, {lookupTimeout, "1"},
{addLookupHeaders, "yes"},
{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"},
{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"},
{rulesEnabled, "yes"}, {rulesEnabled, "True"},
})
}
func TestInvalidBanOrStateValueStopsTheStart(t *testing.T) {
t.Parallel()
wantStartStopped(t, []struct{ name, value string }{
{banResponse, "404"}, {banResponse, "drop"}, {banResponse, ""},
{limitBanDuration, off}, {limitBanDuration, "0s"}, {limitBanDuration, "1"},
{limitBanRepeatWindow, off}, {limitBanRepeatWindow, "-1h"},
{maxBanDuration, off}, {maxBanDuration, "1w"}, {attackBanDuration, off},
{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"},
})
}
// wantStartStopped checks that each setting, set to its value, stops the
// start with an error that names the setting.
func wantStartStopped(t *testing.T, invalid []struct{ name, value string }) {
t.Helper()
for _, tc := range invalid {
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 _, name := range []string{adminToken, metricsToken} {
for _, value := range []string{"", token[1:], strings.Repeat("é", 31)} {
t.Run(name+"="+value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{name: value}.lookupEnv)
want := name + ": is shorter than 32 characters"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
})
}
}
}
func TestTokensAreReadAndLoggedMasked(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{adminToken: otherToken, metricsToken: token})
if cfg.AdminToken != otherToken || cfg.MetricsToken != token {
t.Errorf("admin token %q and metrics token %q, want %q and %q",
cfg.AdminToken, cfg.MetricsToken, otherToken, token)
}
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
logged := out.String()
if strings.Contains(logged, token) || strings.Contains(logged, otherToken) ||
!strings.Contains(logged, `"`+adminToken+`":"********"`) ||
!strings.Contains(logged, `"`+metricsToken+`":"********"`) {
t.Errorf("the tokens are not logged masked: %s", logged)
}
}
func TestSettingFromFileLosesOneNewlineAndNoMore(t *testing.T) {
t.Parallel()
for contents, want := range map[string]string{
token: token,
token + "\n": token,
token + "\n\n": token + "\n",
token + " \n": token + " ",
} {
cfg := fromEnvironment(t, environment{
metricsToken + "_FILE": writeFile(t, contents),
})
if cfg.MetricsToken != want {
t.Errorf("file holding %q gave %s %q, want %q", contents, metricsToken,
cfg.MetricsToken, want)
}
}
}
func TestSettingFromFileIsCheckedAsTheSettingItself(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{
requestMaxBytes + "_FILE": writeFile(t, "lots\n"),
}.lookupEnv)
want := requestMaxBytes + `: "lots" is not a size such as 512K, 100M or 5G, or off`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
_, err = config.FromEnvironment(environment{
logRemoteURL: remoteURL,
instanceName: instance,
logRemoteAppName + "_FILE": writeFile(t, "my app\n"),
}.lookupEnv)
want = logRemoteAppName + `: "my app" is not 1 to 48 printable ASCII ` +
`characters without a space, such as gitea`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestSettingAndItsFileBothSetStopsTheStart(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{
metricsToken: token,
metricsToken + "_FILE": writeFile(t, token),
}.lookupEnv)
want := metricsToken + ": is set, and so is " + metricsToken +
"_FILE; set only one of them"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestUnreadableSettingFileStopsTheStart(t *testing.T) {
t.Parallel()
dir := t.TempDir()
for _, path := range []string{filepath.Join(dir, "missing"), dir} {
_, err := config.FromEnvironment(environment{metricsToken + "_FILE": path}.lookupEnv)
want := metricsToken + "_FILE: cannot be read: "
if err == nil || !strings.HasPrefix(err.Error(), want) {
t.Errorf("%s: error %v, want one starting %s", path, err, want)
}
}
}
func TestTokenFromFileIsLoggedMaskedWithTheFile(t *testing.T) {
t.Parallel()
path := writeFile(t, token+"\n")
cfg := fromEnvironment(t, environment{metricsToken + "_FILE": path})
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)
}
if strings.Contains(out.String(), token) ||
line.Settings[metricsToken] != "********" ||
line.Settings[metricsToken+"_FILE"] != path {
t.Errorf("the token is not logged masked, with its file %s: %s", path,
out.String())
}
}
func TestRemoteLogCAFileIsNotReadAsAFileInItsTurn(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
logRemoteTLSCAFile + "_FILE": writeFile(t, "/nonexistent/ca.pem\n"),
})
if cfg.LogRemoteTLSCAs != nil {
t.Errorf("%s_FILE gave certificates", logRemoteTLSCAFile)
}
}
// writeFile writes contents to a file in a directory of its own, removed
// when the test ends, and returns the file's path.
func writeFile(t *testing.T, contents string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "setting")
err := os.WriteFile(path, []byte(contents), 0o600)
if err != nil {
t.Fatalf("write %s: %v", path, err)
}
return path
}
func TestLogsEachSettingWithItsValue(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{clientRequestTimeout: "45s"})
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",
rateLimitExemptPaths: "",
bytesLimitPerMinute: "10G",
bytesLimitPerHour: "20G",
bytesLimitPerDay: "50G",
bytesCount: "both",
lookupSource: defaultLookupSource,
lookupDBPath: "",
lookupTimeout: "1s",
addLookupHeaders: "false",
deniedCountries: "",
allowedCountries: "",
banResponse: "403",
limitBanDuration: "1h",
limitBanRepeatWindow: "24h",
maxBanDuration: "7d",
attackBanDuration: "7d",
maxBans: "5000",
banScopeV4Prefix: "32",
stateDir: "/var/lib/smallwebwaf",
stateWriteDelay: "10s",
stateCounterInterval: "15m",
adminToken: "",
metricsToken: "",
metricsTopN: "50",
instanceName: hostname,
logRequestHeaders: defaultLogRequestHeaders,
rulesDir: "/etc/smallwebwaf/rules.d",
rulesEnabled: "true",
logRemoteURL: "",
logRemoteTLSCAFile: "",
logRemoteBuffer: "10000",
logRemoteFacility: "local0",
logRemoteAppName: hostname,
alertWebhookURL: "",
alertWebhookHeaders: "",
alertSlackWebhookURL: "",
alertNtfyURL: "",
alertNtfyToken: "",
alertEvents: defaultAlertEvents,
alertCooldown: defaultAlertCooldown,
alertMaxPerHour: "60",
}
if got := loggedSettings(t, cfg); !maps.Equal(got, want) {
t.Errorf("logged settings\n%v\nwant\n%v", got, want)
}
}
// loggedSettings returns the settings as cfg logs them, each by its name.
func loggedSettings(t *testing.T, cfg *config.Config) map[string]string {
t.Helper()
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)
}
return line.Settings
}
// 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)
}
// wantByteLimitSettings checks the settings for the byte limits.
func wantByteLimitSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper()
if got.BytesLimitPerMinute != want.BytesLimitPerMinute ||
got.BytesLimitPerHour != want.BytesLimitPerHour ||
got.BytesLimitPerDay != want.BytesLimitPerDay ||
got.BytesCount != want.BytesCount {
t.Errorf("byte limits %d, %d and %d counting %s, want %d, %d and %d counting %s",
got.BytesLimitPerMinute, got.BytesLimitPerHour, got.BytesLimitPerDay,
got.BytesCount, want.BytesLimitPerMinute, want.BytesLimitPerHour,
want.BytesLimitPerDay, want.BytesCount)
}
}
// wantLookupSettings checks the settings for lookups.
func wantLookupSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper()
if got.LookupSource != want.LookupSource || got.LookupTimeout != want.LookupTimeout ||
got.AddLookupHeaders != want.AddLookupHeaders {
t.Errorf("lookups from %q, waited for %s, headers %t; want %q, %s, %t",
got.LookupSource, got.LookupTimeout, got.AddLookupHeaders,
want.LookupSource, want.LookupTimeout, want.AddLookupHeaders)
}
}
// wantBanSettings checks the settings for bans, the state files, the
// metrics and the rule files.
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.AttackBanDuration != want.AttackBanDuration ||
got.MaxBans != want.MaxBans ||
got.BanScopeV4Prefix != want.BanScopeV4Prefix {
t.Errorf("ban settings\n%+v\nwant\n%+v", got, want)
}
if got.RulesDir != want.RulesDir || got.RulesEnabled != want.RulesEnabled {
t.Errorf("rule files in %q, on: %t, want %q, %t",
got.RulesDir, got.RulesEnabled, want.RulesDir, want.RulesEnabled)
}
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)
}
}