package state_test import ( "context" "encoding/json" "io/fs" "log/slog" "net" "net/http" "net/http/httptest" "net/netip" "os" "path/filepath" "slices" "strconv" "strings" "testing" "testing/synctest" "time" "sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/metrics" "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" // What the process log says once Watch watches the directory, and as // it takes in an edit. watching = "watching the state files for edits" tookIn = "took in an edit of a state file" // maxLogLines is how many lines of the process log wait for a test to // read them. maxLogLines = 64 ) // 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 three 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).Run) // 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).Run) // 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 TestEditJustBeforeAScheduledWriteSurvivesIt(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).Run) // A ban, and bans.json written with it. params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(), bans.Notes{}) time.Sleep(params.WriteDelay) synctest.Wait() // A second ban is to be written WriteDelay later. Just before // then, an admin saves bans.json with the first ban lifted and // another added. params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"), midnight(), bans.Notes{}) time.Sleep(params.WriteDelay - time.Nanosecond) synctest.Wait() edit(t, dir, bansJSON, permanentBansJSON) // The write takes the edit in first, and writes it back. The second // ban, made after the admin opened the file, is lost, as "Edits // while running" in SPEC.md says. time.Sleep(time.Nanosecond) synctest.Wait() if got := readFile(t, filepath.Join(dir, bansJSON)); got != permanentBansJSON { t.Errorf("bans.json holds\n%s\nwant the edit", got) } wantEqual(t, bansJSON, params.Ledger.Snapshot(), []bans.Ban{permanentBan()}) }) } 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 TestWritesAreCountedInTheMetrics(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) } const ( ofBans = `{file="bans.json"}` ofClients = `{file="clients.json"}` ) got := scrape(t, params) wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofBans, 1) wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofClients, 1) wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofBans, 0) wantMetric(t, got, "smallwebwaf_state_file_size_bytes"+ofBans, float64(len(permanentBansJSON))) written := metric(t, got, "smallwebwaf_state_file_last_write_timestamp_seconds"+ofBans) if written < float64(time.Now().Add(-time.Hour).Unix()) { t.Errorf("bans.json was last written at %v, not by that write", written) } // A directory in the way of bans.json's temporary file fails its next // write, which leaves its size as it was, although it has a ban more. 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{}) err = files.WriteAll() if err == nil { t.Fatal("the write did not fail") } got = scrape(t, params) wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofBans, 2) wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofBans, 1) wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofClients, 0) wantMetric(t, got, "smallwebwaf_state_file_size_bytes"+ofBans, float64(len(permanentBansJSON))) } func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) { t.Parallel() dir := t.TempDir() path := filepath.Join(dir, bansJSON) files := load(t, newParams(dir)) // bans.json is a socket, which cannot be opened as a file, even by // root, as the tests run in Docker, but which a rename could replace. // Whether it holds an edit cannot be told, so it is left as it is. socket, err := (&net.ListenConfig{}).Listen(t.Context(), "unix", path) if err != nil { t.Fatalf("listen: %v", err) } defer func() { _ = socket.Close() }() err = files.WriteAll() if err == nil { t.Error("writing with bans.json unreadable did not fail") } info, err := os.Lstat(path) if err != nil || info.Mode().Type() != fs.ModeSocket { t.Errorf("bans.json is now %v (%v), want the socket", info, err) } wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) } func TestEditOfEachFileTakenIn(t *testing.T) { t.Parallel() dir := t.TempDir() params := newParams(dir) lines := logInto(¶ms) fill(params) files := load(t, params) err := files.WriteAll() if err != nil { t.Fatalf("write: %v", err) } watch(t, files, lines) // Each edit holds one entry, for a client the parts did not hold, and // takes the place of everything the part held. client := netip.MustParsePrefix("198.51.100.7/32") edit(t, dir, bansJSON, `{"version": 1, "bans": [{"netblock": "198.51.100.7/32", `+ `"start": "2026-10-06T00:00:00Z", "expires": null}]}`) wantTakenIn(t, lines, dir, bansJSON) wantEqual(t, bansJSON, params.Ledger.Snapshot(), []bans.Ban{{Netblock: client, Start: midnight()}}) edit(t, dir, clientsJSON, `{"version": 1, "clients": [`+ `{"client": "198.51.100.7/32", "history": {"requests": 7}}]}`) wantTakenIn(t, lines, dir, clientsJSON) wantEqual(t, clientsJSON, params.Limiter.Snapshot(), []ratelimit.Client{{Client: client, History: ratelimit.History{Requests: 7}}}) edit(t, dir, lookupsJSON, `{"version": 1, "lookups": [{"client": "198.51.100.7/32", `+ `"country": "FR", "answered": "2026-10-06T00:00:00Z"}]}`) wantTakenIn(t, lines, dir, lookupsJSON) wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(), []lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}}) } func TestOwnWritesAreNotTakenIn(t *testing.T) { t.Parallel() dir := t.TempDir() params := newParams(dir) lines := logInto(¶ms) fill(params) files := load(t, params) watch(t, files, lines) // Every file is written while watched, and then lookups.json edited: // the first edit taken in is that one. err := files.WriteAll() if err != nil { t.Fatalf("write: %v", err) } edit(t, dir, lookupsJSON, `{"version": 1, "lookups": []}`) wantTakenIn(t, lines, dir, lookupsJSON) } func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) { t.Parallel() dir := t.TempDir() params := newParams(dir) lines := logInto(¶ms) watch(t, load(t, params), lines) client := netip.MustParseAddr("203.0.113.9") // An entry added, as an admin writes it, bans its netblock. edit(t, dir, bansJSON, `{"version": 1, "bans": [{"netblock": "203.0.113.0/24", `+ `"start": "2026-10-06T00:00:00Z", "expires": null}]}`) wantTakenIn(t, lines, dir, bansJSON) _, banned := params.Ledger.Check(client, midnight()) if !banned { t.Error("the ban added to bans.json does not refuse") } // The entry removed lifts the ban. edit(t, dir, bansJSON, `{"version": 1, "bans": []}`) wantTakenIn(t, lines, dir, bansJSON) _, banned = params.Ledger.Check(client, midnight()) if banned { t.Error("the ban removed from bans.json still refuses") } } func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) { t.Parallel() // It ends a ban's entry with a comma. const broken = "{\n \"version\": 1,\n \"bans\": [\n" + " {\"netblock\": \"203.0.113.9/32\",}\n ]\n}\n" dir := t.TempDir() path := filepath.Join(dir, bansJSON) params := newParams(dir) lines := logInto(¶ms) params.Ledger.Load([]bans.Ban{permanentBan()}) files := load(t, params) err := files.WriteAll() if err != nil { t.Fatalf("write: %v", err) } watch(t, files, lines) // While smallwebwaf runs, the broken edit is left as it is: an edit // of clients.json, made after it and taken in, shows that it has been // seen. edit(t, dir, bansJSON, broken) edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`) wantTakenIn(t, lines, dir, clientsJSON) wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) // The next write sets it aside, logged with where the error is, and // writes bans.json again from what smallwebwaf still holds. err = files.WriteAll() if err != nil { t.Fatalf("write: %v", err) } line := lines.waitFor(t, "set aside an edit of a state file that does not parse") message, _ := line["error"].(string) if line["file"] != path+".bad" || !strings.HasPrefix(message, path+", line 4, column 39: ") { t.Errorf("set aside with %v", line) } wantFiles(t, dir, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON) if got := readFile(t, path+".bad"); got != broken { t.Errorf("bans.json.bad holds\n%s\nwant the edit", got) } if got := readFile(t, path); got != permanentBansJSON { t.Errorf("bans.json holds\n%s\nwant\n%s", got, permanentBansJSON) } } func TestEditsTakenInAreCountedInTheMetrics(t *testing.T) { t.Parallel() dir := t.TempDir() params := newParams(dir) lines := logInto(¶ms) files := load(t, params) // One edit is taken in by the write of its file, before Watch runs, // and one by Watch. edit(t, dir, bansJSON, `{"version": 1, "bans": []}`) err := files.WriteAll() if err != nil { t.Fatalf("write: %v", err) } watch(t, files, lines) edit(t, dir, bansJSON, permanentBansJSON) wantTakenIn(t, lines, dir, bansJSON) wantMetric(t, scrape(t, params), `smallwebwaf_state_file_edits_taken_in_total{file="bans.json"}`, 2) } func TestEditsSetAsideAreCountedInTheMetrics(t *testing.T) { t.Parallel() dir := t.TempDir() params := newParams(dir) files := load(t, params) edit(t, dir, bansJSON, `{"version": 1, "bans": [`) err := files.WriteAll() if err != nil { t.Fatalf("write: %v", err) } wantMetric(t, scrape(t, params), `smallwebwaf_state_file_edits_set_aside_total{file="bans.json"}`, 1) } func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) { t.Parallel() dir := t.TempDir() params := newParams(dir) lines := logInto(¶ms) files := load(t, params) err := os.Remove(dir) if err != nil { t.Fatalf("remove: %v", err) } // Watch returns at once. files.Watch(t.Context()) line := lines.waitFor(t, "cannot watch the state files for edits") if line["level"] != "ERROR" { t.Errorf("logged as %v", line) } } // 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) m := metrics.New(1) 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, Metrics: m, }), Now: midnight, ProcessLog: discard, Metrics: m, } } // 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 task, the Run or the Watch of state files, until the test // ends. func run(t *testing.T, task func(context.Context)) { t.Helper() ctx, stop := context.WithCancel(t.Context()) stopped := make(chan struct{}) go func() { task(ctx) close(stopped) }() t.Cleanup(func() { stop() <-stopped }) } // watch runs files' Watch until the test ends, and waits until it // watches the directory. func watch(t *testing.T, files *state.Files, lines processLog) { t.Helper() run(t, files.Watch) lines.waitFor(t, watching) } // processLog receives the lines of a process log, each a JSON object, for // a test to wait for. type processLog chan string // logInto has params' process log write its lines into a new processLog, // and returns that. func logInto(params *state.Params) processLog { lines := make(processLog, maxLogLines) params.ProcessLog = slog.New(slog.NewJSONHandler(lines, nil)) return lines } // Write receives a line of the process log. func (l processLog) Write(line []byte) (int, error) { l <- string(line) return len(line), nil } // waitFor returns the next line of the process log whose message is msg, // passing over the lines before it. It waits as long as that takes, so // that a slow test process cannot fail the test. func (l processLog) waitFor(t *testing.T, msg string) map[string]any { t.Helper() for line := range l { var fields map[string]any err := json.Unmarshal([]byte(line), &fields) if err != nil { t.Fatalf("process log line %q is not JSON: %v", line, err) } if fields["msg"] == msg { return fields } } return nil // never reached: nothing closes the log } // wantTakenIn waits for the next edit taken in, and checks that it is of // the state file name in dir. func wantTakenIn(t *testing.T, lines processLog, dir, name string) { t.Helper() line := lines.waitFor(t, tookIn) if line["file"] != filepath.Join(dir, name) { t.Fatalf("took in %v, want an edit of %s", line, name) } } // edit writes content to the state file name in dir, as an admin saves an // edit of it. func edit(t *testing.T, dir, name, content string) { t.Helper() err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o600) if err != nil { t.Fatalf("write %s: %v", name, err) } } // 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) } } } // scrape returns the metrics of params, in the Prometheus text format. func scrape(t *testing.T, params state.Params) string { t.Helper() recorder := httptest.NewRecorder() params.Metrics.ServeHTTP(recorder, httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", http.NoBody)) if recorder.Code != http.StatusOK { t.Fatalf("the metrics were answered %d", recorder.Code) } return recorder.Body.String() } // metric returns the value of series in text, the metrics, such as // smallwebwaf_state_file_writes_total{file="bans.json"}, or fails the test // if there is no such series. func metric(t *testing.T, text, series string) float64 { t.Helper() for line := range strings.Lines(text) { value, found := strings.CutPrefix(strings.TrimSuffix(line, "\n"), series+" ") if !found { continue } number, err := strconv.ParseFloat(value, 64) if err != nil { t.Fatalf("%s has the value %q", series, value) } return number } t.Fatalf("no series %s in the metrics:\n%s", series, text) return 0 } // wantMetric checks the value of series in text, the metrics, as metric // reads it. func wantMetric(t *testing.T, text, series string, want float64) { t.Helper() got := metric(t, text, series) if got != want { t.Errorf("%s is %v, want %v", series, got, want) } }