Ban the netblock of a client that breaks a rate limit, in memory (closes #18)
check / check (push) Successful in 4m11s
check / check (push) Successful in 4m11s
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 is contained in:
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
@@ -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},
|
||||
})
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user