package smallwebwaf_test import ( "bytes" "context" "encoding/json" "io" "net" "net/http" "net/http/httptest" "os" "path/filepath" "strings" "sync" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/smallwebwaf" ) const ( // waitLimit bounds how long a test waits for what should happen. waitLimit = 10 * time.Second // pollInterval is how often a test looks for a line. pollInterval = 10 * time.Millisecond // testVersion is the version the tests give smallwebwaf. testVersion = "test" // localhost is where the tests listen. localhost = "127.0.0.1" listenAddr = "SWWAF_LISTEN_ADDR" upstreamURL = "SWWAF_UPSTREAM_URL" trustedProxies = "SWWAF_TRUSTED_PROXIES" stateDir = "SWWAF_STATE_DIR" stateWriteDelay = "SWWAF_STATE_WRITE_DELAY" stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL" rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" // greeting is what the tests' app answers. greeting = "hello from the app" ) // output collects what smallwebwaf writes on stdout. type output struct { mu sync.Mutex buf bytes.Buffer } // Write adds lines smallwebwaf writes. func (o *output) Write(p []byte) (int, error) { o.mu.Lock() defer o.mu.Unlock() return o.buf.Write(p) } // line returns the first line whose field key is value, waiting for it. func (o *output) line(t *testing.T, key, value string) map[string]any { t.Helper() deadline := time.Now().Add(waitLimit) for time.Now().Before(deadline) { o.mu.Lock() text := o.buf.String() o.mu.Unlock() for line := range strings.Lines(text) { var fields map[string]any err := json.Unmarshal([]byte(line), &fields) if err != nil { t.Fatalf("output line %q is not JSON: %v", line, err) } if fields[key] == value { return fields } } time.Sleep(pollInterval) } t.Fatalf("no line with %s %q in the output:\n%s", key, value, o.text()) return nil } // text returns everything written so far. func (o *output) text() string { o.mu.Lock() defer o.mu.Unlock() return o.buf.String() } // run runs smallwebwaf with the settings in env until ctx is done, and // returns its exit status. func run(ctx context.Context, env map[string]string, out *output) int { return smallwebwaf.Run(ctx, smallwebwaf.Params{ Version: testVersion, LookupEnv: func(name string) (string, bool) { value, ok := env[name] return value, ok }, Stdout: out, }) } func TestInvalidSettingStopsTheStart(t *testing.T) { t.Parallel() out := &output{} status := run(t.Context(), map[string]string{"SWWAF_REQUEST_MAX_BYTES": "lots"}, out) if status != 1 { t.Errorf("exit status %d, want 1", status) } line := out.line(t, "msg", "invalid setting") message, _ := line["error"].(string) if line["type"] != "process" || line["level"] != "ERROR" || !strings.HasPrefix(message, "SWWAF_REQUEST_MAX_BYTES: ") { t.Errorf("start refused with %v", line) } } func TestShortMetricsTokenStopsTheStartUnshown(t *testing.T) { t.Parallel() const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use out := &output{} status := run(t.Context(), map[string]string{"SWWAF_METRICS_TOKEN": token}, out) if status != 1 { t.Errorf("exit status %d, want 1", status) } line := out.line(t, "msg", "invalid setting") if line["error"] != "SWWAF_METRICS_TOKEN: is shorter than 32 characters" { t.Errorf("start refused with %v", line) } if strings.Contains(out.text(), token) { t.Errorf("the output shows the token:\n%s", out.text()) } } func TestAddressInUseStopsTheStart(t *testing.T) { t.Parallel() taken, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0") if err != nil { t.Fatalf("listen: %v", err) } defer func() { _ = taken.Close() }() out := &output{} status := run(t.Context(), map[string]string{ listenAddr: taken.Addr().String(), stateDir: t.TempDir(), }, out) if status != 1 { t.Errorf("exit status %d, want 1", status) } out.line(t, "msg", "cannot listen on SWWAF_LISTEN_ADDR") } func TestServesUntilToldToStop(t *testing.T) { t.Parallel() appURL := startApp(t) dir := t.TempDir() ctx, stop := context.WithCancel(t.Context()) out := &output{} exited := make(chan int, 1) go func() { exited <- run(ctx, map[string]string{ listenAddr: localhost + ":0", upstreamURL: appURL, stateDir: dir, }, out) }() starting := out.line(t, "msg", "starting") wantStartingLine(t, starting, appURL, dir) addr, _ := starting["address"].(string) wantGreeting(t, "http://"+addr+"/") out.line(t, "type", "request") stop() select { case status := <-exited: if status != 0 { t.Errorf("exit status %d, want 0", status) } case <-time.After(waitLimit): t.Fatal("still running after being told to stop") } out.line(t, "msg", "stopped") } func TestStateKeptAcrossRestarts(t *testing.T) { t.Parallel() env := map[string]string{ listenAddr: localhost + ":0", upstreamURL: startApp(t), stateDir: t.TempDir(), rateLimitPerDay: "2", // Neither comes due in the test: the files are written as // smallwebwaf stops. stateWriteDelay: "1h", stateCounterInterval: "1h", } // The two requests a day allows, and a stop. runUntilStopped(t, env, func(url string) { wantGreeting(t, url) wantGreeting(t, url) }) // After a restart the client has no fresh allowance: its third // request breaks the day limit, and bans it. out := runUntilStopped(t, env, func(url string) { wantRefused(t, url) }) out.line(t, "action", "rate_limited") // After another, the ban still refuses it. out = runUntilStopped(t, env, func(url string) { wantRefused(t, url) }) out.line(t, "action", "banned") } func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) { t.Parallel() const scope = "SWWAF_BAN_SCOPE_V4_PREFIX" env := map[string]string{ listenAddr: localhost + ":0", upstreamURL: startApp(t), stateDir: t.TempDir(), trustedProxies: localhost + "/32", rateLimitPerDay: "1", scope: "24", } // 203.0.113.9's second request breaks the day limit, and bans // 203.0.113.0/24. runUntilStopped(t, env, func(url string) { wantStatus(t, url, "203.0.113.9", http.StatusOK) wantStatus(t, url, "203.0.113.9", http.StatusForbidden) }) // With each address a netblock of its own after a restart, that ban // still refuses all of 203.0.113.0/24. 198.51.100.7 is banned alone. env[scope] = "32" runUntilStopped(t, env, func(url string) { wantStatus(t, url, "203.0.113.200", http.StatusForbidden) wantStatus(t, url, "203.0.114.1", http.StatusOK) wantStatus(t, url, "198.51.100.7", http.StatusOK) wantStatus(t, url, "198.51.100.7", http.StatusForbidden) }) // With /24 netblocks again, that ban still refuses 198.51.100.7, and // no other address. env[scope] = "24" runUntilStopped(t, env, func(url string) { wantStatus(t, url, "198.51.100.7", http.StatusForbidden) wantStatus(t, url, "198.51.100.8", http.StatusOK) }) } func TestBanAddedAndLiftedByEditingBansJSON(t *testing.T) { t.Parallel() const ( // bans.json as an admin writes it with a ban, permanent, on // 203.0.113.0/24, and with none. oneBan = `{"version": 1, "bans": [{"netblock": "203.0.113.0/24", ` + `"start": "2026-10-06T00:00:00Z", "expires": null}]}` noBan = `{"version": 1, "bans": []}` ) dir := t.TempDir() env := map[string]string{ listenAddr: localhost + ":0", upstreamURL: startApp(t), stateDir: dir, trustedProxies: localhost + "/32", // No write comes due in the test, so only the watch on the // directory can take the edits in. stateWriteDelay: "1h", stateCounterInterval: "1h", } runUntilStopped(t, env, func(url string) { path := filepath.Join(dir, "bans.json") saveUntilAnswered(t, path, oneBan, url, "203.0.113.9", http.StatusForbidden) wantStatus(t, url, "198.51.100.7", http.StatusOK) saveUntilAnswered(t, path, noBan, url, "203.0.113.9", http.StatusOK) }) } func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) { t.Parallel() dir := t.TempDir() err := os.WriteFile(filepath.Join(dir, "bans.json"), []byte("{\n"), 0o600) if err != nil { t.Fatalf("write bans.json: %v", err) } // The file ends at the newline that is the second byte of its first // line. wantStartRefused(t, dir, filepath.Join(dir, "bans.json")+", line 1, column 2: ") } func TestUnwritableStateDirStopsTheStart(t *testing.T) { t.Parallel() wantStartRefused(t, filepath.Join(t.TempDir(), "missing"), "SWWAF_STATE_DIR cannot be written: ") } // wantStartRefused runs smallwebwaf with its state files in dir, and // checks that it stops at start, with an error that starts with want. If // it starts instead, it is stopped after waitLimit. func wantStartRefused(t *testing.T, dir, want string) { t.Helper() ctx, stop := context.WithTimeout(t.Context(), waitLimit) defer stop() out := &output{} status := run(ctx, map[string]string{listenAddr: localhost + ":0", stateDir: dir}, out) if status != 1 { t.Fatalf("exit status %d, want 1", status) } line := out.line(t, "msg", "cannot use the state files") message, _ := line["error"].(string) if !strings.HasPrefix(message, want) { t.Errorf("start refused with %q, want an error starting %q", message, want) } } // startApp starts an app that answers every request with greeting, and // returns its URL. func startApp(t *testing.T) string { t.Helper() app := httptest.NewServer(http.HandlerFunc( func(w http.ResponseWriter, _ *http.Request) { _, _ = io.WriteString(w, greeting) })) t.Cleanup(app.Close) return app.URL } // runUntilStopped runs smallwebwaf with the settings in env, has use send // it requests at url, then stops it as SIGTERM does, checks that it // stopped in order, and returns its output. func runUntilStopped( t *testing.T, env map[string]string, use func(url string), ) *output { t.Helper() ctx, stop := context.WithCancel(t.Context()) out := &output{} exited := make(chan int, 1) go func() { exited <- run(ctx, env, out) }() addr, _ := out.line(t, "msg", "starting")["address"].(string) use("http://" + addr + "/") stop() select { case status := <-exited: if status != 0 { t.Fatalf("exit status %d, want 0; output:\n%s", status, out.text()) } case <-time.After(waitLimit): t.Fatal("still running after being told to stop") } return out } // wantStartingLine checks that the line at start gives the version and // every setting's value. func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) { t.Helper() settings, _ := line["settings"].(map[string]any) want := map[string]any{ listenAddr: localhost + ":0", upstreamURL: appURL, stateDir: dir, "SWWAF_MODE": "enforce", stateWriteDelay: "10s", stateCounterInterval: "15m", trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16", "SWWAF_CLIENT_REQUEST_TIMEOUT": "60s", "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K", "SWWAF_CLIENT_IDLE_TIMEOUT": "120s", "SWWAF_CLIENT_RESPONSE_TIMEOUT": "30m", "SWWAF_UPSTREAM_REQUEST_TIMEOUT": "60s", "SWWAF_UPSTREAM_RESPONSE_TIMEOUT": "30m", "SWWAF_REQUEST_MAX_BYTES": "100M", "SWWAF_RESPONSE_MAX_BYTES": "5G", "SWWAF_ALLOW_NETS": "", "SWWAF_RATE_LIMIT_EXEMPT_NETS": "", "SWWAF_DENY_NETS": "", "SWWAF_RATE_LIMIT_PER_MINUTE": "1000", "SWWAF_RATE_LIMIT_PER_HOUR": "10000", rateLimitPerDay: "50000", "SWWAF_DENIED_COUNTRIES": "", "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "", "SWWAF_BAN_RESPONSE": "403", "SWWAF_LIMIT_BAN_DURATION": "1h", "SWWAF_LIMIT_BAN_REPEAT_WINDOW": "24h", "SWWAF_MAX_BAN_DURATION": "7d", "SWWAF_MAX_BANS": "5000", "SWWAF_BAN_SCOPE_V4_PREFIX": "32", } for name, value := range want { if settings[name] != value { t.Errorf("starting line gives %s=%v, want %v", name, settings[name], value) } } if line["version"] != testVersion || line["type"] != "process" { t.Errorf("starting line %v", line) } } // wantGreeting checks that a request to url gets the app's answer. func wantGreeting(t *testing.T, url string) { t.Helper() req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, http.NoBody) if err != nil { t.Fatalf("new request: %v", err) } transport := &http.Transport{} defer transport.CloseIdleConnections() res, err := (&http.Client{Transport: transport}).Do(req) if err != nil { t.Fatalf("request: %v", err) } body, err := io.ReadAll(res.Body) _ = res.Body.Close() if err != nil || string(body) != greeting { t.Errorf("got %q (%v), want the app's answer", body, err) } } // wantRefused checks that a request to url is refused with 403, the // default SWWAF_BAN_RESPONSE. func wantRefused(t *testing.T, url string) { t.Helper() req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, http.NoBody) if err != nil { t.Fatalf("new request: %v", err) } transport := &http.Transport{} defer transport.CloseIdleConnections() res, err := (&http.Client{Transport: transport}).Do(req) if err != nil { t.Fatalf("request: %v", err) } _ = res.Body.Close() if res.StatusCode != http.StatusForbidden { t.Errorf("status %d, want %d", res.StatusCode, http.StatusForbidden) } } // wantStatus checks that a request to url from the client at from, as // X-Forwarded-For names it, is answered with status. func wantStatus(t *testing.T, url, from string, status int) { t.Helper() got := statusFrom(t, url, from) if got != status { t.Errorf("request from %s: status %d, want %d", from, got, status) } } // saveUntilAnswered writes content to the state file at path, as an // admin saves an edit of it, until a request to url from the client at // from is answered with status. The file is written again before each // request, since smallwebwaf may not watch its directory yet when it is // first written. It waits as long as that takes, so that a slow test // process cannot fail the test. func saveUntilAnswered(t *testing.T, path, content, url, from string, status int) { t.Helper() for { err := os.WriteFile(path, []byte(content), 0o600) if err != nil { t.Fatalf("write %s: %v", path, err) } if statusFrom(t, url, from) == status { return } time.Sleep(pollInterval) } } // statusFrom returns the status a request to url from the client at // from, as X-Forwarded-For names it, is answered with. func statusFrom(t *testing.T, url, from string) int { t.Helper() req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, http.NoBody) if err != nil { t.Fatalf("new request: %v", err) } req.Header.Set("X-Forwarded-For", from) transport := &http.Transport{} defer transport.CloseIdleConnections() res, err := (&http.Client{Transport: transport}).Do(req) if err != nil { t.Fatalf("request: %v", err) } _ = res.Body.Close() return res.StatusCode }