package state_test import ( "context" "encoding/json" "log/slog" "net/netip" "os" "path/filepath" "slices" "strings" "testing" "testing/synctest" "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 ( // 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() wantRefused(t, tc.file, tc.content, tc.want) }) } } func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) { t.Parallel() const ( // The other fields each entry needs. ban = `"start": "2026-10-06T00:00:00Z", "expires": null` answer = `"answered": "2026-10-06T00:00:00Z"` noNetblock = `: entry 1 has no "netblock"` ) for _, tc := range []struct { name, file, content string // want is what the error says after the file's path. want string }{ { "a ban without a netblock", bansJSON, `{"version": 1, "bans": [{` + ban + `}]}`, noNetblock, }, { "a ban whose netblock is null", bansJSON, `{"version": 1, "bans": [{"netblock": null, ` + ban + `}]}`, noNetblock, }, { "a ban whose netblock is empty", bansJSON, `{"version": 1, "bans": [{"netblock": "", ` + ban + `}]}`, noNetblock, }, { "a ban without a start", bansJSON, `{"version": 1, "bans": [{"netblock": "203.0.113.9/32", "expires": null}]}`, `: entry 1 has no "start"`, }, { // The first ban's expires is null, as a permanent ban's is. "a ban without an expires", bansJSON, `{"version": 1, "bans": [{"netblock": "203.0.113.9/32", ` + ban + `}, ` + `{"netblock": "203.0.113.10/32", "start": "2026-10-06T00:00:00Z"}]}`, `: entry 2 has no "expires"`, }, { "a client without its address", clientsJSON, `{"version": 1, "clients": [{"history": {"requests": 3}}]}`, `: entry 1 has no "client"`, }, { "a client with requests in a window without its start", clientsJSON, `{"version": 1, "clients": [{"client": "203.0.113.9/32", ` + `"hour": {"current": 3}}]}`, `: entry 1 has no "hour.start"`, }, { "an answer without a client", lookupsJSON, `{"version": 1, "lookups": [{"country": "DE", ` + answer + `}]}`, `: entry 1 has no "client"`, }, { // A country of "" is a client GeoJS cannot place. "an answer without a country", lookupsJSON, `{"version": 1, "lookups": [{"client": "192.0.2.1/32", "country": "", ` + answer + `}, {"client": "203.0.113.9/32", ` + answer + `}]}`, `: entry 2 has no "country"`, }, { "an answer without the time GeoJS gave it", lookupsJSON, `{"version": 1, "lookups": [{"client": "203.0.113.9/32", "country": "DE"}]}`, `: entry 1 has no "answered"`, }, } { t.Run(tc.name, func(t *testing.T) { t.Parallel() wantRefused(t, tc.file, tc.content, 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() wantRefused(t, file, content, ": unknown version ") }) } } } 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) } } } // The two tests below run Run in a synctest bubble, where time is a clock // of the test's own: time.Sleep moves it on at once, and synctest.Wait // returns once Run waits for its next write, so that every write due by // then is on disk. func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) { t.Parallel() synctest.Test(t, func(t *testing.T) { dir := t.TempDir() params := newParams(dir) params.WriteDelay = 10 * time.Second run(t, load(t, params)) // A second ban, made while the first waits to be written, puts the // write off no further, and is written with it. first := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(), bans.Notes{}) time.Sleep(5 * time.Second) second := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"), midnight(), bans.Notes{}) time.Sleep(5*time.Second - time.Nanosecond) synctest.Wait() wantFiles(t, dir) time.Sleep(time.Nanosecond) synctest.Wait() wantFiles(t, dir, bansJSON) read := newParams(dir) load(t, read) want := []bans.Ban{first, second} if got := read.Ledger.Snapshot(); !slices.Equal(got, want) { t.Errorf("bans.json holds %+v, want %+v", got, want) } // That write was the only one: bans.json is not written again for // the second ban. The other files wait for the interval, an hour // away. removeFiles(t, dir, bansJSON) time.Sleep(params.WriteDelay) synctest.Wait() wantFiles(t, dir) }) } func TestEveryFileWrittenEveryCounterInterval(t *testing.T) { t.Parallel() synctest.Test(t, func(t *testing.T) { dir := t.TempDir() params := newParams(dir) params.CounterInterval = time.Minute run(t, load(t, params)) // The files are removed once written, so that each interval shows // them written again. for range 3 { time.Sleep(time.Minute - time.Nanosecond) synctest.Wait() wantFiles(t, dir) time.Sleep(time.Nanosecond) synctest.Wait() wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) removeFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) } }) } 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") } } func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) { t.Parallel() dir := t.TempDir() files := load(t, newParams(dir)) // A directory named bans.json cannot be renamed over. err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700) if err != nil { t.Fatalf("mkdir: %v", err) } err = files.WriteAll() if err == nil { t.Error("writing over a directory did not fail") } wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) } // 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) } // wantRefused writes content to the state file named file in a new // directory, and checks that Load refuses it with an error that is the // file's path and then starts with want. func wantRefused(t *testing.T, file, content, want string) { t.Helper() 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+want) { t.Errorf("error %v, want one starting %s%s", err, path, want) } } // removeFiles removes the named files from dir. func removeFiles(t *testing.T, dir string, names ...string) { t.Helper() for _, name := range names { err := os.Remove(filepath.Join(dir, name)) if err != nil { t.Fatalf("remove: %v", err) } } } // 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) } } }