Keep the bans, the clients and GeoJS's answers in state files (closes #17)
check / check (push) Successful in 3m24s

smallwebwaf now copies its state to bans.json, clients.json and
lookups.json in SWWAF_STATE_DIR, as "Persistent state" in SPEC.md
describes, and reads them back at start, so a restart lifts no ban and
gives no client a fresh allowance. Each client gains a history, and a
ban's notes count the netblock's requests. bans.json is written
SWWAF_STATE_WRITE_DELAY after a ban, and every file every
SWWAF_STATE_COUNTER_INTERVAL and at the stop. A ban read back is masked
to its netblock and refuses every client in it. A file that does not
parse, an unknown version, an entry without a field it needs, or an
unwritable directory stops the start.

Deviation: no AS number or name, and no ban cause, reason or lifting yet.

Model: opus-5-5
This commit was merged in pull request #72.
This commit is contained in:
2026-10-06 08:31:52 +02:00
parent 73ca94f850
commit df2c5042d2
27 changed files with 2859 additions and 272 deletions
+142 -46
View File
@@ -1,6 +1,7 @@
// Package bans is the ban ledger: the bans smallwebwaf makes on the
// netblocks of clients that break a rate limit, with their notes, as the
// "Bans" section of SPEC.md describes. The bans are kept in memory only.
// "Bans" section of SPEC.md describes. The bans are kept in memory, and
// written to bans.json and read from it by the state package.
package bans
import (
@@ -58,48 +59,65 @@ func (b Ban) ActiveAt(now time.Time) bool {
return b.Permanent() || now.Before(b.Expires)
}
// Notes are what an admin needs to decide whether to lift a ban.
// Notes are what an admin needs to decide whether to lift a ban. The
// JSON names are those of bans.json.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Notes struct {
// Country is the client's country, when it was looked up.
Country string
Country string `json:"country"`
// Limit, Window and Count are the limit that was broken, its window,
// "minute", "hour" or "day", and the count reached: the client's
// requests in the window, the one that broke the limit included.
// These are the requests that counted toward the ban, and the window
// is the time over which they came.
Limit int64
Window string
Count float64
Limit int64 `json:"limit"`
Window string `json:"window"`
Count float64 `json:"count"`
// Request is the request that broke the limit.
Request Request
// Refused is how many requests the ban has refused so far.
Refused int64
Request Request `json:"request"`
// Requests is how many requests the netblock has sent since it was
// first seen, and Refused how many of them the ban has refused so
// far. Both go up with each request the ban refuses.
Requests int64 `json:"requests"`
Refused int64 `json:"refused"`
// EarlierBans is how many bans the netblock had before this one.
EarlierBans int
EarlierBans int `json:"earlier_bans"`
}
// Request is a request in a ban's notes. Each text is cut to 256 bytes.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Request struct {
Time time.Time
Method string
Host string
Time time.Time `json:"time"`
Method string `json:"method"`
Host string `json:"host"`
// Path is the path with its query string.
Path string
Path string `json:"path"`
// Status is what the client was sent, 0 if nothing was.
Status int
UserAgent string
Status int `json:"status"`
UserAgent string `json:"user_agent"`
}
// Ledger holds the bans. It is safe for concurrent use.
type Ledger struct {
rules Rules
// changed receives a value when a ban is made, unless one is waiting
// already.
changed chan struct{}
mu sync.Mutex
// netblocks holds each banned netblock's bans, oldest first. Each
// request from a netblock makes it the most recently seen.
// netblocks holds each banned netblock's bans, oldest first. Check
// makes each netblock it finds the most recently seen.
netblocks *simplelru.LRU[netip.Prefix, *[]Ban]
// held is how many bans netblocks holds, at most rules.MaxBans.
held int
// v4Lengths and v6Lengths are the lengths of the IPv4 and IPv6
// netblocks that have been banned. Check looks for a ban at each of
// them, so that a ban read from bans.json refuses every client in its
// netblock even when it was made with another SWWAF_BAN_SCOPE_V4_PREFIX,
// or another length of an IPv6 client's netblock.
v4Lengths, v6Lengths []int
}
// New returns a Ledger with no ban yet.
@@ -111,31 +129,49 @@ func New(rules Rules) *Ledger {
panic(err) // NewLRU fails only for a size below one
}
return &Ledger{rules: rules, netblocks: netblocks}
return &Ledger{
rules: rules,
changed: make(chan struct{}, 1),
netblocks: netblocks,
}
}
// Check is called for each request from netblock, at now. It reports
// whether a ban on netblock is active, and returns that ban, with the
// request counted among those it refused.
func (l *Ledger) Check(netblock netip.Prefix, now time.Time) (Ban, bool) {
// Changed receives a value after a ban is made, so that bans.json can be
// written. Several bans made before it is read leave one value.
func (l *Ledger) Changed() <-chan struct{} {
return l.changed
}
// Check is called for each request from client, at now. It reports
// whether a ban on a netblock client is in is active, and returns that
// ban, with the request counted among those it refused.
func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool) {
l.mu.Lock()
defer l.mu.Unlock()
bans, found := l.netblocks.Get(netblock)
if !found {
return Ban{}, false
lengths := l.v6Lengths
if client.Is4() {
lengths = l.v4Lengths
}
// 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 Ban{}, false
for _, length := range lengths {
bans, found := l.netblocks.Get(netip.PrefixFrom(client, length).Masked())
if !found {
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) {
last.Notes.Requests++
last.Notes.Refused++
return *last, true
}
}
last.Notes.Refused++
return *last, true
return Ban{}, false
}
// BanForLimit bans netblock at now for a broken limit, with notes, and
@@ -169,21 +205,13 @@ func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes)
Expires: l.expiry(last, now),
Notes: notes,
}
l.add(ban)
if l.held == l.rules.MaxBans {
l.dropOne()
select {
case l.changed <- struct{}{}:
default: // a value is waiting already
}
// dropOne can have dropped netblock's last ban, and netblock with it.
bans, found = l.netblocks.Peek(netblock)
if !found {
bans = &[]Ban{}
l.netblocks.Add(netblock, bans)
}
*bans = append(*bans, ban)
l.held++
return ban
}
@@ -201,6 +229,74 @@ func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
return slices.Clone(*bans)
}
// Snapshot returns every ban held, sorted by netblock, and each
// netblock's bans oldest first, as bans.json lists them.
func (l *Ledger) Snapshot() []Ban {
l.mu.Lock()
defer l.mu.Unlock()
held := make([]Ban, 0, l.held)
for _, bans := range l.netblocks.Values() {
held = append(held, *bans...)
}
slices.SortStableFunc(held, func(a, b Ban) int {
return a.Netblock.Compare(b.Netblock)
})
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.
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)
})
for _, ban := range bans {
ban.Netblock = ban.Netblock.Masked()
ban.Notes.Request = ban.Notes.Request.cut()
l.add(ban)
}
}
// add adds ban to its netblock's bans, after the last, and makes its
// netblock the most recently seen. With MaxBans held, it drops one first.
func (l *Ledger) add(ban Ban) {
if l.held == l.rules.MaxBans {
l.dropOne()
}
// dropOne can have dropped the netblock's last ban, and the netblock
// with it.
bans, found := l.netblocks.Get(ban.Netblock)
if !found {
bans = &[]Ban{}
l.netblocks.Add(ban.Netblock, bans)
}
*bans = append(*bans, ban)
l.held++
lengths := &l.v6Lengths
if ban.Netblock.Addr().Is4() {
lengths = &l.v4Lengths
}
if !slices.Contains(*lengths, ban.Netblock.Bits()) {
*lengths = append(*lengths, ban.Netblock.Bits())
}
}
// 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.
+12 -10
View File
@@ -39,7 +39,7 @@ func TestRepeatsTripleUntilPermanent(t *testing.T) {
t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires)
}
_, banned := ledger.Check(netblock, now.Add(100*365*day))
_, banned := ledger.Check(netblock.Addr(), now.Add(100*365*day))
if !banned {
t.Error("a permanent ban ended")
}
@@ -135,28 +135,30 @@ func TestCheckRefusesWhileTheBanLastsAndCountsTheRefusals(t *testing.T) {
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
for range 3 {
got, banned := ledger.Check(netblock, ban.Expires.Add(-time.Nanosecond))
got, banned := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
if !banned || got.Start != ban.Start {
t.Fatalf("check during the ban gives %+v and %t", got, banned)
}
}
_, banned := ledger.Check(netip.MustParsePrefix("203.0.113.10/32"), midnight())
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight())
if banned {
t.Error("another netblock is banned")
}
_, banned = ledger.Check(netblock, ban.Expires)
_, banned = ledger.Check(netblock.Addr(), ban.Expires)
if banned {
t.Error("the ban did not end")
}
refused := ledger.Bans(netblock)[0].Notes.Refused
if refused != 3 {
t.Errorf("the notes count %d refused requests, want 3", refused)
// The netblock's requests went from 5 to 8 with the three refused.
notes := ledger.Bans(netblock)[0].Notes
if notes.Refused != 3 || notes.Requests != 8 {
t.Errorf("the notes count %d refused requests of %d, want 3 of 8",
notes.Refused, notes.Requests)
}
}
@@ -178,7 +180,7 @@ func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
// A request from a makes b the netblock seen longest ago, and its ban
// goes to make room for d's.
ledger.Check(a, now)
ledger.Check(a.Addr(), now)
ledger.BanForLimit(d, now, bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 0, c: 1, d: 1})
@@ -188,7 +190,7 @@ func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
// With d seen since, a is seen longest ago, and its earlier ban goes
// first.
ledger.Check(d, first.Expires)
ledger.Check(d.Addr(), first.Expires)
ledger.BanForLimit(b, first.Expires, bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 1, d: 1})
+193
View File
@@ -0,0 +1,193 @@
package bans_test
import (
"net/netip"
"slices"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
)
func TestChangedAfterABanIsMade(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
wantChanged(t, ledger, false)
ledger.BanForLimit(netblock, midnight(), bans.Notes{})
wantChanged(t, ledger, true)
// A limit broken during the ban makes no other, and a refusal changes
// only the counts in the notes, which wait for the interval's write.
ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
ledger.Check(netblock.Addr(), midnight().Add(time.Minute))
wantChanged(t, ledger, false)
// Two bans before the value is read leave one.
ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"), midnight(), bans.Notes{})
ledger.BanForLimit(netip.MustParsePrefix("203.0.113.11/32"), midnight(), bans.Notes{})
wantChanged(t, ledger, true)
wantChanged(t, ledger, false)
}
func TestSnapshotListsEveryBanByNetblock(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
v6 := netip.MustParsePrefix("2001:db8::/64")
high := netip.MustParsePrefix("203.0.113.10/32")
low := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(v6, midnight(), bans.Notes{})
ledger.BanForLimit(high, midnight(), bans.Notes{})
ledger.BanForLimit(low, midnight(), bans.Notes{})
ledger.BanForLimit(v6, first.Expires, bans.Notes{})
snapshot := ledger.Snapshot()
got := make([]string, 0, len(snapshot))
for _, ban := range snapshot {
got = append(got, ban.Netblock.String()+" "+ban.Start.Format(time.Kitchen))
}
want := []string{
"203.0.113.9/32 12:00AM", "203.0.113.10/32 12:00AM",
"2001:db8::/64 12:00AM", "2001:db8::/64 1:00AM",
}
if !slices.Equal(got, want) {
t.Errorf("snapshot %v, want %v", got, want)
}
}
func TestLoadedBansCarryOn(t *testing.T) {
t.Parallel()
before := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
ban := before.BanForLimit(netblock, midnight(), bans.Notes{Limit: 1})
// Loaded into a new ledger, as across a restart, the ban still refuses
// while it lasts, and once it has ended a broken limit bans for three
// times as long, with the loaded ban counted among the earlier ones.
after := bans.New(defaultRules())
after.Load(before.Snapshot())
_, banned := after.Check(netblock.Addr(), ban.Expires.Add(-time.Second))
if !banned {
t.Error("the loaded ban does not refuse")
}
again := after.BanForLimit(netblock, ban.Expires, bans.Notes{})
if again.Expires.Sub(again.Start) != 3*time.Hour || again.Notes.EarlierBans != 1 {
t.Errorf("the next ban lasts %s with %d earlier bans, want 3h and 1",
again.Expires.Sub(again.Start), again.Notes.EarlierBans)
}
}
func TestLoadedBanRefusesEveryClientInItsNetblock(t *testing.T) {
t.Parallel()
// Two entries as an admin might write them, with addresses not masked
// to their lengths, the IPv6 one shorter than the /64 an IPv6 client's
// ban covers, beside a ban the ledger makes on one IPv4 address.
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{
{Netblock: netip.MustParsePrefix("203.0.113.9/24"), Start: midnight()},
{Netblock: netip.MustParsePrefix("2001:db8::1/48"), Start: midnight()},
})
ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(), bans.Notes{})
for client, want := range map[string]bool{
"203.0.113.0": true,
"203.0.113.200": true,
"203.0.114.1": false,
"2001:db8:0:5::1": true,
"2001:db8:1::1": false,
"198.51.100.7": true,
"198.51.100.8": false,
} {
_, banned := ledger.Check(netip.MustParseAddr(client), midnight())
if banned != want {
t.Errorf("%s is refused: %t, want %t", client, banned, want)
}
}
// The loaded netblocks are written back masked.
snapshot := ledger.Snapshot()
got := make([]string, 0, len(snapshot))
for _, ban := range snapshot {
got = append(got, ban.Netblock.String())
}
want := []string{"198.51.100.7/32", "203.0.113.0/24", "2001:db8::/48"}
if !slices.Equal(got, want) {
t.Errorf("the ledger holds bans on %v, want %v", got, want)
}
}
func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
t.Parallel()
// bans.json lists the bans by netblock, not in the order they began.
later := bans.Ban{Netblock: netip.MustParsePrefix("203.0.113.1/32"), Start: midnight()}
earlier := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
Start: midnight().Add(-time.Hour),
}
rules := defaultRules()
rules.MaxBans = 1
ledger := bans.New(rules)
ledger.Load([]bans.Ban{later, earlier})
held := ledger.Snapshot()
if len(held) != 1 || held[0] != later {
t.Errorf("the ledger holds %+v, want only the ban that began later", held)
}
}
func TestLoadCutsTheTextsTo256Bytes(t *testing.T) {
t.Parallel()
long := strings.Repeat("a", 300)
ban := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.9/32"),
Start: midnight(),
Notes: bans.Notes{Request: bans.Request{
Method: long, Host: long, Path: long, UserAgent: long,
}},
}
ledger := bans.New(defaultRules())
ledger.Load([]bans.Ban{ban})
cut := long[:256]
want := bans.Request{Method: cut, Host: cut, Path: cut, UserAgent: cut}
got := ledger.Snapshot()[0].Notes.Request
if got != want {
t.Errorf("the notes keep %+v, want each text cut to 256 bytes", got)
}
}
// wantChanged checks whether the ledger's Changed has a value to read.
func wantChanged(t *testing.T, ledger *bans.Ledger, want bool) {
t.Helper()
got := false
select {
case <-ledger.Changed():
got = true
default:
}
if got != want {
t.Errorf("Changed has a value: %t, want %t", got, want)
}
}
+24
View File
@@ -12,6 +12,7 @@ import (
"net/http"
"net/netip"
"net/url"
"path/filepath"
"slices"
"strconv"
"strings"
@@ -94,6 +95,14 @@ type Config struct {
// BanScopeV4Prefix is the length of the netblock around an IPv4
// client that a ban covers (SWWAF_BAN_SCOPE_V4_PREFIX).
BanScopeV4Prefix int
// StateDir is the directory of the state files, an absolute path
// (SWWAF_STATE_DIR). bans.json is written StateWriteDelay after a ban
// is made (SWWAF_STATE_WRITE_DELAY), and every state file every
// StateCounterInterval (SWWAF_STATE_COUNTER_INTERVAL). Neither can be
// off.
StateDir string
StateWriteDelay time.Duration
StateCounterInterval time.Duration
// settings are the values read, as given or by default, for the
// log line at start.
@@ -139,6 +148,8 @@ var (
errNotBanResponse = errors.New("is not 403, 429 or close")
errNotV4Prefix = errors.New(
"is not the length of an IPv4 netblock, from 0 to 32, such as 24")
errNotAbsolutePath = errors.New(
"is not an absolute path, such as /var/lib/smallwebwaf")
)
// FromEnvironment reads the settings with lookupEnv, normally
@@ -174,6 +185,9 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
MaxBanDuration: env.durationNotOff("SWWAF_MAX_BAN_DURATION", "7d"),
MaxBans: env.numberNotOff("SWWAF_MAX_BANS", "5000"),
BanScopeV4Prefix: env.v4Prefix("SWWAF_BAN_SCOPE_V4_PREFIX", "32"),
StateDir: env.absolutePath("SWWAF_STATE_DIR", "/var/lib/smallwebwaf"),
StateWriteDelay: env.durationNotOff("SWWAF_STATE_WRITE_DELAY", "10s"),
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
}
for _, country := range cfg.ExclusivelyAllowedCountries {
@@ -329,6 +343,16 @@ func (e *environment) v4Prefix(name, defaultValue string) int {
return length
}
// absolutePath reads a setting that is an absolute path.
func (e *environment) absolutePath(name, defaultValue string) string {
path := e.value(name, defaultValue)
if !filepath.IsAbs(path) {
e.check(name, fmt.Errorf("%q %w", path, errNotAbsolutePath))
}
return path
}
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
// whole number of days such as 7d, or off.
func parseDuration(value string) (time.Duration, error) {
+25 -1
View File
@@ -41,6 +41,9 @@ const (
maxBanDuration = "SWWAF_MAX_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS"
banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX"
stateDir = "SWWAF_STATE_DIR"
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
)
// off switches a timeout, a size limit or a rate limit off.
@@ -92,6 +95,9 @@ func TestDefaults(t *testing.T) {
MaxBanDuration: 7 * 24 * time.Hour,
MaxBans: 5000,
BanScopeV4Prefix: 32,
StateDir: "/var/lib/smallwebwaf",
StateWriteDelay: 10 * time.Second,
StateCounterInterval: 15 * time.Minute,
})
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
@@ -136,6 +142,9 @@ func TestValuesAsSet(t *testing.T) {
maxBanDuration: "30d",
maxBans: "100",
banScopeV4Prefix: "24",
stateDir: "/srv/waf-state",
stateWriteDelay: "500ms",
stateCounterInterval: "1h",
})
wantSettings(t, cfg, config.Config{
@@ -157,6 +166,9 @@ func TestValuesAsSet(t *testing.T) {
MaxBanDuration: 30 * 24 * time.Hour,
MaxBans: 100,
BanScopeV4Prefix: 24,
StateDir: "/srv/waf-state",
StateWriteDelay: 500 * time.Millisecond,
StateCounterInterval: time.Hour,
})
if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
@@ -329,6 +341,9 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{maxBanDuration, off}, {maxBanDuration, "1w"},
{maxBans, off}, {maxBans, "0"}, {maxBans, "5K"},
{banScopeV4Prefix, "33"}, {banScopeV4Prefix, "-1"}, {banScopeV4Prefix, "/24"},
{stateDir, ""}, {stateDir, "state"}, {stateDir, "./var/lib/smallwebwaf"},
{stateWriteDelay, off}, {stateWriteDelay, "0s"},
{stateCounterInterval, off}, {stateCounterInterval, "15"},
} {
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
t.Parallel()
@@ -389,6 +404,9 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
maxBanDuration: "7d",
maxBans: "5000",
banScopeV4Prefix: "32",
stateDir: "/var/lib/smallwebwaf",
stateWriteDelay: "10s",
stateCounterInterval: "15m",
}
if !maps.Equal(line.Settings, want) {
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
@@ -417,7 +435,7 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
wantBanSettings(t, got, want)
}
// wantBanSettings checks the settings for bans.
// wantBanSettings checks the settings for bans and the state files.
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper()
@@ -429,6 +447,12 @@ func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
got.BanScopeV4Prefix != want.BanScopeV4Prefix {
t.Errorf("ban settings\n%+v\nwant\n%+v", got, want)
}
if got.StateDir != want.StateDir ||
got.StateWriteDelay != want.StateWriteDelay ||
got.StateCounterInterval != want.StateCounterInterval {
t.Errorf("state settings\n%+v\nwant\n%+v", got, want)
}
}
// wantNetblocks checks a list of netblocks.
+65 -13
View File
@@ -1,6 +1,7 @@
// Package lookup looks up each client's country through the GeoJS web
// service, and keeps the answers in memory, for at most 100,000 clients
// and for 7 days each.
// and for 7 days each. The answers are written to lookups.json and read
// from it by the state package.
package lookup
import (
@@ -12,6 +13,7 @@ import (
"log/slog"
"net/http"
"net/netip"
"slices"
"strings"
"sync"
"time"
@@ -76,7 +78,7 @@ type GeoJS struct {
httpClient *http.Client
mu sync.Mutex
answers *simplelru.LRU[netip.Prefix, answer]
answers *simplelru.LRU[netip.Prefix, *Answer]
// waiting are the clients without an answer: those to ask GeoJS about,
// and those it is being asked about.
waiting map[netip.Prefix]*wait
@@ -88,11 +90,14 @@ type GeoJS struct {
retryAt time.Time
}
// answer is what GeoJS said about a client: its country, "" when GeoJS
// cannot place it, and when GeoJS said so.
type answer struct {
country string
received time.Time
// Answer is what GeoJS said about a client, as lookups.json holds it: its
// country, "" when GeoJS cannot place it, when GeoJS said so, and when
// the answer was last used.
type Answer struct {
Client netip.Prefix `json:"client"`
Country string `json:"country"`
Answered time.Time `json:"answered"`
Used time.Time `json:"used"`
}
// wait is a client waiting for its answer.
@@ -107,7 +112,7 @@ type wait struct {
// New returns a GeoJS with no answer kept yet.
func New(params Params) *GeoJS {
answers, err := simplelru.NewLRU[netip.Prefix, answer](maxAnswers, nil)
answers, err := simplelru.NewLRU[netip.Prefix, *Answer](maxAnswers, nil)
if err != nil {
panic(err) // NewLRU fails only for a size below one
}
@@ -164,6 +169,47 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
return country
}
// Snapshot returns every answer kept, sorted by client, as lookups.json
// lists them.
func (g *GeoJS) Snapshot() []Answer {
g.mu.Lock()
answers := make([]Answer, 0, g.answers.Len())
for _, kept := range g.answers.Values() {
answers = append(answers, *kept)
}
g.mu.Unlock()
slices.SortFunc(answers, func(a, b Answer) int {
return a.Client.Compare(b.Client)
})
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
// 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)
})
now := g.now()
for _, answer := range answers {
if now.Sub(answer.Answered) < keepFor {
g.answers.Add(answer.Client, &answer)
}
}
}
// answerOrWait returns client's kept answer if it has one. Otherwise it
// puts the client among those waiting if there is room, has GeoJS asked
// about them if it can be, and returns what to wait on for the answer, or
@@ -203,15 +249,19 @@ func (g *GeoJS) answerOrWait(
return "", w.asked
}
// kept returns client's answer, if one was received less than keepFor
// ago.
// kept returns client's answer, if GeoJS gave it less than keepFor ago,
// and notes that it was used.
func (g *GeoJS) kept(client netip.Prefix) (string, bool) {
now := g.now()
kept, found := g.answers.Get(client)
if !found || g.now().Sub(kept.received) >= keepFor {
if !found || now.Sub(kept.Answered) >= keepFor {
return "", false
}
return kept.country, true
kept.Used = now
return kept.Country, true
}
// ask starts asking GeoJS about the waiting clients, unless a request to
@@ -293,7 +343,9 @@ func (g *GeoJS) keep(
continue
}
g.answers.Add(client, answer{country: country, received: now})
g.answers.Add(client, &Answer{
Client: client, Country: country, Answered: now, Used: now,
})
close(g.waiting[client].asked)
delete(g.waiting, client)
}
+96
View File
@@ -0,0 +1,96 @@
package lookup_test
import (
"net/netip"
"slices"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/lookup"
)
func TestSnapshotHoldsEachAnswerAndWhenItWasLastUsed(t *testing.T) {
t.Parallel()
_, clock, g := start(t)
placed := netip.MustParsePrefix("203.0.113.9/32")
notPlaced := netip.MustParsePrefix(unplaced + "/32")
asked := clock.Now()
wantCountry(t, g, placed, germany)
wantCountry(t, g, notPlaced, "")
clock.advance(time.Hour)
wantCountry(t, g, placed, germany)
want := []lookup.Answer{
{Client: notPlaced, Country: "", Answered: asked, Used: asked},
{Client: placed, Country: germany, Answered: asked, Used: asked.Add(time.Hour)},
}
if got := g.Snapshot(); !slices.Equal(got, want) {
t.Errorf("snapshot\n%+v\nwant\n%+v", got, want)
}
}
func TestLoadedAnswersAreKeptFor7DaysFromWhenGeoJSGaveThem(t *testing.T) {
t.Parallel()
geojs, clock, g := start(t)
now := clock.Now()
kept := lookup.Answer{
Client: netip.MustParsePrefix("203.0.113.9/32"),
Country: "FR",
Answered: now.Add(-week + time.Second),
Used: now.Add(-time.Hour),
}
stale := lookup.Answer{
Client: netip.MustParsePrefix("203.0.113.10/32"),
Country: "FR",
Answered: now.Add(-week),
Used: now.Add(-time.Hour),
}
g.Load([]lookup.Answer{kept, stale})
if got := g.Snapshot(); !slices.Equal(got, []lookup.Answer{kept}) {
t.Errorf("kept %+v, want only the answer GeoJS gave less than 7 days ago", got)
}
wantCountry(t, g, kept.Client, "FR")
wantRequests(t, geojs, 0)
}
func TestLoadDropsTheAnswerUsedLongestAgoFirst(t *testing.T) {
t.Parallel()
const maxAnswers = 100000
_, clock, g := start(t)
now := clock.Now()
// lookups.json lists the answers by client. Here each was last used a
// second before the one listed before it, so the last listed is the
// one used longest ago, and the one dropped.
answers := make([]lookup.Answer, maxAnswers+1)
addr := netip.MustParseAddr("10.0.0.0")
for i := range answers {
answers[i] = lookup.Answer{
Client: netip.PrefixFrom(addr, addr.BitLen()),
Country: germany,
Answered: now,
Used: now.Add(-time.Duration(i) * time.Second),
}
addr = addr.Next()
}
g.Load(answers)
got := g.Snapshot()
if len(got) != maxAnswers || got[0] != answers[0] ||
got[maxAnswers-1] != answers[maxAnswers-1] {
t.Errorf("%d answers kept, from %s to %s; want %d, from %s to %s",
len(got), got[0].Client, got[len(got)-1].Client, maxAnswers,
answers[0].Client, answers[maxAnswers-1].Client)
}
}
+7 -4
View File
@@ -14,10 +14,10 @@ func (rq *request) banResponse(action string) *refusal {
return &refusal{status: rq.h.config.BanResponse, action: action}
}
// banned reports whether a ban on the client's netblock refuses the
// request at now, and notes for the log line when that ban ends.
// banned reports whether a ban on a netblock the client is in refuses
// the request at now, and notes for the log line when that ban ends.
func (rq *request) banned(now time.Time) bool {
ban, banned := rq.h.ledger.Check(rq.netblock(), now)
ban, banned := rq.h.ledger.Check(rq.client, now)
if banned {
rq.line.BanExpires = banExpires(ban)
}
@@ -36,7 +36,8 @@ func (rq *request) limitBroken(now time.Time) bool {
return false
}
ban := rq.h.ledger.BanForLimit(rq.netblock(), now, bans.Notes{
netblock := rq.netblock()
ban := rq.h.ledger.BanForLimit(netblock, now, bans.Notes{
Country: rq.line.Country,
Limit: hit.Limit,
Window: hit.Window,
@@ -49,6 +50,8 @@ func (rq *request) limitBroken(now time.Time) bool {
Status: rq.h.config.BanResponse,
UserAgent: rq.in.UserAgent(),
},
// The histories count this request only once it has ended.
Requests: rq.h.limiter.Requests(netblock) + 1,
})
rq.h.limiter.Reset(group)
+6 -3
View File
@@ -133,7 +133,7 @@ func TestBanIsNotCountedAndResetsTheCounters(t *testing.T) {
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
banned := proxy.LedgerOf(server).Bans(netip.MustParsePrefix(client + "/32"))
banned := server.Ledger.Bans(netip.MustParsePrefix(client + "/32"))
if len(banned) != 2 || banned[0].Notes.Refused != 3 {
t.Errorf("bans %+v, want two, the first with 3 requests refused", banned)
}
@@ -291,12 +291,15 @@ func TestBanNotes(t *testing.T) {
Status: http.StatusForbidden,
UserAgent: userAgent,
},
// The one let through, the one that broke the limit and the two
// refused under the ban.
Requests: 4,
Refused: 2,
EarlierBans: 0,
},
}
ledger := proxy.LedgerOf(server)
ledger := server.Ledger
got := ledger.Bans(netblock)
if len(got) != 1 || got[0] != want {
@@ -361,7 +364,7 @@ func (c *clock) advance(d time.Duration) {
// set to midnight, the start of a bucket in every window.
func startWithClock(
t *testing.T, geojsURL string, env map[string]string,
) (*sender, *clock, *http.Server) {
) (*sender, *clock, *proxy.Server) {
t.Helper()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
-15
View File
@@ -1,15 +0,0 @@
package proxy
import (
"net/http"
"sneak.berlin/go/smallwebwaf/internal/bans"
)
// LedgerOf returns the ban ledger of a server New returned, so that the
// tests can read the bans' notes.
func LedgerOf(server *http.Server) *bans.Ledger {
h, _ := server.Handler.(*handler)
return h.ledger
}
+103
View File
@@ -0,0 +1,103 @@
package proxy_test
import (
"io"
"net/http"
"net/netip"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
s, clk, server := startWithClock(t, geojsURL, map[string]string{
rateLimitPerMinute: "2",
deniedCountries: "kp",
})
start := clk.Now()
// Two let through, one over the limit, which bans the client, and one
// refused under that ban, for which the country is not looked up.
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
clk.advance(time.Second)
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
clk.advance(time.Second)
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
want := ratelimit.History{
FirstSeen: start,
LastSeen: start.Add(2 * time.Second),
Country: "DE",
LookedUp: start.Add(time.Second),
Requests: 4,
Forwarded: 2,
Refused: 2,
// The app answers with no body, smallwebwaf with its status text.
ResponseBytes: 2 * int64(len("Forbidden\n")),
Responses: ratelimit.Responses{Status2xx: 2, Status4xx: 2},
Offences: ratelimit.Offences{Limit: 1},
}
got := historyOf(t, server, fromDE)
if got != want {
t.Errorf("history\n%+v\nwant\n%+v", got, want)
}
}
func TestHistoryCountsTheBodiesEachWay(t *testing.T) {
t.Parallel()
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
_, _ = io.Copy(io.Discard, r.Body)
_, _ = io.WriteString(w, "hello")
})
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, nil)
got := do(t, newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc")))
wantStatus(t, got, http.StatusOK)
out.requestLine(t)
history := historyOf(t, server, localhost)
if history.RequestBytes != 3 || history.ResponseBytes != 5 {
t.Errorf("history counts %d bytes in and %d out, want 3 and 5",
history.RequestBytes, history.ResponseBytes)
}
}
func TestHealthEndpointIsNotInTheHistory(t *testing.T) {
t.Parallel()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, nil)
wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK)
out.requestLine(t)
if clients := server.Limiter.Snapshot(); len(clients) != 0 {
t.Errorf("the table holds %+v, want no client", clients)
}
}
// historyOf returns the history of the client at addr.
func historyOf(t *testing.T, server *proxy.Server, addr string) ratelimit.History {
t.Helper()
client := netip.MustParsePrefix(addr + "/32")
for _, c := range server.Limiter.Snapshot() {
if c.Client == client {
return c.History
}
}
t.Fatalf("%s is not in the table", client)
return ratelimit.History{}
}
+55 -35
View File
@@ -38,53 +38,70 @@ type Params struct {
// lookup.URL. GeoJS is asked only while a country list is set.
GeoJSURL string
// Now tells the time by which requests are counted for the rate
// limits and bans are made and run out, normally time.Now.
// limits, bans are made and run out, and GeoJS's answers are kept,
// normally time.Now in UTC, the time the state files give.
Now func() time.Time
}
// Server is the server smallwebwaf runs, with the parts of the proxy
// whose state the state files keep.
type Server struct {
*http.Server
Ledger *bans.Ledger
Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS
}
// New returns the server smallwebwaf runs: each request it reads passes
// through the proxy. Go's server itself refuses a request line and
// headers over SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES, with 431, closes a
// connection idle for SWWAF_CLIENT_IDLE_TIMEOUT, and applies
// SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy
// applies the timeouts and size limits from then on.
func New(params Params) *http.Server {
func New(params Params) *Server {
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
h := &handler{
config: params.Config,
requestLog: params.RequestLog,
processLog: params.ProcessLog,
errorLog: errorLog,
transport: newTransport(),
now: params.Now,
limiter: ratelimit.New(ratelimit.Limits{
PerMinute: params.Config.RateLimitPerMinute,
PerHour: params.Config.RateLimitPerHour,
PerDay: params.Config.RateLimitPerDay,
}),
ledger: bans.New(bans.Rules{
LimitBanDuration: params.Config.LimitBanDuration,
LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow,
MaxBanDuration: params.Config.MaxBanDuration,
MaxBans: params.Config.MaxBans,
}),
geojs: lookup.New(lookup.Params{
URL: params.GeoJSURL,
Now: params.Now,
ProcessLog: params.ProcessLog,
}),
}
return &http.Server{
Addr: params.Config.ListenAddr,
Handler: &handler{
config: params.Config,
requestLog: params.RequestLog,
processLog: params.ProcessLog,
errorLog: errorLog,
transport: newTransport(),
now: params.Now,
limiter: ratelimit.New(ratelimit.Limits{
PerMinute: params.Config.RateLimitPerMinute,
PerHour: params.Config.RateLimitPerHour,
PerDay: params.Config.RateLimitPerDay,
}),
ledger: bans.New(bans.Rules{
LimitBanDuration: params.Config.LimitBanDuration,
LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow,
MaxBanDuration: params.Config.MaxBanDuration,
MaxBans: params.Config.MaxBans,
}),
geojs: lookup.New(lookup.Params{
URL: params.GeoJSURL,
Now: time.Now,
ProcessLog: params.ProcessLog,
}),
return &Server{
Server: &http.Server{
Addr: params.Config.ListenAddr,
Handler: h,
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
// Off is an IdleTimeout of 0, which Go's server replaces with
// ReadTimeout: no limit, as long as ReadTimeout stays unset.
IdleTimeout: params.Config.ClientIdleTimeout,
// Go's server reads 4 KiB past MaxHeaderBytes before it
// refuses, so the limit a client meets is the setting.
MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
ErrorLog: errorLog,
},
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
// Off is an IdleTimeout of 0, which Go's server replaces with
// ReadTimeout: no limit, as long as ReadTimeout stays unset.
IdleTimeout: params.Config.ClientIdleTimeout,
// Go's server reads 4 KiB past MaxHeaderBytes before it refuses,
// so the limit a client meets is the setting.
MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
ErrorLog: errorLog,
Ledger: h.ledger,
Limiter: h.limiter,
GeoJS: h.geojs,
}
}
@@ -130,6 +147,9 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return
}
// Once the request has ended, before its log line is written.
defer rq.addToHistory()
refused := rq.check(r.Context())
if refused != nil {
rq.answer(*refused)
+1 -1
View File
@@ -200,7 +200,7 @@ func startProxyWithGeoJS(
func startProxyWithClock(
t *testing.T, appURL, geojsURL string, now func() time.Time,
env map[string]string,
) (string, *output, *http.Server) {
) (string, *output, *proxy.Server) {
t.Helper()
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL}
+19
View File
@@ -12,6 +12,7 @@ import (
"sync/atomic"
"time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
@@ -318,6 +319,24 @@ func (rq *request) finish() {
}
}
// addToHistory adds the request, which has ended, to its client's
// history.
func (rq *request) addToHistory() {
var requestBytes int64
if rq.body != nil {
requestBytes = rq.body.bytes.Load()
}
rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{
Country: rq.line.Country,
Forwarded: !rq.upstreamStart.IsZero(),
Status: rq.out.status,
RequestBytes: requestBytes,
ResponseBytes: rq.out.bytes,
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
})
}
// clientRequestDeadline is when the client must have sent its whole
// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off.
func (rq *request) clientRequestDeadline() time.Time {
+116
View File
@@ -0,0 +1,116 @@
package ratelimit_test
import (
"net/netip"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
func TestHistoryKeepsEveryRequest(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for i, r := range []ratelimit.Request{
{Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100},
{Forwarded: true, Status: 101},
{Forwarded: true, Status: 304, RequestBytes: 5},
{Country: "FR", Status: 403, ResponseBytes: 10, BrokeLimit: true},
{Forwarded: true, Status: 502, ResponseBytes: 12},
// Closed without an answer: refused, and no response.
{Status: 0},
} {
limiter.AddToHistory(client, start.Add(time.Duration(i)*time.Minute), r)
}
want := ratelimit.History{
FirstSeen: start,
LastSeen: start.Add(5 * time.Minute),
Country: "FR",
LookedUp: start.Add(3 * time.Minute),
Requests: 6,
Forwarded: 4,
Refused: 2,
RequestBytes: 15,
ResponseBytes: 122,
Responses: ratelimit.Responses{
Status1xx: 1, Status2xx: 1, Status3xx: 1, Status4xx: 1, Status5xx: 1,
},
Offences: ratelimit.Offences{Limit: 1},
}
got := historyOf(t, limiter, client)
if got != want {
t.Errorf("history\n%+v\nwant\n%+v", got, want)
}
}
func TestResetKeepsTheHistory(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range limit {
wantCount(t, limiter, client, start, "")
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
}
limiter.Reset(client)
if got := historyOf(t, limiter, client).Requests; got != limit {
t.Errorf("the history counts %d requests, want %d", got, limit)
}
}
func TestRequestsAddsUpTheClientsInsideTheNetblock(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{})
for client, requests := range map[string]int{
"198.51.100.9/32": 2,
"198.51.100.10/32": 3,
"192.0.2.1/32": 5,
"2001:db8:5::/64": 7,
} {
for range requests {
limiter.AddToHistory(netip.MustParsePrefix(client), midnight(),
ratelimit.Request{})
}
}
for netblock, want := range map[string]int64{
"198.51.100.9/32": 2,
"198.51.100.0/24": 5,
"2001:db8:5::/64": 7,
"203.0.113.0/24": 0,
} {
got := limiter.Requests(netip.MustParsePrefix(netblock))
if got != want {
t.Errorf("%s has sent %d requests, want %d", netblock, got, want)
}
}
}
// historyOf returns client's history.
func historyOf(
t *testing.T, limiter *ratelimit.Limiter, client netip.Prefix,
) ratelimit.History {
t.Helper()
for _, c := range limiter.Snapshot() {
if c.Client == client {
return c.History
}
}
t.Fatalf("%s is not in the table", client)
return ratelimit.History{}
}
+249 -42
View File
@@ -1,11 +1,15 @@
// Package ratelimit counts each client's requests over a minute, an hour
// and a day, as the "Counting method" section of SPEC.md describes, and
// tells when a request takes a client over a rate limit. The counts are
// kept in memory only, for at most 20,000 clients.
// Package ratelimit keeps the table of clients: each client's requests
// counted over a minute, an hour and a day, as the "Counting method"
// section of SPEC.md describes, which tell when a request takes the client
// over a rate limit, and each client's history since it was first seen.
// At most 20,000 clients are kept, in memory, and written to clients.json
// and read from it by the state package.
package ratelimit
import (
"net/http"
"net/netip"
"slices"
"sync"
"time"
@@ -13,7 +17,8 @@ import (
)
// maxClients is how many clients are kept. Past it, the least recently
// seen client is dropped, and starts afresh if it comes back.
// seen client is dropped, with its history, and starts afresh if it comes
// back.
const maxClients = 20000
const day = 24 * time.Hour
@@ -26,20 +31,94 @@ type Limits struct {
PerDay int64
}
// Limiter counts each client's requests against the limits. It is safe
// for concurrent use.
// Limiter counts each client's requests against the limits, and keeps
// its history. It is safe for concurrent use.
type Limiter struct {
// windows are the minute, the hour and the day, in the order of
// Client.buckets.
windows [3]window
mu sync.Mutex
// clients holds each client's buckets, one pair for each of windows,
// in the same order.
clients *simplelru.LRU[netip.Prefix, *[3]buckets]
mu sync.Mutex
clients *simplelru.LRU[netip.Prefix, *Client]
}
// Client is a client in the table, as clients.json holds it: its buckets
// in each window, and its history.
type Client struct {
Client netip.Prefix `json:"client"`
Minute Buckets `json:"minute"`
Hour Buckets `json:"hour"`
Day Buckets `json:"day"`
History History `json:"history"`
}
// Buckets are a client's two buckets in one window: the requests in the
// bucket under way, which began at Start, and in the bucket before it.
type Buckets struct {
Start time.Time `json:"start"`
Current int64 `json:"current"`
Previous int64 `json:"previous"`
}
// History is what is known of a client since it was first seen.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type History struct {
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
// Country is the client's country as it was last looked up, and
// LookedUp when that was; both are empty while it never was.
Country string `json:"country,omitempty"`
LookedUp time.Time `json:"looked_up,omitzero"`
// Requests are all the client's requests: Forwarded those passed to
// the app, Refused those refused before anything reached it.
Requests int64 `json:"requests"`
Forwarded int64 `json:"forwarded"`
Refused int64 `json:"refused"`
// RequestBytes and ResponseBytes are the body bytes of its requests
// and of the responses it was sent.
RequestBytes int64 `json:"request_bytes"`
ResponseBytes int64 `json:"response_bytes"`
Responses Responses `json:"responses,omitzero"`
Offences Offences `json:"offences,omitzero"`
}
// Responses are the responses a client was sent, by status class;
// Status5xx counts every status from 500 up.
type Responses struct {
Status1xx int64 `json:"1xx,omitempty"`
Status2xx int64 `json:"2xx,omitempty"`
Status3xx int64 `json:"3xx,omitempty"`
Status4xx int64 `json:"4xx,omitempty"`
Status5xx int64 `json:"5xx,omitempty"`
}
// Offences are a client's offences, by kind.
type Offences struct {
// Limit is its requests that broke a rate limit.
Limit int64 `json:"limit"`
}
// Request is what a client's history keeps of one of its requests.
type Request struct {
// Country is the client's country, when the request looked it up.
Country string
// Forwarded is true for a request passed to the app, false for one
// refused before anything reached it.
Forwarded bool
// Status is what the client was sent, 0 if nothing was.
Status int
// RequestBytes and ResponseBytes are the body bytes of the request
// and of its response.
RequestBytes int64
ResponseBytes int64
// BrokeLimit is true for a request that broke a rate limit.
BrokeLimit bool
}
// New returns a Limiter for limits, with no client counted yet.
func New(limits Limits) *Limiter {
clients, err := simplelru.NewLRU[netip.Prefix, *[3]buckets](maxClients, nil)
clients, err := simplelru.NewLRU[netip.Prefix, *Client](maxClients, nil)
if err != nil {
panic(err) // NewLRU fails only for a size below one
}
@@ -73,16 +152,12 @@ func (l *Limiter) Count(client netip.Prefix, now time.Time) (Hit, bool) {
l.mu.Lock()
defer l.mu.Unlock()
counts, seen := l.clients.Get(client)
if !seen {
counts = &[3]buckets{}
l.clients.Add(client, counts)
}
var hit Hit
for i, w := range l.windows {
requests := counts[i].add(now, w.length)
for i, b := range l.get(client).buckets() {
w := l.windows[i]
requests := b.add(now, w.length)
if hit.Window == "" && w.limit > 0 && requests > float64(w.limit) {
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests}
}
@@ -91,12 +166,135 @@ func (l *Limiter) Count(client netip.Prefix, now time.Time) (Hit, bool) {
return hit, hit.Window != ""
}
// Reset sets client's counts in every window back to zero.
// Reset sets client's counts in every window back to zero. Its history
// keeps its totals.
func (l *Limiter) Reset(client netip.Prefix) {
l.mu.Lock()
defer l.mu.Unlock()
l.clients.Remove(client)
c, seen := l.clients.Peek(client)
if seen {
c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{}
}
}
// AddToHistory adds r, a request from client at now, to the client's
// history.
func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
l.mu.Lock()
defer l.mu.Unlock()
h := &l.get(client).History
if h.FirstSeen.IsZero() {
h.FirstSeen = now
}
h.LastSeen = now
if r.Country != "" {
h.Country = r.Country
h.LookedUp = now
}
h.Requests++
if r.Forwarded {
h.Forwarded++
} else {
h.Refused++
}
h.RequestBytes += r.RequestBytes
h.ResponseBytes += r.ResponseBytes
h.Responses.add(r.Status)
if r.BrokeLimit {
h.Offences.Limit++
}
}
// Requests returns how many requests the clients inside netblock have
// sent, as their histories count them.
func (l *Limiter) Requests(netblock netip.Prefix) int64 {
l.mu.Lock()
defer l.mu.Unlock()
// Most often the netblock is one client.
c, seen := l.clients.Peek(netblock)
if seen {
return c.History.Requests
}
var requests int64
for _, c := range l.clients.Values() {
if netblock.Overlaps(c.Client) {
requests += c.History.Requests
}
}
return requests
}
// Snapshot returns every client in the table, sorted by address, as
// clients.json lists them.
func (l *Limiter) Snapshot() []Client {
l.mu.Lock()
clients := make([]Client, 0, l.clients.Len())
for _, c := range l.clients.Values() {
clients = append(clients, *c)
}
l.mu.Unlock()
slices.SortFunc(clients, func(a, b Client) int {
return a.Client.Compare(b.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.
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)
})
for _, c := range clients {
for i, b := range c.buckets() {
// The window that ends at now covers neither bucket once it
// begins after the bucket under way has ended.
length := l.windows[i].length
if !now.Add(-length).Before(b.Start.Add(length)) {
*b = Buckets{}
}
}
l.clients.Add(c.Client, &c)
}
}
// get returns client's entry in the table, a new one if it has none, and
// makes it the most recently seen.
func (l *Limiter) get(client netip.Prefix) *Client {
c, seen := l.clients.Get(client)
if !seen {
c = &Client{Client: client}
l.clients.Add(client, c)
}
return c
}
// buckets returns c's buckets in the minute, the hour and the day.
func (c *Client) buckets() [3]*Buckets {
return [3]*Buckets{&c.Minute, &c.Hour, &c.Day}
}
// window is a length of time over which requests are counted, and the
@@ -107,14 +305,6 @@ type window struct {
limit int64
}
// buckets are a client's two buckets in one window: the requests in the
// bucket under way, which began at start, and in the bucket before it.
type buckets struct {
start time.Time
current int64
previous int64
}
// add counts a request at now in a window of length, and returns the
// client's requests in the window that ends at now: those in the bucket
// under way, and those in the bucket before it weighted by how much of
@@ -125,27 +315,44 @@ type buckets struct {
// bucket. A request dated more than a second before it means the clock
// was set back, and the buckets start afresh: otherwise the bucket before
// would keep its full weight until the clock caught up.
func (b *buckets) add(now time.Time, length time.Duration) float64 {
if now.Before(b.start.Add(-time.Second)) {
*b = buckets{}
func (b *Buckets) add(now time.Time, length time.Duration) float64 {
if now.Before(b.Start.Add(-time.Second)) {
*b = Buckets{}
}
start := now.Truncate(length)
if start.After(b.start) {
if start.Equal(b.start.Add(length)) {
b.previous = b.current
if start.After(b.Start) {
if start.Equal(b.Start.Add(length)) {
b.Previous = b.Current
} else {
b.previous = 0
b.Previous = 0
}
b.start = start
b.current = 0
b.Start = start
b.Current = 0
}
b.current++
b.Current++
elapsed := max(now.Sub(b.start), 0)
elapsed := max(now.Sub(b.Start), 0)
covered := 1 - float64(elapsed)/float64(length)
return float64(b.previous)*covered + float64(b.current)
return float64(b.Previous)*covered + float64(b.Current)
}
// add counts a response with status in its class. A status of 0, for
// nothing sent, is not a response.
func (r *Responses) add(status int) {
switch {
case status >= http.StatusInternalServerError:
r.Status5xx++
case status >= http.StatusBadRequest:
r.Status4xx++
case status >= http.StatusMultipleChoices:
r.Status3xx++
case status >= http.StatusOK:
r.Status2xx++
case status >= http.StatusContinue:
r.Status1xx++
}
}
+122
View File
@@ -0,0 +1,122 @@
package ratelimit_test
import (
"net/netip"
"slices"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
func TestSnapshotListsTheClientsByAddress(t *testing.T) {
t.Parallel()
want := []string{"192.0.2.1/32", "203.0.113.9/32", "203.0.113.10/32", "2001:db8::/64"}
limiter := ratelimit.New(ratelimit.Limits{})
for _, i := range []int{2, 3, 0, 1} {
limiter.Count(netip.MustParsePrefix(want[i]), midnight())
}
snapshot := limiter.Snapshot()
got := make([]string, 0, len(snapshot))
for _, c := range snapshot {
got = append(got, c.Client.String())
}
if !slices.Equal(got, want) {
t.Errorf("snapshot %v, want %v", got, want)
}
counted := ratelimit.Buckets{Start: midnight(), Current: 1}
if snapshot[0].Minute != counted || snapshot[0].Day != counted {
t.Errorf("buckets %+v and %+v, want %+v", snapshot[0].Minute, snapshot[0].Day,
counted)
}
}
func TestLoadedCountsCarryOn(t *testing.T) {
t.Parallel()
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
before := ratelimit.New(ratelimit.Limits{PerHour: limit})
for range limit {
wantCount(t, before, client, start, "")
}
// Loaded into a new limiter, as across a restart, the client has no
// fresh allowance.
later := start.Add(time.Minute)
after := ratelimit.New(ratelimit.Limits{PerHour: limit})
after.Load(before.Snapshot(), later)
wantCount(t, after, client, later, hour)
}
func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
t.Parallel()
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
limiter := ratelimit.New(ratelimit.Limits{})
limiter.Count(client, start)
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
loaded := func(now time.Time) ratelimit.Client {
t.Helper()
after := ratelimit.New(ratelimit.Limits{})
after.Load(limiter.Snapshot(), now)
return after.Snapshot()[0]
}
// Two minutes on, the window that ends then covers neither of the
// minute's buckets, which are emptied; the hour's and the day's stay,
// and so does the history.
got := loaded(start.Add(2 * time.Minute))
if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 ||
got.Day.Current != 1 || got.History.Requests != 1 {
t.Errorf("loaded two minutes on as %+v", got)
}
// A moment before, the window still covers some of the earlier one.
got = loaded(start.Add(2*time.Minute - time.Nanosecond))
if got.Minute.Current != 1 {
t.Errorf("loaded just under two minutes on with minute buckets %+v",
got.Minute)
}
}
func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) {
t.Parallel()
const maxClients = 20000
// clients.json lists the clients by address. Here each was last seen
// a second before the one listed before it, so the last listed is the
// one seen longest ago, and the one dropped.
clients := make([]ratelimit.Client, maxClients+1)
addr := netip.MustParseAddr("10.0.0.0")
for i := range clients {
clients[i].Client = netip.PrefixFrom(addr, addr.BitLen())
clients[i].History.LastSeen = midnight().Add(-time.Duration(i) * time.Second)
addr = addr.Next()
}
limiter := ratelimit.New(ratelimit.Limits{})
limiter.Load(clients, midnight())
got := limiter.Snapshot()
if len(got) != maxClients || got[0].Client != clients[0].Client ||
got[maxClients-1].Client != clients[maxClients-1].Client {
t.Errorf("%d clients kept, from %s to %s; want %d, from %s to %s",
len(got), got[0].Client, got[len(got)-1].Client, maxClients,
clients[0].Client, clients[maxClients-1].Client)
}
}
+6 -4
View File
@@ -24,12 +24,14 @@ func TestHealthCheck(t *testing.T) {
out := &output{}
exited := make(chan int, 1)
settings := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: app.URL,
stateDir: t.TempDir(),
}
go func() {
exited <- run(ctx, map[string]string{
listenAddr: localhost + ":0",
upstreamURL: app.URL,
}, out)
exited <- run(ctx, settings, out)
}()
addr, _ := out.line(t, "msg", "starting")["address"].(string)
+64 -17
View File
@@ -1,6 +1,6 @@
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
// serves requests until it is told to stop, and then stops in an orderly
// way.
// Package smallwebwaf runs the smallwebwaf process: it reads the settings
// and the state files, serves requests until it is told to stop, and then
// stops in an orderly way, writing the state files.
package smallwebwaf
import (
@@ -19,6 +19,7 @@ import (
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/state"
)
// shutdownTimeout is how long requests in progress may take to finish
@@ -55,8 +56,9 @@ func Main(version string) int {
})
}
// Run reads the settings, then serves requests until ctx is done. It
// returns the process's exit status, 1 when smallwebwaf cannot start.
// Run reads the settings and the state files, then serves requests until
// ctx is done. It returns the process's exit status, 1 when smallwebwaf
// cannot start.
func Run(ctx context.Context, params Params) int {
processLog := requestlog.NewProcessLogger(params.Stdout)
@@ -67,6 +69,33 @@ func Run(ctx context.Context, params Params) int {
return 1
}
// The state files give times in UTC.
now := func() time.Time { return time.Now().UTC() }
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: params.Stdout,
ProcessLog: processLog,
GeoJSURL: lookup.URL,
Now: now,
})
files, err := state.Load(state.Params{
Dir: cfg.StateDir,
WriteDelay: cfg.StateWriteDelay,
CounterInterval: cfg.StateCounterInterval,
Ledger: server.Ledger,
Limiter: server.Limiter,
GeoJS: server.GeoJS,
Now: now,
ProcessLog: processLog,
})
if err != nil {
processLog.Error("cannot use the state files", "error", err.Error())
return 1
}
listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr)
if err != nil {
processLog.Error("cannot listen on SWWAF_LISTEN_ADDR",
@@ -75,27 +104,20 @@ func Run(ctx context.Context, params Params) int {
return 1
}
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: params.Stdout,
ProcessLog: processLog,
GeoJSURL: lookup.URL,
Now: time.Now,
})
processLog.Info("starting",
"version", params.Version,
"address", listener.Addr().String(),
"settings", cfg)
return serve(ctx, server, listener, processLog)
return serve(ctx, server.Server, listener, files, processLog)
}
// serve serves requests on listener until ctx is done, then gives the
// requests in progress shutdownTimeout to finish.
// 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.
func serve(
ctx context.Context, server *http.Server, listener net.Listener,
processLog *slog.Logger,
files *state.Files, processLog *slog.Logger,
) int {
served := make(chan error, 1)
@@ -103,6 +125,16 @@ func serve(
served <- server.Serve(listener)
}()
writing, stopWriting := context.WithCancel(ctx)
defer stopWriting()
written := make(chan struct{})
go func() {
files.Run(writing)
close(written)
}()
select {
case err := <-served:
processLog.Error("serving failed", "error", err.Error())
@@ -132,6 +164,21 @@ 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.
<-written
err = files.WriteAll()
if err != nil {
processLog.Error("writing the state files failed", "error", err.Error())
return 1
}
processLog.Info("stopped")
return 0
+253 -15
View File
@@ -8,6 +8,8 @@ import (
"net"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync"
"testing"
@@ -24,9 +26,13 @@ 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"
localhost = "127.0.0.1"
listenAddr = "SWWAF_LISTEN_ADDR"
upstreamURL = "SWWAF_UPSTREAM_URL"
stateDir = "SWWAF_STATE_DIR"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
// greeting is what the tests' app answers.
greeting = "hello from the app"
)
// output collects what smallwebwaf writes on stdout.
@@ -69,11 +75,19 @@ func (o *output) line(t *testing.T, key, value string) map[string]any {
time.Sleep(pollInterval)
}
t.Fatalf("no line with %s %q in the output:\n%s", key, value, o.buf.String())
t.Fatalf("no line with %s %q in the output:\n%s", key, value, o.text())
return nil
}
// text returns everything written so far.
func (o *output) text() string {
o.mu.Lock()
defer o.mu.Unlock()
return o.buf.String()
}
// run runs smallwebwaf with the settings in env until ctx is done, and
// returns its exit status.
func run(ctx context.Context, env map[string]string, out *output) int {
@@ -121,7 +135,10 @@ func TestAddressInUseStopsTheStart(t *testing.T) {
out := &output{}
status := run(t.Context(), map[string]string{listenAddr: taken.Addr().String()}, out)
status := run(t.Context(), map[string]string{
listenAddr: taken.Addr().String(),
stateDir: t.TempDir(),
}, out)
if status != 1 {
t.Errorf("exit status %d, want 1", status)
}
@@ -132,11 +149,8 @@ func TestAddressInUseStopsTheStart(t *testing.T) {
func TestServesUntilToldToStop(t *testing.T) {
t.Parallel()
app := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, "hello from the app")
}))
defer app.Close()
appURL := startApp(t)
dir := t.TempDir()
ctx, stop := context.WithCancel(t.Context())
out := &output{}
@@ -145,12 +159,13 @@ func TestServesUntilToldToStop(t *testing.T) {
go func() {
exited <- run(ctx, map[string]string{
listenAddr: localhost + ":0",
upstreamURL: app.URL,
upstreamURL: appURL,
stateDir: dir,
}, out)
}()
starting := out.line(t, "msg", "starting")
wantStartingLine(t, starting, app.URL)
wantStartingLine(t, starting, appURL, dir)
addr, _ := starting["address"].(string)
wantGreeting(t, "http://"+addr+"/")
@@ -170,15 +185,184 @@ func TestServesUntilToldToStop(t *testing.T) {
out.line(t, "msg", "stopped")
}
func TestStateKeptAcrossRestarts(t *testing.T) {
t.Parallel()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
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",
}
// The two requests a day allows, and a stop.
runUntilStopped(t, env, func(url string) {
wantGreeting(t, url)
wantGreeting(t, url)
})
// After a restart the client has no fresh allowance: its third
// request breaks the day limit, and bans it.
out := runUntilStopped(t, env, func(url string) {
wantRefused(t, url)
})
out.line(t, "action", "rate_limited")
// After another, the ban still refuses it.
out = runUntilStopped(t, env, func(url string) {
wantRefused(t, url)
})
out.line(t, "action", "banned")
}
func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) {
t.Parallel()
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",
}
// 203.0.113.9's second request breaks the day limit, and bans
// 203.0.113.0/24.
runUntilStopped(t, env, func(url string) {
wantStatus(t, url, "203.0.113.9", http.StatusOK)
wantStatus(t, url, "203.0.113.9", http.StatusForbidden)
})
// With each address a netblock of its own after a restart, that ban
// still refuses all of 203.0.113.0/24. 198.51.100.7 is banned alone.
env[scope] = "32"
runUntilStopped(t, env, func(url string) {
wantStatus(t, url, "203.0.113.200", http.StatusForbidden)
wantStatus(t, url, "203.0.114.1", http.StatusOK)
wantStatus(t, url, "198.51.100.7", http.StatusOK)
wantStatus(t, url, "198.51.100.7", http.StatusForbidden)
})
// With /24 netblocks again, that ban still refuses 198.51.100.7, and
// no other address.
env[scope] = "24"
runUntilStopped(t, env, func(url string) {
wantStatus(t, url, "198.51.100.7", http.StatusForbidden)
wantStatus(t, url, "198.51.100.8", http.StatusOK)
})
}
func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) {
t.Parallel()
dir := t.TempDir()
err := os.WriteFile(filepath.Join(dir, "bans.json"), []byte("{\n"), 0o600)
if err != nil {
t.Fatalf("write bans.json: %v", err)
}
// The file ends at the newline that is the second byte of its first
// line.
wantStartRefused(t, dir, filepath.Join(dir, "bans.json")+", line 1, column 2: ")
}
func TestUnwritableStateDirStopsTheStart(t *testing.T) {
t.Parallel()
wantStartRefused(t, filepath.Join(t.TempDir(), "missing"),
"SWWAF_STATE_DIR cannot be written: ")
}
// wantStartRefused runs smallwebwaf with its state files in dir, and
// checks that it stops at start, with an error that starts with want. If
// it starts instead, it is stopped after waitLimit.
func wantStartRefused(t *testing.T, dir, want string) {
t.Helper()
ctx, stop := context.WithTimeout(t.Context(), waitLimit)
defer stop()
out := &output{}
status := run(ctx, map[string]string{listenAddr: localhost + ":0", stateDir: dir}, out)
if status != 1 {
t.Fatalf("exit status %d, want 1", status)
}
line := out.line(t, "msg", "cannot use the state files")
message, _ := line["error"].(string)
if !strings.HasPrefix(message, want) {
t.Errorf("start refused with %q, want an error starting %q", message, want)
}
}
// startApp starts an app that answers every request with greeting, and
// returns its URL.
func startApp(t *testing.T) string {
t.Helper()
app := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, greeting)
}))
t.Cleanup(app.Close)
return app.URL
}
// runUntilStopped runs smallwebwaf with the settings in env, has use send
// it requests at url, then stops it as SIGTERM does, checks that it
// stopped in order, and returns its output.
func runUntilStopped(
t *testing.T, env map[string]string, use func(url string),
) *output {
t.Helper()
ctx, stop := context.WithCancel(t.Context())
out := &output{}
exited := make(chan int, 1)
go func() {
exited <- run(ctx, env, out)
}()
addr, _ := out.line(t, "msg", "starting")["address"].(string)
use("http://" + addr + "/")
stop()
select {
case status := <-exited:
if status != 0 {
t.Fatalf("exit status %d, want 0; output:\n%s", status, out.text())
}
case <-time.After(waitLimit):
t.Fatal("still running after being told to stop")
}
return out
}
// wantStartingLine checks that the line at start gives the version and
// every setting's value.
func wantStartingLine(t *testing.T, line map[string]any, appURL string) {
func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
t.Helper()
settings, _ := line["settings"].(map[string]any)
want := map[string]any{
listenAddr: localhost + ":0",
upstreamURL: appURL,
stateDir: dir,
"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",
"SWWAF_CLIENT_REQUEST_TIMEOUT": "60s",
"SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K",
@@ -193,7 +377,7 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL string) {
"SWWAF_DENY_NETS": "",
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
"SWWAF_RATE_LIMIT_PER_HOUR": "10000",
"SWWAF_RATE_LIMIT_PER_DAY": "50000",
rateLimitPerDay: "50000",
"SWWAF_DENIED_COUNTRIES": "",
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "",
"SWWAF_BAN_RESPONSE": "403",
@@ -236,7 +420,61 @@ func wantGreeting(t *testing.T, url string) {
body, err := io.ReadAll(res.Body)
_ = res.Body.Close()
if err != nil || string(body) != "hello from the app" {
if err != nil || string(body) != greeting {
t.Errorf("got %q (%v), want the app's answer", body, err)
}
}
// wantRefused checks that a request to url is refused with 403, the
// default SWWAF_BAN_RESPONSE.
func wantRefused(t *testing.T, url string) {
t.Helper()
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url,
http.NoBody)
if err != nil {
t.Fatalf("new request: %v", err)
}
transport := &http.Transport{}
defer transport.CloseIdleConnections()
res, err := (&http.Client{Transport: transport}).Do(req)
if err != nil {
t.Fatalf("request: %v", err)
}
_ = res.Body.Close()
if res.StatusCode != http.StatusForbidden {
t.Errorf("status %d, want %d", res.StatusCode, http.StatusForbidden)
}
}
// wantStatus checks that a request to url from the client at from, as
// X-Forwarded-For names it, is answered with status.
func wantStatus(t *testing.T, url, from string, status int) {
t.Helper()
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url,
http.NoBody)
if err != nil {
t.Fatalf("new request: %v", err)
}
req.Header.Set("X-Forwarded-For", from)
transport := &http.Transport{}
defer transport.CloseIdleConnections()
res, err := (&http.Client{Transport: transport}).Do(req)
if err != nil {
t.Fatalf("request: %v", err)
}
_ = res.Body.Close()
if res.StatusCode != status {
t.Errorf("request from %s: status %d, want %d", from, res.StatusCode, status)
}
}
+492
View File
@@ -0,0 +1,492 @@
// Package state keeps smallwebwaf's state in JSON files in
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
// bans.json holds the bans, clients.json each client's counters and
// history, and lookups.json GeoJS's answers. Load reads them at start, and
// Run and WriteAll write them, each from a snapshot its part takes under
// its own lock, so that no request waits on the disk.
package state
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io/fs"
"log/slog"
"net/netip"
"os"
"path/filepath"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
// version is the version of the files' format, the only one read.
const version = 1
// fileMode lets the smallwebwaf user alone read and write the files, which
// hold visitors' addresses.
const fileMode = 0o600
// The state files' names.
const (
bansJSON = "bans.json"
clientsJSON = "clients.json"
lookupsJSON = "lookups.json"
)
var (
errVersion = errors.New("unknown version")
// 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
}
// 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 write(f.params.Dir, bansJSON, append(data, '\n'))
}
// writeClients writes clients.json.
func (f *Files) writeClients() error {
data, err := encodeOnePerLine("clients", f.params.Limiter.Snapshot())
if err != nil {
return fmt.Errorf("encode %s: %w", clientsJSON, err)
}
return write(f.params.Dir, clientsJSON, data)
}
// writeLookups writes lookups.json.
func (f *Files) writeLookups() error {
data, err := encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
if err != nil {
return fmt.Errorf("encode %s: %w", lookupsJSON, err)
}
return write(f.params.Dir, lookupsJSON, data)
}
// newBanEntry returns ban as bans.json holds it.
func newBanEntry(ban bans.Ban) banEntry {
entry := banEntry{Netblock: ban.Netblock, Start: ban.Start, Notes: ban.Notes}
if !ban.Permanent() {
entry.Expires = &ban.Expires
}
return entry
}
// ban returns the ban an entry of bans.json holds.
func (e banEntry) ban() bans.Ban {
ban := bans.Ban{Netblock: e.Netblock, Start: e.Start, Notes: e.Notes}
if e.Expires != nil {
ban.Expires = *e.Expires
}
return ban
}
// 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())
}
+635
View File
@@ -0,0 +1,635 @@
package state_test
import (
"context"
"encoding/json"
"log/slog"
"net/netip"
"os"
"path/filepath"
"slices"
"strings"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/state"
)
const (
// The state files.
bansJSON = "bans.json"
clientsJSON = "clients.json"
lookupsJSON = "lookups.json"
)
// permanentBansJSON is bans.json holding permanentBan.
const permanentBansJSON = `{
"version": 1,
"bans": [
{
"netblock": "2001:db8::/64",
"start": "2026-10-06T00:00:00Z",
"expires": null,
"notes": {
"country": "DE",
"limit": 1000,
"window": "minute",
"count": 1000.5,
"request": {
"time": "2026-10-06T00:00:00Z",
"method": "GET",
"host": "app.example",
"path": "/repo?page=2",
"status": 403,
"user_agent": "scraper/1.0"
},
"requests": 1500,
"refused": 3,
"earlier_bans": 5
}
}
]
}
`
func TestFilesWrittenAndReadBack(t *testing.T) {
t.Parallel()
dir := t.TempDir()
before := newParams(dir)
fill(before)
files, err := state.Load(before)
if err != nil {
t.Fatalf("load: %v", err)
}
err = files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
// Read into new parts, as at the next start, the files give back what
// was written.
after := newParams(dir)
load(t, after)
wantEqual(t, bansJSON, after.Ledger.Snapshot(), before.Ledger.Snapshot())
wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot())
wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot())
// Each one-per-line file lists its entries by client, and nothing
// but the three files is left in the directory.
wantEntries(t, filepath.Join(dir, clientsJSON), "clients",
"192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64")
wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups",
"192.0.2.1/32", "203.0.113.9/32")
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
}
func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
params.Ledger.Load([]bans.Ban{permanentBan()})
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
got := readFile(t, filepath.Join(dir, bansJSON))
if got != permanentBansJSON {
t.Errorf("bans.json\n%s\nwant\n%s", got, permanentBansJSON)
}
}
func TestMissingFilesAreEmptyState(t *testing.T) {
t.Parallel()
params := newParams(t.TempDir())
load(t, params)
if len(params.Ledger.Snapshot()) != 0 || len(params.Limiter.Snapshot()) != 0 ||
len(params.GeoJS.Snapshot()) != 0 {
t.Error("state from no files")
}
}
func TestFileThatDoesNotParseStopsTheStart(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name, file, content string
// want is what the error says after the file's path.
want string
}{
{
"a syntax error", bansJSON,
"{\n \"version\": 1,\n \"bans\": [\n" +
" {\"netblock\": \"203.0.113.9/32\",}\n ]\n}\n",
", line 4, column 39: invalid character '}'",
},
{
"a value of the wrong kind", clientsJSON,
"{\n \"version\": 1,\n \"clients\": [\n" +
" {\"client\":\"203.0.113.9/32\",\"history\":{\"requests\":\"many\"}}\n" +
" ]\n}\n",
", line 4, column ",
},
{
// Found at the newline that ends the file.
"a cut-off file", lookupsJSON,
"{\n \"version\": 1,\n \"lookups\": [\n",
", line 3, column 17: unexpected end of JSON input",
},
{
"an unknown field", lookupsJSON,
`{"version": 1, "lookups": [{"client": "203.0.113.9/32", "contry": "DE"}]}`,
`: json: unknown field "contry"`,
},
{
"a netblock that does not read", bansJSON,
`{"version": 1, "bans": [{"netblock": "203.0.113.300/32"}]}`,
`: netip.ParsePrefix("203.0.113.300/32")`,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
wantRefused(t, tc.file, tc.content, tc.want)
})
}
}
func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
t.Parallel()
const (
// The other fields each entry needs.
ban = `"start": "2026-10-06T00:00:00Z", "expires": null`
answer = `"answered": "2026-10-06T00:00:00Z"`
noNetblock = `: entry 1 has no "netblock"`
)
for _, tc := range []struct {
name, file, content string
// want is what the error says after the file's path.
want string
}{
{
"a ban without a netblock", bansJSON,
`{"version": 1, "bans": [{` + ban + `}]}`,
noNetblock,
},
{
"a ban whose netblock is null", bansJSON,
`{"version": 1, "bans": [{"netblock": null, ` + ban + `}]}`,
noNetblock,
},
{
"a ban whose netblock is empty", bansJSON,
`{"version": 1, "bans": [{"netblock": "", ` + ban + `}]}`,
noNetblock,
},
{
"a ban without a start", bansJSON,
`{"version": 1, "bans": [{"netblock": "203.0.113.9/32", "expires": null}]}`,
`: entry 1 has no "start"`,
},
{
// The first ban's expires is null, as a permanent ban's is.
"a ban without an expires", bansJSON,
`{"version": 1, "bans": [{"netblock": "203.0.113.9/32", ` + ban + `}, ` +
`{"netblock": "203.0.113.10/32", "start": "2026-10-06T00:00:00Z"}]}`,
`: entry 2 has no "expires"`,
},
{
"a client without its address", clientsJSON,
`{"version": 1, "clients": [{"history": {"requests": 3}}]}`,
`: entry 1 has no "client"`,
},
{
"a client with requests in a window without its start", clientsJSON,
`{"version": 1, "clients": [{"client": "203.0.113.9/32", ` +
`"hour": {"current": 3}}]}`,
`: entry 1 has no "hour.start"`,
},
{
"an answer without a client", lookupsJSON,
`{"version": 1, "lookups": [{"country": "DE", ` + answer + `}]}`,
`: entry 1 has no "client"`,
},
{
// A country of "" is a client GeoJS cannot place.
"an answer without a country", lookupsJSON,
`{"version": 1, "lookups": [{"client": "192.0.2.1/32", "country": "", ` +
answer + `}, {"client": "203.0.113.9/32", ` + answer + `}]}`,
`: entry 2 has no "country"`,
},
{
"an answer without the time GeoJS gave it", lookupsJSON,
`{"version": 1, "lookups": [{"client": "203.0.113.9/32", "country": "DE"}]}`,
`: entry 1 has no "answered"`,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
wantRefused(t, tc.file, tc.content, tc.want)
})
}
}
func TestUnknownVersionStopsTheStart(t *testing.T) {
t.Parallel()
for _, file := range []string{bansJSON, clientsJSON, lookupsJSON} {
for _, content := range []string{`{"version": 2}`, `{}`} {
t.Run(file+" "+content, func(t *testing.T) {
t.Parallel()
wantRefused(t, file, content, ": unknown version ")
})
}
}
}
func TestUnwritableDirectoryStopsTheStart(t *testing.T) {
t.Parallel()
notADirectory := filepath.Join(t.TempDir(), "file")
err := os.WriteFile(notADirectory, nil, 0o600)
if err != nil {
t.Fatalf("write: %v", err)
}
for _, dir := range []string{
filepath.Join(t.TempDir(), "missing"),
notADirectory,
} {
const want = "SWWAF_STATE_DIR cannot be written: "
_, err := state.Load(newParams(dir))
if err == nil || !strings.HasPrefix(err.Error(), want) {
t.Errorf("state directory %s: error %v, want one starting %s", dir, err, want)
}
}
}
// The two tests below run Run in a synctest bubble, where time is a clock
// of the test's own: time.Sleep moves it on at once, and synctest.Wait
// returns once Run waits for its next write, so that every write due by
// then is on disk.
func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
params := newParams(dir)
params.WriteDelay = 10 * time.Second
run(t, load(t, params))
// A second ban, made while the first waits to be written, puts the
// write off no further, and is written with it.
first := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"),
midnight(), bans.Notes{})
time.Sleep(5 * time.Second)
second := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"),
midnight(), bans.Notes{})
time.Sleep(5*time.Second - time.Nanosecond)
synctest.Wait()
wantFiles(t, dir)
time.Sleep(time.Nanosecond)
synctest.Wait()
wantFiles(t, dir, bansJSON)
read := newParams(dir)
load(t, read)
want := []bans.Ban{first, second}
if got := read.Ledger.Snapshot(); !slices.Equal(got, want) {
t.Errorf("bans.json holds %+v, want %+v", got, want)
}
// That write was the only one: bans.json is not written again for
// the second ban. The other files wait for the interval, an hour
// away.
removeFiles(t, dir, bansJSON)
time.Sleep(params.WriteDelay)
synctest.Wait()
wantFiles(t, dir)
})
}
func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
params := newParams(dir)
params.CounterInterval = time.Minute
run(t, load(t, params))
// The files are removed once written, so that each interval shows
// them written again.
for range 3 {
time.Sleep(time.Minute - time.Nanosecond)
synctest.Wait()
wantFiles(t, dir)
time.Sleep(time.Nanosecond)
synctest.Wait()
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
removeFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
}
})
}
func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
params.Ledger.Load([]bans.Ban{permanentBan()})
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
// A directory in the way of bans.json's temporary file fails its
// next write, but not the others'.
err = os.Mkdir(filepath.Join(dir, bansJSON+".tmp"), 0o700)
if err != nil {
t.Fatalf("mkdir: %v", err)
}
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
bans.Notes{})
params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight())
err = files.WriteAll()
if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") {
t.Errorf("error %v, want one naming bans.json's temporary file", err)
}
got := readFile(t, filepath.Join(dir, bansJSON))
if got != permanentBansJSON {
t.Errorf("bans.json is now\n%s\nwant it as it was", got)
}
read := newParams(dir)
load(t, read)
if len(read.Limiter.Snapshot()) != 1 {
t.Error("clients.json was not written")
}
}
func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) {
t.Parallel()
dir := t.TempDir()
files := load(t, newParams(dir))
// 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 = files.WriteAll()
if err == nil {
t.Error("writing over a directory did not fail")
}
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
}
// midnight is the time of the tests' clock.
func midnight() time.Time {
return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
}
// newParams returns Params for the state files in dir, with parts that
// hold nothing yet. GeoJS is never asked.
func newParams(dir string) state.Params {
discard := slog.New(slog.DiscardHandler)
return state.Params{
Dir: dir,
WriteDelay: time.Hour,
CounterInterval: time.Hour,
Ledger: bans.New(bans.Rules{
LimitBanDuration: time.Hour,
LimitBanRepeatWindow: 24 * time.Hour,
MaxBanDuration: 7 * 24 * time.Hour,
MaxBans: 5000,
}),
Limiter: ratelimit.New(ratelimit.Limits{}),
GeoJS: lookup.New(lookup.Params{Now: midnight, ProcessLog: discard}),
Now: midnight,
ProcessLog: discard,
}
}
// fill puts a ban that ends and one that does not, clients with counts
// and histories, and GeoJS answers into the parts of params.
func fill(params state.Params) {
now := midnight()
client := netip.MustParsePrefix("203.0.113.9/32")
params.Ledger.Load([]bans.Ban{permanentBan()})
params.Ledger.BanForLimit(client, now, bans.Notes{Country: "DE", Limit: 1})
for _, c := range []string{"2001:db8::/64", "203.0.113.9/32", "192.0.2.1/32"} {
params.Limiter.Count(netip.MustParsePrefix(c), now)
}
params.Limiter.AddToHistory(client, now, ratelimit.Request{
Country: "DE", Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5,
})
params.GeoJS.Load([]lookup.Answer{
{Client: client, Country: "DE", Answered: now.Add(-time.Hour), Used: now},
{
Client: netip.MustParsePrefix("192.0.2.1/32"),
Answered: now.Add(-time.Hour), Used: now.Add(-time.Minute),
},
})
}
// permanentBan is the ban permanentBansJSON holds.
func permanentBan() bans.Ban {
return bans.Ban{
Netblock: netip.MustParsePrefix("2001:db8::/64"),
Start: midnight(),
Notes: bans.Notes{
Country: "DE",
Limit: 1000,
Window: "minute",
Count: 1000.5,
Request: bans.Request{
Time: midnight(),
Method: "GET",
Host: "app.example",
Path: "/repo?page=2",
Status: 403,
UserAgent: "scraper/1.0",
},
Requests: 1500,
Refused: 3,
EarlierBans: 5,
},
}
}
// load reads the state files into the parts of params.
func load(t *testing.T, params state.Params) *state.Files {
t.Helper()
files, err := state.Load(params)
if err != nil {
t.Fatalf("load: %v", err)
}
return files
}
// run runs files' writes until the test ends.
func run(t *testing.T, files *state.Files) {
t.Helper()
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
files.Run(ctx)
close(stopped)
}()
t.Cleanup(func() {
stop()
<-stopped
})
}
// wantEqual checks that the entries read back from file are those
// written.
func wantEqual[E comparable](t *testing.T, file string, got, want []E) {
t.Helper()
if !slices.Equal(got, want) {
t.Errorf("%s read back\n%+v\nwant\n%+v", file, got, want)
}
}
// readFile returns what the file at path holds.
func readFile(t *testing.T, path string) string {
t.Helper()
data, err := os.ReadFile(path) //nolint:gosec // a file the test wrote
if err != nil {
t.Fatalf("read: %v", err)
}
return string(data)
}
// wantRefused writes content to the state file named file in a new
// directory, and checks that Load refuses it with an error that is the
// file's path and then starts with want.
func wantRefused(t *testing.T, file, content, want string) {
t.Helper()
dir := t.TempDir()
path := filepath.Join(dir, file)
err := os.WriteFile(path, []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", file, err)
}
_, err = state.Load(newParams(dir))
if err == nil || !strings.HasPrefix(err.Error(), path+want) {
t.Errorf("error %v, want one starting %s%s", err, path, want)
}
}
// removeFiles removes the named files from dir.
func removeFiles(t *testing.T, dir string, names ...string) {
t.Helper()
for _, name := range names {
err := os.Remove(filepath.Join(dir, name))
if err != nil {
t.Fatalf("remove: %v", err)
}
}
}
// wantFiles checks the names of the files in dir.
func wantFiles(t *testing.T, dir string, want ...string) {
t.Helper()
entries, err := os.ReadDir(dir)
if err != nil {
t.Fatalf("read %s: %v", dir, err)
}
got := make([]string, 0, len(entries))
for _, entry := range entries {
got = append(got, entry.Name())
}
if !slices.Equal(got, want) {
t.Errorf("%s holds %v, want %v", dir, got, want)
}
}
// wantEntries checks that the file at path has its version, then its
// entries under key, each on a line of its own, for the clients want
// names in that order.
func wantEntries(t *testing.T, path, key string, want ...string) {
t.Helper()
data := readFile(t, path)
lines := strings.Split(strings.TrimSuffix(data, "\n"), "\n")
head := []string{"{", ` "version": 1,`, ` "` + key + `": [`}
tail := []string{" ]", "}"}
if len(lines) != len(head)+len(want)+len(tail) ||
!slices.Equal(lines[:len(head)], head) ||
!slices.Equal(lines[len(lines)-len(tail):], tail) {
t.Fatalf("%s is\n%s", path, data)
}
for i, client := range want {
line := strings.TrimSuffix(lines[len(head)+i], ",")
var entry struct {
Client string `json:"client"`
}
err := json.Unmarshal([]byte(line), &entry)
if err != nil || entry.Client != client {
t.Errorf("entry %d of %s is %s (%v), want %s's", i, path, line, err, client)
}
}
}