diff --git a/README.md b/README.md index 2a44452..7544ef9 100644 --- a/README.md +++ b/README.md @@ -15,20 +15,21 @@ Status: the first two milestones are built (https://git.eeqj.de/sneak/smallwebwaf/issues/13 and https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are six parts of 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, `observe` mode, which -comes a little later, and the metrics endpoint and the header size and the idle -time as settings, which come last in it. `smallwebwaf` passes each request to -the app and the app's answer back, unchanged, within its timeouts and size -limits, works out each client's address, bans a client that sends too many -requests, refuses a client that comes from a country you refuse or from a -network you refuse, lets the networks you choose through, keeps its bans, each -client's counters and history, and GeoJS's answers in JSON files across -restarts, writes a JSON log line for every request, serves Prometheus metrics to -a scraper that holds the metrics token, and in `observe` mode passes on the -requests it would refuse, logging what it would have done with them. It comes as -the image the app's own image is built on. The rest of the design comes after -that, in the order of the build order in [`SPEC.md`](SPEC.md). The survey of -existing tools that led to the design is in [`EVALUATION.md`](EVALUATION.md). +JSON state files with your edits taken in while it runs, which come next in the +build order, `observe` mode, which comes a little later, and the metrics +endpoint and the header size and the idle time as settings, which come last in +it. `smallwebwaf` passes each request to the app and the app's answer back, +unchanged, within its timeouts and size limits, works out each client's address, +bans a client that sends too many requests, refuses a client that comes from a +country you refuse or from a network you refuse, lets the networks you choose +through, keeps its bans, each client's counters and history, and GeoJS's answers +in JSON files across restarts, takes in your edits of those files while it runs, +writes a JSON log line for every request, serves Prometheus metrics to a scraper +that holds the metrics token, and in `observe` mode passes on the requests it +would refuse, logging what it would have done with them. It comes as the image +the app's own image is built on. The rest of the design comes after that, in the +order of 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 @@ -100,9 +101,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 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 - request is dropped first. `bans.json` shows the bans and their notes, and a - restart lifts none (see "State files" below); lifting a ban by editing it - comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68. + request is dropped first. `bans.json` shows the bans and their notes, a + restart lifts none, and you add or lift a ban by editing it (see "State files" + below). - 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 is not counted for the rate limits. While one of the country lists below is @@ -335,9 +336,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 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 -`answered`. An edit made while `smallwebwaf` runs is overwritten by its next -write: taking it in comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68. -The AS number and AS name come with their lookup. +`answered`. 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 `.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. ## Metrics @@ -376,7 +410,10 @@ other request. No metric carries a client's address. - `smallwebwaf_state_file_writes_total`, `smallwebwaf_state_file_write_failures_total`, `smallwebwaf_state_file_last_write_timestamp_seconds` and - `smallwebwaf_state_file_size_bytes`, by `file`. + `smallwebwaf_state_file_size_bytes`, by `file`; and, by `file` too, + `smallwebwaf_state_file_edits_taken_in_total`: your edits taken in, and + `smallwebwaf_state_file_edits_set_aside_total`: those renamed to `.bad` + because they would stop the start. - Go's own `go_` metrics and the process's `process_` metrics. The requests Go's HTTP server ends before `smallwebwaf` sees them (see "Request @@ -484,9 +521,9 @@ goes through the candidates one by one. 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 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" - above); the others come with their features, and taking in an edit while - running comes with https://git.eeqj.de/sneak/smallwebwaf/issues/68. + for the bans, the clients and the GeoJS answers are built, with an edit taken + in while running (see "State files" above); the others come with their + features. - 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 `/_smallwebwaf/` on the app's own address, through traefik like any other @@ -678,8 +715,8 @@ addresses are never sent to GeoJS. answers. - `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. -- `internal/state`: reads the state files at start, and writes them when they - are due and at the stop. +- `internal/state`: reads the state files at start, takes in an admin's edit of + 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 process's own messages. - `Dockerfile`: the lint and test phases, then the image, whose last stage @@ -691,8 +728,9 @@ addresses are never sent to GeoJS. 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 netblocks to `SWWAF_MAX_BANS`, dropping the least recently seen, and -`github.com/prometheus/client_golang` keeps the metrics and serves them. The -country codes are the list in `internal/config/config.go`. +`github.com/prometheus/client_golang` keeps the metrics and serves them, and +`github.com/fsnotify/fsnotify` tells `smallwebwaf` when a state file is saved. +The country codes are the list in `internal/config/config.go`. ## Entrypoints @@ -735,10 +773,9 @@ so that they run in minimal containers. ## TODO -- The rest of milestone 3: taking in an admin's edits to the state files - (https://git.eeqj.de/sneak/smallwebwaf/issues/68), exemptions and the rest of - the request log's fields; then the rest of the design, in the order of the - build order in [`SPEC.md`](SPEC.md). +- The rest of milestone 3: exemptions and the rest of the request log's fields; + then the rest of the design, in the order of the build order in + [`SPEC.md`](SPEC.md). ## Documents diff --git a/go.mod b/go.mod index a033322..caf899d 100644 --- a/go.mod +++ b/go.mod @@ -3,6 +3,7 @@ module sneak.berlin/go/smallwebwaf go 1.26.0 require ( + github.com/fsnotify/fsnotify v1.10.1 github.com/hashicorp/golang-lru/v2 v2.0.7 github.com/prometheus/client_golang v1.24.1 ) diff --git a/go.sum b/go.sum index c02c3da..0e9f7bb 100644 --- a/go.sum +++ b/go.sum @@ -4,6 +4,8 @@ github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UF github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho= +github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= diff --git a/internal/bans/bans.go b/internal/bans/bans.go index 68ee022..19b7c19 100644 --- a/internal/bans/bans.go +++ b/internal/bans/bans.go @@ -26,9 +26,9 @@ const maxTextBytes = 256 type Rules struct { // LimitBanDuration is how long a first ban lasts. LimitBanDuration time.Duration - // LimitBanRepeatWindow is how soon after the netblock's last ban - // ended a broken limit counts as a repeat, which bans for - // repeatFactor times as long as that ban. + // LimitBanRepeatWindow is how soon after the end of the netblock's + // ban that ended last a broken limit counts as a repeat, which bans + // for repeatFactor times as long as that ban. LimitBanRepeatWindow time.Duration // MaxBanDuration is the longest ban; a ban that would be longer is // permanent instead. @@ -176,9 +176,23 @@ func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool) { return *ban, true } +// activeBan returns the ban in bans, a netblock's bans oldest first, that +// is active at now, or nil when none is. If several are, it returns the +// one that started last. Every ban is looked at, since a ban an admin adds +// to bans.json can start before the netblock's others and outlast them. +func activeBan(bans []Ban, now time.Time) *Ban { + for i := len(bans) - 1; i >= 0; i-- { + if bans[i].ActiveAt(now) { + return &bans[i] + } + } + + return nil +} + // BanForLimit bans netblock at now for a broken limit, with notes, and // returns the ban. A first ban lasts LimitBanDuration. A ban made within -// LimitBanRepeatWindow after the netblock's last ban ended lasts +// LimitBanRepeatWindow after the netblock's ban that ended last lasts // repeatFactor times as long as that one. A ban that would be longer // than MaxBanDuration is permanent instead. If a ban on netblock is still // active, as when two of its requests break a limit at once, that ban is @@ -192,12 +206,22 @@ func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes) bans, found := l.netblocks.Get(netblock) if found { - last = &(*bans)[len(*bans)-1] - if last.ActiveAt(now) { - return *last + active := activeBan(*bans, now) + if active != nil { + return *active } - notes.EarlierBans = last.Notes.EarlierBans + 1 + // No ban is active, so each has an end. A ban an admin adds to + // bans.json can start after another and end before it, so the + // ban that ended last is looked for among them all. + ended := slices.MaxFunc(*bans, func(a, b Ban) int { + return a.Expires.Compare(b.Expires) + }) + last = &ended + + // The netblock's first ban held counts the bans it had before that + // one, since dropped to make room, and each ban held adds one. + notes.EarlierBans = (*bans)[0].Notes.EarlierBans + len(*bans) } notes.Request = notes.Request.cut() @@ -282,21 +306,25 @@ func (l *Ledger) Snapshot() []Ban { return held } -// Load puts bans read from bans.json into a ledger that holds none yet, -// in the order they started, so that a netblock whose last ban started -// latest counts as the most recently seen. Each netblock is masked to its -// length, so that 203.0.113.9/24 is 203.0.113.0/24, and each text in the -// notes is cut to 256 bytes. Past MaxBans the earliest bans are dropped, -// as when they are made. +// Load puts bans read from bans.json into the ledger, in place of the +// bans it holds, in the order they started, so that a netblock whose last +// ban started latest counts as the most recently seen. Each netblock is +// masked to its length, so that 203.0.113.9/24 is 203.0.113.0/24, and +// each text in the notes is cut to 256 bytes. Past MaxBans the earliest +// bans are dropped, as when they are made. func (l *Ledger) Load(bans []Ban) { - l.mu.Lock() - defer l.mu.Unlock() - bans = slices.Clone(bans) slices.SortStableFunc(bans, func(a, b Ban) int { 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 { ban.Netblock = ban.Netblock.Masked() ban.Notes.Request = ban.Notes.Request.cut() @@ -318,11 +346,9 @@ func (l *Ledger) active(client netip.Addr, now time.Time) *Ban { continue } - // A ban is made only once the one before has ended, so only the - // last can be active. - last := &(*bans)[len(*bans)-1] - if last.ActiveAt(now) { - return last + ban := activeBan(*bans, now) + if ban != nil { + return ban } } @@ -358,8 +384,8 @@ func (l *Ledger) add(ban Ban) { } // expiry returns when a ban for a broken limit made at now ends, or zero -// when it is permanent. last is the netblock's last ban, which has ended, -// or nil when it has none. +// when it is permanent. last is the netblock's ban that ended last, or nil +// when it has none. func (l *Ledger) expiry(last *Ban, now time.Time) time.Time { length := l.rules.LimitBanDuration diff --git a/internal/bans/snapshot_test.go b/internal/bans/snapshot_test.go index 873b196..3a1a651 100644 --- a/internal/bans/snapshot_test.go +++ b/internal/bans/snapshot_test.go @@ -130,6 +130,75 @@ func TestLoadedBanRefusesEveryClientInItsNetblock(t *testing.T) { } } +func TestPermanentBanStartedBeforeAnEndedOneRefuses(t *testing.T) { + t.Parallel() + + // As when an admin adds a permanent ban to bans.json with a start + // before that of the netblock's ban that has ended. + netblock := netip.MustParsePrefix("203.0.113.0/24") + permanent := bans.Ban{Netblock: netblock, Start: midnight().Add(-time.Hour)} + ended := bans.Ban{ + Netblock: netblock, + Start: midnight(), + Expires: midnight().Add(time.Hour), + } + + ledger := bans.New(defaultRules()) + ledger.Load([]bans.Ban{permanent, ended}) + + now := midnight().Add(2 * time.Hour) + client := netip.MustParseAddr("203.0.113.9") + + ban, banned := ledger.Find(client, now) + if !banned || !ban.Permanent() { + t.Errorf("find gives %+v and %t, want the permanent ban", ban, banned) + } + + ban, banned = ledger.Check(client, now) + if !banned || !ban.Permanent() { + t.Errorf("the client is refused: %t, under %+v, want under the permanent ban", + banned, ban) + } + + // A limit broken now makes no shorter ban over the permanent one. + ban = ledger.BanForLimit(netblock, now, bans.Notes{}) + if !ban.Permanent() || len(ledger.Bans(netblock)) != 2 { + t.Errorf("a broken limit returned %+v and left the netblock %d bans, "+ + "want the permanent ban and 2", ban, len(ledger.Bans(netblock))) + } +} + +func TestNextBanWorkedOutFromTheBanThatEndedLast(t *testing.T) { + t.Parallel() + + // A 9-hour ban smallwebwaf made, the third in a row, and an admin's + // 1-hour ban added to bans.json over it, with no notes. + netblock := netip.MustParsePrefix("203.0.113.9/32") + nineHours := bans.Ban{ + Netblock: netblock, + Start: midnight(), + Expires: midnight().Add(9 * time.Hour), + Notes: bans.Notes{EarlierBans: 2}, + } + admins := bans.Ban{ + Netblock: netblock, + Start: midnight().Add(time.Hour), + Expires: midnight().Add(2 * time.Hour), + } + + ledger := bans.New(defaultRules()) + ledger.Load([]bans.Ban{nineHours, admins}) + + // Once both have ended, a limit broken within the repeat window bans + // for three times the 9 hours, and the notes count the two bans + // before the 9-hour one, it, and the admin's. + ban := ledger.BanForLimit(netblock, nineHours.Expires.Add(time.Hour), bans.Notes{}) + if ban.Expires.Sub(ban.Start) != 27*time.Hour || ban.Notes.EarlierBans != 4 { + t.Errorf("the next ban lasts %s with %d earlier bans, want 27h and 4", + ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans) + } +} + func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) { t.Parallel() @@ -151,6 +220,41 @@ func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) { } } +func TestLoadReplacesTheBansHeld(t *testing.T) { + t.Parallel() + + // Room for three bans, so that the second load, were it added to the + // two bans held, would drop none of them to make room. + rules := defaultRules() + rules.MaxBans = 3 + 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. + ledger.Load([]bans.Ban{kept}) + + _, banned := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight()) + if banned { + t.Error("a ban left out of the second load still refuses") + } + + // The ledger holds one ban, so it makes two more without dropping any. + first := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(), + bans.Notes{}) + second := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(), + bans.Notes{}) + + want := []bans.Ban{first, second, kept} + if got := ledger.Snapshot(); !slices.Equal(got, want) { + t.Errorf("the ledger holds %+v, want %+v", got, want) + } +} + func TestLoadCutsTheTextsTo256Bytes(t *testing.T) { t.Parallel() diff --git a/internal/lookup/lookup.go b/internal/lookup/lookup.go index c181007..a2b7139 100644 --- a/internal/lookup/lookup.go +++ b/internal/lookup/lookup.go @@ -197,19 +197,21 @@ func (g *GeoJS) Snapshot() []Answer { return answers } -// Load keeps answers read from lookups.json, in a GeoJS that keeps none -// yet, in the order they were last used, so that the one used longest +// Load keeps answers read from lookups.json, in place of the answers it +// 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 // dropped. func (g *GeoJS) Load(answers []Answer) { - g.mu.Lock() - defer g.mu.Unlock() - answers = slices.Clone(answers) slices.SortStableFunc(answers, func(a, b Answer) int { return a.Used.Compare(b.Used) }) + g.mu.Lock() + defer g.mu.Unlock() + + g.answers.Purge() + now := g.now() for _, answer := range answers { diff --git a/internal/metrics/metrics.go b/internal/metrics/metrics.go index 2bba0ed..f08a714 100644 --- a/internal/metrics/metrics.go +++ b/internal/metrics/metrics.go @@ -44,6 +44,8 @@ type Metrics struct { stateFileWriteFailures *prometheus.CounterVec stateFileLastWrite *prometheus.GaugeVec stateFileSize *prometheus.GaugeVec + stateFileEditsTakenIn *prometheus.CounterVec + stateFileEditsSetAside *prometheus.CounterVec } // New returns the metrics, with the Go runtime's and the process's own. @@ -105,6 +107,11 @@ func New(topN int) *Metrics { "When each state file was last written, in seconds since 1970.", byFile), stateFileSize: gaugeVec("smallwebwaf_state_file_size_bytes", "The size of each state file, as it was last written.", byFile), + stateFileEditsTakenIn: counterVec("smallwebwaf_state_file_edits_taken_in_total", + "Edits of each state file taken in while running.", byFile), + stateFileEditsSetAside: counterVec("smallwebwaf_state_file_edits_set_aside_total", + "Edits of each state file renamed to .bad because they did not parse.", + byFile), } m.handler = promhttp.HandlerFor(m.registry, promhttp.HandlerOpts{}) @@ -120,6 +127,7 @@ func New(topN int) *Metrics { m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered, m.stateFileWrites, m.stateFileWriteFailures, m.stateFileLastWrite, m.stateFileSize, + m.stateFileEditsTakenIn, m.stateFileEditsSetAside, ) return m @@ -230,6 +238,18 @@ func (m *Metrics) StateFileWritten(name string, size int, err error) { m.stateFileSize.WithLabelValues(name).Set(float64(size)) } +// StateFileEditTakenIn counts an admin's edit of the state file name +// taken in while smallwebwaf runs. +func (m *Metrics) StateFileEditTakenIn(name string) { + m.stateFileEditsTakenIn.WithLabelValues(name).Inc() +} + +// StateFileEditSetAside counts an admin's edit of the state file name +// renamed to name.bad because it did not parse. +func (m *Metrics) StateFileEditSetAside(name string) { + m.stateFileEditsSetAside.WithLabelValues(name).Inc() +} + // statusClass returns the class of status, such as 2xx, or none when no // status was sent. func statusClass(status int) string { diff --git a/internal/ratelimit/ratelimit.go b/internal/ratelimit/ratelimit.go index 2220c46..fb25801 100644 --- a/internal/ratelimit/ratelimit.go +++ b/internal/ratelimit/ratelimit.go @@ -269,18 +269,21 @@ func (l *Limiter) Snapshot() []Client { return clients } -// Load puts clients read from clients.json into a table that holds none -// yet, in the order they were last seen, so that the least recently seen -// is dropped first. Buckets whose time has passed at now are emptied. +// Load puts clients read from clients.json into the table, in place of +// the clients it holds, in the order they were last seen, so that the +// least recently seen is dropped first. Buckets whose time has passed at +// now are emptied. func (l *Limiter) Load(clients []Client, now time.Time) { - l.mu.Lock() - defer l.mu.Unlock() - clients = slices.Clone(clients) slices.SortStableFunc(clients, func(a, b Client) int { return a.History.LastSeen.Compare(b.History.LastSeen) }) + l.mu.Lock() + defer l.mu.Unlock() + + l.clients.Purge() + for _, c := range clients { for i, b := range c.buckets() { // The window that ends at now covers neither bucket once it diff --git a/internal/smallwebwaf/smallwebwaf.go b/internal/smallwebwaf/smallwebwaf.go index 0fe29e0..06b446b 100644 --- a/internal/smallwebwaf/smallwebwaf.go +++ b/internal/smallwebwaf/smallwebwaf.go @@ -113,9 +113,10 @@ func Run(ctx context.Context, params Params) int { return serve(ctx, server.Server, listener, files, processLog) } -// serve serves requests on listener, and writes the state files as they -// are due, until ctx is done. Then it gives the requests in progress -// shutdownTimeout to finish, and writes every state file. +// serve serves requests on listener, writes the state files as they are +// due, and takes in an admin's edits of them, until ctx is done. Then it +// gives the requests in progress shutdownTimeout to finish, and writes +// every state file. func serve( ctx context.Context, server *http.Server, listener net.Listener, files *state.Files, processLog *slog.Logger, @@ -130,12 +131,18 @@ func serve( defer stopWriting() written := make(chan struct{}) + watched := make(chan struct{}) go func() { files.Run(writing) close(written) }() + go func() { + files.Watch(writing) + close(watched) + }() + select { case err := <-served: processLog.Error("serving failed", "error", err.Error()) @@ -165,13 +172,15 @@ func serve( return 1 } - // Run's last write has ended, so nothing else writes the files. Every - // request has ended too, but for two kinds that Go's server does not - // wait for: one cut off because Shutdown timed out, and one whose - // connection switched protocols, such as a WebSocket. Such a request - // adds to its client's history only as it ends, which can be after - // this write, and then that request is missing from clients.json. + // Run and Watch have ended, so nothing else reads or writes the + // files. Every request has ended too, but for two kinds + // that Go's server does not wait for: one cut off because Shutdown + // timed out, and one whose connection switched protocols, such as a + // WebSocket. Such a request adds to its client's history only as it + // ends, which can be after this write, and then that request is + // missing from clients.json. <-written + <-watched err = files.WriteAll() if err != nil { diff --git a/internal/smallwebwaf/smallwebwaf_test.go b/internal/smallwebwaf/smallwebwaf_test.go index 76f6a22..befbab9 100644 --- a/internal/smallwebwaf/smallwebwaf_test.go +++ b/internal/smallwebwaf/smallwebwaf_test.go @@ -26,11 +26,14 @@ const ( // testVersion is the version the tests give smallwebwaf. testVersion = "test" // localhost is where the tests listen. - localhost = "127.0.0.1" - listenAddr = "SWWAF_LISTEN_ADDR" - upstreamURL = "SWWAF_UPSTREAM_URL" - stateDir = "SWWAF_STATE_DIR" - rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" + localhost = "127.0.0.1" + listenAddr = "SWWAF_LISTEN_ADDR" + upstreamURL = "SWWAF_UPSTREAM_URL" + trustedProxies = "SWWAF_TRUSTED_PROXIES" + 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 = "hello from the app" ) @@ -217,8 +220,8 @@ func TestStateKeptAcrossRestarts(t *testing.T) { rateLimitPerDay: "2", // Neither comes due in the test: the files are written as // smallwebwaf stops. - "SWWAF_STATE_WRITE_DELAY": "1h", - "SWWAF_STATE_COUNTER_INTERVAL": "1h", + stateWriteDelay: "1h", + stateCounterInterval: "1h", } // The two requests a day allows, and a stop. @@ -247,12 +250,12 @@ func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) { const scope = "SWWAF_BAN_SCOPE_V4_PREFIX" env := map[string]string{ - listenAddr: localhost + ":0", - upstreamURL: startApp(t), - stateDir: t.TempDir(), - "SWWAF_TRUSTED_PROXIES": localhost + "/32", - rateLimitPerDay: "1", - scope: "24", + listenAddr: localhost + ":0", + upstreamURL: startApp(t), + stateDir: t.TempDir(), + trustedProxies: localhost + "/32", + rateLimitPerDay: "1", + scope: "24", } // 203.0.113.9's second request breaks the day limit, and bans @@ -281,6 +284,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) { t.Parallel() @@ -384,9 +419,9 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) { upstreamURL: appURL, stateDir: dir, "SWWAF_MODE": "enforce", - "SWWAF_STATE_WRITE_DELAY": "10s", - "SWWAF_STATE_COUNTER_INTERVAL": "15m", - "SWWAF_TRUSTED_PROXIES": "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16", + stateWriteDelay: "10s", + stateCounterInterval: "15m", + trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16", "SWWAF_CLIENT_REQUEST_TIMEOUT": "60s", "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K", "SWWAF_CLIENT_IDLE_TIMEOUT": "120s", @@ -479,6 +514,40 @@ func wantRefused(t *testing.T, url string) { func wantStatus(t *testing.T, url, from string, status int) { 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, http.NoBody) if err != nil { @@ -497,7 +566,5 @@ func wantStatus(t *testing.T, url, from string, status int) { _ = res.Body.Close() - if res.StatusCode != status { - t.Errorf("request from %s: status %d, want %d", from, res.StatusCode, status) - } + return res.StatusCode } diff --git a/internal/state/state.go b/internal/state/state.go index f0d40dc..35551cb 100644 --- a/internal/state/state.go +++ b/internal/state/state.go @@ -1,14 +1,17 @@ // 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. +// history, and lookups.json GeoJS's answers. Load reads them at start, +// Watch takes in an admin's edit of one while smallwebwaf runs, and Run +// and WriteAll write them. The disk is read and written outside the +// parts' locks, which are held only to take a snapshot or to put in what +// a file holds, so that no request waits on the disk. package state import ( "bytes" "context" + "crypto/sha256" "encoding/json" "errors" "fmt" @@ -17,8 +20,10 @@ import ( "net/netip" "os" "path/filepath" + "sync" "time" + "github.com/fsnotify/fsnotify" "sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/metrics" @@ -61,15 +66,26 @@ type Params struct { // 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 receives what was read and taken in, the edits set aside, + // and the writes that fail. ProcessLog *slog.Logger - // Metrics count each file's writes. + // Metrics count each file's writes, and the edits taken in and set + // aside. Metrics *metrics.Metrics } // Files are the state files of a running smallwebwaf. type Files struct { 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. @@ -120,41 +136,28 @@ func Load(params Params) (*Files, error) { return nil, fmt.Errorf("SWWAF_STATE_DIR cannot be written: %w", err) } - var ( - bansIn bansFile - clientsIn clientsFile - lookupsIn lookupsFile - ) + f := &Files{params: params, sums: map[string][sha256.Size]byte{}} - err = errors.Join( - read(params.Dir, bansJSON, &bansIn), - read(params.Dir, clientsJSON, &clientsIn), - read(params.Dir, lookupsJSON, &lookupsIn), - ) + bansRead, bansErr := f.read(bansJSON) + clientsRead, clientsErr := f.read(clientsJSON) + lookupsRead, lookupsErr := f.read(lookupsJSON) + + err = errors.Join(bansErr, clientsErr, lookupsErr) 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)) + "bans", bansRead, "clients", clientsRead, "lookups", lookupsRead) - return &Files{params: params}, nil + return f, 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. +// 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) { interval := time.NewTicker(f.params.CounterInterval) defer interval.Stop() @@ -172,7 +175,7 @@ func (f *Files) Run(ctx context.Context) { case <-bansDue: bansDue = nil - f.logFailure(f.writeBans()) + f.logFailure(f.writeFile(bansJSON)) case <-interval.C: f.logFailure(f.WriteAll()) } @@ -182,7 +185,50 @@ func (f *Files) Run(ctx context.Context) { // 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()) + 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.fileChanged(name) + } + case err = <-watcher.Errors: + f.params.ProcessLog.Warn("watching the state files failed", + "error", err.Error()) + } + } } // logFailure logs a write that failed. @@ -193,52 +239,209 @@ func (f *Files) logFailure(err error) { } } -// writeBans writes bans.json. -func (f *Files) writeBans() error { - held := f.params.Ledger.Snapshot() +// fileChanged takes in what the state file name holds, as Watch sees it +// change, if that is an edit made since smallwebwaf last read or wrote +// the file. A file that cannot be read or does not parse is left for its +// next write. +func (f *Files) fileChanged(name string) { + f.mu.Lock() + defer f.mu.Unlock() - file := bansFile{Version: version, Bans: make([]banEntry, 0, len(held))} - for _, ban := range held { - file.Bans = append(file.Bans, newBanEntry(ban)) + data, changed, err := f.readChanged(name) + if err != nil || !changed { + return } - data, err := json.MarshalIndent(file, "", " ") - if err != nil { - return fmt.Errorf("encode %s: %w", bansJSON, err) - } - - return f.writeCounted(bansJSON, append(data, '\n')) + _ = f.takeInEdit(name, data) } -// writeClients writes clients.json. -func (f *Files) writeClients() error { - data, err := encodeOnePerLine("clients", f.params.Limiter.Snapshot()) +// takeInEdit takes in data, an edit of the state file name, as takeIn +// does, and counts and logs it. Every edit taken in while smallwebwaf +// runs, by Watch or by a write, is taken in here. An edit that does not +// parse is neither counted nor logged, and takeIn's error returned. +func (f *Files) takeInEdit(name string, data []byte) error { + _, err := f.takeIn(name, data) if err != nil { - return fmt.Errorf("encode %s: %w", clientsJSON, err) + return err } - return f.writeCounted(clientsJSON, data) + // Counted before it is logged, so that the count is there once the + // log line is. + f.params.Metrics.StateFileEditTakenIn(name) + f.params.ProcessLog.Info("took in an edit of a state file", + "file", filepath.Join(f.params.Dir, name)) + + return nil } -// 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) +// read takes in the state file name at start, and returns how many +// entries it holds. A missing file holds none. +func (f *Files) read(name string) (int, error) { + data, changed, err := f.readChanged(name) + if err != nil || !changed { + return 0, err } - return f.writeCounted(lookupsJSON, data) + return f.takeIn(name, data) } -// writeCounted writes data to the state file name, as write does, and -// counts the write in the metrics. -func (f *Files) writeCounted(name string, data []byte) error { - err := write(f.params.Dir, name, data) +// readChanged returns what the state file name holds, and whether that +// has changed since smallwebwaf last read or wrote the file, as it has +// for a file smallwebwaf never read or wrote. A missing file has not +// changed: it is written again at its next write. +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 + } + + 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, or set aside if it does not +// parse. A file that cannot be read, or an edit that cannot be set +// aside, is left as it is, and the write given up. Every write is counted +// in the metrics, and one that fails or is given up as a failure. +func (f *Files) writeFile(name string) error { + f.mu.Lock() + defer f.mu.Unlock() + + data, changed, err := f.readChanged(name) + if err == nil && changed { + err = f.takeInEdit(name, data) + if err != nil { + err = f.setAside(name, err) + } + } + + if err == nil { + data, err = f.encode(name) + if err != nil { + err = fmt.Errorf("encode %s: %w", name, err) + } + } + + if err == nil { + err = write(f.params.Dir, name, data) + } + + if err == nil { + // The file holds data from here on, even if the directory sync + // fails, so that its next read does not take it for an admin's + // edit. + f.sums[name] = sha256.Sum256(data) + err = syncDirectory(f.params.Dir) + } + f.params.Metrics.StateFileWritten(name, len(data), err) return err } +// setAside renames the state file name, an edit that does not parse with +// parseErr, to name.bad, for the admin to mend, and logs it with where in +// the file the error is. If the rename fails, the edit is left as it is, +// and the error returned is parseErr joined with the rename's. +func (f *Files) setAside(name string, parseErr error) error { + path := filepath.Join(f.params.Dir, name) + + err := os.Rename(path, path+".bad") + if err != nil { + return errors.Join(parseErr, err) + } + + f.params.ProcessLog.Error("set aside an edit of a state file that does not parse", + "file", path+".bad", "error", parseErr.Error()) + f.params.Metrics.StateFileEditSetAside(name) + + 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. func newBanEntry(ban bans.Ban) banEntry { entry := banEntry{Netblock: ban.Netblock, Start: ban.Start, Notes: ban.Notes} @@ -389,28 +592,16 @@ func checkWritable(dir string) error { 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, and checks its entries. A missing file leaves file as it -// is. -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 - } - +// parse reads data, what the state file at path holds, into file, a +// pointer to that file's struct, and checks its entries. +func parse(path string, data []byte, file stateFile) error { // 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) + 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) @@ -463,7 +654,7 @@ func position(data []byte, err error) string { // 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. +// name. syncDirectory must follow, so that the rename lasts. func write(dir, name string, data []byte) error { path := filepath.Join(dir, name) temporary := path + ".tmp" @@ -475,10 +666,13 @@ func write(dir, name string, data []byte) error { if err != nil { _ = os.Remove(temporary) - - return err } + return err +} + +// syncDirectory syncs dir to the disk, so that a rename in it lasts. +func syncDirectory(dir string) error { directory, err := os.Open(dir) //nolint:gosec // SWWAF_STATE_DIR itself if err != nil { return err diff --git a/internal/state/state_test.go b/internal/state/state_test.go index 7306687..b508b96 100644 --- a/internal/state/state_test.go +++ b/internal/state/state_test.go @@ -3,7 +3,10 @@ package state_test import ( "context" "encoding/json" + "io/fs" "log/slog" + "maps" + "net" "net/http" "net/http/httptest" "net/netip" @@ -28,6 +31,13 @@ const ( 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. @@ -290,7 +300,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 // returns once Run waits for its next write, so that every write due by // then is on disk. @@ -302,7 +312,7 @@ func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) { dir := t.TempDir() params := newParams(dir) 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 // write off no further, and is written with it. @@ -347,7 +357,7 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) { dir := t.TempDir() params := newParams(dir) 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 // them written again. @@ -364,6 +374,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) { t.Parallel() @@ -459,24 +507,359 @@ func TestWritesAreCountedInTheMetrics(t *testing.T) { float64(len(permanentBansJSON))) } -func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) { +func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) { t.Parallel() dir := t.TempDir() - files := load(t, newParams(dir)) + path := filepath.Join(dir, bansJSON) + params := newParams(dir) + files := load(t, params) - // A directory named bans.json cannot be renamed over. - err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700) + // 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) + wantWriteFailed(t, params, bansJSON) +} + +func TestBrokenEditThatCannotBeSetAsideIsNotWrittenOver(t *testing.T) { + t.Parallel() + + const broken = `{"version": 1, "bans": [` + + dir := t.TempDir() + path := filepath.Join(dir, bansJSON) + params := newParams(dir) + files := load(t, params) + + // A directory named bans.json.bad cannot be renamed over, so the + // broken edit cannot be set aside, and is left as it is. + edit(t, dir, bansJSON, broken) + + err := os.Mkdir(path+".bad", 0o700) if err != nil { t.Fatalf("mkdir: %v", err) } err = files.WriteAll() if err == nil { - t.Error("writing over a directory did not fail") + t.Error("writing with bans.json.bad in the way did not fail") } + if got := readFile(t, path); got != broken { + t.Errorf("bans.json holds\n%s\nwant the edit", got) + } + + wantWriteFailed(t, params, bansJSON) +} + +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 TestFileRenamedOverAStateFileTakenIn(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + path := filepath.Join(dir, bansJSON) + params := newParams(dir) + lines := logInto(¶ms) + files := load(t, params) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + watch(t, files, lines) + + // The admin mends bans.json.bad and moves it back, as editors that + // save by renaming do with a file of their own: nothing is written + // into bans.json itself. An edit of clients.json after it must be + // taken in second. + edit(t, dir, bansJSON+".bad", permanentBansJSON) + + err = os.Rename(path+".bad", path) + if err != nil { + t.Fatalf("rename: %v", err) + } + + edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`) + + wantTakenIn(t, lines, dir, bansJSON) + wantEqual(t, bansJSON, params.Ledger.Snapshot(), []bans.Ban{permanentBan()}) +} + +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 TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + lines := logInto(¶ms) + files := load(t, params) + + // An edit taken in by Watch, which is then stopped. + ctx, stop := context.WithCancel(t.Context()) + stopped := make(chan struct{}) + + go func() { + files.Watch(ctx) + close(stopped) + }() + + lines.waitFor(t, watching) + edit(t, dir, bansJSON, `{"version": 1, "bans": []}`) + byWatch := lines.waitFor(t, tookIn) + + stop() + <-stopped + + // An edit taken in by the write of its file. Nothing logs after the + // write, so the log is closed, and a write that does not log the edit + // fails the test at once instead of waiting for the line. + edit(t, dir, bansJSON, permanentBansJSON) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + close(lines) + + byWrite := lines.waitFor(t, tookIn) + + // The two lines differ only in their time. + delete(byWatch, "time") + delete(byWrite, "time") + + if !maps.Equal(byWrite, byWatch) { + t.Errorf("the write logged %v, where Watch logged %v", byWrite, byWatch) + } +} + +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. @@ -573,15 +956,16 @@ func load(t *testing.T, params state.Params) *state.Files { return files } -// run runs files' writes until the test ends. -func run(t *testing.T, files *state.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() { - files.Run(ctx) + task(ctx) close(stopped) }() @@ -591,6 +975,80 @@ 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, or nil if the log is closed first. 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 +} + +// 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) { @@ -735,6 +1193,18 @@ func metric(t *testing.T, text, series string) float64 { return 0 } +// wantWriteFailed checks that the metrics of params count one write of the +// state file name, and that it failed. +func wantWriteFailed(t *testing.T, params state.Params, name string) { + t.Helper() + + got := scrape(t, params) + file := `{file="` + name + `"}` + + wantMetric(t, got, "smallwebwaf_state_file_writes_total"+file, 1) + wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+file, 1) +} + // wantMetric checks the value of series in text, the metrics, as metric // reads it. func wantMetric(t *testing.T, text, series string, want float64) { diff --git a/internal/state/write_internal_test.go b/internal/state/write_internal_test.go new file mode 100644 index 0000000..698fb3b --- /dev/null +++ b/internal/state/write_internal_test.go @@ -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) + } +}