Keep the bans, the clients and GeoJS's answers in state files (closes #17)
check / check (push) Successful in 3m53s
check / check (push) Successful in 3m53s
smallwebwaf now copies its state to bans.json, clients.json and lookups.json in SWWAF_STATE_DIR, as "Persistent state" in SPEC.md describes, and reads them back at start, so a restart lifts no ban and gives no client a fresh allowance. Each client gains a history, and a ban's notes count the netblock's requests. bans.json is written SWWAF_STATE_WRITE_DELAY after a ban, every file every SWWAF_STATE_COUNTER_INTERVAL and at the stop, each through a synced temporary file renamed over it. A file that does not parse, an unknown version or an unwritable directory stops the start. The image gets /var/lib/smallwebwaf, which the run script gives to the smallwebwaf user. Deviation: no AS number or name, and no ban cause, reason or lifting yet. Model: opus-5-5
This commit is contained in:
@@ -0,0 +1,380 @@
|
||||
// Package state keeps smallwebwaf's state in JSON files in
|
||||
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
|
||||
// bans.json holds the bans, clients.json each client's counters and
|
||||
// history, and lookups.json GeoJS's answers. Load reads them at start, and
|
||||
// Run and WriteAll write them, each from a snapshot its part takes under
|
||||
// its own lock, so that no request waits on the disk.
|
||||
package state
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
)
|
||||
|
||||
// version is the version of the files' format, the only one read.
|
||||
const version = 1
|
||||
|
||||
// fileMode lets the smallwebwaf user alone read and write the files, which
|
||||
// hold visitors' addresses.
|
||||
const fileMode = 0o600
|
||||
|
||||
// The state files' names.
|
||||
const (
|
||||
bansJSON = "bans.json"
|
||||
clientsJSON = "clients.json"
|
||||
lookupsJSON = "lookups.json"
|
||||
)
|
||||
|
||||
var errVersion = errors.New("unknown version")
|
||||
|
||||
// Params are what Load needs.
|
||||
type Params struct {
|
||||
// Dir is the directory of the state files (SWWAF_STATE_DIR).
|
||||
Dir string
|
||||
// WriteDelay is how long after a ban is made bans.json is written
|
||||
// (SWWAF_STATE_WRITE_DELAY), and CounterInterval how often every file
|
||||
// is (SWWAF_STATE_COUNTER_INTERVAL).
|
||||
WriteDelay time.Duration
|
||||
CounterInterval time.Duration
|
||||
// Ledger, Limiter and GeoJS hold the state.
|
||||
Ledger *bans.Ledger
|
||||
Limiter *ratelimit.Limiter
|
||||
GeoJS *lookup.GeoJS
|
||||
// Now tells the time by which the counters' buckets run out, normally
|
||||
// time.Now in UTC.
|
||||
Now func() time.Time
|
||||
// ProcessLog receives what was read, and the writes that fail.
|
||||
ProcessLog *slog.Logger
|
||||
}
|
||||
|
||||
// Files are the state files of a running smallwebwaf.
|
||||
type Files struct {
|
||||
params Params
|
||||
}
|
||||
|
||||
// bansFile is bans.json, indented for an admin to read and edit.
|
||||
type bansFile struct {
|
||||
Version int `json:"version"`
|
||||
Bans []banEntry `json:"bans"`
|
||||
}
|
||||
|
||||
// banEntry is a ban as bans.json holds it: a permanent ban's expires is
|
||||
// null.
|
||||
type banEntry struct {
|
||||
Netblock netip.Prefix `json:"netblock"`
|
||||
Start time.Time `json:"start"`
|
||||
Expires *time.Time `json:"expires"`
|
||||
Notes bans.Notes `json:"notes"`
|
||||
}
|
||||
|
||||
// clientsFile is clients.json, with each client on a line of its own.
|
||||
type clientsFile struct {
|
||||
Version int `json:"version"`
|
||||
Clients []ratelimit.Client `json:"clients"`
|
||||
}
|
||||
|
||||
// lookupsFile is lookups.json, with each answer on a line of its own.
|
||||
type lookupsFile struct {
|
||||
Version int `json:"version"`
|
||||
Lookups []lookup.Answer `json:"lookups"`
|
||||
}
|
||||
|
||||
// Load checks that files can be written in Dir, and reads the state files
|
||||
// in it into the ledger, the limiter and GeoJS. A missing file is empty
|
||||
// state, as on a first start. A file that does not parse, or has an
|
||||
// unknown version, is an error that names the file and, where the JSON
|
||||
// decoder tells it, the line and column.
|
||||
func Load(params Params) (*Files, error) {
|
||||
err := checkWritable(params.Dir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SWWAF_STATE_DIR cannot be written: %w", err)
|
||||
}
|
||||
|
||||
var (
|
||||
bansIn bansFile
|
||||
clientsIn clientsFile
|
||||
lookupsIn lookupsFile
|
||||
)
|
||||
|
||||
err = errors.Join(
|
||||
read(params.Dir, bansJSON, &bansIn),
|
||||
read(params.Dir, clientsJSON, &clientsIn),
|
||||
read(params.Dir, lookupsJSON, &lookupsIn),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
held := make([]bans.Ban, 0, len(bansIn.Bans))
|
||||
for _, entry := range bansIn.Bans {
|
||||
held = append(held, entry.ban())
|
||||
}
|
||||
|
||||
params.Ledger.Load(held)
|
||||
params.Limiter.Load(clientsIn.Clients, params.Now())
|
||||
params.GeoJS.Load(lookupsIn.Lookups)
|
||||
|
||||
params.ProcessLog.Info("read the state files", "directory", params.Dir,
|
||||
"bans", len(bansIn.Bans), "clients", len(clientsIn.Clients),
|
||||
"lookups", len(lookupsIn.Lookups))
|
||||
|
||||
return &Files{params: params}, nil
|
||||
}
|
||||
|
||||
// Run writes bans.json WriteDelay after a ban is made, with every ban
|
||||
// made in between, and every file every CounterInterval, until ctx is
|
||||
// done. A write that fails is logged, and the file is written again at
|
||||
// its next write.
|
||||
func (f *Files) Run(ctx context.Context) {
|
||||
interval := time.NewTicker(f.params.CounterInterval)
|
||||
defer interval.Stop()
|
||||
|
||||
var bansDue <-chan time.Time // nil while no ban waits to be written
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-f.params.Ledger.Changed():
|
||||
if bansDue == nil {
|
||||
bansDue = time.After(f.params.WriteDelay)
|
||||
}
|
||||
case <-bansDue:
|
||||
bansDue = nil
|
||||
|
||||
f.logFailure(f.writeBans())
|
||||
case <-interval.C:
|
||||
f.logFailure(f.WriteAll())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// WriteAll writes every state file, as smallwebwaf stops. A file that
|
||||
// fails does not keep the others from being written.
|
||||
func (f *Files) WriteAll() error {
|
||||
return errors.Join(f.writeBans(), f.writeClients(), f.writeLookups())
|
||||
}
|
||||
|
||||
// logFailure logs a write that failed.
|
||||
func (f *Files) logFailure(err error) {
|
||||
if err != nil {
|
||||
f.params.ProcessLog.Error("writing the state files failed",
|
||||
"error", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
// writeBans writes bans.json.
|
||||
func (f *Files) writeBans() error {
|
||||
held := f.params.Ledger.Snapshot()
|
||||
|
||||
file := bansFile{Version: version, Bans: make([]banEntry, 0, len(held))}
|
||||
for _, ban := range held {
|
||||
file.Bans = append(file.Bans, newBanEntry(ban))
|
||||
}
|
||||
|
||||
data, err := json.MarshalIndent(file, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("encode %s: %w", bansJSON, err)
|
||||
}
|
||||
|
||||
return write(f.params.Dir, bansJSON, append(data, '\n'))
|
||||
}
|
||||
|
||||
// writeClients writes clients.json.
|
||||
func (f *Files) writeClients() error {
|
||||
data, err := encodeOnePerLine("clients", f.params.Limiter.Snapshot())
|
||||
if err != nil {
|
||||
return fmt.Errorf("encode %s: %w", clientsJSON, err)
|
||||
}
|
||||
|
||||
return write(f.params.Dir, clientsJSON, data)
|
||||
}
|
||||
|
||||
// writeLookups writes lookups.json.
|
||||
func (f *Files) writeLookups() error {
|
||||
data, err := encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
|
||||
if err != nil {
|
||||
return fmt.Errorf("encode %s: %w", lookupsJSON, err)
|
||||
}
|
||||
|
||||
return write(f.params.Dir, lookupsJSON, data)
|
||||
}
|
||||
|
||||
// newBanEntry returns ban as bans.json holds it.
|
||||
func newBanEntry(ban bans.Ban) banEntry {
|
||||
entry := banEntry{Netblock: ban.Netblock, Start: ban.Start, Notes: ban.Notes}
|
||||
if !ban.Permanent() {
|
||||
entry.Expires = &ban.Expires
|
||||
}
|
||||
|
||||
return entry
|
||||
}
|
||||
|
||||
// ban returns the ban an entry of bans.json holds.
|
||||
func (e banEntry) ban() bans.Ban {
|
||||
ban := bans.Ban{Netblock: e.Netblock, Start: e.Start, Notes: e.Notes}
|
||||
if e.Expires != nil {
|
||||
ban.Expires = *e.Expires
|
||||
}
|
||||
|
||||
return ban
|
||||
}
|
||||
|
||||
// encodeOnePerLine encodes a state file whose entries, under key, are one
|
||||
// to a line, so that grep shows everything about one client.
|
||||
func encodeOnePerLine[E any](key string, entries []E) ([]byte, error) {
|
||||
var b bytes.Buffer
|
||||
|
||||
fmt.Fprintf(&b, "{\n \"version\": %d,\n %q: [", version, key)
|
||||
|
||||
for i, entry := range entries {
|
||||
line, err := json.Marshal(entry)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if i > 0 {
|
||||
b.WriteString(",")
|
||||
}
|
||||
|
||||
b.WriteString("\n ")
|
||||
b.Write(line)
|
||||
}
|
||||
|
||||
b.WriteString("\n ]\n}\n")
|
||||
|
||||
return b.Bytes(), nil
|
||||
}
|
||||
|
||||
// checkWritable makes a file in dir and removes it again.
|
||||
func checkWritable(dir string) error {
|
||||
file, err := os.CreateTemp(dir, "write-check-*")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return errors.Join(file.Close(), os.Remove(file.Name()))
|
||||
}
|
||||
|
||||
// read reads the state file name in dir into file, a pointer to that
|
||||
// file's struct. A missing file leaves file as it is.
|
||||
func read(dir, name string, file any) error {
|
||||
path := filepath.Join(dir, name)
|
||||
|
||||
data, err := os.ReadFile(path) //nolint:gosec // a state file, in SWWAF_STATE_DIR
|
||||
if errors.Is(err, fs.ErrNotExist) {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// The version is read first, so that a file of another version is
|
||||
// refused for that, and not for an entry this version cannot read.
|
||||
var header struct {
|
||||
Version int `json:"version"`
|
||||
}
|
||||
|
||||
err = json.Unmarshal(data, &header)
|
||||
if err == nil && header.Version != version {
|
||||
err = fmt.Errorf("%w %d, where this smallwebwaf reads version %d",
|
||||
errVersion, header.Version, version)
|
||||
}
|
||||
|
||||
if err == nil {
|
||||
decoder := json.NewDecoder(bytes.NewReader(data))
|
||||
// A field this version does not know is most likely misspelt, and
|
||||
// its value would be lost without a word.
|
||||
decoder.DisallowUnknownFields()
|
||||
err = decoder.Decode(file)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s%s: %w", path, position(data, err), err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// position returns where in data err was found, as ", line L, column C"
|
||||
// of the last byte the JSON decoder read, or "" when err does not tell.
|
||||
func position(data []byte, err error) string {
|
||||
var (
|
||||
syntaxErr *json.SyntaxError
|
||||
typeErr *json.UnmarshalTypeError
|
||||
read int64
|
||||
)
|
||||
|
||||
switch {
|
||||
case errors.As(err, &syntaxErr):
|
||||
read = syntaxErr.Offset
|
||||
case errors.As(err, &typeErr):
|
||||
read = typeErr.Offset
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
|
||||
before := data[:max(min(read, int64(len(data)))-1, 0)]
|
||||
line := bytes.Count(before, []byte("\n")) + 1
|
||||
column := len(before) - bytes.LastIndexByte(before, '\n')
|
||||
|
||||
return fmt.Sprintf(", line %d, column %d", line, column)
|
||||
}
|
||||
|
||||
// write writes data to the file name in dir so that a crash at any
|
||||
// moment leaves either the old file or the new one, whole: data goes to a
|
||||
// temporary file in the same directory, which is synced and renamed over
|
||||
// name, and then the directory is synced, so that the rename lasts.
|
||||
func write(dir, name string, data []byte) error {
|
||||
path := filepath.Join(dir, name)
|
||||
temporary := path + ".tmp"
|
||||
|
||||
err := writeSynced(temporary, data)
|
||||
if err != nil {
|
||||
_ = os.Remove(temporary)
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
err = os.Rename(temporary, path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
directory, err := os.Open(dir) //nolint:gosec // SWWAF_STATE_DIR itself
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return errors.Join(directory.Sync(), directory.Close())
|
||||
}
|
||||
|
||||
// writeSynced writes data to the file at path, and syncs it to the disk.
|
||||
func writeSynced(path string, data []byte) error {
|
||||
//nolint:gosec // a state file's temporary file, in SWWAF_STATE_DIR
|
||||
file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, fileMode)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, err = file.Write(data)
|
||||
if err == nil {
|
||||
err = file.Sync()
|
||||
}
|
||||
|
||||
return errors.Join(err, file.Close())
|
||||
}
|
||||
@@ -0,0 +1,508 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user