package state_test import ( "context" "encoding/json" "log/slog" "net/netip" "os" "path/filepath" "slices" "strings" "testing" "time" "sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/state" ) const ( // waitLimit bounds how long a test waits for what should happen. waitLimit = 10 * time.Second // pollInterval is how often a test looks again. pollInterval = 10 * time.Millisecond // The state files. bansJSON = "bans.json" clientsJSON = "clients.json" lookupsJSON = "lookups.json" ) // permanentBansJSON is bans.json holding permanentBan. const permanentBansJSON = `{ "version": 1, "bans": [ { "netblock": "2001:db8::/64", "start": "2026-10-06T00:00:00Z", "expires": null, "notes": { "country": "DE", "limit": 1000, "window": "minute", "count": 1000.5, "request": { "time": "2026-10-06T00:00:00Z", "method": "GET", "host": "app.example", "path": "/repo?page=2", "status": 403, "user_agent": "scraper/1.0" }, "requests": 1500, "refused": 3, "earlier_bans": 5 } } ] } ` func TestFilesWrittenAndReadBack(t *testing.T) { t.Parallel() dir := t.TempDir() before := newParams(dir) fill(before) files, err := state.Load(before) if err != nil { t.Fatalf("load: %v", err) } err = files.WriteAll() if err != nil { t.Fatalf("write: %v", err) } // Read into new parts, as at the next start, the files give back what // was written. after := newParams(dir) load(t, after) wantEqual(t, bansJSON, after.Ledger.Snapshot(), before.Ledger.Snapshot()) wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot()) wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot()) // Each one-per-line file lists its entries by client, and nothing // but the three files is left in the directory. wantEntries(t, filepath.Join(dir, clientsJSON), "clients", "192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64") wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups", "192.0.2.1/32", "203.0.113.9/32") wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) } func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) { t.Parallel() dir := t.TempDir() params := newParams(dir) params.Ledger.Load([]bans.Ban{permanentBan()}) files := load(t, params) err := files.WriteAll() if err != nil { t.Fatalf("write: %v", err) } got := readFile(t, filepath.Join(dir, bansJSON)) if got != permanentBansJSON { t.Errorf("bans.json\n%s\nwant\n%s", got, permanentBansJSON) } } func TestMissingFilesAreEmptyState(t *testing.T) { t.Parallel() params := newParams(t.TempDir()) load(t, params) if len(params.Ledger.Snapshot()) != 0 || len(params.Limiter.Snapshot()) != 0 || len(params.GeoJS.Snapshot()) != 0 { t.Error("state from no files") } } func TestFileThatDoesNotParseStopsTheStart(t *testing.T) { t.Parallel() for _, tc := range []struct { name, file, content string // want is what the error says after the file's path. want string }{ { "a syntax error", bansJSON, "{\n \"version\": 1,\n \"bans\": [\n" + " {\"netblock\": \"203.0.113.9/32\",}\n ]\n}\n", ", line 4, column 39: invalid character '}'", }, { "a value of the wrong kind", clientsJSON, "{\n \"version\": 1,\n \"clients\": [\n" + " {\"client\":\"203.0.113.9/32\",\"history\":{\"requests\":\"many\"}}\n" + " ]\n}\n", ", line 4, column ", }, { // Found at the newline that ends the file. "a cut-off file", lookupsJSON, "{\n \"version\": 1,\n \"lookups\": [\n", ", line 3, column 17: unexpected end of JSON input", }, { "an unknown field", lookupsJSON, `{"version": 1, "lookups": [{"client": "203.0.113.9/32", "contry": "DE"}]}`, `: json: unknown field "contry"`, }, { "a netblock that does not read", bansJSON, `{"version": 1, "bans": [{"netblock": "203.0.113.300/32"}]}`, `: netip.ParsePrefix("203.0.113.300/32")`, }, } { t.Run(tc.name, func(t *testing.T) { t.Parallel() dir := t.TempDir() path := filepath.Join(dir, tc.file) err := os.WriteFile(path, []byte(tc.content), 0o600) if err != nil { t.Fatalf("write %s: %v", tc.file, err) } _, err = state.Load(newParams(dir)) if err == nil || !strings.HasPrefix(err.Error(), path+tc.want) { t.Errorf("error %v, want one starting %s%s", err, path, tc.want) } }) } } func TestUnknownVersionStopsTheStart(t *testing.T) { t.Parallel() for _, file := range []string{bansJSON, clientsJSON, lookupsJSON} { for _, content := range []string{`{"version": 2}`, `{}`} { t.Run(file+" "+content, func(t *testing.T) { t.Parallel() dir := t.TempDir() path := filepath.Join(dir, file) err := os.WriteFile(path, []byte(content), 0o600) if err != nil { t.Fatalf("write %s: %v", file, err) } _, err = state.Load(newParams(dir)) if err == nil || !strings.HasPrefix(err.Error(), path+": unknown version ") { t.Errorf("error %v, want one naming %s and its version", err, path) } }) } } } func TestUnwritableDirectoryStopsTheStart(t *testing.T) { t.Parallel() notADirectory := filepath.Join(t.TempDir(), "file") err := os.WriteFile(notADirectory, nil, 0o600) if err != nil { t.Fatalf("write: %v", err) } for _, dir := range []string{ filepath.Join(t.TempDir(), "missing"), notADirectory, } { const want = "SWWAF_STATE_DIR cannot be written: " _, err := state.Load(newParams(dir)) if err == nil || !strings.HasPrefix(err.Error(), want) { t.Errorf("state directory %s: error %v, want one starting %s", dir, err, want) } } } func TestBanWrittenAfterTheWriteDelay(t *testing.T) { t.Parallel() dir := t.TempDir() params := newParams(dir) params.WriteDelay = 10 * time.Millisecond run(t, load(t, params)) ban := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(), bans.Notes{Limit: 1}) waitForFile(t, filepath.Join(dir, bansJSON)) read := newParams(dir) load(t, read) if got := read.Ledger.Snapshot(); !slices.Equal(got, []bans.Ban{ban}) { t.Errorf("bans.json holds %+v, want %+v", got, ban) } // The other files wait for the interval, an hour away. wantFiles(t, dir, bansJSON) } func TestEveryFileWrittenEveryCounterInterval(t *testing.T) { t.Parallel() dir := t.TempDir() params := newParams(dir) params.CounterInterval = 10 * time.Millisecond run(t, load(t, params)) for _, file := range []string{bansJSON, clientsJSON, lookupsJSON} { waitForFile(t, filepath.Join(dir, file)) } } func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) { t.Parallel() dir := t.TempDir() params := newParams(dir) params.Ledger.Load([]bans.Ban{permanentBan()}) files := load(t, params) err := files.WriteAll() if err != nil { t.Fatalf("write: %v", err) } // A directory in the way of bans.json's temporary file fails its // next write, but not the others'. err = os.Mkdir(filepath.Join(dir, bansJSON+".tmp"), 0o700) if err != nil { t.Fatalf("mkdir: %v", err) } params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(), bans.Notes{}) params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight()) err = files.WriteAll() if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") { t.Errorf("error %v, want one naming bans.json's temporary file", err) } got := readFile(t, filepath.Join(dir, bansJSON)) if got != permanentBansJSON { t.Errorf("bans.json is now\n%s\nwant it as it was", got) } read := newParams(dir) load(t, read) if len(read.Limiter.Snapshot()) != 1 { t.Error("clients.json was not written") } } // midnight is the time of the tests' clock. func midnight() time.Time { return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC) } // newParams returns Params for the state files in dir, with parts that // hold nothing yet. GeoJS is never asked. func newParams(dir string) state.Params { discard := slog.New(slog.DiscardHandler) return state.Params{ Dir: dir, WriteDelay: time.Hour, CounterInterval: time.Hour, Ledger: bans.New(bans.Rules{ LimitBanDuration: time.Hour, LimitBanRepeatWindow: 24 * time.Hour, MaxBanDuration: 7 * 24 * time.Hour, MaxBans: 5000, }), Limiter: ratelimit.New(ratelimit.Limits{}), GeoJS: lookup.New(lookup.Params{Now: midnight, ProcessLog: discard}), Now: midnight, ProcessLog: discard, } } // fill puts a ban that ends and one that does not, clients with counts // and histories, and GeoJS answers into the parts of params. func fill(params state.Params) { now := midnight() client := netip.MustParsePrefix("203.0.113.9/32") params.Ledger.Load([]bans.Ban{permanentBan()}) params.Ledger.BanForLimit(client, now, bans.Notes{Country: "DE", Limit: 1}) for _, c := range []string{"2001:db8::/64", "203.0.113.9/32", "192.0.2.1/32"} { params.Limiter.Count(netip.MustParsePrefix(c), now) } params.Limiter.AddToHistory(client, now, ratelimit.Request{ Country: "DE", Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5, }) params.GeoJS.Load([]lookup.Answer{ {Client: client, Country: "DE", Answered: now.Add(-time.Hour), Used: now}, { Client: netip.MustParsePrefix("192.0.2.1/32"), Answered: now.Add(-time.Hour), Used: now.Add(-time.Minute), }, }) } // permanentBan is the ban permanentBansJSON holds. func permanentBan() bans.Ban { return bans.Ban{ Netblock: netip.MustParsePrefix("2001:db8::/64"), Start: midnight(), Notes: bans.Notes{ Country: "DE", Limit: 1000, Window: "minute", Count: 1000.5, Request: bans.Request{ Time: midnight(), Method: "GET", Host: "app.example", Path: "/repo?page=2", Status: 403, UserAgent: "scraper/1.0", }, Requests: 1500, Refused: 3, EarlierBans: 5, }, } } // load reads the state files into the parts of params. func load(t *testing.T, params state.Params) *state.Files { t.Helper() files, err := state.Load(params) if err != nil { t.Fatalf("load: %v", err) } return files } // run runs files' writes until the test ends. func run(t *testing.T, files *state.Files) { t.Helper() ctx, stop := context.WithCancel(t.Context()) stopped := make(chan struct{}) go func() { files.Run(ctx) close(stopped) }() t.Cleanup(func() { stop() <-stopped }) } // wantEqual checks that the entries read back from file are those // written. func wantEqual[E comparable](t *testing.T, file string, got, want []E) { t.Helper() if !slices.Equal(got, want) { t.Errorf("%s read back\n%+v\nwant\n%+v", file, got, want) } } // readFile returns what the file at path holds. func readFile(t *testing.T, path string) string { t.Helper() data, err := os.ReadFile(path) //nolint:gosec // a file the test wrote if err != nil { t.Fatalf("read: %v", err) } return string(data) } // waitForFile waits for the file at path to exist. func waitForFile(t *testing.T, path string) { t.Helper() deadline := time.Now().Add(waitLimit) for time.Now().Before(deadline) { _, err := os.Stat(path) if err == nil { return } time.Sleep(pollInterval) } t.Fatalf("no %s after %s", path, waitLimit) } // wantFiles checks the names of the files in dir. func wantFiles(t *testing.T, dir string, want ...string) { t.Helper() entries, err := os.ReadDir(dir) if err != nil { t.Fatalf("read %s: %v", dir, err) } got := make([]string, 0, len(entries)) for _, entry := range entries { got = append(got, entry.Name()) } if !slices.Equal(got, want) { t.Errorf("%s holds %v, want %v", dir, got, want) } } // wantEntries checks that the file at path has its version, then its // entries under key, each on a line of its own, for the clients want // names in that order. func wantEntries(t *testing.T, path, key string, want ...string) { t.Helper() data := readFile(t, path) lines := strings.Split(strings.TrimSuffix(data, "\n"), "\n") head := []string{"{", ` "version": 1,`, ` "` + key + `": [`} tail := []string{" ]", "}"} if len(lines) != len(head)+len(want)+len(tail) || !slices.Equal(lines[:len(head)], head) || !slices.Equal(lines[len(lines)-len(tail):], tail) { t.Fatalf("%s is\n%s", path, data) } for i, client := range want { line := strings.TrimSuffix(lines[len(head)+i], ",") var entry struct { Client string `json:"client"` } err := json.Unmarshal([]byte(line), &entry) if err != nil || entry.Client != client { t.Errorf("entry %d of %s is %s (%v), want %s's", i, path, line, err, client) } } }