Ban the netblock of a client that breaks a rate limit, in memory (closes #18)
check / check (push) Successful in 3m48s

A request over a rate limit is refused with SWWAF_BAN_RESPONSE and bans
the client's netblock: an hour at first, three times the last ban when
broken again within a day of its end, permanent past seven days. The
ban ledger in internal/bans is checked after the static lists and
before the lookup, and the requests it refuses are not counted. A ban
resets the client's counters and carries notes holding the request
that broke the limit, as SPEC.md now says. At most SWWAF_MAX_BANS are
held. SWWAF_BAN_RESPONSE also answers SWWAF_DENY_NETS and the country
lists.

Judgement call: the six ban settings cannot be off.
Judgement call: a permanent ban's ban_expires is "permanent".

Model: opus-5-5
This commit was merged in pull request #69.
This commit is contained in:
2026-10-06 05:29:03 +02:00
parent 0f85c9ae07
commit 73ca94f850
20 changed files with 1522 additions and 127 deletions
+255
View File
@@ -0,0 +1,255 @@
// 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.
package bans
import (
"net/netip"
"slices"
"strings"
"sync"
"time"
"github.com/hashicorp/golang-lru/v2/simplelru"
)
// repeatFactor is how many times as long as the netblock's last ban a ban
// for a limit broken again within the repeat window lasts.
const repeatFactor = 3
// maxTextBytes is how much of each text in a ban's notes is kept.
const maxTextBytes = 256
// Rules are how long a ban for a broken limit lasts, and how many bans
// are held.
type Rules struct {
// LimitBanDuration is how long a first ban lasts.
LimitBanDuration time.Duration
// LimitBanRepeatWindow is how soon after the netblock's last ban
// ended a broken limit counts as a repeat, which bans for
// repeatFactor times as long as that ban.
LimitBanRepeatWindow time.Duration
// MaxBanDuration is the longest ban; a ban that would be longer is
// permanent instead.
MaxBanDuration time.Duration
// MaxBans is the most bans held, at least one. Past it, the earliest
// ban of the netblock that has gone longest without a request is
// dropped.
MaxBans int
}
// Ban is a ban on a netblock for a broken limit, the only kind of ban
// smallwebwaf makes so far.
type Ban struct {
Netblock netip.Prefix
Start time.Time
// Expires is when the ban ends, zero for a permanent ban.
Expires time.Time
Notes Notes
}
// Permanent reports whether the ban never runs out.
func (b Ban) Permanent() bool {
return b.Expires.IsZero()
}
// ActiveAt reports whether the ban refuses requests at now.
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.
type Notes struct {
// Country is the client's country, when it was looked up.
Country string
// 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
// Request is the request that broke the limit.
Request Request
// Refused is how many requests the ban has refused so far.
Refused int64
// EarlierBans is how many bans the netblock had before this one.
EarlierBans int
}
// Request is a request in a ban's notes. Each text is cut to 256 bytes.
type Request struct {
Time time.Time
Method string
Host string
// Path is the path with its query string.
Path string
// Status is what the client was sent, 0 if nothing was.
Status int
UserAgent string
}
// Ledger holds the bans. It is safe for concurrent use.
type Ledger struct {
rules Rules
mu sync.Mutex
// netblocks holds each banned netblock's bans, oldest first. Each
// request from a netblock makes it the most recently seen.
netblocks *simplelru.LRU[netip.Prefix, *[]Ban]
// held is how many bans netblocks holds, at most rules.MaxBans.
held int
}
// New returns a Ledger with no ban yet.
func New(rules Rules) *Ledger {
// Every netblock held has a ban, so there are never more netblocks
// than rules.MaxBans, and the LRU never drops one itself.
netblocks, err := simplelru.NewLRU[netip.Prefix, *[]Ban](rules.MaxBans, nil)
if err != nil {
panic(err) // NewLRU fails only for a size below one
}
return &Ledger{rules: rules, 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) {
l.mu.Lock()
defer l.mu.Unlock()
bans, found := l.netblocks.Get(netblock)
if !found {
return Ban{}, false
}
// 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
}
last.Notes.Refused++
return *last, true
}
// BanForLimit bans netblock at now for a broken limit, with notes, and
// returns the ban. A first ban lasts LimitBanDuration. A ban made within
// LimitBanRepeatWindow after the netblock's last ban ended lasts
// repeatFactor times as long as that one. A ban that would be longer
// than MaxBanDuration is permanent instead. If a ban on netblock is still
// active, as when two of its requests break a limit at once, that ban is
// returned and no other is made. The ledger fills in the notes' Refused
// and EarlierBans itself.
func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes) Ban {
l.mu.Lock()
defer l.mu.Unlock()
var last *Ban
bans, found := l.netblocks.Get(netblock)
if found {
last = &(*bans)[len(*bans)-1]
if last.ActiveAt(now) {
return *last
}
notes.EarlierBans = last.Notes.EarlierBans + 1
}
notes.Request = notes.Request.cut()
ban := Ban{
Netblock: netblock,
Start: now,
Expires: l.expiry(last, now),
Notes: notes,
}
if l.held == l.rules.MaxBans {
l.dropOne()
}
// 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
}
// Bans returns the bans held on netblock, oldest first. It is not a
// request from netblock, and leaves when it was last seen unchanged.
func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
l.mu.Lock()
defer l.mu.Unlock()
bans, found := l.netblocks.Peek(netblock)
if !found {
return nil
}
return slices.Clone(*bans)
}
// 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.
func (l *Ledger) expiry(last *Ban, now time.Time) time.Time {
length := l.rules.LimitBanDuration
if last != nil && now.Sub(last.Expires) <= l.rules.LimitBanRepeatWindow {
lastLength := last.Expires.Sub(last.Start)
// This is repeatFactor * lastLength > MaxBanDuration, written so
// that it cannot overflow.
if lastLength > l.rules.MaxBanDuration/repeatFactor {
return time.Time{}
}
length = repeatFactor * lastLength
}
if length > l.rules.MaxBanDuration {
return time.Time{}
}
return now.Add(length)
}
// dropOne drops the earliest ban of the netblock that has gone longest
// without a request, and the netblock with it if that was its only ban.
func (l *Ledger) dropOne() {
netblock, bans, _ := l.netblocks.GetOldest()
if len(*bans) == 1 {
l.netblocks.Remove(netblock)
} else {
*bans = slices.Delete(*bans, 0, 1)
}
l.held--
}
// cut returns r with each text cut to maxTextBytes and copied, so that
// the notes do not keep the rest of the request in memory.
func (r Request) cut() Request {
r.Method = cutText(r.Method)
r.Host = cutText(r.Host)
r.Path = cutText(r.Path)
r.UserAgent = cutText(r.UserAgent)
return r
}
// cutText returns a copy of the first maxTextBytes of text.
func cutText(text string) string {
return strings.Clone(text[:min(len(text), maxTextBytes)])
}
+267
View File
@@ -0,0 +1,267 @@
package bans_test
import (
"net/netip"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
)
const day = 24 * time.Hour
func TestRepeatsTripleUntilPermanent(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
now := midnight()
// Each ban is followed by another as soon as it ends: 1, 3, 9, 27 and
// 81 hours.
for i, hours := range []int{1, 3, 9, 27, 81} {
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
length := time.Duration(hours) * time.Hour
if !ban.Expires.Equal(now.Add(length)) || ban.Notes.EarlierBans != i {
t.Fatalf("ban %d lasts %s with %d earlier bans, want %d hours and %d",
i+1, ban.Expires.Sub(now), ban.Notes.EarlierBans, hours, i)
}
now = ban.Expires
}
// The sixth would last 243 hours, more than seven days: it is
// permanent, and never ends.
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() {
t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires)
}
_, banned := ledger.Check(netblock, now.Add(100*365*day))
if !banned {
t.Error("a permanent ban ended")
}
}
func TestRepeatWindowRunsOut(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
// gap is the time between the end of the first ban and the second.
gap time.Duration
want time.Duration
}{
{"broken again as the window ends", day, 3 * time.Hour},
{"broken again after the window", day + time.Nanosecond, time.Hour},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{})
if second.Expires.Sub(second.Start) != tc.want || second.Notes.EarlierBans != 1 {
t.Errorf("second ban lasts %s with %d earlier bans, want %s and 1",
second.Expires.Sub(second.Start), second.Notes.EarlierBans, tc.want)
}
})
}
}
func TestFirstBanLongerThanTheMaximumIsPermanent(t *testing.T) {
t.Parallel()
rules := defaultRules()
rules.LimitBanDuration = rules.MaxBanDuration + time.Hour
ledger := bans.New(rules)
ban := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
bans.Notes{})
if !ban.Permanent() {
t.Errorf("first ban ends at %s, want a permanent one", ban.Expires)
}
}
func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) {
t.Parallel()
// With bans of up to 100,000 days, the 14th ban in a row, of 3^13
// hours, is within the maximum, and three times as long would not fit
// in a time.Duration. The 15th is permanent.
rules := defaultRules()
rules.MaxBanDuration = 100000 * day
ledger := bans.New(rules)
netblock := netip.MustParsePrefix("203.0.113.9/32")
now := midnight()
for i := range 14 {
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Expires.After(ban.Start) {
t.Fatalf("ban %d starts at %s and ends at %s", i+1, ban.Start, ban.Expires)
}
now = ban.Expires
}
ban := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() {
t.Errorf("15th ban ends at %s, want a permanent one", ban.Expires)
}
}
func TestBrokenLimitDuringABanMakesNoOther(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
again := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
if again != first || len(ledger.Bans(netblock)) != 1 {
t.Errorf("a limit broken during a ban gave %+v and %d bans, want %+v and 1",
again, len(ledger.Bans(netblock)), first)
}
}
func TestCheckRefusesWhileTheBanLastsAndCountsTheRefusals(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
for range 3 {
got, banned := ledger.Check(netblock, 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())
if banned {
t.Error("another netblock is banned")
}
_, banned = ledger.Check(netblock, 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)
}
}
func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
t.Parallel()
rules := defaultRules()
rules.MaxBans = 3
ledger := bans.New(rules)
a := netip.MustParsePrefix("203.0.113.1/32")
b := netip.MustParsePrefix("203.0.113.2/32")
c := netip.MustParsePrefix("203.0.113.3/32")
d := netip.MustParsePrefix("2001:db8::/64")
now := midnight()
first := ledger.BanForLimit(a, now, bans.Notes{})
ledger.BanForLimit(b, now, bans.Notes{})
ledger.BanForLimit(c, now, bans.Notes{})
// 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.BanForLimit(d, now, bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 0, c: 1, d: 1})
// a is banned again once its ban has ended; c, seen longest ago, goes.
ledger.BanForLimit(a, first.Expires, bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{a: 2, c: 0, d: 1})
// With d seen since, a is seen longest ago, and its earlier ban goes
// first.
ledger.Check(d, first.Expires)
ledger.BanForLimit(b, first.Expires, bans.Notes{})
wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 1, d: 1})
if !ledger.Bans(a)[0].Start.Equal(first.Expires) {
t.Errorf("a kept its ban of %s, want the later one", ledger.Bans(a)[0].Start)
}
}
func TestFullLedgerDropsTheEarlierBanOfTheNetblockBannedAgain(t *testing.T) {
t.Parallel()
// With room for one ban, the netblock's ended ban goes to make room for
// its new one, whose notes still count it.
rules := defaultRules()
rules.MaxBans = 1
ledger := bans.New(rules)
netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second := ledger.BanForLimit(netblock, first.Expires, bans.Notes{})
held := ledger.Bans(netblock)
if len(held) != 1 || held[0] != second || held[0].Notes.EarlierBans != 1 {
t.Errorf("the ledger holds %+v, want only the second ban, with 1 earlier ban",
held)
}
}
func TestRequestTextsAreCutTo256Bytes(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
long := strings.Repeat("a", 300)
request := bans.Request{
Time: midnight(), Method: long, Host: long, Path: long, Status: 403, UserAgent: long,
}
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request})
cut := long[:256]
want := bans.Request{
Time: midnight(), Method: cut, Host: cut, Path: cut, Status: 403, UserAgent: cut,
}
if ban.Notes.Request != want || ledger.Bans(netblock)[0].Notes.Request != want {
t.Errorf("the notes keep %+v, want each text cut to 256 bytes", ban.Notes.Request)
}
}
// defaultRules are the rules at the settings' defaults.
func defaultRules() bans.Rules {
return bans.Rules{
LimitBanDuration: time.Hour,
LimitBanRepeatWindow: day,
MaxBanDuration: 7 * day,
MaxBans: 5000,
}
}
// midnight is when the tests' first bans are made.
func midnight() time.Time {
return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
}
// wantBans checks how many bans the ledger holds on each netblock.
func wantBans(t *testing.T, ledger *bans.Ledger, want map[netip.Prefix]int) {
t.Helper()
for netblock, count := range want {
got := len(ledger.Bans(netblock))
if got != count {
t.Errorf("%s has %d bans, want %d", netblock, got, count)
}
}
}
+116 -2
View File
@@ -9,6 +9,7 @@ import (
"log/slog"
"math"
"net"
"net/http"
"net/netip"
"net/url"
"slices"
@@ -74,6 +75,25 @@ type Config struct {
// capitals, as GeoJS gives them.
DeniedCountries []string
ExclusivelyAllowedCountries []string
// BanResponse is the status a refused client is answered with, 403
// or 429, or 0 to close the connection without an answer
// (SWWAF_BAN_RESPONSE). It answers a banned client, a request that
// breaks a rate limit, SWWAF_DENY_NETS and the country lists.
BanResponse int
// LimitBanDuration is the ban for a first broken rate limit
// (SWWAF_LIMIT_BAN_DURATION). A limit broken again within
// LimitBanRepeatWindow after the last ban ended bans for three times
// as long as that ban (SWWAF_LIMIT_BAN_REPEAT_WINDOW), and a ban that
// would be longer than MaxBanDuration is permanent instead
// (SWWAF_MAX_BAN_DURATION). None of them can be off.
LimitBanDuration time.Duration
LimitBanRepeatWindow time.Duration
MaxBanDuration time.Duration
// MaxBans is the most bans held (SWWAF_MAX_BANS).
MaxBans int
// BanScopeV4Prefix is the length of the netblock around an IPv4
// client that a ban covers (SWWAF_BAN_SCOPE_V4_PREFIX).
BanScopeV4Prefix int
// settings are the values read, as given or by default, for the
// log line at start.
@@ -89,6 +109,7 @@ const (
kibibyte = 1 << 10
mebibyte = 1 << 20
gibibyte = 1 << 30
ipv4Bits = 32
)
var (
@@ -109,8 +130,15 @@ var (
"such as http://127.0.0.1:8081")
errNotCountry = errors.New(
"is not a two-letter country code such as de or kp")
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
errNotOver4K = errors.New("is not a size of more than 4K, such as 32K")
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
errNotOver4K = errors.New("is not a size of more than 4K, such as 32K")
errNotDurationAboveZero = errors.New(
"is not a duration above zero, such as 1h or 7d")
errNotNumberAboveZero = errors.New(
"is not a whole number above zero, such as 5000")
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")
)
// FromEnvironment reads the settings with lookupEnv, normally
@@ -140,6 +168,12 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
ExclusivelyAllowedCountries: env.countries(
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""),
BanResponse: env.banResponse("SWWAF_BAN_RESPONSE", "403"),
LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"),
LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"),
MaxBanDuration: env.durationNotOff("SWWAF_MAX_BAN_DURATION", "7d"),
MaxBans: env.numberNotOff("SWWAF_MAX_BANS", "5000"),
BanScopeV4Prefix: env.v4Prefix("SWWAF_BAN_SCOPE_V4_PREFIX", "32"),
}
for _, country := range cfg.ExclusivelyAllowedCountries {
@@ -261,6 +295,40 @@ func (e *environment) countries(name, defaultValue string) []string {
return countries
}
// durationNotOff reads a setting that is a duration and, unlike a
// timeout, cannot be off.
func (e *environment) durationNotOff(name, defaultValue string) time.Duration {
duration, err := parseDurationNotOff(e.value(name, defaultValue))
e.check(name, err)
return duration
}
// numberNotOff reads a setting that is a whole number above zero, which
// cannot be off.
func (e *environment) numberNotOff(name, defaultValue string) int {
number, err := parseNumberNotOff(e.value(name, defaultValue))
e.check(name, err)
return number
}
// banResponse reads a setting that is how a refused client is answered.
func (e *environment) banResponse(name, defaultValue string) int {
status, err := parseBanResponse(e.value(name, defaultValue))
e.check(name, err)
return status
}
// v4Prefix reads a setting that is the length of an IPv4 netblock.
func (e *environment) v4Prefix(name, defaultValue string) int {
length, err := parseV4Prefix(e.value(name, defaultValue))
e.check(name, err)
return length
}
// 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) {
@@ -362,6 +430,52 @@ func parseCount(value string) (int64, error) {
return n, nil
}
// parseDurationNotOff reads a duration above zero, as parseDuration does,
// but not off.
func parseDurationNotOff(value string) (time.Duration, error) {
duration, err := parseDuration(value)
if err != nil || duration == 0 {
return 0, fmt.Errorf("%q %w", value, errNotDurationAboveZero)
}
return duration, nil
}
// parseNumberNotOff reads a whole number above zero.
func parseNumberNotOff(value string) (int, error) {
n, err := strconv.Atoi(value)
if err != nil || n <= 0 {
return 0, fmt.Errorf("%q %w", value, errNotNumberAboveZero)
}
return n, nil
}
// parseBanResponse reads how a refused client is answered: 403, 429, or
// close, which is 0.
func parseBanResponse(value string) (int, error) {
switch value {
case "403":
return http.StatusForbidden, nil
case "429":
return http.StatusTooManyRequests, nil
case "close":
return 0, nil
default:
return 0, fmt.Errorf("%q %w", value, errNotBanResponse)
}
}
// parseV4Prefix reads the length of an IPv4 netblock, from 0 to 32.
func parseV4Prefix(value string) (int, error) {
n, err := strconv.Atoi(value)
if err != nil || n < 0 || n > ipv4Bits {
return 0, fmt.Errorf("%q %w", value, errNotV4Prefix)
}
return n, nil
}
// parseList splits a comma-separated list and trims the spaces around
// each item. An empty value is an empty list.
func parseList(value string) ([]string, error) {
+61
View File
@@ -35,6 +35,12 @@ const (
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
deniedCountries = "SWWAF_DENIED_COUNTRIES"
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
banResponse = "SWWAF_BAN_RESPONSE"
limitBanDuration = "SWWAF_LIMIT_BAN_DURATION"
limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW"
maxBanDuration = "SWWAF_MAX_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS"
banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX"
)
// off switches a timeout, a size limit or a rate limit off.
@@ -80,6 +86,12 @@ func TestDefaults(t *testing.T) {
RateLimitPerMinute: 1000,
RateLimitPerHour: 10000,
RateLimitPerDay: 50000,
BanResponse: 403,
LimitBanDuration: time.Hour,
LimitBanRepeatWindow: 24 * time.Hour,
MaxBanDuration: 7 * 24 * time.Hour,
MaxBans: 5000,
BanScopeV4Prefix: 32,
})
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
@@ -118,6 +130,12 @@ func TestValuesAsSet(t *testing.T) {
rateLimitPerDay: "6000",
deniedCountries: "cn, RU,kp,Xk",
allowedCountries: "de",
banResponse: "429",
limitBanDuration: "15m",
limitBanRepeatWindow: "2d",
maxBanDuration: "30d",
maxBans: "100",
banScopeV4Prefix: "24",
})
wantSettings(t, cfg, config.Config{
@@ -133,6 +151,12 @@ func TestValuesAsSet(t *testing.T) {
RateLimitPerMinute: 60,
RateLimitPerHour: 600,
RateLimitPerDay: 6000,
BanResponse: 429,
LimitBanDuration: 15 * time.Minute,
LimitBanRepeatWindow: 48 * time.Hour,
MaxBanDuration: 30 * 24 * time.Hour,
MaxBans: 100,
BanScopeV4Prefix: 24,
})
if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
@@ -226,6 +250,15 @@ func TestRateLimitsOff(t *testing.T) {
}
}
func TestBanResponseCloseIsZero(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{banResponse: "close"})
if cfg.BanResponse != 0 {
t.Errorf("close read as %d, want 0", cfg.BanResponse)
}
}
func TestTrustedProxiesSetButEmptyTrustNothing(t *testing.T) {
t.Parallel()
@@ -290,6 +323,12 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{allowedCountries, "uk"},
{allowedCountries, "zz"},
{allowedCountries, "de,germany"},
{banResponse, "404"}, {banResponse, "drop"}, {banResponse, ""},
{limitBanDuration, off}, {limitBanDuration, "0s"}, {limitBanDuration, "1"},
{limitBanRepeatWindow, off}, {limitBanRepeatWindow, "-1h"},
{maxBanDuration, off}, {maxBanDuration, "1w"},
{maxBans, off}, {maxBans, "0"}, {maxBans, "5K"},
{banScopeV4Prefix, "33"}, {banScopeV4Prefix, "-1"}, {banScopeV4Prefix, "/24"},
} {
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
t.Parallel()
@@ -344,6 +383,12 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
rateLimitPerDay: "50000",
deniedCountries: "",
allowedCountries: "",
banResponse: "403",
limitBanDuration: "1h",
limitBanRepeatWindow: "24h",
maxBanDuration: "7d",
maxBans: "5000",
banScopeV4Prefix: "32",
}
if !maps.Equal(line.Settings, want) {
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
@@ -368,6 +413,22 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
got.RateLimitPerDay != want.RateLimitPerDay {
t.Errorf("settings\n%+v\nwant\n%+v", got, want)
}
wantBanSettings(t, got, want)
}
// wantBanSettings checks the settings for bans.
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper()
if got.BanResponse != want.BanResponse ||
got.LimitBanDuration != want.LimitBanDuration ||
got.LimitBanRepeatWindow != want.LimitBanRepeatWindow ||
got.MaxBanDuration != want.MaxBanDuration ||
got.MaxBans != want.MaxBans ||
got.BanScopeV4Prefix != want.BanScopeV4Prefix {
t.Errorf("ban settings\n%+v\nwant\n%+v", got, want)
}
}
// wantNetblocks checks a list of netblocks.
+82
View File
@@ -0,0 +1,82 @@
package proxy
import (
"net/netip"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// banResponse is a refusal answered with SWWAF_BAN_RESPONSE, and logged
// with action.
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.
func (rq *request) banned(now time.Time) bool {
ban, banned := rq.h.ledger.Check(rq.netblock(), now)
if banned {
rq.line.BanExpires = banExpires(ban)
}
return banned
}
// limitBroken counts the request for the rate limits at now, and reports
// whether it takes the client over one. Such a request bans the client's
// netblock, and sets the client's counters back to zero.
func (rq *request) limitBroken(now time.Time) bool {
group := clientGroup(rq.client)
hit, over := rq.h.limiter.Count(group, now)
if !over {
return false
}
ban := rq.h.ledger.BanForLimit(rq.netblock(), now, bans.Notes{
Country: rq.line.Country,
Limit: hit.Limit,
Window: hit.Window,
Count: hit.Requests,
Request: bans.Request{
Time: now,
Method: rq.in.Method,
Host: rq.in.Host,
Path: rq.in.URL.RequestURI(),
Status: rq.h.config.BanResponse,
UserAgent: rq.in.UserAgent(),
},
})
rq.h.limiter.Reset(group)
rq.line.LimitHit = hit.Window
rq.line.Offence = requestlog.OffenceLimit
rq.line.BanExpires = banExpires(ban)
return true
}
// netblock is the netblock a ban on the client covers: its IPv4 address,
// widened to SWWAF_BAN_SCOPE_V4_PREFIX, or the IPv6 group clientGroup
// counts it in.
func (rq *request) netblock() netip.Prefix {
addr := rq.client.Unmap()
if addr.Is4() {
return netip.PrefixFrom(addr, rq.h.config.BanScopeV4Prefix).Masked()
}
return clientGroup(addr)
}
// banExpires is when ban ends, as the log line gives it: a time, or
// permanent.
func banExpires(ban bans.Ban) string {
if ban.Permanent() {
return "permanent"
}
return requestlog.FormatTime(ban.Expires)
}
+434
View File
@@ -0,0 +1,434 @@
package proxy_test
import (
"bufio"
"errors"
"io"
"maps"
"net/http"
"net/netip"
"slices"
"sync"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
// otherClient is a client next to client.
otherClient = "203.0.113.10"
// userAgent is the user agent of every request a sender sends.
userAgent = "ban-test/1.0"
// permanent is the log line's ban_expires for a permanent ban.
permanent = "permanent"
)
func TestBrokenLimitBansTheClient(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{rateLimitPerMinute: "1"})
expires := requestlog.FormatTime(clk.Now().Add(time.Hour))
// The request over the limit of one a minute is refused, and bans the
// client for an hour, the default.
s.get(client, http.StatusOK, requestlog.ActionForward)
line := s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
if line.LimitHit != minute || line.Offence != requestlog.OffenceLimit ||
line.BanExpires != expires {
t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+
"want minute, limit and %s", line.LimitHit, line.Offence, line.BanExpires,
expires)
}
// Every request while the ban lasts is refused.
clk.advance(time.Hour - time.Second)
line = s.get(client, http.StatusForbidden, requestlog.ActionBanned)
if line.BanExpires != expires || line.Offence != "" || line.LimitHit != "" {
t.Errorf("log line has ban_expires %q, offence %q and limit_hit %q, "+
"want %s and neither of the others", line.BanExpires, line.Offence,
line.LimitHit, expires)
}
// Once it ends, the client is let through.
clk.advance(time.Second)
s.get(client, http.StatusOK, requestlog.ActionForward)
}
func TestBanLengthsFollowTheSettings(t *testing.T) {
t.Parallel()
s, clk, _ := startWithClock(t, "", map[string]string{
rateLimitPerMinute: "1",
limitBanDuration: "10m",
limitBanRepeatWindow: "1h",
maxBanDuration: "1h",
})
// breakLimit has client go over the limit of one a minute, and
// returns when the ban that makes ends.
breakLimit := func() string {
s.get(client, http.StatusOK, requestlog.ActionForward)
return s.get(client, http.StatusForbidden, requestlog.ActionRateLimited).BanExpires
}
wantExpires := func(got string, length time.Duration) {
t.Helper()
want := requestlog.FormatTime(clk.Now().Add(length))
if got != want {
t.Errorf("ban ends at %s, want %s", got, want)
}
}
// A first ban lasts SWWAF_LIMIT_BAN_DURATION; one within
// SWWAF_LIMIT_BAN_REPEAT_WINDOW after it ended, three times as long.
wantExpires(breakLimit(), 10*time.Minute)
clk.advance(10*time.Minute + time.Hour)
wantExpires(breakLimit(), 30*time.Minute)
// Later than that, SWWAF_LIMIT_BAN_DURATION again.
clk.advance(30*time.Minute + time.Hour + time.Second)
wantExpires(breakLimit(), 10*time.Minute)
clk.advance(10 * time.Minute)
wantExpires(breakLimit(), 30*time.Minute)
// 90 minutes would be longer than SWWAF_MAX_BAN_DURATION: the ban is
// permanent.
clk.advance(30 * time.Minute)
got := breakLimit()
if got != permanent {
t.Errorf("ban ends at %s, want a permanent one", got)
}
clk.advance(365 * 24 * time.Hour)
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
}
func TestBanIsNotCountedAndResetsTheCounters(t *testing.T) {
t.Parallel()
s, clk, server := startWithClock(t, "", map[string]string{rateLimitPerDay: "2"})
// The third request in a day is over the limit of two, and bans the
// client for an hour.
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
for range 3 {
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
}
// Later the same day the client has its whole allowance again: the
// ban set its counters back to zero, and the requests it refused were
// not counted for the rate limits, only in its notes.
clk.advance(time.Hour)
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
banned := proxy.LedgerOf(server).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)
}
}
func TestBanCoversTheClientsNetblock(t *testing.T) {
t.Parallel()
// In the IPv4 cases, client breaks the limit; these two are next to it.
const (
allowed = "203.0.113.60" // in SWWAF_ALLOW_NETS
exempt = "203.0.113.50" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
)
for _, tc := range []struct {
name string
env map[string]string
breaker string // the client that breaks the limit
refused []string
let []string // let through
}{
{
"an IPv4 address, by default", nil, client,
nil, []string{otherClient, exempt},
},
{
"the IPv4 netblock SWWAF_BAN_SCOPE_V4_PREFIX sets",
map[string]string{banScopeV4Prefix: "24"}, client,
[]string{otherClient, exempt}, []string{"203.0.112.9", allowed},
},
{
"an IPv6 /64", nil, "2001:db8:5::1",
[]string{"2001:db8:5::ffff:1"}, []string{"2001:db8:5:1::1"},
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{
rateLimitPerMinute: "1",
allowNets: allowed,
rateLimitExemptNets: exempt,
}
maps.Copy(env, tc.env)
s, _, _ := startWithClock(t, "", env)
s.get(tc.breaker, http.StatusOK, requestlog.ActionForward)
s.get(tc.breaker, http.StatusForbidden, requestlog.ActionRateLimited)
for _, sent := range tc.refused {
s.get(sent, http.StatusForbidden, requestlog.ActionBanned)
}
for _, sent := range tc.let {
s.get(sent, http.StatusOK, requestlog.ActionForward)
}
})
}
}
func TestBannedClientIsRefusedBeforeItsCountryIsLookedUp(t *testing.T) {
t.Parallel()
geojsURL, asked := startGeoJS(t)
s, _, _ := startWithClock(t, geojsURL, map[string]string{
rateLimitPerMinute: "1",
banScopeV4Prefix: "24",
deniedCountries: "kp",
})
// fromDE's ban covers otherClient, which is refused unasked about.
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
line := s.get(otherClient, http.StatusForbidden, requestlog.ActionBanned)
if line.Country != "" {
t.Errorf("log line has country %q, want none", line.Country)
}
if !slices.Equal(asked(), []string{fromDE}) {
t.Errorf("GeoJS was asked about %v, want %s alone", asked(), fromDE)
}
}
func TestBanResponseAnswersEveryRefusalButTheSizeLimits(t *testing.T) {
t.Parallel()
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
for _, tc := range []struct {
setting string // "" leaves SWWAF_BAN_RESPONSE at its default
status int // 0 is the connection closed without an answer
}{
{"", http.StatusForbidden},
{"403", http.StatusForbidden},
{"429", http.StatusTooManyRequests},
{"close", 0},
} {
t.Run(banResponse+"="+tc.setting, func(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
env := map[string]string{
rateLimitPerMinute: "1",
denyNets: denied,
deniedCountries: "kp",
}
if tc.setting != "" {
env[banResponse] = tc.setting
}
s, _, _ := startWithClock(t, geojsURL, env)
s.get(denied, tc.status, requestlog.ActionDenied)
s.get(fromKP, tc.status, requestlog.ActionCountryDenied)
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromDE, tc.status, requestlog.ActionRateLimited)
s.get(fromDE, tc.status, requestlog.ActionBanned)
})
}
}
func TestBanNotes(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
s, clk, server := startWithClock(t, geojsURL, map[string]string{
rateLimitPerMinute: "1",
deniedCountries: "kp",
})
start := clk.Now()
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.request(fromDE, "/repo/commits?page=2",
http.StatusForbidden, requestlog.ActionRateLimited)
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
netblock := netip.MustParsePrefix(fromDE + "/32")
want := bans.Ban{
Netblock: netblock,
Start: start,
Expires: start.Add(time.Hour),
Notes: bans.Notes{
Country: "DE",
Limit: 1,
Window: minute,
Count: 2,
Request: bans.Request{
Time: start,
Method: http.MethodGet,
Host: appHost,
Path: "/repo/commits?page=2",
Status: http.StatusForbidden,
UserAgent: userAgent,
},
Refused: 2,
EarlierBans: 0,
},
}
ledger := proxy.LedgerOf(server)
got := ledger.Bans(netblock)
if len(got) != 1 || got[0] != want {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
}
// The next ban counts this one among the earlier.
clk.advance(time.Hour)
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
got = ledger.Bans(netblock)
if len(got) != 2 || got[1].Notes.EarlierBans != 1 {
t.Errorf("bans %+v, want two, the second with one earlier ban", got)
}
}
func TestMaxBansDropsTheBanOfTheNetblockSeenLongestAgo(t *testing.T) {
t.Parallel()
s, _, _ := startWithClock(t, "", map[string]string{
rateLimitPerMinute: "1",
maxBans: "1",
})
// One ban is held, so otherClient's ban drops client's.
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
s.get(otherClient, http.StatusOK, requestlog.ActionForward)
s.get(otherClient, http.StatusForbidden, requestlog.ActionRateLimited)
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(otherClient, http.StatusForbidden, requestlog.ActionBanned)
}
// clock is the time a test sets, by which smallwebwaf counts requests and
// makes bans.
type clock struct {
mu sync.Mutex
now time.Time
}
// Now tells the time.
func (c *clock) Now() time.Time {
c.mu.Lock()
defer c.mu.Unlock()
return c.now
}
// advance moves the clock on by d.
func (c *clock) advance(d time.Duration) {
c.mu.Lock()
defer c.mu.Unlock()
c.now = c.now.Add(d)
}
// startWithClock starts smallwebwaf in front of an app that answers 200,
// with the settings in env on top of trusting localhost's
// X-Forwarded-For, clients' countries looked up at geojsURL, and a clock
// 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) {
t.Helper()
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
settings := map[string]string{trustedProxies: trustLocalhost}
maps.Copy(settings, env)
addr, out, server := startProxyWithClock(t, app.URL, geojsURL, clk.Now, settings)
return &sender{t: t, addr: addr, out: out}, clk, server
}
// sender sends requests to smallwebwaf one after another, each on a
// connection of its own, and checks each one's answer and log line. They
// must be the only requests smallwebwaf is sent, since the log lines are
// matched to them in order.
type sender struct {
t *testing.T
addr string
out *output
sent int
}
// get sends a GET request for / from the client at from.
func (s *sender) get(from string, status int, action string) logLine {
s.t.Helper()
return s.request(from, "/", status, action)
}
// request sends a GET request for path from the client at from, as
// X-Forwarded-For names it, and checks that its answer and its log line
// have status, 0 for the connection closed without an answer, and that
// the line has action. It returns the log line.
func (s *sender) request(from, path string, status int, action string) logLine {
s.t.Helper()
conn := dial(s.t, s.addr)
send(s.t, conn, "GET "+path+" HTTP/1.1\r\nHost: "+appHost+
"\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n\r\n")
err := conn.SetReadDeadline(time.Now().Add(waitLimit))
if err != nil {
s.t.Fatalf("set read deadline: %v", err)
}
got := 0
res, err := http.ReadResponse(bufio.NewReader(conn), nil)
switch {
case err == nil:
got = readAnswer(res).status
case !errors.Is(err, io.ErrUnexpectedEOF):
s.t.Fatalf("read response: %v", err)
}
_ = conn.Close()
if got != status {
s.t.Errorf("request %d, from %s: status %d, want %d", s.sent+1, from, got,
status)
}
line := s.out.requestLines(s.t, s.sent+1)[s.sent]
s.sent++
wantLine(s.t, line, status, action)
return line
}
+15
View File
@@ -0,0 +1,15 @@
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
}
+13
View File
@@ -10,6 +10,7 @@ import (
"net/http"
"time"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
@@ -36,6 +37,9 @@ type Params struct {
// GeoJSURL is where clients' countries are looked up, normally
// 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.
Now func() time.Time
}
// New returns the server smallwebwaf runs: each request it reads passes
@@ -55,11 +59,18 @@ func New(params Params) *http.Server {
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,
@@ -85,7 +96,9 @@ type handler struct {
processLog *slog.Logger
errorLog *log.Logger
transport http.RoundTripper
now func() time.Time
limiter *ratelimit.Limiter
ledger *bans.Ledger
geojs *lookup.GeoJS
}
+22 -1
View File
@@ -57,8 +57,15 @@ const (
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
denyNets = "SWWAF_DENY_NETS"
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
deniedCountries = "SWWAF_DENIED_COUNTRIES"
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
banResponse = "SWWAF_BAN_RESPONSE"
limitBanDuration = "SWWAF_LIMIT_BAN_DURATION"
limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW"
maxBanDuration = "SWWAF_MAX_BAN_DURATION"
maxBans = "SWWAF_MAX_BANS"
banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX"
)
// output collects what smallwebwaf writes on stdout.
@@ -183,6 +190,19 @@ func startProxyWithGeoJS(
) (string, *output) {
t.Helper()
addr, out, _ := startProxyWithClock(t, appURL, geojsURL, time.Now, env)
return addr, out
}
// startProxyWithClock is startProxyWithGeoJS with requests counted and
// bans made by the time now tells, and returns the server as well.
func startProxyWithClock(
t *testing.T, appURL, geojsURL string, now func() time.Time,
env map[string]string,
) (string, *output, *http.Server) {
t.Helper()
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL}
maps.Copy(settings, env)
@@ -201,6 +221,7 @@ func startProxyWithGeoJS(
RequestLog: out,
ProcessLog: requestlog.NewProcessLogger(out),
GeoJSURL: geojsURL,
Now: now,
})
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
@@ -216,7 +237,7 @@ func startProxyWithGeoJS(
_ = server.Close()
})
return listener.Addr().String(), out
return listener.Addr().String(), out, server
}
// newClient returns an HTTP client that sends requests as they are made,
+12 -8
View File
@@ -8,7 +8,11 @@ import (
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
func TestRateLimitRefusesWith429BeforeTheApp(t *testing.T) {
// minute is the window of SWWAF_RATE_LIMIT_PER_MINUTE, as a log line's
// limit_hit names it.
const minute = "minute"
func TestRateLimitRefusesBeforeTheApp(t *testing.T) {
t.Parallel()
var calls atomic.Int32
@@ -24,19 +28,19 @@ func TestRateLimitRefusesWith429BeforeTheApp(t *testing.T) {
const otherClient = "203.0.113.10"
// With a limit of one request a minute, a client's second request is
// refused. A client is one IPv4 address, or one IPv6 /64; an IPv4
// address in IPv6 form is that IPv4 address.
// refused, with 403 by default. A client is one IPv4 address, or one
// IPv6 /64; an IPv4 address in IPv6 form is that IPv4 address.
requests := []struct {
client string // as X-Forwarded-For names it
logged string // as the log line's client_ip names it
want int
}{
{client, client, http.StatusOK},
{client, client, http.StatusTooManyRequests},
{client, client, http.StatusForbidden},
{otherClient, otherClient, http.StatusOK},
{"::ffff:" + otherClient, otherClient, http.StatusTooManyRequests},
{"::ffff:" + otherClient, otherClient, http.StatusForbidden},
{"2001:db8::1", "2001:db8::1", http.StatusOK},
{"2001:db8::8000:0:0:1", "2001:db8::8000:0:0:1", http.StatusTooManyRequests},
{"2001:db8::8000:0:0:1", "2001:db8::8000:0:0:1", http.StatusForbidden},
{"2001:db8:0:1::1", "2001:db8:0:1::1", http.StatusOK},
}
@@ -53,9 +57,9 @@ func TestRateLimitRefusesWith429BeforeTheApp(t *testing.T) {
if sent.want == http.StatusOK {
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
} else {
wantLine(t, line, http.StatusTooManyRequests, requestlog.ActionRateLimited)
wantLine(t, line, http.StatusForbidden, requestlog.ActionRateLimited)
if line.LimitHit != "minute" {
if line.LimitHit != minute {
t.Errorf("log line has limit_hit %q, want minute", line.LimitHit)
}
}
+25 -23
View File
@@ -21,7 +21,8 @@ const flushAfterEachWrite time.Duration = -1
// refusal is smallwebwaf refusing a request, or refusing to go on with it:
// the status the client is answered if the response has not started yet,
// and the action the log line names.
// 0 to close the connection without an answer, and the action the log
// line names.
type refusal struct {
status int
action string
@@ -105,39 +106,33 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
// is known, before its body is read or anything reaches the app. It
// returns nil to let the request through. A client in SWWAF_ALLOW_NETS
// skips every check but the size limit. For any other client,
// SWWAF_DENY_NETS comes first, so that a client it refuses is not looked
// up, and then the country lists; a request either refuses is not counted
// for the rate limits. Then come the rate limits, unless the client is in
// SWWAF_DENY_NETS comes first, then a ban on its netblock, so that a
// client either refuses is not looked up, and then the country lists; a
// request any of them refuses is not counted for the rate limits. Then
// come the rate limits, unless the client is in
// SWWAF_RATE_LIMIT_EXEMPT_NETS, so that every other request is counted,
// one refused for its size too. ctx is the request's own context.
// one refused for its size too. Every refusal but the size limit's is
// answered with SWWAF_BAN_RESPONSE. ctx is the request's own context.
func (rq *request) check(ctx context.Context) *refusal {
cfg := rq.h.config
allowed := isInside(rq.client, cfg.AllowNets)
exempt := isInside(rq.client, cfg.RateLimitExemptNets)
now := rq.h.now()
if !allowed && isInside(rq.client, cfg.DenyNets) {
return &refusal{
status: http.StatusForbidden,
action: requestlog.ActionDenied,
}
return rq.banResponse(requestlog.ActionDenied)
}
if !allowed && rq.banned(now) {
return rq.banResponse(requestlog.ActionBanned)
}
if !allowed && rq.countryDenied(ctx) {
return &refusal{
status: http.StatusForbidden,
action: requestlog.ActionCountryDenied,
}
return rq.banResponse(requestlog.ActionCountryDenied)
}
if !allowed && !isInside(rq.client, cfg.RateLimitExemptNets) {
limitHit := rq.h.limiter.Count(clientGroup(rq.client), rq.start)
if limitHit != "" {
rq.line.LimitHit = limitHit
return &refusal{
status: http.StatusTooManyRequests,
action: requestlog.ActionRateLimited,
}
}
if !allowed && !exempt && rq.limitBroken(now) {
return rq.banResponse(requestlog.ActionRateLimited)
}
maxBytes := cfg.RequestMaxBytes
@@ -250,6 +245,13 @@ func (rq *request) answer(r refusal) {
return // too late to answer: the connection can only be cut
}
if r.status == 0 {
// SWWAF_BAN_RESPONSE is close. This panic has Go's server close
// the connection without an answer, and log nothing; the log line
// is still written as the handler returns.
panic(http.ErrAbortHandler)
}
// A client found too slow is read no more; any other may go on
// sending until its time is up, so that Go's server can read the
// rest of the body and end the request cleanly.
+3 -3
View File
@@ -78,7 +78,7 @@ func TestRequestFromAllowNetsIsNotCounted(t *testing.T) {
{listedAddr, http.StatusOK, requestlog.ActionForward},
{listedAddr, http.StatusOK, requestlog.ActionForward},
{unlistedAddr, http.StatusOK, requestlog.ActionForward},
{unlistedAddr, http.StatusTooManyRequests, requestlog.ActionRateLimited},
{unlistedAddr, http.StatusForbidden, requestlog.ActionRateLimited},
})
}
@@ -133,7 +133,7 @@ func TestRequestRefusedByDenyNetsIsNotCounted(t *testing.T) {
{listedAddr, http.StatusForbidden, requestlog.ActionDenied},
{listedAddr, http.StatusForbidden, requestlog.ActionDenied},
{unlistedAddr, http.StatusOK, requestlog.ActionForward},
{unlistedAddr, http.StatusTooManyRequests, requestlog.ActionRateLimited},
{unlistedAddr, http.StatusForbidden, requestlog.ActionRateLimited},
})
}
@@ -156,7 +156,7 @@ func TestRateLimitExemptNetsAreNeitherCountedNorRefused(t *testing.T) {
{listedAddr, http.StatusOK, requestlog.ActionForward},
{listedAddr, http.StatusOK, requestlog.ActionForward},
{unlistedAddr, http.StatusOK, requestlog.ActionForward},
{unlistedAddr, http.StatusTooManyRequests, requestlog.ActionRateLimited},
{unlistedAddr, http.StatusForbidden, requestlog.ActionRateLimited},
{fromKP, http.StatusForbidden, requestlog.ActionCountryDenied},
})
}
+27 -8
View File
@@ -54,11 +54,22 @@ func New(limits Limits) *Limiter {
}
}
// Hit is a request that takes a client over a rate limit.
type Hit struct {
// Window is "minute", "hour" or "day".
Window string
// Limit is the window's limit.
Limit int64
// Requests is the client's requests counted in the window, this one
// included.
Requests float64
}
// Count counts a request from client at now, in every window, whether or
// not it is refused. It returns the window whose limit the request takes
// the client over, "minute", "hour" or "day", the shortest if it is over
// several, or "" if it is within every limit.
func (l *Limiter) Count(client netip.Prefix, now time.Time) string {
// not it is refused. It reports whether the request takes the client over
// a limit, and the window whose limit it goes over, the shortest if it is
// over several.
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Hit, bool) {
l.mu.Lock()
defer l.mu.Unlock()
@@ -68,16 +79,24 @@ func (l *Limiter) Count(client netip.Prefix, now time.Time) string {
l.clients.Add(client, counts)
}
limitHit := ""
var hit Hit
for i, w := range l.windows {
requests := counts[i].add(now, w.length)
if limitHit == "" && w.limit > 0 && requests > float64(w.limit) {
limitHit = w.name
if hit.Window == "" && w.limit > 0 && requests > float64(w.limit) {
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests}
}
}
return limitHit
return hit, hit.Window != ""
}
// Reset sets client's counts in every window back to zero.
func (l *Limiter) Reset(client netip.Prefix) {
l.mu.Lock()
defer l.mu.Unlock()
l.clients.Remove(client)
}
// window is a length of time over which requests are counted, and the
+49 -3
View File
@@ -54,6 +54,52 @@ func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) {
}
}
func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: limit})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range limit {
_, over := limiter.Count(client, start)
if over {
t.Fatal("a request within the limit is over it")
}
}
// Over both limits; the minute's is named, with the four requests.
hit, over := limiter.Count(client, start)
want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1}
if !over || hit != want {
t.Errorf("request over the limit gives %+v and %t, want %+v and true",
hit, over, want)
}
}
func TestResetSetsTheCountsBackToZero(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerDay: limit})
client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight()
for range limit {
wantCount(t, limiter, client, start, "")
}
wantCount(t, limiter, client, start, minute)
limiter.Reset(client)
// At the same moment, the client has its whole allowance again.
for range limit {
wantCount(t, limiter, client, start, "")
}
wantCount(t, limiter, client, start, minute)
}
func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) {
t.Parallel()
@@ -192,9 +238,9 @@ func wantCount(
) {
t.Helper()
got := limiter.Count(client, now)
if got != want {
hit, _ := limiter.Count(client, now)
if hit.Window != want {
t.Errorf("request from %s at %s is over %q, want %q",
client, now.Format(time.RFC3339), got, want)
client, now.Format(time.RFC3339), hit.Window, want)
}
}
+12 -1
View File
@@ -24,8 +24,10 @@ const (
// for, or whose answer could not be passed on.
ActionUpstreamError = "upstream_error"
// ActionRateLimited is a request refused because it took its client
// over a rate limit, or came while the client was over one.
// over a rate limit, which bans the client.
ActionRateLimited = "rate_limited"
// ActionBanned is a request refused because a ban covers its client.
ActionBanned = "banned"
// ActionDenied is a request refused because its client is in
// SWWAF_DENY_NETS.
ActionDenied = "denied"
@@ -36,6 +38,10 @@ const (
ActionAdmin = "admin"
)
// OffenceLimit is the offence a request line names for a request that
// broke a rate limit.
const OffenceLimit = "limit"
// timeLayout is RFC 3339 with milliseconds.
const timeLayout = "2006-01-02T15:04:05.000Z07:00"
@@ -64,6 +70,11 @@ type Line struct {
// LimitHit is the window whose rate limit the request went over:
// minute, hour or day.
LimitHit string `json:"limit_hit,omitempty"`
// Offence is the offence the request was held as, OffenceLimit.
Offence string `json:"offence,omitempty"`
// BanExpires is when the ban the request made, or was refused under,
// ends: a time, or "permanent".
BanExpires string `json:"ban_expires,omitempty"`
// Aborted is true when the client went away early.
Aborted bool `json:"aborted,omitempty"`
// DurationTotal and DurationUpstreamTotal are in milliseconds.
+2 -1
View File
@@ -50,7 +50,8 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
}
unset := []string{
"upstream_status", "limit_hit", "aborted", "duration_upstream_total",
"upstream_status", "limit_hit", "offence", "ban_expires", "aborted",
"duration_upstream_total",
}
for _, name := range unset {
_, present := fields[name]
+1
View File
@@ -80,6 +80,7 @@ func Run(ctx context.Context, params Params) int {
RequestLog: params.Stdout,
ProcessLog: processLog,
GeoJSURL: lookup.URL,
Now: time.Now,
})
processLog.Info("starting",
+6
View File
@@ -196,6 +196,12 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL string) {
"SWWAF_RATE_LIMIT_PER_DAY": "50000",
"SWWAF_DENIED_COUNTRIES": "",
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "",
"SWWAF_BAN_RESPONSE": "403",
"SWWAF_LIMIT_BAN_DURATION": "1h",
"SWWAF_LIMIT_BAN_REPEAT_WINDOW": "24h",
"SWWAF_MAX_BAN_DURATION": "7d",
"SWWAF_MAX_BANS": "5000",
"SWWAF_BAN_SCOPE_V4_PREFIX": "32",
}
for name, value := range want {