Files
smallwebwaf/internal/smallwebwaf/smallwebwaf_test.go
T
clawbot c87bcca6e6
check / check (push) Waiting to run
Take in an admin's edits of the state files while running (closes #68)
smallwebwaf watches SWWAF_STATE_DIR with fsnotify and takes in a saved
edit of a state file in place of what it held. It knows its own writes
by the SHA-256 of what it last read or wrote; each write first takes in
an edit made since. An edit that does not parse is renamed to
<name>.bad at the next write. Each edit taken in or set aside is logged
and counted. Every ban on a netblock is checked, and the next ban is
worked out from the one that ended last. README.md says how to add and
lift a ban.

Judgement call: a broken edit is set aside at the next write, since an
editor's file can be read half written.

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

571 lines
15 KiB
Go

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
}