Take in an admin's edits of the state files while running (closes #68)
check / check (push) Successful in 3m33s

smallwebwaf watches SWWAF_STATE_DIR with fsnotify and takes in an edit of
bans.json, clients.json or lookups.json as soon as it is saved, in place
of what it held. It tells its own writes from an admin's by the SHA-256
of what it last read or wrote, and each write takes in an edit made since
first. An edit that does not parse is renamed to <name>.bad at the
file's next write, which writes the file again from memory and logs the
file and where the error is. README.md says how to add and lift a ban.

Judgement call: a broken edit is set aside at the file's next write, not
when seen, since an editor's file can be read half written.
Judgement call: a state file that cannot be read is not written over.

Model: opus-5-5
This commit is contained in:
2026-10-06 07:23:39 +00:00
parent df2c5042d2
commit bc0152445a
12 changed files with 813 additions and 154 deletions
+61 -28
View File
@@ -15,17 +15,18 @@ Status: the first two milestones are built
(https://git.eeqj.de/sneak/smallwebwaf/issues/13 and (https://git.eeqj.de/sneak/smallwebwaf/issues/13 and
https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are four parts of https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are four parts of
milestone 3: the static lists, the bans that broken rate limits lead to and the milestone 3: the static lists, the bans that broken rate limits lead to and the
JSON state files, which come next in the build order, and the header size and JSON state files with your edits taken in while it runs, which come next in the
the idle time as settings, which come last in it. `smallwebwaf` passes each build order, and the header size and the idle time as settings, which come last
request to the app and the app's answer back, unchanged, within its timeouts and in it. `smallwebwaf` passes each request to the app and the app's answer back,
size limits, works out each client's address, bans a client that sends too many unchanged, within its timeouts and size limits, works out each client's address,
requests, refuses a client that comes from a country you refuse or from a bans a client that sends too many requests, refuses a client that comes from a
network you refuse, lets the networks you choose through, keeps its bans, each country you refuse or from a network you refuse, lets the networks you choose
client's counters and history, and GeoJS's answers in JSON files across through, keeps its bans, each client's counters and history, and GeoJS's answers
restarts, and writes a JSON log line for every request. It comes as the image in JSON files across restarts, takes in your edits of those files while it runs,
the app's own image is built on. The rest of the design comes after that, in the and writes a JSON log line for every request. It comes as the image the app's
order of the build order in [`SPEC.md`](SPEC.md). The survey of existing tools own image is built on. The rest of the design comes after that, in the order of
that led to the design is in [`EVALUATION.md`](EVALUATION.md). the build order in [`SPEC.md`](SPEC.md). The survey of existing tools that led
to the design is in [`EVALUATION.md`](EVALUATION.md).
## Getting started ## Getting started
@@ -97,9 +98,9 @@ in `bin/state` unless `SWWAF_STATE_DIR` is set.
seen, how many of them the ban has refused, and how many bans the netblock had seen, how many of them the ban has refused, and how many bans the netblock had
before. At most `SWWAF_MAX_BANS` bans are kept, past, active and permanent; before. At most `SWWAF_MAX_BANS` bans are kept, past, active and permanent;
past that, the earliest ban of the netblock that has gone longest without a past that, the earliest ban of the netblock that has gone longest without a
request is dropped first. `bans.json` shows the bans and their notes, and a request is dropped first. `bans.json` shows the bans and their notes, a
restart lifts none (see "State files" below); lifting a ban by editing it restart lifts none, and you add or lift a ban by editing it (see "State files"
comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68. below).
- Refuses a request from a country you refuse with `SWWAF_BAN_RESPONSE`, as soon - Refuses a request from a country you refuse with `SWWAF_BAN_RESPONSE`, as soon
as the client's country is known and before its body is read; such a request as the client's country is known and before its body is read; such a request
is not counted for the rate limits. While one of the country lists below is is not counted for the rate limits. While one of the country lists below is
@@ -295,9 +296,42 @@ without a field it needs, named with the entry's place in the file: a ban's
`netblock`, `start` or `expires`, which is `null` for a permanent ban; a `netblock`, `start` or `expires`, which is `null` for a permanent ban; a
client's `client`, or the `start` of a window in which it has requests; an client's `client`, or the `start` of a window in which it has requests; an
answer's `client`, `country`, which is `""` for a client GeoJS cannot place, or answer's `client`, `country`, which is `""` for a client GeoJS cannot place, or
`answered`. An edit made while `smallwebwaf` runs is overwritten by its next `answered`. The AS number and AS name come with their lookup.
write: taking it in comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68.
The AS number and AS name come with their lookup. While it runs, `smallwebwaf` watches `SWWAF_STATE_DIR` and takes in your edit of
a state file as soon as you save it: what the file then holds replaces what
`smallwebwaf` held for it, as if read at start. It tells its own writes from
yours by comparing the file with what it last read or wrote, and before it
writes a file it takes in any edit made since, so your edit is not overwritten;
a change `smallwebwaf` made after you opened the file, such as a new ban, is
lost when you save over it. An edit that would stop the start, because it does
not parse, has another `version` or leaves out a field an entry needs, does not
stop the running `smallwebwaf`: it keeps what it holds, and at the file's next
write renames your file to `<name>.bad`, such as `bans.json.bad`, writes the
file again from memory, and logs the file and where the error is. It waits for
that write because an editor's file can be read before the editor has finished
writing it. Mend the `.bad` file and move it back. A file you remove is written
again at its next write.
To ban a netblock, add an entry to `bans.json` with its `netblock`, its `start`
and its `expires`, `null` for a ban that never ends; its `notes` may be left
out. This `bans.json` bans `203.0.113.0/24` for good:
```json
{
"version": 1,
"bans": [
{
"netblock": "203.0.113.0/24",
"start": "2026-10-06T12:00:00Z",
"expires": null
}
]
}
```
To lift a ban, delete its entry. `smallwebwaf` then forgets the ban, so it does
not make the netblock's next ban longer.
## Why ## Why
@@ -400,9 +434,9 @@ goes through the candidates one by one.
readable JSON files, written regularly and at every stop, so a restart loses readable JSON files, written regularly and at every stop, so a restart loses
nothing. Edit a file, or add a rule file, and the running `smallwebwaf` picks nothing. Edit a file, or add a rule file, and the running `smallwebwaf` picks
up the change. Nothing is read from disk while serving a request. The files up the change. Nothing is read from disk while serving a request. The files
for the bans, the clients and the GeoJS answers are built (see "State files" for the bans, the clients and the GeoJS answers are built, with an edit taken
above); the others come with their features, and taking in an edit while in while running (see "State files" above); the others come with their
running comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68. features.
- Health checks, the metrics, and listing, adding and lifting bans or asking why - Health checks, the metrics, and listing, adding and lifting bans or asking why
a given address was refused, all on the one port every request uses: under a given address was refused, all on the one port every request uses: under
`/_smallwebwaf/` on the app's own address, through traefik like any other `/_smallwebwaf/` on the app's own address, through traefik like any other
@@ -590,8 +624,8 @@ addresses are never sent to GeoJS.
answers. answers.
- `internal/ratelimit`: the table of clients: counts each client's requests, - `internal/ratelimit`: the table of clients: counts each client's requests,
tells when one takes it over a rate limit, and keeps each client's history. tells when one takes it over a rate limit, and keeps each client's history.
- `internal/state`: reads the state files at start, and writes them when they - `internal/state`: reads the state files at start, takes in an admin's edit of
are due and at the stop. one while running, and writes them when they are due and at the stop.
- `internal/requestlog`: the lines on stdout: the request log line and the - `internal/requestlog`: the lines on stdout: the request log line and the
process's own messages. process's own messages.
- `Dockerfile`: the lint and test phases, then the image, whose last stage - `Dockerfile`: the lint and test phases, then the image, whose last stage
@@ -602,8 +636,9 @@ addresses are never sent to GeoJS.
Besides the Go standard library, `github.com/hashicorp/golang-lru/v2` keeps the Besides the Go standard library, `github.com/hashicorp/golang-lru/v2` keeps the
table of clients to 20,000, the GeoJS answers to 100,000 and the banned table of clients to 20,000, the GeoJS answers to 100,000 and the banned
netblocks to `SWWAF_MAX_BANS`, dropping the least recently seen. The country netblocks to `SWWAF_MAX_BANS`, dropping the least recently seen, and
codes are the list in `internal/config/config.go`. `github.com/fsnotify/fsnotify` tells `smallwebwaf` when a state file is saved.
The country codes are the list in `internal/config/config.go`.
## Entrypoints ## Entrypoints
@@ -646,10 +681,8 @@ so that they run in minimal containers.
## TODO ## TODO
- The rest of milestone 3, from taking in an admin's edits to the state files - The rest of milestone 3, from exemptions up to the metrics endpoint, and the
(https://git.eeqj.de/sneak/smallwebwaf/issues/68) up to the metrics endpoint, rest of the design, in the order of the build order in [`SPEC.md`](SPEC.md).
and the rest of the design, in the order of the build order in
[`SPEC.md`](SPEC.md).
## Documents ## Documents
+6 -1
View File
@@ -2,4 +2,9 @@ module sneak.berlin/go/smallwebwaf
go 1.26.0 go 1.26.0
require github.com/hashicorp/golang-lru/v2 v2.0.7 require (
github.com/fsnotify/fsnotify v1.10.1
github.com/hashicorp/golang-lru/v2 v2.0.7
)
require golang.org/x/sys v0.13.0 // indirect
+4
View File
@@ -1,2 +1,6 @@
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
golang.org/x/sys v0.13.0 h1:Af8nKPmuFypiUBjVoU9V20FiaFXOcuZI21p0ycVYYGE=
golang.org/x/sys v0.13.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
+13 -9
View File
@@ -247,21 +247,25 @@ func (l *Ledger) Snapshot() []Ban {
return held return held
} }
// Load puts bans read from bans.json into a ledger that holds none yet, // Load puts bans read from bans.json into the ledger, in place of the
// in the order they started, so that a netblock whose last ban started // bans it holds, in the order they started, so that a netblock whose last
// latest counts as the most recently seen. Each netblock is masked to its // ban started latest counts as the most recently seen. Each netblock is
// length, so that 203.0.113.9/24 is 203.0.113.0/24, and each text in the // masked to its length, so that 203.0.113.9/24 is 203.0.113.0/24, and
// notes is cut to 256 bytes. Past MaxBans the earliest bans are dropped, // each text in the notes is cut to 256 bytes. Past MaxBans the earliest
// as when they are made. // bans are dropped, as when they are made.
func (l *Ledger) Load(bans []Ban) { func (l *Ledger) Load(bans []Ban) {
l.mu.Lock()
defer l.mu.Unlock()
bans = slices.Clone(bans) bans = slices.Clone(bans)
slices.SortStableFunc(bans, func(a, b Ban) int { slices.SortStableFunc(bans, func(a, b Ban) int {
return a.Start.Compare(b.Start) return a.Start.Compare(b.Start)
}) })
l.mu.Lock()
defer l.mu.Unlock()
l.netblocks.Purge()
l.held = 0
l.v4Lengths, l.v6Lengths = nil, nil
for _, ban := range bans { for _, ban := range bans {
ban.Netblock = ban.Netblock.Masked() ban.Netblock = ban.Netblock.Masked()
ban.Notes.Request = ban.Notes.Request.cut() ban.Notes.Request = ban.Notes.Request.cut()
+29
View File
@@ -151,6 +151,35 @@ func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
} }
} }
func TestLoadReplacesTheBansHeld(t *testing.T) {
t.Parallel()
rules := defaultRules()
rules.MaxBans = 2
ledger := bans.New(rules)
kept := bans.Ban{Netblock: netip.MustParsePrefix("2001:db8::/64"), Start: midnight()}
ledger.Load([]bans.Ban{
{Netblock: netip.MustParsePrefix("203.0.113.0/24"), Start: midnight()},
kept,
})
// Loaded again without the first ban, as when an admin's edit of
// bans.json is taken in: that ban is lifted, and the ledger, holding
// one ban, makes another without dropping any.
ledger.Load([]bans.Ban{kept})
made := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(),
bans.Notes{})
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight())
if banned {
t.Error("a ban left out of the second load still refuses")
}
if got, want := ledger.Snapshot(), []bans.Ban{made, kept}; !slices.Equal(got, want) {
t.Errorf("the ledger holds %+v, want %+v", got, want)
}
}
func TestLoadCutsTheTextsTo256Bytes(t *testing.T) { func TestLoadCutsTheTextsTo256Bytes(t *testing.T) {
t.Parallel() t.Parallel()
+7 -5
View File
@@ -188,19 +188,21 @@ func (g *GeoJS) Snapshot() []Answer {
return answers return answers
} }
// Load keeps answers read from lookups.json, in a GeoJS that keeps none // Load keeps answers read from lookups.json, in place of the answers it
// yet, in the order they were last used, so that the one used longest // keeps, in the order they were last used, so that the one used longest
// ago is dropped first. Answers GeoJS gave keepFor ago or more are // ago is dropped first. Answers GeoJS gave keepFor ago or more are
// dropped. // dropped.
func (g *GeoJS) Load(answers []Answer) { func (g *GeoJS) Load(answers []Answer) {
g.mu.Lock()
defer g.mu.Unlock()
answers = slices.Clone(answers) answers = slices.Clone(answers)
slices.SortStableFunc(answers, func(a, b Answer) int { slices.SortStableFunc(answers, func(a, b Answer) int {
return a.Used.Compare(b.Used) return a.Used.Compare(b.Used)
}) })
g.mu.Lock()
defer g.mu.Unlock()
g.answers.Purge()
now := g.now() now := g.now()
for _, answer := range answers { for _, answer := range answers {
+9 -6
View File
@@ -254,18 +254,21 @@ func (l *Limiter) Snapshot() []Client {
return clients return clients
} }
// Load puts clients read from clients.json into a table that holds none // Load puts clients read from clients.json into the table, in place of
// yet, in the order they were last seen, so that the least recently seen // the clients it holds, in the order they were last seen, so that the
// is dropped first. Buckets whose time has passed at now are emptied. // least recently seen is dropped first. Buckets whose time has passed at
// now are emptied.
func (l *Limiter) Load(clients []Client, now time.Time) { func (l *Limiter) Load(clients []Client, now time.Time) {
l.mu.Lock()
defer l.mu.Unlock()
clients = slices.Clone(clients) clients = slices.Clone(clients)
slices.SortStableFunc(clients, func(a, b Client) int { slices.SortStableFunc(clients, func(a, b Client) int {
return a.History.LastSeen.Compare(b.History.LastSeen) return a.History.LastSeen.Compare(b.History.LastSeen)
}) })
l.mu.Lock()
defer l.mu.Unlock()
l.clients.Purge()
for _, c := range clients { for _, c := range clients {
for i, b := range c.buckets() { for i, b := range c.buckets() {
// The window that ends at now covers neither bucket once it // The window that ends at now covers neither bucket once it
+18 -9
View File
@@ -112,9 +112,10 @@ func Run(ctx context.Context, params Params) int {
return serve(ctx, server.Server, listener, files, processLog) return serve(ctx, server.Server, listener, files, processLog)
} }
// serve serves requests on listener, and writes the state files as they // serve serves requests on listener, writes the state files as they are
// are due, until ctx is done. Then it gives the requests in progress // due, and takes in an admin's edits of them, until ctx is done. Then it
// shutdownTimeout to finish, and writes every state file. // gives the requests in progress shutdownTimeout to finish, and writes
// every state file.
func serve( func serve(
ctx context.Context, server *http.Server, listener net.Listener, ctx context.Context, server *http.Server, listener net.Listener,
files *state.Files, processLog *slog.Logger, files *state.Files, processLog *slog.Logger,
@@ -129,12 +130,18 @@ func serve(
defer stopWriting() defer stopWriting()
written := make(chan struct{}) written := make(chan struct{})
watched := make(chan struct{})
go func() { go func() {
files.Run(writing) files.Run(writing)
close(written) close(written)
}() }()
go func() {
files.Watch(writing)
close(watched)
}()
select { select {
case err := <-served: case err := <-served:
processLog.Error("serving failed", "error", err.Error()) processLog.Error("serving failed", "error", err.Error())
@@ -164,13 +171,15 @@ func serve(
return 1 return 1
} }
// Run's last write has ended, so nothing else writes the files. Every // Run and Watch have ended, so nothing else reads or writes the
// request has ended too, but for two kinds that Go's server does not // files. Every request has ended too, but for two kinds
// wait for: one cut off because Shutdown timed out, and one whose // that Go's server does not wait for: one cut off because Shutdown
// connection switched protocols, such as a WebSocket. Such a request // timed out, and one whose connection switched protocols, such as a
// adds to its client's history only as it ends, which can be after // WebSocket. Such a request adds to its client's history only as it
// this write, and then that request is missing from clients.json. // ends, which can be after this write, and then that request is
// missing from clients.json.
<-written <-written
<-watched
err = files.WriteAll() err = files.WriteAll()
if err != nil { if err != nil {
+86 -19
View File
@@ -26,11 +26,14 @@ const (
// testVersion is the version the tests give smallwebwaf. // testVersion is the version the tests give smallwebwaf.
testVersion = "test" testVersion = "test"
// localhost is where the tests listen. // localhost is where the tests listen.
localhost = "127.0.0.1" localhost = "127.0.0.1"
listenAddr = "SWWAF_LISTEN_ADDR" listenAddr = "SWWAF_LISTEN_ADDR"
upstreamURL = "SWWAF_UPSTREAM_URL" upstreamURL = "SWWAF_UPSTREAM_URL"
stateDir = "SWWAF_STATE_DIR" trustedProxies = "SWWAF_TRUSTED_PROXIES"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" 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 is what the tests' app answers.
greeting = "hello from the app" greeting = "hello from the app"
) )
@@ -195,8 +198,8 @@ func TestStateKeptAcrossRestarts(t *testing.T) {
rateLimitPerDay: "2", rateLimitPerDay: "2",
// Neither comes due in the test: the files are written as // Neither comes due in the test: the files are written as
// smallwebwaf stops. // smallwebwaf stops.
"SWWAF_STATE_WRITE_DELAY": "1h", stateWriteDelay: "1h",
"SWWAF_STATE_COUNTER_INTERVAL": "1h", stateCounterInterval: "1h",
} }
// The two requests a day allows, and a stop. // The two requests a day allows, and a stop.
@@ -225,12 +228,12 @@ func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) {
const scope = "SWWAF_BAN_SCOPE_V4_PREFIX" const scope = "SWWAF_BAN_SCOPE_V4_PREFIX"
env := map[string]string{ env := map[string]string{
listenAddr: localhost + ":0", listenAddr: localhost + ":0",
upstreamURL: startApp(t), upstreamURL: startApp(t),
stateDir: t.TempDir(), stateDir: t.TempDir(),
"SWWAF_TRUSTED_PROXIES": localhost + "/32", trustedProxies: localhost + "/32",
rateLimitPerDay: "1", rateLimitPerDay: "1",
scope: "24", scope: "24",
} }
// 203.0.113.9's second request breaks the day limit, and bans // 203.0.113.9's second request breaks the day limit, and bans
@@ -259,6 +262,38 @@ func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) {
}) })
} }
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) { func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
@@ -361,9 +396,9 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
listenAddr: localhost + ":0", listenAddr: localhost + ":0",
upstreamURL: appURL, upstreamURL: appURL,
stateDir: dir, stateDir: dir,
"SWWAF_STATE_WRITE_DELAY": "10s", stateWriteDelay: "10s",
"SWWAF_STATE_COUNTER_INTERVAL": "15m", stateCounterInterval: "15m",
"SWWAF_TRUSTED_PROXIES": "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16", trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
"SWWAF_CLIENT_REQUEST_TIMEOUT": "60s", "SWWAF_CLIENT_REQUEST_TIMEOUT": "60s",
"SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K", "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K",
"SWWAF_CLIENT_IDLE_TIMEOUT": "120s", "SWWAF_CLIENT_IDLE_TIMEOUT": "120s",
@@ -456,6 +491,40 @@ func wantRefused(t *testing.T, url string) {
func wantStatus(t *testing.T, url, from string, status int) { func wantStatus(t *testing.T, url, from string, status int) {
t.Helper() 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, req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url,
http.NoBody) http.NoBody)
if err != nil { if err != nil {
@@ -474,7 +543,5 @@ func wantStatus(t *testing.T, url, from string, status int) {
_ = res.Body.Close() _ = res.Body.Close()
if res.StatusCode != status { return res.StatusCode
t.Errorf("request from %s: status %d, want %d", from, res.StatusCode, status)
}
} }
+234 -66
View File
@@ -1,14 +1,16 @@
// Package state keeps smallwebwaf's state in JSON files in // Package state keeps smallwebwaf's state in JSON files in
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes: // SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
// bans.json holds the bans, clients.json each client's counters and // 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 // history, and lookups.json GeoJS's answers. Load reads them at start,
// Run and WriteAll write them, each from a snapshot its part takes under // Watch takes in an admin's edit of one while smallwebwaf runs, and Run
// its own lock, so that no request waits on the disk. // and WriteAll write them. Each part takes a snapshot, or what a file
// holds, under its own lock, so that no request waits on the disk.
package state package state
import ( import (
"bytes" "bytes"
"context" "context"
"crypto/sha256"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
@@ -17,8 +19,10 @@ import (
"net/netip" "net/netip"
"os" "os"
"path/filepath" "path/filepath"
"sync"
"time" "time"
"github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
@@ -60,13 +64,23 @@ type Params struct {
// Now tells the time by which the counters' buckets run out, normally // Now tells the time by which the counters' buckets run out, normally
// time.Now in UTC. // time.Now in UTC.
Now func() time.Time Now func() time.Time
// ProcessLog receives what was read, and the writes that fail. // ProcessLog receives what was read and taken in, the edits set aside,
// and the writes that fail.
ProcessLog *slog.Logger ProcessLog *slog.Logger
} }
// Files are the state files of a running smallwebwaf. // Files are the state files of a running smallwebwaf.
type Files struct { type Files struct {
params Params params Params
// mu is held while a file is read for an edit, and while it is
// written, so that Watch and the writes take turns. No request takes
// it.
mu sync.Mutex
// sums are the SHA-256 sums of what each file held, by name, when
// smallwebwaf last read or wrote it. A file that holds anything else
// has been edited since.
sums map[string][sha256.Size]byte
} }
// bansFile is bans.json, indented for an admin to read and edit. // bansFile is bans.json, indented for an admin to read and edit.
@@ -117,41 +131,28 @@ func Load(params Params) (*Files, error) {
return nil, fmt.Errorf("SWWAF_STATE_DIR cannot be written: %w", err) return nil, fmt.Errorf("SWWAF_STATE_DIR cannot be written: %w", err)
} }
var ( f := &Files{params: params, sums: map[string][sha256.Size]byte{}}
bansIn bansFile
clientsIn clientsFile
lookupsIn lookupsFile
)
err = errors.Join( bansRead, bansErr := f.read(bansJSON)
read(params.Dir, bansJSON, &bansIn), clientsRead, clientsErr := f.read(clientsJSON)
read(params.Dir, clientsJSON, &clientsIn), lookupsRead, lookupsErr := f.read(lookupsJSON)
read(params.Dir, lookupsJSON, &lookupsIn),
) err = errors.Join(bansErr, clientsErr, lookupsErr)
if err != nil { if err != nil {
return nil, err 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, params.ProcessLog.Info("read the state files", "directory", params.Dir,
"bans", len(bansIn.Bans), "clients", len(clientsIn.Clients), "bans", bansRead, "clients", clientsRead, "lookups", lookupsRead)
"lookups", len(lookupsIn.Lookups))
return &Files{params: params}, nil return f, nil
} }
// Run writes bans.json WriteDelay after a ban is made, with every ban // Run writes bans.json WriteDelay after a ban is made, with every ban
// made in between, and every file every CounterInterval, until ctx is // 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 // done. A write that fails is logged, and the file is written again at
// its next write. // its next write. Each write takes in an admin's edit of its file first,
// as writeFile describes.
func (f *Files) Run(ctx context.Context) { func (f *Files) Run(ctx context.Context) {
interval := time.NewTicker(f.params.CounterInterval) interval := time.NewTicker(f.params.CounterInterval)
defer interval.Stop() defer interval.Stop()
@@ -169,7 +170,7 @@ func (f *Files) Run(ctx context.Context) {
case <-bansDue: case <-bansDue:
bansDue = nil bansDue = nil
f.logFailure(f.writeBans()) f.logFailure(f.writeFile(bansJSON))
case <-interval.C: case <-interval.C:
f.logFailure(f.WriteAll()) f.logFailure(f.WriteAll())
} }
@@ -179,7 +180,50 @@ func (f *Files) Run(ctx context.Context) {
// WriteAll writes every state file, as smallwebwaf stops. A file that // WriteAll writes every state file, as smallwebwaf stops. A file that
// fails does not keep the others from being written. // fails does not keep the others from being written.
func (f *Files) WriteAll() error { func (f *Files) WriteAll() error {
return errors.Join(f.writeBans(), f.writeClients(), f.writeLookups()) return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON),
f.writeFile(lookupsJSON))
}
// Watch watches Dir until ctx is done, and takes in an admin's edit of a
// state file as soon as it is saved: what the file holds replaces what
// smallwebwaf held for it. An edit that does not parse is left for the
// file's next write, which sets it aside, since a file can be read while
// an editor is still writing it. If Dir cannot be watched, that is
// logged, and an edit is taken in only before its file is written.
func (f *Files) Watch(ctx context.Context) {
watcher, err := fsnotify.NewWatcher()
if err == nil {
defer func() {
_ = watcher.Close()
}()
err = watcher.Add(f.params.Dir)
}
if err != nil {
f.params.ProcessLog.Error("cannot watch the state files for edits",
"error", err.Error())
return
}
f.params.ProcessLog.Info("watching the state files for edits",
"directory", f.params.Dir)
for {
select {
case <-ctx.Done():
return
case event := <-watcher.Events:
switch name := filepath.Base(event.Name); name {
case bansJSON, clientsJSON, lookupsJSON:
f.takeInEdit(name)
}
case err = <-watcher.Errors:
f.params.ProcessLog.Warn("watching the state files failed",
"error", err.Error())
}
}
} }
// logFailure logs a write that failed. // logFailure logs a write that failed.
@@ -190,41 +234,177 @@ func (f *Files) logFailure(err error) {
} }
} }
// writeBans writes bans.json. // takeInEdit takes in an edit of the state file name and logs it, if the
func (f *Files) writeBans() error { // file has changed since smallwebwaf last read or wrote it and parses. A
held := f.params.Ledger.Snapshot() // file that cannot be read or does not parse is left for its next write.
func (f *Files) takeInEdit(name string) {
f.mu.Lock()
defer f.mu.Unlock()
file := bansFile{Version: version, Bans: make([]banEntry, 0, len(held))} data, changed, err := f.readChanged(name)
for _, ban := range held { if err != nil || !changed {
file.Bans = append(file.Bans, newBanEntry(ban)) return
} }
data, err := json.MarshalIndent(file, "", " ") _, err = f.takeIn(name, data)
if err != nil { if err != nil {
return fmt.Errorf("encode %s: %w", bansJSON, err) return
} }
return write(f.params.Dir, bansJSON, append(data, '\n')) f.params.ProcessLog.Info("took in an edit of a state file",
"file", filepath.Join(f.params.Dir, name))
} }
// writeClients writes clients.json. // read takes in the state file name, as at start, and returns how many
func (f *Files) writeClients() error { // entries it holds. A missing file holds none, and is not read.
data, err := encodeOnePerLine("clients", f.params.Limiter.Snapshot()) func (f *Files) read(name string) (int, error) {
if err != nil { data, changed, err := f.readChanged(name)
return fmt.Errorf("encode %s: %w", clientsJSON, err) if err != nil || !changed {
return 0, err
} }
return write(f.params.Dir, clientsJSON, data) return f.takeIn(name, data)
} }
// writeLookups writes lookups.json. // readChanged returns what the state file name holds, and whether that
func (f *Files) writeLookups() error { // has changed since smallwebwaf last read or wrote the file, as it has
data, err := encodeOnePerLine("lookups", f.params.GeoJS.Snapshot()) // for a file smallwebwaf never read or wrote. A missing file has not
if err != nil { // changed: it is written again at its next write.
return fmt.Errorf("encode %s: %w", lookupsJSON, err) func (f *Files) readChanged(name string) ([]byte, bool, error) {
path := filepath.Join(f.params.Dir, name)
data, err := os.ReadFile(path) //nolint:gosec // a state file, in SWWAF_STATE_DIR
if errors.Is(err, fs.ErrNotExist) {
return nil, false, nil
} }
return write(f.params.Dir, lookupsJSON, data) if err != nil {
return nil, false, err
}
return data, sha256.Sum256(data) != f.sums[name], nil
}
// takeIn parses data, what the state file name holds, puts it into the
// part that keeps that state, in place of what the part held, and returns
// how many entries the file holds. An error names the file and, where the
// JSON decoder tells it, the line and column, or else the entry.
func (f *Files) takeIn(name string, data []byte) (int, error) {
path := filepath.Join(f.params.Dir, name)
var entries int
switch name {
case bansJSON:
var file bansFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
held := make([]bans.Ban, 0, len(file.Bans))
for _, entry := range file.Bans {
held = append(held, entry.ban())
}
f.params.Ledger.Load(held)
entries = len(held)
case clientsJSON:
var file clientsFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
f.params.Limiter.Load(file.Clients, f.params.Now())
entries = len(file.Clients)
case lookupsJSON:
var file lookupsFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
f.params.GeoJS.Load(file.Lookups)
entries = len(file.Lookups)
}
f.sums[name] = sha256.Sum256(data)
return entries, nil
}
// writeFile writes the state file name from what smallwebwaf holds. An
// edit made since smallwebwaf last read or wrote the file is taken in
// first, so that it is not overwritten. An edit that does not parse is
// renamed to name.bad, for the admin to mend, and logged with where in
// the file the error is.
func (f *Files) writeFile(name string) error {
f.mu.Lock()
defer f.mu.Unlock()
path := filepath.Join(f.params.Dir, name)
data, changed, err := f.readChanged(name)
if err != nil {
return err
}
if changed {
_, err = f.takeIn(name, data)
}
if err != nil {
renameErr := os.Rename(path, path+".bad")
if renameErr != nil {
return errors.Join(err, renameErr)
}
f.params.ProcessLog.Error("set aside an edit of a state file that does not parse",
"file", path+".bad", "error", err.Error())
}
data, err = f.encode(name)
if err != nil {
return fmt.Errorf("encode %s: %w", name, err)
}
err = write(f.params.Dir, name, data)
if err != nil {
return err
}
f.sums[name] = sha256.Sum256(data)
return nil
}
// encode returns the state file name as smallwebwaf writes it, from a
// snapshot of the part that keeps that state.
func (f *Files) encode(name string) ([]byte, error) {
switch name {
case bansJSON:
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 nil, err
}
return append(data, '\n'), nil
case clientsJSON:
return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
default: // lookups.json
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
}
} }
// newBanEntry returns ban as bans.json holds it. // newBanEntry returns ban as bans.json holds it.
@@ -377,28 +557,16 @@ func checkWritable(dir string) error {
return errors.Join(file.Close(), os.Remove(file.Name())) return errors.Join(file.Close(), os.Remove(file.Name()))
} }
// read reads the state file name in dir into file, a pointer to that // parse reads data, what the state file at path holds, into file, a
// file's struct, and checks its entries. A missing file leaves file as it // pointer to that file's struct, and checks its entries.
// is. func parse(path string, data []byte, file stateFile) error {
func read(dir, name string, file stateFile) 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 // 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. // refused for that, and not for an entry this version cannot read.
var header struct { var header struct {
Version int `json:"version"` Version int `json:"version"`
} }
err = json.Unmarshal(data, &header) err := json.Unmarshal(data, &header)
if err == nil && header.Version != version { if err == nil && header.Version != version {
err = fmt.Errorf("%w %d, where this smallwebwaf reads version %d", err = fmt.Errorf("%w %d, where this smallwebwaf reads version %d",
errVersion, header.Version, version) errVersion, header.Version, version)
+311 -11
View File
@@ -3,7 +3,9 @@ package state_test
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"io/fs"
"log/slog" "log/slog"
"net"
"net/netip" "net/netip"
"os" "os"
"path/filepath" "path/filepath"
@@ -24,6 +26,13 @@ const (
bansJSON = "bans.json" bansJSON = "bans.json"
clientsJSON = "clients.json" clientsJSON = "clients.json"
lookupsJSON = "lookups.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. // permanentBansJSON is bans.json holding permanentBan.
@@ -286,7 +295,7 @@ func TestUnwritableDirectoryStopsTheStart(t *testing.T) {
} }
} }
// The two tests below run Run in a synctest bubble, where time is a clock // 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 // 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 // returns once Run waits for its next write, so that every write due by
// then is on disk. // then is on disk.
@@ -298,7 +307,7 @@ func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) {
dir := t.TempDir() dir := t.TempDir()
params := newParams(dir) params := newParams(dir)
params.WriteDelay = 10 * time.Second params.WriteDelay = 10 * time.Second
run(t, load(t, params)) run(t, load(t, params).Run)
// A second ban, made while the first waits to be written, puts the // A second ban, made while the first waits to be written, puts the
// write off no further, and is written with it. // write off no further, and is written with it.
@@ -343,7 +352,7 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
dir := t.TempDir() dir := t.TempDir()
params := newParams(dir) params := newParams(dir)
params.CounterInterval = time.Minute params.CounterInterval = time.Minute
run(t, load(t, params)) run(t, load(t, params).Run)
// The files are removed once written, so that each interval shows // The files are removed once written, so that each interval shows
// them written again. // them written again.
@@ -360,6 +369,44 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
}) })
} }
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) { func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
t.Parallel() t.Parallel()
@@ -402,26 +449,205 @@ func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
} }
} }
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) { func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) {
t.Parallel() t.Parallel()
dir := t.TempDir() dir := t.TempDir()
path := filepath.Join(dir, bansJSON)
files := load(t, newParams(dir)) files := load(t, newParams(dir))
// A directory named bans.json cannot be renamed over. // A socket cannot be opened as a file, even by root, as which the
err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700) // tests run in Docker, but a rename can replace it. 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 { if err != nil {
t.Fatalf("mkdir: %v", err) t.Fatalf("listen: %v", err)
} }
defer func() {
_ = socket.Close()
}()
err = files.WriteAll() err = files.WriteAll()
if err == nil { if err == nil {
t.Error("writing over a directory did not fail") 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) wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
} }
func TestEditOfEachFileTakenIn(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
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(&params)
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(&params)
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(&params)
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 TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
lines := logInto(&params)
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. // midnight is the time of the tests' clock.
func midnight() time.Time { func midnight() time.Time {
return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC) return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
@@ -512,15 +738,16 @@ func load(t *testing.T, params state.Params) *state.Files {
return files return files
} }
// run runs files' writes until the test ends. // run runs task, the Run or the Watch of state files, until the test
func run(t *testing.T, files *state.Files) { // ends.
func run(t *testing.T, task func(context.Context)) {
t.Helper() t.Helper()
ctx, stop := context.WithCancel(t.Context()) ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{}) stopped := make(chan struct{})
go func() { go func() {
files.Run(ctx) task(ctx)
close(stopped) close(stopped)
}() }()
@@ -530,6 +757,79 @@ func run(t *testing.T, files *state.Files) {
}) })
} }
// 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 // wantEqual checks that the entries read back from file are those
// written. // written.
func wantEqual[E comparable](t *testing.T, file string, got, want []E) { func wantEqual[E comparable](t *testing.T, file string, got, want []E) {
+35
View File
@@ -0,0 +1,35 @@
package state
import (
"os"
"path/filepath"
"testing"
)
// The test is on write itself: a state file is read before it is
// written, and a directory in its place fails that read first.
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) {
t.Parallel()
dir := t.TempDir()
// 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 = write(dir, bansJSON, []byte("{}\n"))
if err == nil {
t.Error("writing over a directory did not fail")
}
entries, err := os.ReadDir(dir)
if err != nil {
t.Fatalf("read %s: %v", dir, err)
}
if len(entries) != 1 || entries[0].Name() != bansJSON {
t.Errorf("%s holds %v, want only bans.json", dir, entries)
}
}