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