Compare commits
11
Commits
d4f5b3e404
..
next
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
04e66d2069 | ||
|
|
a6634454cd | ||
|
|
ca787985f8 | ||
|
|
0b0f207423 | ||
|
|
2b8c98ba1f | ||
|
|
82e20e0cb5 | ||
|
|
2421cdc273 | ||
|
|
f35e3ddfe8 | ||
|
|
0dc26041dc | ||
|
|
c80753c56e | ||
|
|
26f4abef7f |
@@ -61,6 +61,10 @@ linters:
|
||||
desc: >-
|
||||
Test-support code belongs in test files and in packages whose
|
||||
directory name ends in test, not in the shipped binary.
|
||||
- pkg: sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest
|
||||
desc: >-
|
||||
Test-support code belongs in test files and in packages whose
|
||||
directory name ends in test, not in the shipped binary.
|
||||
# Only decisions already recorded in the Go package defaults are
|
||||
# listed here. Every entry matches the module path exactly.
|
||||
gomodguard_v2:
|
||||
|
||||
+25
-1
@@ -29,6 +29,12 @@ RUN go mod download
|
||||
|
||||
COPY . .
|
||||
|
||||
# go.mod and go.sum must be as `go mod tidy` writes them, which is what
|
||||
# `make tidy` does. Checked before the tests, which a missing go.sum line
|
||||
# fails with a message that does not name `make tidy`.
|
||||
RUN go mod tidy -diff || \
|
||||
{ echo "go.mod or go.sum is not tidy: run make tidy" >&2; exit 1; }
|
||||
|
||||
# Go's build cache is kept on a tmpfs, out of the image: nothing uses it
|
||||
# after this step, and writing it into the image takes seconds.
|
||||
RUN --mount=type=tmpfs,target=/root/.cache/go-build \
|
||||
@@ -36,7 +42,25 @@ RUN --mount=type=tmpfs,target=/root/.cache/go-build \
|
||||
{ echo "--- Rerunning with -v for details ---"; \
|
||||
go test -timeout 90s -race -v ./...; exit 1; }
|
||||
|
||||
# Build stage. Nothing is wanted from the two phases above; the copies
|
||||
# Tidy stage: `go mod tidy` in the test phase's Go, so that the files it
|
||||
# writes pass the test phase's check. Nothing else depends on it, so only
|
||||
# script/tidy, which names the stage after it, builds it.
|
||||
#
|
||||
# golang 1.27.1-trixie, 2026-09-19
|
||||
FROM golang@sha256:3b77fc618ec235a1ab412de7737f120dd507c57e8d87de4cbb7994fb94275ed5 AS tidy
|
||||
|
||||
WORKDIR /src
|
||||
|
||||
COPY . .
|
||||
|
||||
RUN go mod tidy
|
||||
|
||||
# go.mod and go.sum alone, which script/tidy writes into the working tree.
|
||||
FROM scratch AS tidy-files
|
||||
|
||||
COPY --from=tidy /src/go.mod /src/go.sum /
|
||||
|
||||
# Build stage. Nothing is wanted from the lint and test phases; the copies
|
||||
# are what make BuildKit build them first, so the image, which needs this
|
||||
# stage, cannot be produced unless lint and test passed.
|
||||
#
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
.PHONY: bootstrap setup test lint fmt fmt-check check docker hooks build run \
|
||||
example-app
|
||||
.PHONY: bootstrap setup test lint fmt fmt-check tidy check docker hooks build \
|
||||
run example-app
|
||||
|
||||
# Makefile targets are thin shims; the implementations live in script/
|
||||
# per the scripts-to-rule-them-all pattern (see the Entrypoints section
|
||||
# of README.md). build and run are for working on the code by hand;
|
||||
# of README.md). tidy writes go.mod and go.sum as `go mod tidy` does,
|
||||
# which test checks. build and run are for working on the code by hand;
|
||||
# example-app checks the image with an app built on it.
|
||||
|
||||
bootstrap:
|
||||
@@ -24,6 +25,9 @@ fmt:
|
||||
fmt-check:
|
||||
@script/fmt-check
|
||||
|
||||
tidy:
|
||||
@script/tidy
|
||||
|
||||
check:
|
||||
@script/check
|
||||
|
||||
|
||||
@@ -5,6 +5,8 @@ go 1.26.0
|
||||
require (
|
||||
github.com/fsnotify/fsnotify v1.10.1
|
||||
github.com/hashicorp/golang-lru/v2 v2.0.7
|
||||
github.com/maxmind/mmdbwriter v1.2.0
|
||||
github.com/oschwald/maxminddb-golang/v2 v2.7.0
|
||||
github.com/prometheus/client_golang v1.24.1
|
||||
)
|
||||
|
||||
@@ -16,6 +18,7 @@ require (
|
||||
github.com/prometheus/client_model v0.6.2 // indirect
|
||||
github.com/prometheus/common v0.70.1 // indirect
|
||||
github.com/prometheus/procfs v0.21.1 // indirect
|
||||
golang.org/x/sys v0.47.0 // indirect
|
||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba // indirect
|
||||
golang.org/x/sys v0.48.0 // indirect
|
||||
google.golang.org/protobuf v1.36.11 // indirect
|
||||
)
|
||||
|
||||
@@ -2,8 +2,6 @@ github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
|
||||
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
|
||||
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
||||
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
|
||||
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
|
||||
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||
@@ -14,10 +12,12 @@ github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJn
|
||||
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
|
||||
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
|
||||
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
|
||||
github.com/maxmind/mmdbwriter v1.2.0 h1:hyvDopImmgvle3aR8AaddxXnT0iQH2KWJX3vNfkwzYM=
|
||||
github.com/maxmind/mmdbwriter v1.2.0/go.mod h1:EQmKHhk2y9DRVvyNxwCLKC5FrkXZLx4snc5OlLY5XLE=
|
||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
|
||||
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||
github.com/oschwald/maxminddb-golang/v2 v2.7.0 h1:ZcAr3GYc2LYC8aec2mCMX9+QOF0EolH3jDFKRV/Z1+U=
|
||||
github.com/oschwald/maxminddb-golang/v2 v2.7.0/go.mod h1:DuKJLbbug6TXC0yJXgs1MWifvXHmudRWzMobMIUu04g=
|
||||
github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU=
|
||||
github.com/prometheus/client_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE=
|
||||
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
|
||||
@@ -26,15 +26,17 @@ github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi
|
||||
github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc=
|
||||
github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI=
|
||||
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
|
||||
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
|
||||
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
|
||||
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
|
||||
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
|
||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
||||
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg=
|
||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBseWJUpBw5I82+2U4M=
|
||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
|
||||
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
|
||||
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
|
||||
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
||||
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
|
||||
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
|
||||
|
||||
+92
-45
@@ -1,5 +1,6 @@
|
||||
// Package alerts sends alerts on bans, on a source that fails and on a
|
||||
// file with an error to each destination set: to the webhook
|
||||
// Package alerts sends alerts on bans, on traffic over an anomaly
|
||||
// threshold, on a source that fails and on a file with an error to each
|
||||
// destination set: to the webhook
|
||||
// SWWAF_ALERT_WEBHOOK_URL names, each as one JSON object, as the "Alert
|
||||
// webhook schema" section of SPEC.md describes, to the Slack incoming
|
||||
// webhook SWWAF_ALERT_SLACK_WEBHOOK_URL names, as a message, and to the
|
||||
@@ -40,19 +41,27 @@ const (
|
||||
// EventPermanentBan is a permanent ban smallwebwaf made, or a ban it
|
||||
// made permanent.
|
||||
EventPermanentBan = "permanent_ban"
|
||||
// EventWAFBlock, EventAnomaly and EventReputationHit come with the
|
||||
// Core Rule Set, the anomaly thresholds and the reputation sources;
|
||||
// nothing raises them yet.
|
||||
EventWAFBlock = "waf_block"
|
||||
// EventAnomaly is a count of requests or bytes over an anomaly
|
||||
// threshold.
|
||||
EventAnomaly = "anomaly"
|
||||
// EventWAFBlock comes with the Core Rule Set; nothing raises it yet.
|
||||
EventWAFBlock = "waf_block"
|
||||
// EventReputationHit is a request whose client a blocklist or a DNSBL
|
||||
// zone lists, or whose AbuseIPDB score is a hit.
|
||||
EventReputationHit = "reputation_hit"
|
||||
// EventSourceFailure is GeoJS failing or refusing smallwebwaf.
|
||||
// EventSourceFailure is GeoJS failing or refusing smallwebwaf, a fetch
|
||||
// of a list failing, a query to a DNSBL zone or a check with AbuseIPDB
|
||||
// failing or refused, or the day's AbuseIPDB checks used up.
|
||||
EventSourceFailure = "source_failure"
|
||||
// EventFileError is a rule file or state file edited while smallwebwaf
|
||||
// runs that does not parse, or a state file that cannot be written.
|
||||
// runs that does not parse, a replacement of the lookup database that
|
||||
// cannot be read, or a state file that cannot be written.
|
||||
EventFileError = "file_error"
|
||||
// EventSummary is the summary of the alerts an hour held back past
|
||||
// SWWAF_ALERT_MAX_PER_HOUR. SWWAF_ALERT_EVENTS does not name it.
|
||||
// EventSummary is the summary sent as an hour ends: of the alerts held
|
||||
// back in it past SWWAF_ALERT_MAX_PER_HOUR, and of the repeats held
|
||||
// back by the cooldowns dropped as it ends, which no alert let through
|
||||
// has given. It is sent with SWWAF_ALERT_MAX_PER_HOUR off too, for
|
||||
// those repeats. SWWAF_ALERT_EVENTS does not name it.
|
||||
EventSummary = "summary"
|
||||
)
|
||||
|
||||
@@ -152,18 +161,22 @@ type Alert struct {
|
||||
ASName string `json:"as_name"`
|
||||
Country string `json:"country"`
|
||||
// Reason is a short sentence, and Detail what is particular to the
|
||||
// event: for a file_error, its "file", and for a source_failure, its
|
||||
// "source", which the cooldown tells repeats by.
|
||||
// event: for a file_error, its "file", for a source_failure, its
|
||||
// "source", and for an anomaly, its "scope", with the "asn" or the
|
||||
// "name" of some scopes, which the cooldown tells repeats by.
|
||||
Reason string `json:"reason"`
|
||||
Detail map[string]any `json:"detail"`
|
||||
// SuppressedRepeats is how many repeats of the alert the cooldown
|
||||
// held back since the last one let through.
|
||||
// held back since the last one let through. For a summary, it is how
|
||||
// many the cooldowns dropped as the hour ended had held back that no
|
||||
// alert let through gave.
|
||||
SuppressedRepeats int `json:"suppressed_repeats"`
|
||||
}
|
||||
|
||||
// Cooldown is, for an event on a netblock, or about a file or a source,
|
||||
// when the last alert let through was raised, and how many repeats the
|
||||
// cooldown has held back since, as alerts.json holds it.
|
||||
// Cooldown is, for an event on a netblock, about a file or a source, or
|
||||
// for an anomaly in a scope, when the last alert let through was raised,
|
||||
// and how many repeats the cooldown has held back since, as alerts.json
|
||||
// holds it.
|
||||
//
|
||||
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||
type Cooldown struct {
|
||||
@@ -171,6 +184,9 @@ type Cooldown struct {
|
||||
Netblock netip.Prefix `json:"netblock"`
|
||||
File string `json:"file,omitempty"`
|
||||
Source string `json:"source,omitempty"`
|
||||
Scope string `json:"scope,omitempty"`
|
||||
ASN string `json:"asn,omitempty"`
|
||||
Name string `json:"name,omitempty"`
|
||||
Sent time.Time `json:"sent"`
|
||||
SuppressedRepeats int `json:"suppressed_repeats"`
|
||||
}
|
||||
@@ -213,7 +229,7 @@ type Queue struct {
|
||||
|
||||
mu sync.Mutex
|
||||
// cooldowns are the alerts last let through, by event and netblock,
|
||||
// file or source.
|
||||
// file, source or scope.
|
||||
cooldowns map[cooldownKey]*Cooldown
|
||||
hour Hour
|
||||
|
||||
@@ -247,21 +263,28 @@ type destination struct {
|
||||
}
|
||||
|
||||
// cooldownKey is what makes an alert a repeat of another: the same event
|
||||
// on the same netblock, and about the same file or source, as its detail
|
||||
// names them. Each is empty for an alert without one.
|
||||
// on the same netblock, and about the same file or source, or in the same
|
||||
// scope with the same AS number or name, as its detail names them. Each
|
||||
// is empty for an alert without one.
|
||||
type cooldownKey struct {
|
||||
event string
|
||||
netblock netip.Prefix
|
||||
file string
|
||||
source string
|
||||
scope string
|
||||
asn string
|
||||
name string
|
||||
}
|
||||
|
||||
// cooldownKeyOf returns what makes another alert a repeat of alert.
|
||||
func cooldownKeyOf(alert *Alert) cooldownKey {
|
||||
file, _ := alert.Detail["file"].(string)
|
||||
source, _ := alert.Detail["source"].(string)
|
||||
scope, _ := alert.Detail["scope"].(string)
|
||||
asn, _ := alert.Detail["asn"].(string)
|
||||
name, _ := alert.Detail["name"].(string)
|
||||
|
||||
return cooldownKey{alert.Event, alert.Netblock, file, source}
|
||||
return cooldownKey{alert.Event, alert.Netblock, file, source, scope, asn, name}
|
||||
}
|
||||
|
||||
// New returns a Queue with no alert yet.
|
||||
@@ -294,14 +317,14 @@ func New(params Params) *Queue {
|
||||
// unless no destination is set or SWWAF_ALERT_EVENTS leaves its event
|
||||
// out. It gives alert the instance and the time. An alert that repeats
|
||||
// the last one let through less than Cooldown before is held back and
|
||||
// counted, and the next one let through gives that count. Past MaxPerHour
|
||||
// alerts let through in the hour under way, by the clock, an alert is
|
||||
// held back for that hour's summary instead, which is sent once the hour
|
||||
// has ended; it starts no cooldown, and the repeats held back before it
|
||||
// are given by the next alert let through. Raise never waits: an alert
|
||||
// let through joins the queue of each destination, from which Run sends
|
||||
// it, and with queueSize alerts waiting for a destination, the oldest is
|
||||
// dropped.
|
||||
// counted. The next one let through gives that count, unless an hour of
|
||||
// the clock ends first after the cooldown has run out: the cooldown is
|
||||
// then dropped, and that hour's summary gives the count. Past MaxPerHour
|
||||
// alerts let through in the hour under way, an alert is held back for
|
||||
// that hour's summary instead, which is sent once the hour has ended; it
|
||||
// starts no cooldown. Raise never waits: an alert let through joins the
|
||||
// queue of each destination, from which Run sends it, and with queueSize
|
||||
// alerts waiting for a destination, the oldest is dropped.
|
||||
func (q *Queue) Raise(alert Alert) {
|
||||
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, alert.Event) {
|
||||
return
|
||||
@@ -424,7 +447,8 @@ func (q *Queue) Suppressed() int64 {
|
||||
}
|
||||
|
||||
// Snapshot returns the queue's state, as alerts.json holds it, with the
|
||||
// cooldowns sorted by netblock, then by event, file and source.
|
||||
// cooldowns sorted by netblock, then by event, file, source, scope, AS
|
||||
// number and name.
|
||||
func (q *Queue) Snapshot() State {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
@@ -442,7 +466,9 @@ func (q *Queue) Snapshot() State {
|
||||
|
||||
slices.SortFunc(state.Cooldowns, func(a, b Cooldown) int {
|
||||
return cmp.Or(a.Netblock.Compare(b.Netblock), cmp.Compare(a.Event, b.Event),
|
||||
cmp.Compare(a.File, b.File), cmp.Compare(a.Source, b.Source))
|
||||
cmp.Compare(a.File, b.File), cmp.Compare(a.Source, b.Source),
|
||||
cmp.Compare(a.Scope, b.Scope), cmp.Compare(a.ASN, b.ASN),
|
||||
cmp.Compare(a.Name, b.Name))
|
||||
})
|
||||
|
||||
for _, d := range q.destinations {
|
||||
@@ -465,7 +491,10 @@ func (q *Queue) Load(state State) {
|
||||
|
||||
for _, cooldown := range state.Cooldowns {
|
||||
cooldown.Netblock = cooldown.Netblock.Masked()
|
||||
key := cooldownKey{cooldown.Event, cooldown.Netblock, cooldown.File, cooldown.Source}
|
||||
key := cooldownKey{
|
||||
cooldown.Event, cooldown.Netblock, cooldown.File, cooldown.Source,
|
||||
cooldown.Scope, cooldown.ASN, cooldown.Name,
|
||||
}
|
||||
q.cooldowns[key] = &cooldown
|
||||
}
|
||||
|
||||
@@ -537,46 +566,64 @@ func (q *Queue) startCooldown(alert *Alert, now time.Time) {
|
||||
|
||||
q.cooldowns[key] = &Cooldown{
|
||||
Event: alert.Event, Netblock: alert.Netblock, File: key.file, Source: key.source,
|
||||
Sent: now,
|
||||
Scope: key.scope, ASN: key.asn, Name: key.name, Sent: now,
|
||||
}
|
||||
}
|
||||
|
||||
// endHour ends the hour under way, if now is past it: it queues that
|
||||
// hour's summary when alerts were held back in it past MaxPerHour, and
|
||||
// forgets the cooldowns that have run out with no repeat held back, which
|
||||
// no alert needs any more.
|
||||
// endHour ends the hour under way, if now is past it. It drops the
|
||||
// cooldowns that have run out, whatever repeats they held back, so that
|
||||
// they do not pile up, and queues that hour's summary when alerts were
|
||||
// held back in it past MaxPerHour, or when a cooldown dropped had held
|
||||
// back repeats, which no alert let through has given: the summary gives
|
||||
// them.
|
||||
func (q *Queue) endHour(now time.Time) {
|
||||
start := now.Truncate(time.Hour)
|
||||
if !start.After(q.hour.Start) {
|
||||
return
|
||||
}
|
||||
|
||||
repeats := 0
|
||||
|
||||
for key, cooldown := range q.cooldowns {
|
||||
if now.Sub(cooldown.Sent) >= q.params.Cooldown {
|
||||
repeats += cooldown.SuppressedRepeats
|
||||
|
||||
delete(q.cooldowns, key)
|
||||
}
|
||||
}
|
||||
|
||||
heldBack := 0
|
||||
for _, count := range q.hour.HeldBack {
|
||||
heldBack += count
|
||||
}
|
||||
|
||||
var reasons []string
|
||||
|
||||
if heldBack > 0 {
|
||||
reasons = append(reasons, fmt.Sprintf("%d alerts held back in the hour from %s, "+
|
||||
"past the %d an hour SWWAF_ALERT_MAX_PER_HOUR allows", heldBack,
|
||||
q.hour.Start.Format(time.RFC3339), q.params.MaxPerHour))
|
||||
}
|
||||
|
||||
if repeats > 0 {
|
||||
reasons = append(reasons, fmt.Sprintf("%d repeats held back by "+
|
||||
"SWWAF_ALERT_COOLDOWN that no later alert gives", repeats))
|
||||
}
|
||||
|
||||
if len(reasons) > 0 {
|
||||
q.queue(&Alert{
|
||||
Instance: q.params.Instance,
|
||||
Time: now,
|
||||
Event: EventSummary,
|
||||
Reason: fmt.Sprintf("%d alerts held back in the hour from %s, past the %d "+
|
||||
"an hour SWWAF_ALERT_MAX_PER_HOUR allows", heldBack,
|
||||
q.hour.Start.Format(time.RFC3339), q.params.MaxPerHour),
|
||||
Reason: strings.Join(reasons, "; "),
|
||||
Detail: map[string]any{
|
||||
"hour": q.hour.Start, "count": heldBack, "events": q.hour.HeldBack,
|
||||
},
|
||||
SuppressedRepeats: repeats,
|
||||
})
|
||||
}
|
||||
|
||||
q.hour = Hour{Start: start, HeldBack: map[string]int{}}
|
||||
|
||||
for key, cooldown := range q.cooldowns {
|
||||
if now.Sub(cooldown.Sent) >= q.params.Cooldown && cooldown.SuppressedRepeats == 0 {
|
||||
delete(q.cooldowns, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// queue adds alert to the alerts waiting for each destination.
|
||||
|
||||
@@ -381,7 +381,7 @@ func TestAlertsPastTheHourlyLimitAreRolledIntoOneSummary(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheNextSent(t *testing.T) {
|
||||
func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheSummary(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
@@ -401,8 +401,8 @@ func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheNextSent(t *testing.
|
||||
time.Sleep(cooldown)
|
||||
raise()
|
||||
|
||||
// The next hour's first alert gives the two repeats, and the summary
|
||||
// the alert past the limit.
|
||||
// The summary gives the alert past the limit and the two repeats,
|
||||
// and the next hour's first alert none.
|
||||
time.Sleep(time.Hour - cooldown)
|
||||
synctest.Wait()
|
||||
raise()
|
||||
@@ -413,11 +413,14 @@ func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheNextSent(t *testing.
|
||||
got := webhook.received()
|
||||
if len(got) == 3 {
|
||||
detail, _ := got[1].alert["detail"].(map[string]any)
|
||||
repeats := got[2].alert["suppressed_repeats"]
|
||||
summaryRepeats := got[1].alert["suppressed_repeats"]
|
||||
lastRepeats := got[2].alert["suppressed_repeats"]
|
||||
|
||||
if detail["count"] != float64(1) || repeats != float64(2) {
|
||||
t.Errorf("the summary counts %v alerts, and the last alert gives %v "+
|
||||
"repeats, want 1 and 2", detail["count"], repeats)
|
||||
if detail["count"] != float64(1) || summaryRepeats != float64(2) ||
|
||||
lastRepeats != float64(0) {
|
||||
t.Errorf("the summary counts %v alerts and %v repeats, and the last "+
|
||||
"alert gives %v repeats, want 1, 2 and 0", detail["count"],
|
||||
summaryRepeats, lastRepeats)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -425,6 +428,70 @@ func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheNextSent(t *testing.
|
||||
})
|
||||
}
|
||||
|
||||
func TestCooldownsThatHaveRunOutAreDroppedAndTheirRepeatsSummedUp(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for name, maxPerHour := range map[string]int{"limit off": 0, "limit set": 60} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
params := newParams()
|
||||
params.MaxPerHour = maxPerHour
|
||||
webhook, q := start(t, params)
|
||||
|
||||
// For four hours, a netblock of its own each minute is over an
|
||||
// anomaly threshold twice: an alert, and a repeat the cooldown
|
||||
// holds back.
|
||||
netblocks := 0
|
||||
|
||||
for range 4 {
|
||||
for range 60 {
|
||||
anomaly := alerts.Alert{
|
||||
Event: alerts.EventAnomaly, Netblock: netblock(netblocks),
|
||||
Detail: map[string]any{"scope": "net"},
|
||||
}
|
||||
q.Raise(anomaly)
|
||||
q.Raise(anomaly)
|
||||
|
||||
netblocks++
|
||||
|
||||
time.Sleep(time.Minute)
|
||||
}
|
||||
|
||||
// As the hour ends, only the cooldowns started less than the
|
||||
// cooldown before are kept, in memory and for alerts.json.
|
||||
synctest.Wait()
|
||||
|
||||
kept := len(q.Snapshot().Cooldowns)
|
||||
if kept > int(cooldown/time.Minute) {
|
||||
t.Errorf("after %d netblocks, %d cooldowns are kept, want at most %d",
|
||||
netblocks, kept, int(cooldown/time.Minute))
|
||||
}
|
||||
}
|
||||
|
||||
// An hour on, every cooldown has been dropped, and the summaries
|
||||
// have given every repeat.
|
||||
time.Sleep(time.Hour)
|
||||
synctest.Wait()
|
||||
|
||||
repeats := 0.0
|
||||
|
||||
for _, request := range webhook.received() {
|
||||
count, _ := request.alert["suppressed_repeats"].(float64)
|
||||
repeats += count
|
||||
}
|
||||
|
||||
kept := len(q.Snapshot().Cooldowns)
|
||||
if kept != 0 || repeats != float64(netblocks) {
|
||||
t.Errorf("%d cooldowns are kept and the webhook was given %v repeats, "+
|
||||
"want 0 and %d", kept, repeats, netblocks)
|
||||
}
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFailedRequestIsSentAgainWithBackoff(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -635,7 +702,8 @@ func TestStateLoadedIntoANewQueueCarriesOn(t *testing.T) {
|
||||
after.Load(roundTrip(t, before.Snapshot()))
|
||||
|
||||
// The new queue sends the alert waiting, holds back the repeat as
|
||||
// the cooldown still runs, and sends the summary of the hour.
|
||||
// the cooldown still runs, and sends the summary of the hour, which
|
||||
// gives both repeats, as the cooldown has run out.
|
||||
after.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
|
||||
synctest.Wait()
|
||||
wantEvents(t, webhook, alerts.EventBan)
|
||||
@@ -644,18 +712,12 @@ func TestStateLoadedIntoANewQueueCarriesOn(t *testing.T) {
|
||||
synctest.Wait()
|
||||
wantEvents(t, webhook, alerts.EventBan, alerts.EventSummary)
|
||||
|
||||
detail, _ := webhook.received()[1].alert["detail"].(map[string]any)
|
||||
if detail["count"] != float64(1) {
|
||||
t.Errorf("the summary counts %v alerts, want 1", detail["count"])
|
||||
}
|
||||
summary := webhook.received()[1].alert
|
||||
detail, _ := summary["detail"].(map[string]any)
|
||||
|
||||
// The cooldown has run out, and the next one gives both repeats.
|
||||
after.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
|
||||
synctest.Wait()
|
||||
|
||||
got := webhook.received()
|
||||
if repeats := got[len(got)-1].alert["suppressed_repeats"]; repeats != float64(2) {
|
||||
t.Errorf("the last alert gives %v repeats, want 2", repeats)
|
||||
if detail["count"] != float64(1) || summary["suppressed_repeats"] != float64(2) {
|
||||
t.Errorf("the summary counts %v alerts and %v repeats, want 1 and 2",
|
||||
detail["count"], summary["suppressed_repeats"])
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -701,8 +763,8 @@ func TestSlackAndNtfyAreSentTheSummaryAndTheRepeatsHeldBack(t *testing.T) {
|
||||
ban := alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1), Reason: "a ban"}
|
||||
|
||||
// The hour's one alert, a repeat of it the cooldown holds back, and
|
||||
// an alert past the limit; once the hour has ended, its summary, and
|
||||
// the next alert, which gives the repeat.
|
||||
// an alert past the limit; once the hour has ended, its summary,
|
||||
// which gives the repeat, and the next alert, which gives none.
|
||||
q.Raise(ban)
|
||||
q.Raise(ban)
|
||||
q.Raise(alerts.Alert{Event: alerts.EventFileError, Reason: "a file error"})
|
||||
@@ -720,14 +782,14 @@ func TestSlackAndNtfyAreSentTheSummaryAndTheRepeatsHeldBack(t *testing.T) {
|
||||
}
|
||||
|
||||
const summary = "1 alerts held back in the hour from 2000-01-01T00:00:00Z, " +
|
||||
"past the 1 an hour SWWAF_ALERT_MAX_PER_HOUR allows"
|
||||
"past the 1 an hour SWWAF_ALERT_MAX_PER_HOUR allows; 1 repeats held back " +
|
||||
"by SWWAF_ALERT_COOLDOWN that no later alert gives\nsuppressed repeats: 1"
|
||||
|
||||
wantSlackMessage(t, slack[1], "*"+instance+": summary*\n"+summary)
|
||||
wantNtfyMessage(t, ntfy[1], instance+": summary", "default bar_chart", summary)
|
||||
wantSlackMessage(t, slack[2],
|
||||
"*"+instance+": ban*\na ban\nnetblock: 203.0.113.1/32\nsuppressed repeats: 1")
|
||||
wantSlackMessage(t, slack[2], "*"+instance+": ban*\na ban\nnetblock: 203.0.113.1/32")
|
||||
wantNtfyMessage(t, ntfy[2], instance+": ban", "default no_entry",
|
||||
"a ban\nnetblock: 203.0.113.1/32\nsuppressed repeats: 1")
|
||||
"a ban\nnetblock: 203.0.113.1/32")
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,408 @@
|
||||
// Package anomaly counts requests and bytes over a minute and an hour, per
|
||||
// client, per surrounding netblock, per AS number, for the whole service
|
||||
// and per named netblock, and raises an anomaly alert for a count over its
|
||||
// threshold, as "Anomaly thresholds" under "Configuration surface" in
|
||||
// SPEC.md describes. It refuses and bans nothing. At most 20,000 counters
|
||||
// are kept, in memory, and written to alerts.json and read from it by the
|
||||
// state package.
|
||||
package anomaly
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/hashicorp/golang-lru/v2/simplelru"
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
)
|
||||
|
||||
// maxCounters is how many counters are kept. Past it, the counter counted
|
||||
// least recently is dropped, and starts afresh if it is counted again.
|
||||
const maxCounters = 20000
|
||||
|
||||
// The scopes, what a counter counts, as the settings, alerts.json and the
|
||||
// alerts name them.
|
||||
const (
|
||||
// ScopeClient is one client: an IPv4 address, or an IPv6 netblock of
|
||||
// SWWAF_IPV6_GROUP_PREFIX.
|
||||
ScopeClient = "client"
|
||||
// ScopeNet is the netblock around a client, SWWAF_ANOMALY_NET_V4_PREFIX
|
||||
// or SWWAF_ANOMALY_NET_V6_PREFIX long.
|
||||
ScopeNet = "net"
|
||||
// ScopeASN is an AS number.
|
||||
ScopeASN = "asn"
|
||||
// ScopeTotal is the whole service.
|
||||
ScopeTotal = "total"
|
||||
// ScopeWatch is a named netblock of SWWAF_WATCH_NETS.
|
||||
ScopeWatch = "watch"
|
||||
)
|
||||
|
||||
// Scopes returns every scope.
|
||||
func Scopes() []string {
|
||||
return []string{ScopeClient, ScopeNet, ScopeASN, ScopeTotal, ScopeWatch}
|
||||
}
|
||||
|
||||
// The windows a counter counts in, as the alerts name them.
|
||||
const (
|
||||
minute = "minute"
|
||||
hour = "hour"
|
||||
)
|
||||
|
||||
// Thresholds are the most requests and the most bytes a scope may have
|
||||
// counted in a minute and in an hour before an alert is raised. Zero is
|
||||
// off.
|
||||
type Thresholds struct {
|
||||
RequestsPerMinute int64
|
||||
RequestsPerHour int64
|
||||
BytesPerMinute int64
|
||||
BytesPerHour int64
|
||||
}
|
||||
|
||||
// NamedNetblock is a netblock SWWAF_WATCH_NETS names.
|
||||
type NamedNetblock struct {
|
||||
Name string
|
||||
Netblock netip.Prefix
|
||||
}
|
||||
|
||||
// Params are what New needs.
|
||||
type Params struct {
|
||||
// The thresholds of each scope: SWWAF_ANOMALY_CLIENT_*,
|
||||
// SWWAF_ANOMALY_NET_*, SWWAF_ANOMALY_ASN_*, SWWAF_ANOMALY_TOTAL_* and
|
||||
// SWWAF_WATCH_*.
|
||||
Client, Net, ASN, Total, Watch Thresholds
|
||||
// NetV4Prefix and NetV6Prefix are the lengths of the netblock around a
|
||||
// client (SWWAF_ANOMALY_NET_V4_PREFIX and SWWAF_ANOMALY_NET_V6_PREFIX).
|
||||
NetV4Prefix, NetV6Prefix int
|
||||
// NamedNetblocks are SWWAF_WATCH_NETS.
|
||||
NamedNetblocks []NamedNetblock
|
||||
// Alerts receive the anomaly alerts.
|
||||
Alerts *alerts.Queue
|
||||
}
|
||||
|
||||
// Counter is one scope's counts, as alerts.json holds them: the scope,
|
||||
// with the netblock, the AS number or the name that tells it from the
|
||||
// others in that scope, and its two buckets of requests and of bytes in
|
||||
// the minute and in the hour. A bucket whose threshold is off counts
|
||||
// nothing, and is left out.
|
||||
//
|
||||
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||
type Counter struct {
|
||||
Scope string `json:"scope"`
|
||||
Netblock netip.Prefix `json:"netblock,omitzero"`
|
||||
ASN string `json:"asn,omitempty"`
|
||||
Name string `json:"name,omitempty"`
|
||||
Minute ratelimit.Buckets `json:"minute,omitzero"`
|
||||
Hour ratelimit.Buckets `json:"hour,omitzero"`
|
||||
MinuteBytes ratelimit.Buckets `json:"minute_bytes,omitzero"`
|
||||
HourBytes ratelimit.Buckets `json:"hour_bytes,omitzero"`
|
||||
}
|
||||
|
||||
// Request is a request that has ended, as the counters count it.
|
||||
type Request struct {
|
||||
// Client is the client's address, and ClientGroup the client it is
|
||||
// counted as: its IPv4 address, or the IPv6 netblock of
|
||||
// SWWAF_IPV6_GROUP_PREFIX its address is in.
|
||||
Client netip.Addr
|
||||
ClientGroup netip.Prefix
|
||||
// ASN, ASName and Country are the client's as looked up, each "" when
|
||||
// unknown.
|
||||
ASN, ASName, Country string
|
||||
// Bytes are the request's bytes, as SWWAF_BYTES_COUNT counts them.
|
||||
Bytes int64
|
||||
}
|
||||
|
||||
// Counters counts each request in the scopes it is in. It is safe for
|
||||
// concurrent use.
|
||||
type Counters struct {
|
||||
params Params
|
||||
|
||||
mu sync.Mutex
|
||||
counters *simplelru.LRU[key, *Counter]
|
||||
}
|
||||
|
||||
// key is what tells a counter from the others: its scope, with its
|
||||
// netblock, AS number or name.
|
||||
type key struct {
|
||||
scope string
|
||||
netblock netip.Prefix
|
||||
asn string
|
||||
name string
|
||||
}
|
||||
|
||||
// New returns Counters for params, with nothing counted yet.
|
||||
func New(params Params) *Counters {
|
||||
counters, err := simplelru.NewLRU[key, *Counter](maxCounters, nil)
|
||||
if err != nil {
|
||||
panic(err) // NewLRU fails only for a size below one
|
||||
}
|
||||
|
||||
return &Counters{params: params, counters: counters}
|
||||
}
|
||||
|
||||
// Count counts r, a request that has ended, at now, in each scope it is
|
||||
// in whose thresholds are not all off: its client, the netblock around
|
||||
// it, its AS number once known, the whole service, and each named
|
||||
// netblock it is in. Only the counts whose threshold is set are counted.
|
||||
// For each scope whose count is over a threshold, it raises an anomaly
|
||||
// alert, for the first such count in the order requests and bytes in the
|
||||
// minute, then in the hour; the alert queue's cooldown holds back the
|
||||
// repeats. Nothing is refused or banned.
|
||||
func (c *Counters) Count(now time.Time, r Request) {
|
||||
var raised []alerts.Alert
|
||||
|
||||
c.mu.Lock()
|
||||
|
||||
for _, scope := range c.scopesOf(r) {
|
||||
counter, found := c.counters.Get(scope.key)
|
||||
if !found {
|
||||
counter = scope.key.counter()
|
||||
c.counters.Add(scope.key, counter)
|
||||
}
|
||||
|
||||
over, passed := counter.add(now, r.Bytes, scope.thresholds)
|
||||
if passed {
|
||||
raised = append(raised, alertFor(r, scope.key, over))
|
||||
}
|
||||
}
|
||||
|
||||
c.mu.Unlock()
|
||||
|
||||
for _, alert := range raised {
|
||||
c.params.Alerts.Raise(alert)
|
||||
}
|
||||
}
|
||||
|
||||
// Snapshot returns every counter, sorted by scope, then by netblock, AS
|
||||
// number and name, as alerts.json lists them.
|
||||
func (c *Counters) Snapshot() []Counter {
|
||||
c.mu.Lock()
|
||||
|
||||
counters := make([]Counter, 0, c.counters.Len())
|
||||
for _, counter := range c.counters.Values() {
|
||||
counters = append(counters, *counter)
|
||||
}
|
||||
|
||||
c.mu.Unlock()
|
||||
|
||||
slices.SortFunc(counters, func(a, b Counter) int {
|
||||
return cmp.Or(cmp.Compare(a.Scope, b.Scope), a.Netblock.Compare(b.Netblock),
|
||||
cmp.Compare(a.ASN, b.ASN), cmp.Compare(a.Name, b.Name))
|
||||
})
|
||||
|
||||
return counters
|
||||
}
|
||||
|
||||
// Load puts counters, read from alerts.json, in place of those held, in
|
||||
// the order they were last counted, as the starts of their buckets tell,
|
||||
// so that the one counted least recently is dropped first. Each netblock
|
||||
// is masked to its length, so that 203.0.113.9/24 is 203.0.113.0/24.
|
||||
// Buckets whose time has passed at now are emptied, and a counter left
|
||||
// with every bucket empty is dropped.
|
||||
func (c *Counters) Load(counters []Counter, now time.Time) {
|
||||
counters = slices.Clone(counters)
|
||||
slices.SortStableFunc(counters, func(a, b Counter) int {
|
||||
return a.lastStart().Compare(b.lastStart())
|
||||
})
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
c.counters.Purge()
|
||||
|
||||
for _, counter := range counters {
|
||||
counter.Netblock = counter.Netblock.Masked()
|
||||
empty := true
|
||||
|
||||
for _, count := range counter.counts() {
|
||||
if count.buckets.Passed(now, count.length) {
|
||||
*count.buckets = ratelimit.Buckets{}
|
||||
}
|
||||
|
||||
empty = empty && *count.buckets == ratelimit.Buckets{}
|
||||
}
|
||||
|
||||
if !empty {
|
||||
c.counters.Add(counter.key(), &counter)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// scope is a scope a request is counted in, and its thresholds.
|
||||
type scope struct {
|
||||
key key
|
||||
thresholds Thresholds
|
||||
}
|
||||
|
||||
// scopesOf returns the scopes r is in whose thresholds are not all off.
|
||||
func (c *Counters) scopesOf(r Request) []scope {
|
||||
p := c.params
|
||||
client := r.Client.Unmap()
|
||||
|
||||
all := []scope{
|
||||
{key{scope: ScopeClient, netblock: r.ClientGroup}, p.Client},
|
||||
{key{scope: ScopeNet, netblock: c.netAround(client)}, p.Net},
|
||||
{key{scope: ScopeTotal}, p.Total},
|
||||
}
|
||||
|
||||
if r.ASN != "" {
|
||||
all = append(all, scope{key{scope: ScopeASN, asn: r.ASN}, p.ASN})
|
||||
}
|
||||
|
||||
for _, named := range p.NamedNetblocks {
|
||||
if named.Netblock.Contains(client) {
|
||||
all = append(all, scope{
|
||||
key{scope: ScopeWatch, netblock: named.Netblock, name: named.Name}, p.Watch,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
return slices.DeleteFunc(all, func(s scope) bool {
|
||||
return s.thresholds == Thresholds{}
|
||||
})
|
||||
}
|
||||
|
||||
// netAround returns the netblock around client that ScopeNet counts it
|
||||
// in: NetV4Prefix or NetV6Prefix long.
|
||||
func (c *Counters) netAround(client netip.Addr) netip.Prefix {
|
||||
length := c.params.NetV6Prefix
|
||||
if client.Is4() {
|
||||
length = c.params.NetV4Prefix
|
||||
}
|
||||
|
||||
return netip.PrefixFrom(client, length).Masked()
|
||||
}
|
||||
|
||||
// overThreshold is a count over its threshold: what it counts, requests or
|
||||
// bytes, its window, the count and the threshold.
|
||||
type overThreshold struct {
|
||||
kind, window string
|
||||
count float64
|
||||
threshold int64
|
||||
}
|
||||
|
||||
// add counts a request of bytes at now in each of c's counts whose
|
||||
// threshold, in thresholds, is set, and returns the first count over its
|
||||
// threshold, and whether there is one.
|
||||
func (c *Counter) add(
|
||||
now time.Time, bytes int64, thresholds Thresholds,
|
||||
) (overThreshold, bool) {
|
||||
// In the order of counts.
|
||||
inOrder := [4]int64{
|
||||
thresholds.RequestsPerMinute, thresholds.BytesPerMinute,
|
||||
thresholds.RequestsPerHour, thresholds.BytesPerHour,
|
||||
}
|
||||
|
||||
var (
|
||||
first overThreshold
|
||||
passed bool
|
||||
)
|
||||
|
||||
for i, count := range c.counts() {
|
||||
threshold := inOrder[i]
|
||||
if threshold == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
n := int64(1)
|
||||
if count.kind == ratelimit.KindBytes {
|
||||
n = bytes
|
||||
}
|
||||
|
||||
counted := count.buckets.Add(now, count.length, n)
|
||||
if !passed && counted > float64(threshold) {
|
||||
first = overThreshold{count.kind, count.window, counted, threshold}
|
||||
passed = true
|
||||
}
|
||||
}
|
||||
|
||||
return first, passed
|
||||
}
|
||||
|
||||
// bucketCount is one of a counter's four counts: requests or bytes, in a
|
||||
// window of length, and the buckets they are counted in.
|
||||
type bucketCount struct {
|
||||
kind, window string
|
||||
length time.Duration
|
||||
buckets *ratelimit.Buckets
|
||||
}
|
||||
|
||||
// counts returns c's counts: requests and bytes in the minute, then in
|
||||
// the hour.
|
||||
func (c *Counter) counts() [4]bucketCount {
|
||||
return [4]bucketCount{
|
||||
{ratelimit.KindRequests, minute, time.Minute, &c.Minute},
|
||||
{ratelimit.KindBytes, minute, time.Minute, &c.MinuteBytes},
|
||||
{ratelimit.KindRequests, hour, time.Hour, &c.Hour},
|
||||
{ratelimit.KindBytes, hour, time.Hour, &c.HourBytes},
|
||||
}
|
||||
}
|
||||
|
||||
// lastStart returns the start of c's latest bucket, which tells, to the
|
||||
// minute or to the hour, when c was last counted.
|
||||
func (c *Counter) lastStart() time.Time {
|
||||
var latest time.Time
|
||||
|
||||
for _, count := range c.counts() {
|
||||
if count.buckets.Start.After(latest) {
|
||||
latest = count.buckets.Start
|
||||
}
|
||||
}
|
||||
|
||||
return latest
|
||||
}
|
||||
|
||||
// key returns what tells c from the other counters.
|
||||
func (c *Counter) key() key {
|
||||
return key{scope: c.Scope, netblock: c.Netblock, asn: c.ASN, name: c.Name}
|
||||
}
|
||||
|
||||
// counter returns a counter for k, with nothing counted yet.
|
||||
func (k key) counter() *Counter {
|
||||
return &Counter{Scope: k.scope, Netblock: k.netblock, ASN: k.asn, Name: k.name}
|
||||
}
|
||||
|
||||
// alertFor returns the anomaly alert for o, a count over its threshold in
|
||||
// the scope k, which r took over it. It gives r's client, with its AS
|
||||
// number, AS name and country, and the netblock counted, of a client, the
|
||||
// netblock around it or a named netblock. Its detail gives the scope, the
|
||||
// AS number or the name of a scope that has one, the window, what is
|
||||
// counted, the count and the threshold.
|
||||
func alertFor(r Request, k key, o overThreshold) alerts.Alert {
|
||||
detail := map[string]any{
|
||||
"scope": k.scope, "window": o.window, "kind": o.kind, "count": o.count,
|
||||
"threshold": o.threshold,
|
||||
}
|
||||
|
||||
var counted string
|
||||
|
||||
switch k.scope {
|
||||
case ScopeClient:
|
||||
counted = "the client " + k.netblock.String()
|
||||
case ScopeNet:
|
||||
counted = "the netblock " + k.netblock.String()
|
||||
case ScopeASN:
|
||||
counted = k.asn
|
||||
detail["asn"] = k.asn
|
||||
case ScopeTotal:
|
||||
counted = "the whole service"
|
||||
default: // watch
|
||||
counted = "the named netblock " + k.name + ", " + k.netblock.String()
|
||||
detail["name"] = k.name
|
||||
}
|
||||
|
||||
return alerts.Alert{
|
||||
Event: alerts.EventAnomaly,
|
||||
Client: r.Client,
|
||||
Netblock: k.netblock,
|
||||
ASN: r.ASN,
|
||||
ASName: r.ASName,
|
||||
Country: r.Country,
|
||||
Reason: fmt.Sprintf("%s per %s of %s over the threshold of %d", o.kind, o.window,
|
||||
counted, o.threshold),
|
||||
Detail: detail,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,238 @@
|
||||
package anomaly_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"reflect"
|
||||
"slices"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/anomaly"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
)
|
||||
|
||||
// maxCounters is how many counters are kept.
|
||||
const maxCounters = 20000
|
||||
|
||||
func TestEachScopeHasACooldownOfItsOwn(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
queue := newQueue()
|
||||
office := netip.MustParsePrefix("203.0.113.0/24")
|
||||
overAtTheSecond := anomaly.Thresholds{RequestsPerMinute: 1}
|
||||
counters := anomaly.New(anomaly.Params{
|
||||
Client: overAtTheSecond, Net: overAtTheSecond, ASN: overAtTheSecond,
|
||||
Total: overAtTheSecond, Watch: overAtTheSecond,
|
||||
// The netblock around a client is the client's own, and two names
|
||||
// name one netblock.
|
||||
NetV4Prefix: 32,
|
||||
NamedNetblocks: []anomaly.NamedNetblock{
|
||||
{Name: "office", Netblock: office}, {Name: "hq", Netblock: office},
|
||||
},
|
||||
Alerts: queue,
|
||||
})
|
||||
|
||||
// The first client's second request is over the threshold in the six
|
||||
// scopes it is in. The other client's two are both over it in the whole
|
||||
// service and in each named netblock, three repeats each, and its
|
||||
// second is over it in the scopes of its own, its client, its netblock
|
||||
// and its AS number, which are no repeats.
|
||||
for _, r := range []anomaly.Request{
|
||||
{Client: netip.MustParseAddr("203.0.113.9"), ASN: "AS64496"},
|
||||
{Client: netip.MustParseAddr("203.0.113.10"), ASN: "AS64511"},
|
||||
} {
|
||||
r.ClientGroup = netip.PrefixFrom(r.Client, 32)
|
||||
|
||||
for range 2 {
|
||||
counters.Count(midnight(), r)
|
||||
}
|
||||
}
|
||||
|
||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||
if len(waiting) != 9 || queue.Suppressed() != 6 {
|
||||
t.Fatalf("%d alerts wait and %d are held back, want 9 and 6: %+v",
|
||||
len(waiting), queue.Suppressed(), waiting)
|
||||
}
|
||||
|
||||
// alerts.json keeps each scope's cooldown: each alert raised again
|
||||
// after a restart is a repeat.
|
||||
data, err := json.Marshal(queue.Snapshot())
|
||||
if err != nil {
|
||||
t.Fatalf("encode: %v", err)
|
||||
}
|
||||
|
||||
var read alerts.State
|
||||
|
||||
err = json.Unmarshal(data, &read)
|
||||
if err != nil {
|
||||
t.Fatalf("decode: %v", err)
|
||||
}
|
||||
|
||||
after := newQueue()
|
||||
after.Load(read)
|
||||
|
||||
for _, alert := range read.Waiting[alerts.DestinationWebhook] {
|
||||
after.Raise(alert)
|
||||
}
|
||||
|
||||
if after.Suppressed() != 9 {
|
||||
t.Errorf("after loading, %d alerts are held back, want 9", after.Suppressed())
|
||||
}
|
||||
}
|
||||
|
||||
func TestKeepsAtMost20000CountersDroppingTheLeastRecentlyCounted(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
counters := newCounters(anomaly.Params{
|
||||
Client: anomaly.Thresholds{RequestsPerMinute: 1000},
|
||||
})
|
||||
|
||||
for i := range maxCounters {
|
||||
counters.Count(midnight(), request(i))
|
||||
}
|
||||
|
||||
// Counted again, the first client is the most recently counted, and
|
||||
// the second is dropped for a new one.
|
||||
counters.Count(midnight(), request(0))
|
||||
counters.Count(midnight(), request(maxCounters))
|
||||
|
||||
got := counters.Snapshot()
|
||||
if len(got) != maxCounters || !holds(got, 0) || holds(got, 1) ||
|
||||
!holds(got, maxCounters) {
|
||||
t.Errorf("%d counters, holding the first client %v, the second %v and the "+
|
||||
"new one %v, want %d, the first and the new one", len(got), holds(got, 0),
|
||||
holds(got, 1), holds(got, maxCounters), maxCounters)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadEmptiesBucketsWhoseTimeHasPassedAndDropsEmptyCounters(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
counters := newCounters(anomaly.Params{
|
||||
Net: anomaly.Thresholds{RequestsPerMinute: 1000, RequestsPerHour: 1000},
|
||||
Total: anomaly.Thresholds{RequestsPerMinute: 1000},
|
||||
NetV4Prefix: 24,
|
||||
})
|
||||
halfAnHourOn := midnight().Add(30 * time.Minute)
|
||||
|
||||
// Half an hour on, the hour's buckets count still, and the minute's
|
||||
// do not.
|
||||
counters.Load([]anomaly.Counter{
|
||||
{
|
||||
Scope: anomaly.ScopeNet,
|
||||
Netblock: netip.MustParsePrefix("203.0.113.9/24"),
|
||||
Minute: ratelimit.Buckets{Start: midnight(), Current: 5},
|
||||
Hour: ratelimit.Buckets{Start: midnight(), Current: 7},
|
||||
},
|
||||
{
|
||||
Scope: anomaly.ScopeTotal,
|
||||
Minute: ratelimit.Buckets{Start: midnight(), Current: 1},
|
||||
},
|
||||
}, halfAnHourOn)
|
||||
|
||||
// The whole service's counter, left empty, is dropped, and the
|
||||
// netblock read is masked to its length.
|
||||
netblock := anomaly.Counter{
|
||||
Scope: anomaly.ScopeNet,
|
||||
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
|
||||
Hour: ratelimit.Buckets{Start: midnight(), Current: 7},
|
||||
}
|
||||
if got, want := counters.Snapshot(), []anomaly.Counter{netblock}; !reflect.DeepEqual(
|
||||
got, want) {
|
||||
t.Errorf("counters read\n%+v\nwant\n%+v", got, want)
|
||||
}
|
||||
|
||||
// A request from the netblock is counted with the requests read.
|
||||
counters.Count(halfAnHourOn, anomaly.Request{
|
||||
Client: netip.MustParseAddr("203.0.113.9"),
|
||||
ClientGroup: netip.MustParsePrefix("203.0.113.9/32"),
|
||||
})
|
||||
|
||||
netblock.Minute = ratelimit.Buckets{Start: halfAnHourOn, Current: 1}
|
||||
netblock.Hour.Current = 8
|
||||
want := []anomaly.Counter{netblock, {
|
||||
Scope: anomaly.ScopeTotal,
|
||||
Minute: ratelimit.Buckets{Start: halfAnHourOn, Current: 1},
|
||||
}}
|
||||
|
||||
if got := counters.Snapshot(); !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("counters after a request\n%+v\nwant\n%+v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadDropsTheLeastRecentlyCountedFirst(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
counters := newCounters(anomaly.Params{
|
||||
Client: anomaly.Thresholds{RequestsPerMinute: 1000},
|
||||
})
|
||||
now := midnight().Add(time.Minute)
|
||||
|
||||
// The second half of the file was counted in the minute before the
|
||||
// first half.
|
||||
read := make([]anomaly.Counter, 0, maxCounters)
|
||||
for i := range maxCounters {
|
||||
start := now
|
||||
if i >= maxCounters/2 {
|
||||
start = midnight()
|
||||
}
|
||||
|
||||
read = append(read, anomaly.Counter{
|
||||
Scope: anomaly.ScopeClient, Netblock: request(i).ClientGroup,
|
||||
Minute: ratelimit.Buckets{Start: start, Current: 1},
|
||||
})
|
||||
}
|
||||
|
||||
counters.Load(read, now)
|
||||
counters.Count(now, request(maxCounters))
|
||||
|
||||
got := counters.Snapshot()
|
||||
if !holds(got, 0) || holds(got, maxCounters/2) {
|
||||
t.Errorf("holding the first client of the file %v, and the first counted in "+
|
||||
"the minute before %v, want only the first", holds(got, 0),
|
||||
holds(got, maxCounters/2))
|
||||
}
|
||||
}
|
||||
|
||||
// midnight is the time of the tests' requests.
|
||||
func midnight() time.Time {
|
||||
return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
|
||||
}
|
||||
|
||||
// newCounters returns Counters for params, whose alerts go nowhere.
|
||||
func newCounters(params anomaly.Params) *anomaly.Counters {
|
||||
params.Alerts = alerts.New(alerts.Params{})
|
||||
|
||||
return anomaly.New(params)
|
||||
}
|
||||
|
||||
// newQueue returns a queue of alerts to a webhook, with the default
|
||||
// cooldown, which keeps them waiting, since it is never run.
|
||||
func newQueue() *alerts.Queue {
|
||||
return alerts.New(alerts.Params{
|
||||
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
||||
Events: alerts.Events(),
|
||||
Cooldown: 15 * time.Minute,
|
||||
MaxPerHour: 60,
|
||||
Now: midnight,
|
||||
})
|
||||
}
|
||||
|
||||
// request returns a request from client number i, an address in
|
||||
// 10.0.0.0/8.
|
||||
func request(i int) anomaly.Request {
|
||||
client := netip.MustParseAddr(fmt.Sprintf("10.%d.%d.%d", i>>16, i>>8&255, i&255))
|
||||
|
||||
return anomaly.Request{Client: client, ClientGroup: netip.PrefixFrom(client, 32)}
|
||||
}
|
||||
|
||||
// holds reports whether counters hold the counter of client number i.
|
||||
func holds(counters []anomaly.Counter, i int) bool {
|
||||
return slices.ContainsFunc(counters, func(counter anomaly.Counter) bool {
|
||||
return counter.Netblock == request(i).ClientGroup
|
||||
})
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package bans_test
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -66,12 +67,15 @@ func TestReasonOfTheBansSmallwebwafMakes(t *testing.T) {
|
||||
ledger := bans.New(defaultRules())
|
||||
|
||||
limit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
|
||||
bans.Notes{Limit: 1000, Window: "minute"})
|
||||
bans.Notes{Kind: "requests", Limit: 1000, Window: "minute"})
|
||||
byteLimit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.3/32"),
|
||||
midnight(), bans.Notes{Kind: "bytes", Limit: 10 << 30, Window: "hour"})
|
||||
attack, _ := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
|
||||
bans.Notes{RuleID: "git-dir", Target: "path"})
|
||||
|
||||
for _, tc := range []struct{ got, want string }{
|
||||
{limit.Reason, "requests per minute over the limit of 1000"},
|
||||
{byteLimit.Reason, "bytes per hour over the limit of 10737418240"},
|
||||
{attack.Reason, "matched the rule git-dir"},
|
||||
} {
|
||||
if tc.got != tc.want {
|
||||
@@ -114,7 +118,7 @@ func TestLiftedBanForALimitRefusesNothingAndMakesNoBanLonger(t *testing.T) {
|
||||
}
|
||||
|
||||
held := ledger.Bans(netblock)
|
||||
if len(held) != 2 || held[0] != lifted {
|
||||
if len(held) != 2 || !reflect.DeepEqual(held[0], lifted) {
|
||||
t.Errorf("the ledger holds %+v, want the lifted ban and the new one", held)
|
||||
}
|
||||
}
|
||||
@@ -209,7 +213,7 @@ func TestAdminsBanIsMadeWhileAnotherLasts(t *testing.T) {
|
||||
|
||||
got := ledger.BanForAdmin(netip.MustParsePrefix("203.0.113.9/24"), now, time.Time{},
|
||||
"probes for logins")
|
||||
if got != want {
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("the admin's ban is\n%+v\nwant\n%+v", got, want)
|
||||
}
|
||||
|
||||
@@ -221,7 +225,7 @@ func TestAdminsBanIsMadeWhileAnotherLasts(t *testing.T) {
|
||||
|
||||
// It refuses once the ban for the limit has ended.
|
||||
ban, banned, _ := ledger.Find(netblock.Addr(), midnight().Add(2*time.Hour))
|
||||
if !banned || ban != want {
|
||||
if !banned || !reflect.DeepEqual(ban, want) {
|
||||
t.Errorf("after the limit's ban the netblock is under %+v (%t), want %+v",
|
||||
ban, banned, want)
|
||||
}
|
||||
|
||||
+41
-17
@@ -1,8 +1,8 @@
|
||||
// Package bans is the ban ledger: the bans smallwebwaf makes on the
|
||||
// netblocks of clients that break a rate limit or show a clear sign of
|
||||
// attack, and those an admin makes, with their notes, as the "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.
|
||||
// netblocks of clients that break a rate limit or a byte limit or show a
|
||||
// clear sign of attack, and those an admin makes, with their notes, as
|
||||
// the "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 (
|
||||
@@ -97,20 +97,33 @@ type Notes struct {
|
||||
ASN string `json:"asn"`
|
||||
ASName string `json:"as_name"`
|
||||
Country string `json:"country"`
|
||||
// Limit, Window and Count are, for a ban for a broken limit, 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.
|
||||
// Kind, Limit, Window and Count are, for a ban for a broken limit,
|
||||
// what the limit was on, "requests" for a rate limit or "bytes" for a
|
||||
// byte limit, the limit that was broken, its window, "minute", "hour"
|
||||
// or "day", and the count reached: the client's requests, or bytes, in
|
||||
// the window, those of the request that broke the limit included.
|
||||
// These are what counted toward the ban, and the window is the time
|
||||
// over which they came.
|
||||
Kind string `json:"kind,omitempty"`
|
||||
Limit int64 `json:"limit,omitempty"`
|
||||
Window string `json:"window,omitempty"`
|
||||
Count float64 `json:"count,omitempty"`
|
||||
// LimitPercent and LimitPercentSetting are, for a ban for a limit a
|
||||
// biased threshold lowered, the client's percentage of that kind of
|
||||
// limit, of which Limit is the result, and the setting that gave it.
|
||||
// Both are left out for a limit that was not lowered.
|
||||
LimitPercent *int64 `json:"limit_percent,omitempty"`
|
||||
LimitPercentSetting string `json:"limit_percent_setting,omitempty"`
|
||||
// RuleID and Target are, for a ban for a clear sign of attack, the id
|
||||
// of the rule file rule that matched, and its target.
|
||||
RuleID string `json:"rule_id,omitempty"`
|
||||
Target string `json:"target,omitempty"`
|
||||
// Request is the request that broke the limit, or that was the clear
|
||||
// sign of attack.
|
||||
// Reputation is the reputation sources that listed the client when
|
||||
// the request that caused the ban was made, in the order the request
|
||||
// log's reputation names them. It is left out when none did.
|
||||
Reputation []ReputationHit `json:"reputation,omitempty"`
|
||||
// Request is the request that broke the limit, or whose bytes broke
|
||||
// it, or that was the clear sign of attack.
|
||||
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
|
||||
@@ -122,6 +135,15 @@ type Notes struct {
|
||||
EarlierBans EarlierBans `json:"earlier_bans"`
|
||||
}
|
||||
|
||||
// ReputationHit is a reputation source that listed a client, as a
|
||||
// reputation_hit alert's detail gives it: Source is the blocklist's URL,
|
||||
// the DNSBL zone with its key masked, or "abuseipdb", and Score, for
|
||||
// AbuseIPDB alone, its score of the client.
|
||||
type ReputationHit struct {
|
||||
Source string `json:"source"`
|
||||
Score *int64 `json:"score,omitempty"`
|
||||
}
|
||||
|
||||
// EarlierBans counts a netblock's bans before a ban, by cause.
|
||||
type EarlierBans struct {
|
||||
Limit int `json:"limit"`
|
||||
@@ -260,7 +282,8 @@ func activeBan(bans []Ban, now time.Time) *Ban {
|
||||
// active, as when two of its requests break a limit at once, that ban is
|
||||
// returned with false, and no other is made. The ledger fills in the
|
||||
// notes' Refused and EarlierBans itself, and gives the ban the reason
|
||||
// "requests per <Window> over the limit of <Limit>", from the notes.
|
||||
// "<Kind> per <Window> over the limit of <Limit>", from the notes, such
|
||||
// as "requests per minute over the limit of 1000".
|
||||
func (l *Ledger) BanForLimit(
|
||||
netblock netip.Prefix, now time.Time, notes Notes,
|
||||
) (Ban, bool) {
|
||||
@@ -317,7 +340,8 @@ func (l *Ledger) WouldBePermanent(
|
||||
|
||||
// limitReason is the reason of a ban for a broken limit, with notes.
|
||||
func limitReason(notes Notes) string {
|
||||
return fmt.Sprintf("requests per %s over the limit of %d", notes.Window, notes.Limit)
|
||||
return fmt.Sprintf("%s per %s over the limit of %d",
|
||||
notes.Kind, notes.Window, notes.Limit)
|
||||
}
|
||||
|
||||
// attackReason is the reason of a ban for a clear sign of attack, with
|
||||
@@ -416,10 +440,10 @@ func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
|
||||
}
|
||||
|
||||
// AddLookup gives the notes of netblock's bans that have no AS number, AS
|
||||
// name or country yet those of a client in it, as GeoJS answered about
|
||||
// it. It is not a request from netblock, and leaves when it was last seen
|
||||
// unchanged. It does not have bans.json written at once: the notes are
|
||||
// written with its next write, as the counts in them are.
|
||||
// name or country yet those of a client in it, as the lookup answered
|
||||
// about it. It is not a request from netblock, and leaves when it was last
|
||||
// seen unchanged. It does not have bans.json written at once: the notes
|
||||
// are written with its next write, as the counts in them are.
|
||||
func (l *Ledger) AddLookup(netblock netip.Prefix, asn, asName, country string) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
@@ -2,6 +2,7 @@ package bans_test
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -130,13 +131,13 @@ func TestBrokenLimitDuringABanMakesNoOther(t *testing.T) {
|
||||
|
||||
again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
|
||||
|
||||
if made || again != first || len(ledger.Bans(netblock)) != 1 {
|
||||
if made || !reflect.DeepEqual(again, first) || len(ledger.Bans(netblock)) != 1 {
|
||||
t.Errorf("a limit broken during a ban gave %+v, made %t, and %d bans, "+
|
||||
"want %+v, not made, and 1", again, made, len(ledger.Bans(netblock)), first)
|
||||
}
|
||||
|
||||
again, made = ledger.BanForAttack(netblock, midnight().Add(time.Minute), bans.Notes{})
|
||||
if made || again != first {
|
||||
if made || !reflect.DeepEqual(again, first) {
|
||||
t.Errorf("an attack during a ban gave %+v, made %t, want %+v, not made",
|
||||
again, made, first)
|
||||
}
|
||||
@@ -182,7 +183,7 @@ func TestFindCountsNothing(t *testing.T) {
|
||||
ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
|
||||
|
||||
got, banned, _ := ledger.Find(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
|
||||
if !banned || got != ban {
|
||||
if !banned || !reflect.DeepEqual(got, ban) {
|
||||
t.Errorf("find during the ban gives %+v and %t, want %+v", got, banned, ban)
|
||||
}
|
||||
|
||||
@@ -191,7 +192,7 @@ func TestFindCountsNothing(t *testing.T) {
|
||||
t.Error("the ban did not end")
|
||||
}
|
||||
|
||||
if notes := ledger.Bans(netblock)[0].Notes; notes != ban.Notes {
|
||||
if notes := ledger.Bans(netblock)[0].Notes; !reflect.DeepEqual(notes, ban.Notes) {
|
||||
t.Errorf("the notes are %+v, want them unchanged, %+v", notes, ban.Notes)
|
||||
}
|
||||
}
|
||||
@@ -247,7 +248,7 @@ func TestFullLedgerDropsTheEarlierBanOfTheNetblockBannedAgain(t *testing.T) {
|
||||
second, _ := ledger.BanForLimit(netblock, first.Expires, bans.Notes{})
|
||||
|
||||
held := ledger.Bans(netblock)
|
||||
if len(held) != 1 || held[0] != second ||
|
||||
if len(held) != 1 || !reflect.DeepEqual(held[0], second) ||
|
||||
held[0].Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
|
||||
t.Errorf("the ledger holds %+v, want only the second ban, "+
|
||||
"with 1 earlier ban for a limit", held)
|
||||
@@ -342,14 +343,14 @@ func TestWouldBanGivesTheBanWithoutMakingIt(t *testing.T) {
|
||||
|
||||
// While the first ban lasts, none would be made.
|
||||
during, would := ledger.WouldBanForAttack(netblock, midnight(), bans.Notes{})
|
||||
if would || during != first {
|
||||
if would || !reflect.DeepEqual(during, first) {
|
||||
t.Errorf("during the first ban, would ban %t with %+v, want false with %+v",
|
||||
would, during, first)
|
||||
}
|
||||
|
||||
// As it ends, a clear sign of attack would ban for seven days, and a
|
||||
// limit broken again for three hours, but neither is made.
|
||||
limitNotes := bans.Notes{Limit: 1, Window: "minute"}
|
||||
limitNotes := bans.Notes{Kind: "requests", Limit: 1, Window: "minute"}
|
||||
attack, wouldAttack := ledger.WouldBanForAttack(netblock, first.Expires,
|
||||
bans.Notes{RuleID: "git-dir"})
|
||||
limit, wouldLimit := ledger.WouldBanForLimit(netblock, first.Expires, limitNotes)
|
||||
@@ -371,7 +372,7 @@ func TestWouldBanGivesTheBanWithoutMakingIt(t *testing.T) {
|
||||
|
||||
// The ban made is the one that would have been.
|
||||
made, _ := ledger.BanForLimit(netblock, first.Expires, limitNotes)
|
||||
if made != limit {
|
||||
if !reflect.DeepEqual(made, limit) {
|
||||
t.Errorf("the ban made is %+v, want %+v", made, limit)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,6 +2,7 @@ package bans_test
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -224,7 +225,7 @@ func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
|
||||
ledger.Load([]bans.Ban{later, earlier})
|
||||
|
||||
held := ledger.Snapshot()
|
||||
if len(held) != 1 || held[0] != later {
|
||||
if len(held) != 1 || !reflect.DeepEqual(held[0], later) {
|
||||
t.Errorf("the ledger holds %+v, want only the ban that began later", held)
|
||||
}
|
||||
}
|
||||
@@ -267,7 +268,7 @@ func TestLoadReplacesTheBansHeld(t *testing.T) {
|
||||
bans.Notes{})
|
||||
|
||||
want := []bans.Ban{first, second, kept}
|
||||
if got := ledger.Snapshot(); !slices.Equal(got, want) {
|
||||
if got := ledger.Snapshot(); !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("the ledger holds %+v, want %+v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
+669
-31
@@ -24,6 +24,7 @@ import (
|
||||
"unicode/utf8"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/anomaly"
|
||||
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||
)
|
||||
|
||||
@@ -47,6 +48,12 @@ type Config struct {
|
||||
// TrustedProxies are the netblocks whose X-Forwarded-For is
|
||||
// believed (SWWAF_TRUSTED_PROXIES).
|
||||
TrustedProxies []netip.Prefix
|
||||
// IPv6GroupPrefix is the length of the IPv6 netblock that is one client
|
||||
// (SWWAF_IPV6_GROUP_PREFIX), from 32 to 128.
|
||||
IPv6GroupPrefix int
|
||||
// MaxTrackedClients is the most clients the table of clients holds, in
|
||||
// memory and in clients.json (SWWAF_MAX_TRACKED_CLIENTS).
|
||||
MaxTrackedClients int
|
||||
// ClientRequestTimeout bounds reading the whole request from the
|
||||
// client (SWWAF_CLIENT_REQUEST_TIMEOUT).
|
||||
ClientRequestTimeout time.Duration
|
||||
@@ -91,13 +98,26 @@ type Config struct {
|
||||
// limits neither count nor refuse (SWWAF_RATE_LIMIT_EXEMPT_PATHS).
|
||||
// Each starts with /.
|
||||
RateLimitExemptPaths []string
|
||||
// BytesLimitPerMinute, BytesLimitPerHour and BytesLimitPerDay are the
|
||||
// most bytes a client's requests may carry in a minute, an hour and a
|
||||
// day (SWWAF_BYTES_LIMIT_PER_MINUTE, SWWAF_BYTES_LIMIT_PER_HOUR and
|
||||
// SWWAF_BYTES_LIMIT_PER_DAY). BytesCount is which body bytes count
|
||||
// toward them (SWWAF_BYTES_COUNT): response, request or both.
|
||||
BytesLimitPerMinute int64
|
||||
BytesLimitPerHour int64
|
||||
BytesLimitPerDay int64
|
||||
BytesCount string
|
||||
// LookupSource is where each client's AS number and country are
|
||||
// looked up (SWWAF_LOOKUP_SOURCE): geojs, or off for nowhere. A request
|
||||
// waits up to LookupTimeout for its client's first answer while a
|
||||
// setting needs it (SWWAF_LOOKUP_TIMEOUT), which cannot be off.
|
||||
// AddLookupHeaders is true when the app is passed the client's AS
|
||||
// number and country in headers (SWWAF_ADD_LOOKUP_HEADERS).
|
||||
// looked up (SWWAF_LOOKUP_SOURCE): geojs, file, or off for nowhere.
|
||||
// LookupDBPath is the lookup database, the IPinfo Lite file looked up
|
||||
// in while LookupSource is file (SWWAF_LOOKUP_DB_PATH), and "" for any
|
||||
// other source. A request waits up to LookupTimeout for its client's
|
||||
// first answer from GeoJS while a setting needs it
|
||||
// (SWWAF_LOOKUP_TIMEOUT), which cannot be off. AddLookupHeaders is true
|
||||
// when the app is passed the client's AS number and country in headers
|
||||
// (SWWAF_ADD_LOOKUP_HEADERS).
|
||||
LookupSource string
|
||||
LookupDBPath string
|
||||
LookupTimeout time.Duration
|
||||
AddLookupHeaders bool
|
||||
// DeniedCountries are the countries whose clients are refused
|
||||
@@ -107,6 +127,60 @@ type Config struct {
|
||||
// capitals, as GeoJS gives them.
|
||||
DeniedCountries []string
|
||||
ExclusivelyAllowedCountries []string
|
||||
// The biased thresholds. ASNLimitPercent and CountryLimitPercent give
|
||||
// the clients of the AS numbers and the countries they list that
|
||||
// percentage of every rate limit and byte limit
|
||||
// (SWWAF_ASN_LIMIT_PERCENT and SWWAF_COUNTRY_LIMIT_PERCENT).
|
||||
// ASNBytesPercent and CountryBytesPercent give those they list a
|
||||
// percentage of the byte limits in place of that one
|
||||
// (SWWAF_ASN_BYTES_PERCENT and SWWAF_COUNTRY_BYTES_PERCENT). Each holds
|
||||
// percentages from 0 to 100, by AS number, written as AS64496, or by
|
||||
// country, a two-letter code in capitals, as the lookup gives them.
|
||||
// UnknownLimitPercent is the percentage of every limit a client without
|
||||
// a country gets (SWWAF_UNKNOWN_LIMIT_PERCENT). ASNLimitPercentURL is
|
||||
// where a file of AS:percent lines is fetched from, whose percentages
|
||||
// count as those of ASNLimitPercent do (SWWAF_ASN_LIMIT_PERCENT_URL), ""
|
||||
// while it is unset.
|
||||
ASNLimitPercent map[string]int64
|
||||
CountryLimitPercent map[string]int64
|
||||
ASNBytesPercent map[string]int64
|
||||
CountryBytesPercent map[string]int64
|
||||
UnknownLimitPercent int64
|
||||
ASNLimitPercentURL string
|
||||
// BlocklistURLs are where the blocklists are fetched from
|
||||
// (SWWAF_BLOCKLIST_URLS). Each list, and ASNLimitPercentURL's, is
|
||||
// fetched again BlocklistRefresh after it was last fetched or tried
|
||||
// (SWWAF_BLOCKLIST_REFRESH), which is never less than an hour.
|
||||
// BlocklistAction is what is done with a client a blocklist lists
|
||||
// (SWWAF_BLOCKLIST_ACTION): deny, limit or log; for limit,
|
||||
// BlocklistLimitPercent is the percentage of every limit it gets.
|
||||
BlocklistURLs []string
|
||||
BlocklistRefresh time.Duration
|
||||
BlocklistAction string
|
||||
BlocklistLimitPercent int64
|
||||
// DNSBLZones are the DNSBL zones clients are asked about
|
||||
// (SWWAF_DNSBL_ZONES), through DNSBLResolver (SWWAF_DNSBL_RESOLVER), or
|
||||
// the host's resolver while that is the zero AddrPort.
|
||||
// AbuseIPDBKey is the key of the AbuseIPDB account clients are checked
|
||||
// with (SWWAF_ABUSEIPDB_KEY), "" while it is unset and none is. A score
|
||||
// of AbuseIPDBMinScore or more is a hit (SWWAF_ABUSEIPDB_MIN_SCORE), and
|
||||
// at most AbuseIPDBDailyBudget checks are made a day
|
||||
// (SWWAF_ABUSEIPDB_DAILY_BUDGET).
|
||||
// ReputationAction is what is done with a client a zone's verdict lists,
|
||||
// or whose score is a hit (SWWAF_REPUTATION_ACTION): deny, limit or log;
|
||||
// for limit, ReputationLimitPercent is the percentage of every limit it
|
||||
// gets. A verdict or a score is used for ReputationCacheTTL after it was
|
||||
// fetched (SWWAF_REPUTATION_CACHE_TTL), and a query or a check may take
|
||||
// ReputationTimeout (SWWAF_REPUTATION_TIMEOUT). Neither can be off.
|
||||
DNSBLZones []string
|
||||
DNSBLResolver netip.AddrPort
|
||||
AbuseIPDBKey string
|
||||
AbuseIPDBMinScore int64
|
||||
AbuseIPDBDailyBudget int
|
||||
ReputationAction string
|
||||
ReputationLimitPercent int64
|
||||
ReputationCacheTTL time.Duration
|
||||
ReputationTimeout time.Duration
|
||||
// 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
|
||||
@@ -141,6 +215,9 @@ type Config struct {
|
||||
// LogRequestHeaders are the request headers whose values the request
|
||||
// log gives, in lower case (SWWAF_LOG_REQUEST_HEADERS).
|
||||
LogRequestHeaders []string
|
||||
// LogLevel is the least severe of the process's own messages that are
|
||||
// written (SWWAF_LOG_LEVEL). It holds back no request log line.
|
||||
LogLevel slog.Level
|
||||
// AdminToken is the bearer token an admin sends for the ban endpoints
|
||||
// and /_smallwebwaf/clients/<ip> (SWWAF_ADMIN_TOKEN), "" while it is
|
||||
// unset and they are off.
|
||||
@@ -191,6 +268,23 @@ type Config struct {
|
||||
AlertEvents []string
|
||||
AlertCooldown time.Duration
|
||||
AlertMaxPerHour int
|
||||
// The anomaly thresholds, which only raise alerts: the most requests
|
||||
// and bytes a minute and an hour per client (SWWAF_ANOMALY_CLIENT_*),
|
||||
// per netblock around a client (SWWAF_ANOMALY_NET_*), per AS number
|
||||
// (SWWAF_ANOMALY_ASN_*), for the whole service (SWWAF_ANOMALY_TOTAL_*)
|
||||
// and per named netblock (SWWAF_WATCH_*), each 0 while it is off.
|
||||
// AnomalyNetV4Prefix and AnomalyNetV6Prefix are the lengths of the
|
||||
// netblock around a client (SWWAF_ANOMALY_NET_V4_PREFIX and
|
||||
// SWWAF_ANOMALY_NET_V6_PREFIX), and WatchNets the named netblocks
|
||||
// (SWWAF_WATCH_NETS).
|
||||
AnomalyClient anomaly.Thresholds
|
||||
AnomalyNet anomaly.Thresholds
|
||||
AnomalyASN anomaly.Thresholds
|
||||
AnomalyTotal anomaly.Thresholds
|
||||
AnomalyWatch anomaly.Thresholds
|
||||
AnomalyNetV4Prefix int
|
||||
AnomalyNetV6Prefix int
|
||||
WatchNets []anomaly.NamedNetblock
|
||||
|
||||
// settings are the values read, as given or by default, and the
|
||||
// files they were read from, for the log line at start.
|
||||
@@ -201,12 +295,21 @@ type Config struct {
|
||||
// off.
|
||||
const off = "off"
|
||||
|
||||
// fileSource is the SWWAF_LOOKUP_SOURCE that looks clients up in the
|
||||
// lookup database, the file SWWAF_LOOKUP_DB_PATH names.
|
||||
const fileSource = "file"
|
||||
|
||||
const (
|
||||
day = 24 * time.Hour
|
||||
kibibyte = 1 << 10
|
||||
mebibyte = 1 << 20
|
||||
gibibyte = 1 << 30
|
||||
ipv4Bits = 32
|
||||
ipv6Bits = 128
|
||||
// minIPv6GroupPrefix is the shortest SWWAF_IPV6_GROUP_PREFIX, the
|
||||
// netblock a provider is usually given: a shorter one would make one
|
||||
// client of the customers of several providers.
|
||||
minIPv6GroupPrefix = 32
|
||||
// minTokenLength is the fewest characters a token may have.
|
||||
minTokenLength = 32
|
||||
// masked is what the log shows for a token that is set, and in place of
|
||||
@@ -242,8 +345,10 @@ var (
|
||||
"is taken out of every request by Go's HTTP server, so it can never " +
|
||||
"be logged")
|
||||
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
|
||||
errNotLookupSource = errors.New("is not geojs or off")
|
||||
errNotLookupSource = errors.New("is not geojs, file or off")
|
||||
errNeedsLookups = errors.New("it needs each client looked up")
|
||||
errNeedsDBPath = errors.New("it names the file to look clients up in")
|
||||
errDBPathUnused = errors.New("only file reads it")
|
||||
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")
|
||||
@@ -252,10 +357,18 @@ 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")
|
||||
errNotV6Prefix = errors.New(
|
||||
"is not the length of an IPv6 netblock, from 0 to 128, such as 48")
|
||||
errNotIPv6GroupPrefix = errors.New(
|
||||
"is not the length of an IPv6 netblock, from 32 to 128, such as 64")
|
||||
errNotLogLevel = errors.New("is not debug, info, warn or error")
|
||||
errNotNamedNetblock = errors.New(
|
||||
"is not a name, = and a netblock, such as office=203.0.113.0/24")
|
||||
errNotAbsolutePath = errors.New(
|
||||
"is not an absolute path, such as /var/lib/smallwebwaf")
|
||||
errShortToken = errors.New("is shorter than 32 characters")
|
||||
errNotMode = errors.New("is not enforce or observe")
|
||||
errNotBytesCount = errors.New("is not response, request or both")
|
||||
errNotPathPrefix = errors.New(
|
||||
"is not a path prefix starting with /, such as /assets/")
|
||||
errNotBoolean = errors.New("is not true or false")
|
||||
@@ -280,6 +393,23 @@ var (
|
||||
"source_failure or file_error")
|
||||
errNotNumberOrOff = errors.New("is not a whole number above zero, such as 60, or off")
|
||||
errNotUTF8 = errors.New("is not valid UTF-8")
|
||||
errNotASN = errors.New("is not an AS number such as AS64496")
|
||||
errNotPercentItem = errors.New(
|
||||
"is not a code, : and a percentage, such as AS64496:50 or cn:25")
|
||||
errNotPercent = errors.New("is not a percentage, a whole number from 0 to 100")
|
||||
errListedTwice = errors.New("is listed twice")
|
||||
errNotListURL = errors.New(
|
||||
"is not an http or https URL without a user or a fragment, " +
|
||||
"such as https://www.spamhaus.org/drop/drop.txt")
|
||||
errInBlocklistURLs = errors.New("is in SWWAF_BLOCKLIST_URLS too")
|
||||
errNotAnHourOrMore = errors.New("is not a duration of 1h or more, such as 24h")
|
||||
errNotAction = errors.New(
|
||||
"is not deny, limit:<percent> such as limit:25, or log")
|
||||
errNotZone = errors.New("is not a DNS zone such as dnsbl.dronebl.org")
|
||||
errZoneTooLong = errors.New("is longer than 189 characters, too long for the " +
|
||||
"names IPv6 clients are asked about by")
|
||||
errNotResolver = errors.New("is not an IP address with an optional port, " +
|
||||
"such as 192.0.2.53 or [2001:db8::53]:5353")
|
||||
)
|
||||
|
||||
// FromEnvironment reads the settings with lookupEnv, normally
|
||||
@@ -287,6 +417,8 @@ var (
|
||||
// named by the setting's name with _FILE added names the file, which is
|
||||
// read now (see lookup). A setting that is not set takes its default. A
|
||||
// setting that is set but invalid is an error that names it.
|
||||
//
|
||||
//nolint:funlen // one line for each setting, a list that grows with them
|
||||
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
||||
env := &environment{lookupEnv: lookupEnv}
|
||||
cfg := &Config{
|
||||
@@ -295,6 +427,8 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
||||
InstanceName: env.instanceName(),
|
||||
Observe: env.observe("SWWAF_MODE", "enforce"),
|
||||
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
|
||||
IPv6GroupPrefix: env.ipv6GroupPrefix("SWWAF_IPV6_GROUP_PREFIX", "64"),
|
||||
MaxTrackedClients: env.numberNotOff("SWWAF_MAX_TRACKED_CLIENTS", "20000"),
|
||||
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
|
||||
ClientRequestHeaderMaxBytes: env.headerSize(
|
||||
"SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES", "32K"),
|
||||
@@ -311,12 +445,32 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
||||
RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"),
|
||||
RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"),
|
||||
RateLimitExemptPaths: env.pathPrefixes("SWWAF_RATE_LIMIT_EXEMPT_PATHS", ""),
|
||||
BytesLimitPerMinute: env.size("SWWAF_BYTES_LIMIT_PER_MINUTE", "10G"),
|
||||
BytesLimitPerHour: env.size("SWWAF_BYTES_LIMIT_PER_HOUR", "20G"),
|
||||
BytesLimitPerDay: env.size("SWWAF_BYTES_LIMIT_PER_DAY", "50G"),
|
||||
BytesCount: env.bytesCount("SWWAF_BYTES_COUNT", "both"),
|
||||
LookupSource: env.lookupSource("SWWAF_LOOKUP_SOURCE", "geojs"),
|
||||
LookupDBPath: env.value("SWWAF_LOOKUP_DB_PATH", ""),
|
||||
LookupTimeout: env.durationNotOff("SWWAF_LOOKUP_TIMEOUT", "1s"),
|
||||
AddLookupHeaders: env.boolean("SWWAF_ADD_LOOKUP_HEADERS", "false"),
|
||||
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
|
||||
ExclusivelyAllowedCountries: env.countries(
|
||||
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""),
|
||||
ASNLimitPercent: env.percents("SWWAF_ASN_LIMIT_PERCENT", ParseASN),
|
||||
CountryLimitPercent: env.percents("SWWAF_COUNTRY_LIMIT_PERCENT", parseCountry),
|
||||
ASNBytesPercent: env.percents("SWWAF_ASN_BYTES_PERCENT", ParseASN),
|
||||
CountryBytesPercent: env.percents("SWWAF_COUNTRY_BYTES_PERCENT", parseCountry),
|
||||
UnknownLimitPercent: env.percent("SWWAF_UNKNOWN_LIMIT_PERCENT", "100"),
|
||||
ASNLimitPercentURL: env.listURL("SWWAF_ASN_LIMIT_PERCENT_URL"),
|
||||
BlocklistURLs: env.listURLs("SWWAF_BLOCKLIST_URLS"),
|
||||
BlocklistRefresh: env.refresh("SWWAF_BLOCKLIST_REFRESH", "24h"),
|
||||
DNSBLZones: env.zones("SWWAF_DNSBL_ZONES"),
|
||||
DNSBLResolver: env.resolver("SWWAF_DNSBL_RESOLVER"),
|
||||
AbuseIPDBKey: env.secret("SWWAF_ABUSEIPDB_KEY"),
|
||||
AbuseIPDBMinScore: env.percent("SWWAF_ABUSEIPDB_MIN_SCORE", "75"),
|
||||
AbuseIPDBDailyBudget: env.numberNotOff("SWWAF_ABUSEIPDB_DAILY_BUDGET", "900"),
|
||||
ReputationCacheTTL: env.durationNotOff("SWWAF_REPUTATION_CACHE_TTL", "24h"),
|
||||
ReputationTimeout: env.durationNotOff("SWWAF_REPUTATION_TIMEOUT", "2s"),
|
||||
BanResponse: env.banResponse("SWWAF_BAN_RESPONSE", "403"),
|
||||
LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"),
|
||||
LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"),
|
||||
@@ -329,6 +483,7 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
||||
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
|
||||
LogRequestHeaders: env.headerNames("SWWAF_LOG_REQUEST_HEADERS",
|
||||
"accept,accept-language,accept-encoding,content-type,origin,range"),
|
||||
LogLevel: env.logLevel("SWWAF_LOG_LEVEL", "info"),
|
||||
AdminToken: env.token("SWWAF_ADMIN_TOKEN"),
|
||||
MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
|
||||
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
|
||||
@@ -347,13 +502,27 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
||||
strings.Join(alerts.Events(), ",")),
|
||||
AlertCooldown: env.duration("SWWAF_ALERT_COOLDOWN", "15m"),
|
||||
AlertMaxPerHour: env.numberOrOff("SWWAF_ALERT_MAX_PER_HOUR", "60"),
|
||||
AnomalyClient: env.thresholds("SWWAF_ANOMALY_CLIENT_"),
|
||||
AnomalyNet: env.thresholds("SWWAF_ANOMALY_NET_"),
|
||||
AnomalyASN: env.thresholds("SWWAF_ANOMALY_ASN_"),
|
||||
AnomalyTotal: env.thresholds("SWWAF_ANOMALY_TOTAL_"),
|
||||
AnomalyWatch: env.thresholds("SWWAF_WATCH_"),
|
||||
AnomalyNetV4Prefix: env.v4Prefix("SWWAF_ANOMALY_NET_V4_PREFIX", "24"),
|
||||
AnomalyNetV6Prefix: env.v6Prefix("SWWAF_ANOMALY_NET_V6_PREFIX", "48"),
|
||||
WatchNets: env.namedNetblocks("SWWAF_WATCH_NETS"),
|
||||
}
|
||||
|
||||
cfg.LogRemoteAppName = env.appName("SWWAF_LOG_REMOTE_APP_NAME",
|
||||
cfg.InstanceName, cfg.LogRemoteURL != nil)
|
||||
cfg.BlocklistAction, cfg.BlocklistLimitPercent = env.action(
|
||||
"SWWAF_BLOCKLIST_ACTION", "deny")
|
||||
cfg.ReputationAction, cfg.ReputationLimitPercent = env.action(
|
||||
"SWWAF_REPUTATION_ACTION", "limit:25")
|
||||
|
||||
env.checkInstanceNameForNtfy(cfg.InstanceName, cfg.AlertNtfyURL != nil)
|
||||
env.checkLookupDBPath(cfg)
|
||||
env.checkCountriesAndLookups(cfg)
|
||||
env.checkASNLimitPercentURL(cfg)
|
||||
|
||||
if env.err != nil {
|
||||
return nil, env.err
|
||||
@@ -539,6 +708,17 @@ func (e *environment) count(name, defaultValue string) int64 {
|
||||
return count
|
||||
}
|
||||
|
||||
// bytesCount reads the setting that is which body bytes count toward the
|
||||
// byte limits: response, request or both.
|
||||
func (e *environment) bytesCount(name, defaultValue string) string {
|
||||
value := e.value(name, defaultValue)
|
||||
if value != "response" && value != "request" && value != "both" {
|
||||
e.check(name, fmt.Errorf("%q %w", value, errNotBytesCount))
|
||||
}
|
||||
|
||||
return value
|
||||
}
|
||||
|
||||
// pathPrefixes reads a setting that is a list of path prefixes.
|
||||
func (e *environment) pathPrefixes(name, defaultValue string) []string {
|
||||
prefixes, err := parsePathPrefixes(e.value(name, defaultValue))
|
||||
@@ -555,20 +735,137 @@ func (e *environment) countries(name, defaultValue string) []string {
|
||||
return countries
|
||||
}
|
||||
|
||||
// percents reads a setting that is a list of AS numbers or countries,
|
||||
// which parseCode reads, each with a percentage. It is empty by default.
|
||||
func (e *environment) percents(
|
||||
name string, parseCode func(string) (string, error),
|
||||
) map[string]int64 {
|
||||
percents, err := parsePercents(e.value(name, ""), parseCode)
|
||||
e.check(name, err)
|
||||
|
||||
return percents
|
||||
}
|
||||
|
||||
// percent reads a setting that is a percentage, from 0 to 100.
|
||||
func (e *environment) percent(name, defaultValue string) int64 {
|
||||
percent, err := ParsePercent(e.value(name, defaultValue))
|
||||
e.check(name, err)
|
||||
|
||||
return percent
|
||||
}
|
||||
|
||||
// listURL reads a setting that is the URL a list is fetched from, "" while
|
||||
// it is unset or empty.
|
||||
func (e *environment) listURL(name string) string {
|
||||
value := e.value(name, "")
|
||||
if value != "" && !isListURL(value) {
|
||||
e.check(name, fmt.Errorf("%q %w", value, errNotListURL))
|
||||
}
|
||||
|
||||
return value
|
||||
}
|
||||
|
||||
// listURLs reads a setting that is a list of the URLs lists are fetched
|
||||
// from. It is empty by default.
|
||||
func (e *environment) listURLs(name string) []string {
|
||||
urls, err := parseListURLs(e.value(name, ""))
|
||||
e.check(name, err)
|
||||
|
||||
return urls
|
||||
}
|
||||
|
||||
// refresh reads the setting that is how long after a list was last
|
||||
// fetched or tried it is fetched again: a duration of an hour or more,
|
||||
// since the Spamhaus lists may be fetched no more often, which cannot be
|
||||
// off.
|
||||
func (e *environment) refresh(name, defaultValue string) time.Duration {
|
||||
value := e.value(name, defaultValue)
|
||||
|
||||
duration, err := parseDuration(value)
|
||||
if err != nil || duration < time.Hour {
|
||||
e.check(name, fmt.Errorf("%q %w", value, errNotAnHourOrMore))
|
||||
}
|
||||
|
||||
return duration
|
||||
}
|
||||
|
||||
// action reads a setting that is what is done with a client a blocklist
|
||||
// or a DNSBL zone lists: deny, log, or limit:<percent>, which it returns
|
||||
// as limit and the percentage.
|
||||
func (e *environment) action(name, defaultValue string) (string, int64) {
|
||||
value := e.value(name, defaultValue)
|
||||
if value == "deny" || value == "log" {
|
||||
return value, 0
|
||||
}
|
||||
|
||||
percentText, isLimit := strings.CutPrefix(value, "limit:")
|
||||
|
||||
percent, err := ParsePercent(percentText)
|
||||
if !isLimit || err != nil {
|
||||
e.check(name, fmt.Errorf("%q %w", value, errNotAction))
|
||||
}
|
||||
|
||||
return "limit", percent
|
||||
}
|
||||
|
||||
// zones reads the setting that is the list of DNSBL zones. It is empty by
|
||||
// default. The log shows each zone with its key masked, as MaskZoneKey
|
||||
// masks it.
|
||||
func (e *environment) zones(name string) []string {
|
||||
value, _ := e.lookup(name)
|
||||
zones, err := parseZones(value)
|
||||
e.check(name, err)
|
||||
|
||||
logged := make([]string, len(zones))
|
||||
for i, zone := range zones {
|
||||
logged[i] = MaskZoneKey(zone)
|
||||
}
|
||||
|
||||
e.settings = append(e.settings, slog.String(name, strings.Join(logged, ",")))
|
||||
|
||||
return zones
|
||||
}
|
||||
|
||||
// resolver reads the setting that is the resolver the DNSBL zones are
|
||||
// asked through, the zero AddrPort while it is unset or empty.
|
||||
func (e *environment) resolver(name string) netip.AddrPort {
|
||||
resolver, err := parseResolver(e.value(name, ""))
|
||||
e.check(name, err)
|
||||
|
||||
return resolver
|
||||
}
|
||||
|
||||
// lookupSource reads the setting that is where clients are looked up:
|
||||
// geojs, or off.
|
||||
// geojs, file, or off.
|
||||
func (e *environment) lookupSource(name, defaultValue string) string {
|
||||
source := e.value(name, defaultValue)
|
||||
if source != "geojs" && source != off {
|
||||
if source != "geojs" && source != fileSource && source != off {
|
||||
e.check(name, fmt.Errorf("%q %w", source, errNotLookupSource))
|
||||
}
|
||||
|
||||
return source
|
||||
}
|
||||
|
||||
// checkLookupDBPath refuses SWWAF_LOOKUP_SOURCE=file without
|
||||
// SWWAF_LOOKUP_DB_PATH, and SWWAF_LOOKUP_DB_PATH with any other source:
|
||||
// one source at a time.
|
||||
func (e *environment) checkLookupDBPath(cfg *Config) {
|
||||
switch {
|
||||
case cfg.LookupSource == fileSource && cfg.LookupDBPath == "":
|
||||
e.check("SWWAF_LOOKUP_SOURCE", fmt.Errorf(
|
||||
"is file while SWWAF_LOOKUP_DB_PATH is unset; %w", errNeedsDBPath))
|
||||
case cfg.LookupSource != fileSource && cfg.LookupDBPath != "":
|
||||
e.check("SWWAF_LOOKUP_DB_PATH", fmt.Errorf(
|
||||
"is set while SWWAF_LOOKUP_SOURCE is %s; %w", cfg.LookupSource, errDBPathUnused))
|
||||
}
|
||||
}
|
||||
|
||||
// checkCountriesAndLookups refuses a country on both country lists, and,
|
||||
// while SWWAF_LOOKUP_SOURCE is off, each setting that needs clients looked
|
||||
// up: the country lists and SWWAF_ADD_LOOKUP_HEADERS.
|
||||
// up: the country lists, SWWAF_ADD_LOOKUP_HEADERS, the biased thresholds,
|
||||
// SWWAF_ASN_LIMIT_PERCENT_URL among them, of which
|
||||
// SWWAF_UNKNOWN_LIMIT_PERCENT needs them only below 100, where it lowers a
|
||||
// limit, and the anomaly thresholds per AS number.
|
||||
func (e *environment) checkCountriesAndLookups(cfg *Config) {
|
||||
for _, country := range cfg.ExclusivelyAllowedCountries {
|
||||
if slices.Contains(cfg.DeniedCountries, country) {
|
||||
@@ -588,6 +885,16 @@ func (e *environment) checkCountriesAndLookups(cfg *Config) {
|
||||
{"SWWAF_DENIED_COUNTRIES", len(cfg.DeniedCountries) > 0},
|
||||
{"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", len(cfg.ExclusivelyAllowedCountries) > 0},
|
||||
{"SWWAF_ADD_LOOKUP_HEADERS", cfg.AddLookupHeaders},
|
||||
{"SWWAF_ASN_LIMIT_PERCENT", len(cfg.ASNLimitPercent) > 0},
|
||||
{"SWWAF_COUNTRY_LIMIT_PERCENT", len(cfg.CountryLimitPercent) > 0},
|
||||
{"SWWAF_ASN_BYTES_PERCENT", len(cfg.ASNBytesPercent) > 0},
|
||||
{"SWWAF_COUNTRY_BYTES_PERCENT", len(cfg.CountryBytesPercent) > 0},
|
||||
{"SWWAF_UNKNOWN_LIMIT_PERCENT", cfg.UnknownLimitPercent < 100},
|
||||
{"SWWAF_ASN_LIMIT_PERCENT_URL", cfg.ASNLimitPercentURL != ""},
|
||||
{"SWWAF_ANOMALY_ASN_REQUESTS_PER_MINUTE", cfg.AnomalyASN.RequestsPerMinute > 0},
|
||||
{"SWWAF_ANOMALY_ASN_REQUESTS_PER_HOUR", cfg.AnomalyASN.RequestsPerHour > 0},
|
||||
{"SWWAF_ANOMALY_ASN_BYTES_PER_MINUTE", cfg.AnomalyASN.BytesPerMinute > 0},
|
||||
{"SWWAF_ANOMALY_ASN_BYTES_PER_HOUR", cfg.AnomalyASN.BytesPerHour > 0},
|
||||
} {
|
||||
if setting.set {
|
||||
e.check(setting.name, fmt.Errorf("is set while SWWAF_LOOKUP_SOURCE is off; %w",
|
||||
@@ -596,6 +903,15 @@ func (e *environment) checkCountriesAndLookups(cfg *Config) {
|
||||
}
|
||||
}
|
||||
|
||||
// checkASNLimitPercentURL refuses SWWAF_ASN_LIMIT_PERCENT_URL naming a
|
||||
// blocklist too: the file at a URL is fetched as one list or the other.
|
||||
func (e *environment) checkASNLimitPercentURL(cfg *Config) {
|
||||
if slices.Contains(cfg.BlocklistURLs, cfg.ASNLimitPercentURL) {
|
||||
e.check("SWWAF_ASN_LIMIT_PERCENT_URL",
|
||||
fmt.Errorf("%q %w", cfg.ASNLimitPercentURL, errInBlocklistURLs))
|
||||
}
|
||||
}
|
||||
|
||||
// headerNames reads a setting that is a list of header names, and
|
||||
// returns them in lower case.
|
||||
func (e *environment) headerNames(name, defaultValue string) []string {
|
||||
@@ -639,6 +955,66 @@ func (e *environment) v4Prefix(name, defaultValue string) int {
|
||||
return length
|
||||
}
|
||||
|
||||
// v6Prefix reads a setting that is the length of an IPv6 netblock.
|
||||
func (e *environment) v6Prefix(name, defaultValue string) int {
|
||||
length, err := parseV6Prefix(e.value(name, defaultValue))
|
||||
e.check(name, err)
|
||||
|
||||
return length
|
||||
}
|
||||
|
||||
// ipv6GroupPrefix reads the setting that is the length of the IPv6
|
||||
// netblock that is one client, from minIPv6GroupPrefix to 128.
|
||||
func (e *environment) ipv6GroupPrefix(name, defaultValue string) int {
|
||||
value := e.value(name, defaultValue)
|
||||
|
||||
length, err := strconv.Atoi(value)
|
||||
if err != nil || length < minIPv6GroupPrefix || length > ipv6Bits {
|
||||
e.check(name, fmt.Errorf("%q %w", value, errNotIPv6GroupPrefix))
|
||||
}
|
||||
|
||||
return length
|
||||
}
|
||||
|
||||
// logLevel reads the setting that is the least severe of the process's
|
||||
// own messages that are written: debug, info, warn or error.
|
||||
func (e *environment) logLevel(name, defaultValue string) slog.Level {
|
||||
value := e.value(name, defaultValue)
|
||||
|
||||
level, known := map[string]slog.Level{
|
||||
"debug": slog.LevelDebug,
|
||||
"info": slog.LevelInfo,
|
||||
"warn": slog.LevelWarn,
|
||||
"error": slog.LevelError,
|
||||
}[value]
|
||||
if !known {
|
||||
e.check(name, fmt.Errorf("%q %w", value, errNotLogLevel))
|
||||
}
|
||||
|
||||
return level
|
||||
}
|
||||
|
||||
// thresholds reads the four anomaly thresholds whose settings' names
|
||||
// start with prefix: requests and bytes per minute and per hour. Each is
|
||||
// off by default.
|
||||
func (e *environment) thresholds(prefix string) anomaly.Thresholds {
|
||||
return anomaly.Thresholds{
|
||||
RequestsPerMinute: e.count(prefix+"REQUESTS_PER_MINUTE", off),
|
||||
RequestsPerHour: e.count(prefix+"REQUESTS_PER_HOUR", off),
|
||||
BytesPerMinute: e.size(prefix+"BYTES_PER_MINUTE", off),
|
||||
BytesPerHour: e.size(prefix+"BYTES_PER_HOUR", off),
|
||||
}
|
||||
}
|
||||
|
||||
// namedNetblocks reads a setting that is a list of named netblocks. It is
|
||||
// empty by default.
|
||||
func (e *environment) namedNetblocks(name string) []anomaly.NamedNetblock {
|
||||
named, err := parseNamedNetblocks(e.value(name, ""))
|
||||
e.check(name, err)
|
||||
|
||||
return named
|
||||
}
|
||||
|
||||
// absolutePath reads a setting that is an absolute path.
|
||||
func (e *environment) absolutePath(name, defaultValue string) string {
|
||||
path := e.value(name, defaultValue)
|
||||
@@ -795,10 +1171,10 @@ func (e *environment) webhookHeaders(name string) http.Header {
|
||||
}
|
||||
|
||||
// secret reads a setting that is a secret another service gave, such as
|
||||
// an ntfy token, "" while it is unset. It is sent in a header, which
|
||||
// cannot hold a control character, so one in it is an error. The log
|
||||
// shows ******** in place of a value that is not empty, and an error
|
||||
// shows none of it.
|
||||
// an ntfy token or an AbuseIPDB key, "" while it is unset. It is sent in
|
||||
// a header, which cannot hold a control character, so one in it is an
|
||||
// error. The log shows ******** in place of a value that is not empty, and
|
||||
// an error shows none of it.
|
||||
func (e *environment) secret(name string) string {
|
||||
value, _ := e.lookup(name)
|
||||
|
||||
@@ -980,6 +1356,16 @@ func parseV4Prefix(value string) (int, error) {
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// parseV6Prefix reads the length of an IPv6 netblock, from 0 to 128.
|
||||
func parseV6Prefix(value string) (int, error) {
|
||||
n, err := strconv.Atoi(value)
|
||||
if err != nil || n < 0 || n > ipv6Bits {
|
||||
return 0, fmt.Errorf("%q %w", value, errNotV6Prefix)
|
||||
}
|
||||
|
||||
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) {
|
||||
@@ -1008,7 +1394,7 @@ func parseNetblocks(value string) ([]netip.Prefix, error) {
|
||||
netblocks := make([]netip.Prefix, 0, len(items))
|
||||
|
||||
for _, item := range items {
|
||||
netblock, err := parseNetblock(item)
|
||||
netblock, err := ParseNetblock(item)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -1019,9 +1405,10 @@ func parseNetblocks(value string) ([]netip.Prefix, error) {
|
||||
return netblocks, nil
|
||||
}
|
||||
|
||||
// parseNetblock reads a netblock in CIDR form, such as 10.0.0.0/8. A bare
|
||||
// address is a netblock of that address alone, a /32 or a /128.
|
||||
func parseNetblock(value string) (netip.Prefix, error) {
|
||||
// ParseNetblock reads a netblock in CIDR form, such as 10.0.0.0/8. A bare
|
||||
// address is a netblock of that address alone, a /32 or a /128. A
|
||||
// blocklist's lines are read with it too.
|
||||
func ParseNetblock(value string) (netip.Prefix, error) {
|
||||
if strings.Contains(value, "/") {
|
||||
netblock, err := netip.ParsePrefix(value)
|
||||
if err != nil {
|
||||
@@ -1039,6 +1426,42 @@ func parseNetblock(value string) (netip.Prefix, error) {
|
||||
return netip.PrefixFrom(addr, addr.BitLen()), nil
|
||||
}
|
||||
|
||||
// parseNamedNetblocks reads a comma-separated list of named netblocks,
|
||||
// each a name, = and a netblock, such as office=203.0.113.0/24. An empty
|
||||
// value is an empty list. A name listed twice is an error.
|
||||
func parseNamedNetblocks(value string) ([]anomaly.NamedNetblock, error) {
|
||||
items, err := parseList(value)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
named := make([]anomaly.NamedNetblock, 0, len(items))
|
||||
|
||||
for _, item := range items {
|
||||
name, netblockText, found := strings.Cut(item, "=")
|
||||
name = strings.TrimSpace(name)
|
||||
|
||||
if !found || name == "" {
|
||||
return nil, fmt.Errorf("%q %w", item, errNotNamedNetblock)
|
||||
}
|
||||
|
||||
netblock, err := ParseNetblock(strings.TrimSpace(netblockText))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if slices.ContainsFunc(named, func(n anomaly.NamedNetblock) bool {
|
||||
return n.Name == name
|
||||
}) {
|
||||
return nil, fmt.Errorf("%q %w", name, errListedTwice)
|
||||
}
|
||||
|
||||
named = append(named, anomaly.NamedNetblock{Name: name, Netblock: netblock})
|
||||
}
|
||||
|
||||
return named, nil
|
||||
}
|
||||
|
||||
// parsePathPrefixes reads a comma-separated list of path prefixes, each
|
||||
// starting with /.
|
||||
func parsePathPrefixes(value string) ([]string, error) {
|
||||
@@ -1097,13 +1520,12 @@ func parseCountries(value string) ([]string, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
known := strings.Fields(countryCodes)
|
||||
countries := make([]string, 0, len(items))
|
||||
|
||||
for _, item := range items {
|
||||
country := strings.ToUpper(item)
|
||||
if !slices.Contains(known, country) {
|
||||
return nil, fmt.Errorf("%q %w", item, errNotCountry)
|
||||
country, err := parseCountry(item)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
countries = append(countries, country)
|
||||
@@ -1112,6 +1534,83 @@ func parseCountries(value string) ([]string, error) {
|
||||
return countries, nil
|
||||
}
|
||||
|
||||
// parseCountry reads a country code in either case, and returns it in
|
||||
// capitals.
|
||||
func parseCountry(value string) (string, error) {
|
||||
country := strings.ToUpper(value)
|
||||
if !slices.Contains(strings.Fields(countryCodes), country) {
|
||||
return "", fmt.Errorf("%q %w", value, errNotCountry)
|
||||
}
|
||||
|
||||
return country, nil
|
||||
}
|
||||
|
||||
// ParseASN reads an AS number such as AS64496, in either case, and
|
||||
// returns it as the lookup gives it: AS and the number, in capitals and
|
||||
// without leading zeros. The file SWWAF_ASN_LIMIT_PERCENT_URL names is
|
||||
// read with it too.
|
||||
func ParseASN(value string) (string, error) {
|
||||
digits, hasAS := strings.CutPrefix(strings.ToUpper(value), "AS")
|
||||
|
||||
number, err := strconv.ParseUint(digits, 10, 32)
|
||||
if !hasAS || err != nil {
|
||||
return "", fmt.Errorf("%q %w", value, errNotASN)
|
||||
}
|
||||
|
||||
return "AS" + strconv.FormatUint(number, 10), nil
|
||||
}
|
||||
|
||||
// parsePercents reads a comma-separated list of items, each an AS number
|
||||
// or a country, which parseCode reads, then : and a percentage, such as
|
||||
// AS64496:50 or cn:25, and returns each one's percentage. An empty value
|
||||
// is an empty list. An AS number or country listed twice is an error.
|
||||
func parsePercents(
|
||||
value string, parseCode func(string) (string, error),
|
||||
) (map[string]int64, error) {
|
||||
items, err := parseList(value)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
percents := make(map[string]int64, len(items))
|
||||
|
||||
for _, item := range items {
|
||||
codeText, percentText, found := strings.Cut(item, ":")
|
||||
if !found {
|
||||
return nil, fmt.Errorf("%q %w", item, errNotPercentItem)
|
||||
}
|
||||
|
||||
code, err := parseCode(codeText)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
percent, err := ParsePercent(percentText)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if _, listed := percents[code]; listed {
|
||||
return nil, fmt.Errorf("%q %w", codeText, errListedTwice)
|
||||
}
|
||||
|
||||
percents[code] = percent
|
||||
}
|
||||
|
||||
return percents, nil
|
||||
}
|
||||
|
||||
// ParsePercent reads a percentage, a whole number from 0 to 100. The file
|
||||
// SWWAF_ASN_LIMIT_PERCENT_URL names is read with it too.
|
||||
func ParsePercent(value string) (int64, error) {
|
||||
percent, err := strconv.ParseInt(value, 10, 64)
|
||||
if err != nil || percent < 0 || percent > 100 {
|
||||
return 0, fmt.Errorf("%q %w", value, errNotPercent)
|
||||
}
|
||||
|
||||
return percent, nil
|
||||
}
|
||||
|
||||
// headerNameChars are the characters RFC 9110 allows in a header name:
|
||||
// letters, digits and these marks.
|
||||
const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" +
|
||||
@@ -1259,16 +1758,7 @@ func parseWebhookURL(value string) (*url.URL, string, error) {
|
||||
}
|
||||
|
||||
webhook, err := url.Parse(value)
|
||||
if err != nil {
|
||||
return nil, "", errNotWebhookURL
|
||||
}
|
||||
|
||||
port, err := strconv.ParseUint(webhook.Port(), 10, 16)
|
||||
|
||||
valid := (webhook.Scheme == "http" || webhook.Scheme == "https") &&
|
||||
webhook.Hostname() != "" && (webhook.Port() == "" || (err == nil && port != 0)) &&
|
||||
webhook.User == nil && webhook.Opaque == "" && webhook.Fragment == ""
|
||||
if !valid {
|
||||
if err != nil || !isHTTPURL(webhook) {
|
||||
return nil, "", errNotWebhookURL
|
||||
}
|
||||
|
||||
@@ -1280,6 +1770,154 @@ func parseWebhookURL(value string) (*url.URL, string, error) {
|
||||
return webhook, logged, nil
|
||||
}
|
||||
|
||||
// isHTTPURL reports whether u is http or https, with a host, and an
|
||||
// optional port from 1 to 65535, path and query, without a user or a
|
||||
// fragment.
|
||||
func isHTTPURL(u *url.URL) bool {
|
||||
port, err := strconv.ParseUint(u.Port(), 10, 16)
|
||||
|
||||
return (u.Scheme == "http" || u.Scheme == "https") && u.Hostname() != "" &&
|
||||
(u.Port() == "" || (err == nil && port != 0)) &&
|
||||
u.User == nil && u.Opaque == "" && u.Fragment == ""
|
||||
}
|
||||
|
||||
// isListURL reports whether value is a URL a list can be fetched from, as
|
||||
// isHTTPURL says.
|
||||
func isListURL(value string) bool {
|
||||
u, err := url.Parse(value)
|
||||
|
||||
return err == nil && isHTTPURL(u)
|
||||
}
|
||||
|
||||
// parseListURLs reads a comma-separated list of the URLs lists are fetched
|
||||
// from. A URL listed twice is an error: it would be fetched twice as
|
||||
// often.
|
||||
func parseListURLs(value string) ([]string, error) {
|
||||
urls, err := parseList(value)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for i, listURL := range urls {
|
||||
if !isListURL(listURL) {
|
||||
return nil, fmt.Errorf("%q %w", listURL, errNotListURL)
|
||||
}
|
||||
|
||||
if slices.Contains(urls[:i], listURL) {
|
||||
return nil, fmt.Errorf("%q %w", listURL, errListedTwice)
|
||||
}
|
||||
}
|
||||
|
||||
return urls, nil
|
||||
}
|
||||
|
||||
const (
|
||||
// maxZoneLength is the most characters a DNSBL zone may have: 253, the
|
||||
// most a DNS name may have, less the 64 that come before the zone in
|
||||
// the name an IPv6 client is asked about by, its 32 hex digits each
|
||||
// followed by a dot.
|
||||
maxZoneLength = 189
|
||||
// maxLabelLength is the most characters a label of a DNS name may have.
|
||||
maxLabelLength = 63
|
||||
// labelChars are the characters a label of a DNS zone may hold.
|
||||
labelChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-"
|
||||
// dnsPort is the port a resolver is asked on when SWWAF_DNSBL_RESOLVER
|
||||
// gives none.
|
||||
dnsPort = 53
|
||||
)
|
||||
|
||||
// parseZones reads a comma-separated list of DNSBL zones, each a DNS name
|
||||
// such as dnsbl.dronebl.org: labels separated by dots, each of 1 to 63
|
||||
// letters, digits and hyphens, neither starting nor ending with a hyphen,
|
||||
// and at most maxZoneLength characters in all. Go's resolver takes any
|
||||
// other name for one that does not exist, so that the zone would list no
|
||||
// client. A zone listed twice is an error, whatever the case of its
|
||||
// letters, which DNS names ignore, and whatever its key, since
|
||||
// MaskZoneKey shows two keys of one zone alike. An error shows a zone as
|
||||
// MaskZoneKey does.
|
||||
func parseZones(value string) ([]string, error) {
|
||||
zones, err := parseList(value)
|
||||
if err != nil {
|
||||
// parseList's error, for an empty item, shows the whole value, keys
|
||||
// included.
|
||||
return nil, errEmptyItem
|
||||
}
|
||||
|
||||
for i, zone := range zones {
|
||||
shown := MaskZoneKey(zone)
|
||||
listedBefore := slices.ContainsFunc(zones[:i], func(earlier string) bool {
|
||||
return strings.EqualFold(MaskZoneKey(earlier), shown)
|
||||
})
|
||||
|
||||
switch {
|
||||
case len(zone) > maxZoneLength:
|
||||
return nil, fmt.Errorf("%q %w", shown, errZoneTooLong)
|
||||
case !isZone(zone):
|
||||
return nil, fmt.Errorf("%q %w", shown, errNotZone)
|
||||
case listedBefore:
|
||||
return nil, fmt.Errorf("%q %w", shown, errListedTwice)
|
||||
}
|
||||
}
|
||||
|
||||
return zones, nil
|
||||
}
|
||||
|
||||
// MaskZoneKey returns zone with ******** in place of its key, if it is a
|
||||
// zone of Spamhaus's keyed query service, a name under dq.spamhaus.net,
|
||||
// such as <key>.xbl.dq.spamhaus.net, whose first label is the key. Any
|
||||
// other zone it returns as it is. A zone is shown so wherever it leaves
|
||||
// the process: in the log, the alerts and the metrics.
|
||||
func MaskZoneKey(zone string) string {
|
||||
// DNS names ignore case, and a name may be written with a dot at its
|
||||
// end.
|
||||
name := strings.TrimSuffix(strings.ToLower(zone), ".")
|
||||
if !strings.HasSuffix(name, ".dq.spamhaus.net") {
|
||||
return zone
|
||||
}
|
||||
|
||||
_, rest, _ := strings.Cut(zone, ".")
|
||||
|
||||
return masked + "." + rest
|
||||
}
|
||||
|
||||
// isZone reports whether each label of zone is as parseZones takes it.
|
||||
func isZone(zone string) bool {
|
||||
for label := range strings.SplitSeq(zone, ".") {
|
||||
badChar := strings.ContainsFunc(label, func(char rune) bool {
|
||||
return !strings.ContainsRune(labelChars, char)
|
||||
})
|
||||
|
||||
if label == "" || len(label) > maxLabelLength || badChar ||
|
||||
strings.HasPrefix(label, "-") || strings.HasSuffix(label, "-") {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// parseResolver reads the resolver the DNSBL zones are asked through: an
|
||||
// IP address with a port from 1 to 65535, such as 192.0.2.53:5353 or
|
||||
// [2001:db8::53]:5353, or without one, such as 192.0.2.53 or 2001:db8::53,
|
||||
// for port 53. An empty value is none, the zero AddrPort.
|
||||
func parseResolver(value string) (netip.AddrPort, error) {
|
||||
if value == "" {
|
||||
return netip.AddrPort{}, nil
|
||||
}
|
||||
|
||||
resolver, err := netip.ParseAddrPort(value)
|
||||
if err == nil && resolver.Port() != 0 {
|
||||
return resolver, nil
|
||||
}
|
||||
|
||||
addr, err := netip.ParseAddr(value)
|
||||
if err != nil {
|
||||
return netip.AddrPort{}, fmt.Errorf("%q %w", value, errNotResolver)
|
||||
}
|
||||
|
||||
return netip.AddrPortFrom(addr, dnsPort), nil
|
||||
}
|
||||
|
||||
// parseWebhookHeaders reads a comma-separated list of headers, each its
|
||||
// name, :, and its value, and returns them, and how the log shows them,
|
||||
// with each value as ********. An error names the item by its place in
|
||||
|
||||
+910
-21
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,233 @@
|
||||
package lookup
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/fsnotify/fsnotify"
|
||||
"github.com/oschwald/maxminddb-golang/v2"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
)
|
||||
|
||||
// quietTime is how long the lookup database must go without a change
|
||||
// before it is read again, so that a file still being copied in is read
|
||||
// only once whole.
|
||||
const quietTime = 2 * time.Second
|
||||
|
||||
// FileParams are what OpenFile needs.
|
||||
type FileParams struct {
|
||||
// Path is the lookup database, the IPinfo Lite file in its .mmdb form
|
||||
// (SWWAF_LOOKUP_DB_PATH).
|
||||
Path string
|
||||
// Now tells the time, normally time.Now.
|
||||
Now func() time.Time
|
||||
// ProcessLog receives each reading of the file, and why a replacement
|
||||
// of it cannot be read.
|
||||
ProcessLog *slog.Logger
|
||||
// Alerts receive a file_error alert for each replacement that cannot
|
||||
// be read.
|
||||
Alerts *alerts.Queue
|
||||
}
|
||||
|
||||
// File looks up clients' AS numbers and countries in the lookup database,
|
||||
// held in memory, and reads it again when it is replaced. It is safe for
|
||||
// concurrent use.
|
||||
type File struct {
|
||||
params FileParams
|
||||
|
||||
mu sync.Mutex
|
||||
// reader is the database in use, and lastRead when it was read.
|
||||
// readFailures are the replacements that could not be read.
|
||||
reader *maxminddb.Reader
|
||||
lastRead time.Time
|
||||
readFailures int
|
||||
}
|
||||
|
||||
// record is what the lookup database holds about a network, of the fields
|
||||
// smallwebwaf reads.
|
||||
type record struct {
|
||||
ASN string `maxminddb:"asn"`
|
||||
ASName string `maxminddb:"as_name"`
|
||||
CountryCode string `maxminddb:"country_code"`
|
||||
}
|
||||
|
||||
// OpenFile reads the lookup database. A file that cannot be read, or that
|
||||
// is not a .mmdb file, is an error.
|
||||
func OpenFile(params FileParams) (*File, error) {
|
||||
reader, err := read(params.Path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
f := &File{params: params}
|
||||
f.use(reader)
|
||||
|
||||
return f, nil
|
||||
}
|
||||
|
||||
// LookUp returns what the lookup database says about client: its AS
|
||||
// number, such as AS64496, the AS's name, and its country, such as DE,
|
||||
// each "" when the database does not give it, as for an address missing
|
||||
// from it. The database is asked about the client's first address, as
|
||||
// GeoJS is.
|
||||
func (f *File) LookUp(client netip.Prefix) Answer {
|
||||
f.mu.Lock()
|
||||
reader := f.reader
|
||||
f.mu.Unlock()
|
||||
|
||||
var found record
|
||||
|
||||
// A record that cannot be decoded places the client nowhere, as a
|
||||
// missing one does.
|
||||
err := reader.Lookup(client.Addr()).Decode(&found)
|
||||
if err != nil {
|
||||
found = record{}
|
||||
}
|
||||
|
||||
return Answer{
|
||||
Client: client,
|
||||
ASN: found.ASN,
|
||||
ASName: found.ASName,
|
||||
Country: found.CountryCode,
|
||||
Answered: f.params.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
// LastRead returns when the lookup database in use was read.
|
||||
func (f *File) LastRead() time.Time {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
|
||||
return f.lastRead
|
||||
}
|
||||
|
||||
// ReadFailures returns how many replacements of the lookup database could
|
||||
// not be read.
|
||||
func (f *File) ReadFailures() int {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
|
||||
return f.readFailures
|
||||
}
|
||||
|
||||
// Watch watches the directory of the lookup database until ctx is done,
|
||||
// and reads the file again once it has gone without a change for
|
||||
// quietTime, after it is replaced, written or removed, and after Watch
|
||||
// starts watching. If the directory cannot be watched, that is logged, and
|
||||
// the database read at start stays in use.
|
||||
func (f *File) Watch(ctx context.Context) {
|
||||
watcher, err := fsnotify.NewWatcher()
|
||||
if err == nil {
|
||||
defer func() {
|
||||
_ = watcher.Close()
|
||||
}()
|
||||
|
||||
err = watcher.Add(filepath.Dir(f.params.Path))
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
f.params.ProcessLog.Error("cannot watch the lookup database for replacements",
|
||||
"error", err.Error())
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
f.params.ProcessLog.Info("watching the lookup database for replacements",
|
||||
"file", f.params.Path)
|
||||
|
||||
f.readAfterChanges(ctx, watcher.Events, watcher.Errors)
|
||||
}
|
||||
|
||||
// readAfterChanges reads the lookup database again once quietTime has
|
||||
// passed without a change to it from events, until ctx is done, and logs
|
||||
// the errors from errs. A change to another file in its directory does not
|
||||
// count. The wait starts at once, as if for a change, so that a file
|
||||
// replaced after OpenFile read it, and before its directory was watched,
|
||||
// is read too.
|
||||
func (f *File) readAfterChanges(
|
||||
ctx context.Context, events <-chan fsnotify.Event, errs <-chan error,
|
||||
) {
|
||||
path := filepath.Clean(f.params.Path)
|
||||
|
||||
quiet := time.NewTimer(quietTime)
|
||||
defer quiet.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case event := <-events:
|
||||
if filepath.Clean(event.Name) == path {
|
||||
quiet.Reset(quietTime)
|
||||
}
|
||||
case <-quiet.C:
|
||||
f.readAgain()
|
||||
case err := <-errs:
|
||||
f.params.ProcessLog.Warn("watching the lookup database failed",
|
||||
"error", err.Error())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// readAgain reads the lookup database again, in place of the one in use,
|
||||
// or, if it cannot be read, counts that, raises a file_error alert for it
|
||||
// and logs it, and the one in use stays in use.
|
||||
func (f *File) readAgain() {
|
||||
reader, err := read(f.params.Path)
|
||||
if err != nil {
|
||||
const kept = "the lookup database cannot be read, " +
|
||||
"and the one read before stays in use"
|
||||
|
||||
f.mu.Lock()
|
||||
f.readFailures++
|
||||
f.mu.Unlock()
|
||||
|
||||
// Raised before it is logged, so that the alert is there once the
|
||||
// log line is.
|
||||
f.params.Alerts.Raise(alerts.Alert{
|
||||
Event: alerts.EventFileError,
|
||||
Reason: kept,
|
||||
Detail: map[string]any{"file": f.params.Path, "error": err.Error()},
|
||||
})
|
||||
f.params.ProcessLog.Error(kept, "error", err.Error())
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
f.use(reader)
|
||||
}
|
||||
|
||||
// use puts reader in use, in place of the database read before, and logs
|
||||
// that the file was read.
|
||||
func (f *File) use(reader *maxminddb.Reader) {
|
||||
f.mu.Lock()
|
||||
f.reader = reader
|
||||
f.lastRead = f.params.Now()
|
||||
f.mu.Unlock()
|
||||
|
||||
f.params.ProcessLog.Info("read the lookup database", "file", f.params.Path)
|
||||
}
|
||||
|
||||
// read reads the lookup database at path. The whole file is read into
|
||||
// memory, rather than mapped into it as the reader can, so that a file
|
||||
// overwritten in place cannot change, or end, under a lookup.
|
||||
func read(path string) (*maxminddb.Reader, error) {
|
||||
data, err := os.ReadFile(path) //nolint:gosec // the file the admin names
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SWWAF_LOOKUP_DB_PATH cannot be read: %w", err)
|
||||
}
|
||||
|
||||
reader, err := maxminddb.OpenBytes(data)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SWWAF_LOOKUP_DB_PATH %s is not a .mmdb file: %w", path, err)
|
||||
}
|
||||
|
||||
return reader, nil
|
||||
}
|
||||
@@ -0,0 +1,379 @@
|
||||
package lookup
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"testing"
|
||||
"testing/synctest"
|
||||
"time"
|
||||
|
||||
"github.com/fsnotify/fsnotify"
|
||||
"github.com/maxmind/mmdbwriter/mmdbtype"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
|
||||
)
|
||||
|
||||
// testNetblock is the netblock the tests' lookup databases place, and
|
||||
// testClient a client in it.
|
||||
const (
|
||||
testNetblock = "203.0.113.0/24"
|
||||
testClient = "203.0.113.9/32"
|
||||
)
|
||||
|
||||
func TestFilePlacesClientsAndCountsAnAddressMissingFromItAsUnknown(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
germany := lookuptest.Network{ASN: "AS64496", ASName: "Example Net", Country: "DE"}
|
||||
northKorea := lookuptest.Network{ASN: "AS64511", ASName: "Other Net", Country: "KP"}
|
||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
||||
testNetblock: germany,
|
||||
"2001:db8::/32": northKorea,
|
||||
})
|
||||
|
||||
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
|
||||
|
||||
f, err := OpenFile(FileParams{
|
||||
Path: path,
|
||||
Now: func() time.Time { return now },
|
||||
ProcessLog: slog.New(slog.DiscardHandler),
|
||||
Alerts: newQueue(),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open %s: %v", path, err)
|
||||
}
|
||||
|
||||
for client, want := range map[string]lookuptest.Network{
|
||||
testClient: germany,
|
||||
// An IPv6 client is its /64.
|
||||
"2001:db8:1:2::/64": northKorea,
|
||||
"198.51.100.7/32": {},
|
||||
} {
|
||||
prefix := netip.MustParsePrefix(client)
|
||||
|
||||
got := f.LookUp(prefix)
|
||||
if got != (Answer{
|
||||
Client: prefix, ASN: want.ASN, ASName: want.ASName, Country: want.Country,
|
||||
Answered: now,
|
||||
}) {
|
||||
t.Errorf("%s has the answer %+v, want %+v, answered %s", client, got, want, now)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecordThatCannotBeReadPlacesTheClientNowhere(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// The AS number is a number, where a string belongs. The writer writes
|
||||
// a record's fields in the order of their names, so as_name is read
|
||||
// before the AS number fails.
|
||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||
lookuptest.WriteRecords(t, path, map[string]mmdbtype.Map{
|
||||
testNetblock: {
|
||||
"asn": mmdbtype.Uint32(64496),
|
||||
"as_name": mmdbtype.String("Example Net"),
|
||||
"country_code": mmdbtype.String("DE"),
|
||||
},
|
||||
})
|
||||
|
||||
f := openFile(t, path, newQueue())
|
||||
|
||||
answer := f.LookUp(netip.MustParsePrefix(testClient))
|
||||
if answer.ASN != "" || answer.ASName != "" || answer.Country != "" {
|
||||
t.Errorf("%s is placed %+v, want nowhere", testClient, answer)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFileThatCannotBeReadIsAnError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
missing := filepath.Join(dir, "missing.mmdb")
|
||||
notDatabase := filepath.Join(dir, "not.mmdb")
|
||||
writeFile(t, notDatabase, "not a lookup database\n")
|
||||
|
||||
for path, want := range map[string]string{
|
||||
missing: "SWWAF_LOOKUP_DB_PATH cannot be read: open " + missing +
|
||||
": no such file or directory",
|
||||
notDatabase: "SWWAF_LOOKUP_DB_PATH " + notDatabase +
|
||||
" is not a .mmdb file: error opening database: invalid MaxMind DB file",
|
||||
} {
|
||||
_, err := OpenFile(FileParams{
|
||||
Path: path,
|
||||
Now: time.Now,
|
||||
ProcessLog: slog.New(slog.DiscardHandler),
|
||||
Alerts: newQueue(),
|
||||
})
|
||||
if err == nil || err.Error() != want {
|
||||
t.Errorf("opening %s failed with %v, want %s", path, err, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The tests below run readAfterChanges 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 readAfterChanges waits again, so that every
|
||||
// reading due by then is done. The test sends the changes itself, as the
|
||||
// watch of a directory cannot run in a bubble.
|
||||
|
||||
func TestReplacementCopiedOverTheFileInTwoPartsIsReadOnlyWhole(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "ipinfo_lite.mmdb")
|
||||
writeDatabase(t, path, "DE")
|
||||
|
||||
queue := newQueue()
|
||||
f := openFile(t, path, queue)
|
||||
changes := watch(t, f)
|
||||
|
||||
other := filepath.Join(dir, "replacement.mmdb")
|
||||
writeDatabase(t, other, "KP")
|
||||
|
||||
replacement, err := os.ReadFile(other) //nolint:gosec // a file the test wrote
|
||||
if err != nil {
|
||||
t.Fatalf("read %s: %v", other, err)
|
||||
}
|
||||
|
||||
// The file in use is overwritten in place, and keeps giving what
|
||||
// it gave. Its first part alone is not a .mmdb file.
|
||||
file, err := os.Create(path) //nolint:gosec // a file the test wrote
|
||||
if err != nil {
|
||||
t.Fatalf("create %s: %v", path, err)
|
||||
}
|
||||
|
||||
defer func() {
|
||||
_ = file.Close()
|
||||
}()
|
||||
|
||||
half := len(replacement) / 2
|
||||
write(t, file, replacement[:half])
|
||||
|
||||
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
||||
|
||||
time.Sleep(quietTime - time.Nanosecond)
|
||||
synctest.Wait()
|
||||
wantCountry(t, f, "DE")
|
||||
|
||||
// The second part starts the wait again.
|
||||
write(t, file, replacement[half:])
|
||||
|
||||
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
||||
|
||||
time.Sleep(quietTime - time.Nanosecond)
|
||||
synctest.Wait()
|
||||
wantCountry(t, f, "DE")
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
synctest.Wait()
|
||||
wantCountry(t, f, "KP")
|
||||
|
||||
if !f.LastRead().Equal(time.Now()) || f.ReadFailures() != 0 {
|
||||
t.Errorf("read at %s, with %d failures; want read now, with none",
|
||||
f.LastRead(), f.ReadFailures())
|
||||
}
|
||||
|
||||
wantAlerts(t, queue)
|
||||
})
|
||||
}
|
||||
|
||||
func TestReplacementThatCannotBeReadLeavesTheFileInUseWithOneAlert(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||
writeDatabase(t, path, "DE")
|
||||
|
||||
queue := newQueue()
|
||||
f := openFile(t, path, queue)
|
||||
read := f.LastRead()
|
||||
changes := watch(t, f)
|
||||
|
||||
writeFile(t, path, "not a lookup database\n")
|
||||
|
||||
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
||||
|
||||
// Long after, the replacement has been read once.
|
||||
time.Sleep(time.Hour)
|
||||
synctest.Wait()
|
||||
wantCountry(t, f, "DE")
|
||||
|
||||
if !f.LastRead().Equal(read) || f.ReadFailures() != 1 {
|
||||
t.Errorf("read at %s, with %d failures; want read at %s, with one",
|
||||
f.LastRead(), f.ReadFailures(), read)
|
||||
}
|
||||
|
||||
wantAlerts(t, queue, alerts.Alert{
|
||||
Time: read.Add(quietTime),
|
||||
Event: alerts.EventFileError,
|
||||
Reason: "the lookup database cannot be read, and the one read before stays in use",
|
||||
Detail: map[string]any{
|
||||
"file": path,
|
||||
"error": "SWWAF_LOOKUP_DB_PATH " + path + " is not a .mmdb file: " +
|
||||
"error opening database: invalid MaxMind DB file",
|
||||
},
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
func TestChangeOfAnotherFileInTheDirectoryIsNoReplacement(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "ipinfo_lite.mmdb")
|
||||
writeDatabase(t, path, "DE")
|
||||
|
||||
f := openFile(t, path, newQueue())
|
||||
changes := watch(t, f)
|
||||
|
||||
// The wait that starts with the watch ends with a reading.
|
||||
time.Sleep(quietTime)
|
||||
synctest.Wait()
|
||||
writeDatabase(t, path, "KP")
|
||||
|
||||
changes <- fsnotify.Event{Name: filepath.Join(dir, "other.mmdb"), Op: fsnotify.Create}
|
||||
|
||||
time.Sleep(quietTime)
|
||||
synctest.Wait()
|
||||
wantCountry(t, f, "DE")
|
||||
|
||||
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
||||
|
||||
time.Sleep(quietTime)
|
||||
synctest.Wait()
|
||||
wantCountry(t, f, "KP")
|
||||
})
|
||||
}
|
||||
|
||||
func TestReplacementSavedBeforeTheWatchStartsIsRead(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||
writeDatabase(t, path, "DE")
|
||||
|
||||
f := openFile(t, path, newQueue())
|
||||
|
||||
// Saved after OpenFile read the file, and before its directory was
|
||||
// watched, so that no change is seen for it.
|
||||
writeDatabase(t, path, "KP")
|
||||
watch(t, f)
|
||||
time.Sleep(quietTime)
|
||||
synctest.Wait()
|
||||
wantCountry(t, f, "KP")
|
||||
})
|
||||
}
|
||||
|
||||
// newQueue returns a queue of alerts for a webhook that is never sent
|
||||
// them, so that they wait in it for the test to look at.
|
||||
func newQueue() *alerts.Queue {
|
||||
return alerts.New(alerts.Params{
|
||||
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
||||
Events: alerts.Events(),
|
||||
Cooldown: 15 * time.Minute,
|
||||
Now: time.Now,
|
||||
})
|
||||
}
|
||||
|
||||
// writeDatabase writes a lookup database at path that places testNetblock
|
||||
// in country, and no other address.
|
||||
func writeDatabase(t *testing.T, path, country string) {
|
||||
t.Helper()
|
||||
|
||||
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
||||
testNetblock: {ASN: "AS64496", ASName: "Example Net", Country: country},
|
||||
})
|
||||
}
|
||||
|
||||
// openFile opens the lookup database at path, which raises its alerts to
|
||||
// queue.
|
||||
func openFile(t *testing.T, path string, queue *alerts.Queue) *File {
|
||||
t.Helper()
|
||||
|
||||
f, err := OpenFile(FileParams{
|
||||
Path: path,
|
||||
Now: time.Now,
|
||||
ProcessLog: slog.New(slog.DiscardHandler),
|
||||
Alerts: queue,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open %s: %v", path, err)
|
||||
}
|
||||
|
||||
return f
|
||||
}
|
||||
|
||||
// watch runs f's readAfterChanges until the test ends, and returns the
|
||||
// channel that sends it changes.
|
||||
func watch(t *testing.T, f *File) chan<- fsnotify.Event {
|
||||
t.Helper()
|
||||
|
||||
changes := make(chan fsnotify.Event)
|
||||
ctx, stop := context.WithCancel(t.Context())
|
||||
stopped := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
f.readAfterChanges(ctx, changes, nil)
|
||||
close(stopped)
|
||||
}()
|
||||
|
||||
t.Cleanup(func() {
|
||||
stop()
|
||||
<-stopped
|
||||
})
|
||||
|
||||
return changes
|
||||
}
|
||||
|
||||
// wantCountry checks the country f gives testClient.
|
||||
func wantCountry(t *testing.T, f *File, want string) {
|
||||
t.Helper()
|
||||
|
||||
got := f.LookUp(netip.MustParsePrefix(testClient)).Country
|
||||
if got != want {
|
||||
t.Errorf("%s is in %q, want %q", testClient, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// wantAlerts checks the alerts waiting in queue, and that it held none
|
||||
// back.
|
||||
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
|
||||
t.Helper()
|
||||
|
||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||
if len(waiting) != len(want) || (len(want) > 0 && !reflect.DeepEqual(waiting, want)) {
|
||||
t.Errorf("alerts waiting %+v, want %+v", waiting, want)
|
||||
}
|
||||
|
||||
if queue.Suppressed() != 0 {
|
||||
t.Errorf("%d alerts held back, want none", queue.Suppressed())
|
||||
}
|
||||
}
|
||||
|
||||
// writeFile writes content to the file at path.
|
||||
func writeFile(t *testing.T, path, content string) {
|
||||
t.Helper()
|
||||
|
||||
err := os.WriteFile(path, []byte(content), 0o600)
|
||||
if err != nil {
|
||||
t.Fatalf("write %s: %v", path, err)
|
||||
}
|
||||
}
|
||||
|
||||
// write writes data to the end of file.
|
||||
func write(t *testing.T, file *os.File, data []byte) {
|
||||
t.Helper()
|
||||
|
||||
_, err := file.Write(data)
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
// Package lookup looks up each client's AS number and country through
|
||||
// the GeoJS web service, and keeps the answers in memory, for at most
|
||||
// 100,000 clients and for 7 days each. The answers are written to
|
||||
// Package lookup looks up each client's AS number and country, through
|
||||
// the GeoJS web service or in the lookup database, the IPinfo Lite file
|
||||
// SWWAF_LOOKUP_DB_PATH names. GeoJS's answers are kept in memory, for at
|
||||
// most 100,000 clients and for 7 days each, and are written to
|
||||
// lookups.json and read from it by the state package.
|
||||
package lookup
|
||||
|
||||
@@ -113,11 +114,12 @@ type GeoJS struct {
|
||||
retryAt time.Time
|
||||
}
|
||||
|
||||
// Answer is what GeoJS said about a client, as lookups.json holds it: its
|
||||
// AS number, such as AS64496, and the AS's name, both "" when GeoJS knows
|
||||
// no AS number for it; its country, "" when GeoJS cannot place it; when
|
||||
// GeoJS said so, and when the answer was last used. The zero Answer is
|
||||
// that of a client with no answer.
|
||||
// Answer is what GeoJS or the lookup database said about a client: its AS
|
||||
// number, such as AS64496, and the AS's name, both "" when the source knows
|
||||
// no AS number for it; its country, "" when the source cannot place it;
|
||||
// when the source said so; and, for GeoJS's answers, which lookups.json
|
||||
// holds, when the answer was last used. The zero Answer is that of a
|
||||
// client with no answer.
|
||||
//
|
||||
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||
type Answer struct {
|
||||
@@ -174,7 +176,8 @@ func New(params Params) *GeoJS {
|
||||
// when it ends.
|
||||
//
|
||||
// GeoJS is asked about the client's first address, which is the client's
|
||||
// own address for IPv4, and an address in the same place for an IPv6 /64.
|
||||
// own address for IPv4, and an address in the same place for an IPv6
|
||||
// netblock.
|
||||
func (g *GeoJS) LookUp(ctx context.Context, client netip.Prefix) Answer {
|
||||
answer, asked := g.answerOrWait(ctx, client)
|
||||
if asked == nil {
|
||||
|
||||
@@ -119,6 +119,58 @@ func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestRequestWaitsAsLongAsTheTimeoutSaysAndGeoJSIsAbandonedAfterIt(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
// A timeout longer than the default second, and a GeoJS that does
|
||||
// not answer.
|
||||
const longerTimeout = 3 * time.Second
|
||||
|
||||
m := metrics.New(1, "app")
|
||||
g := lookup.New(lookup.Params{
|
||||
URL: lookup.URL,
|
||||
Timeout: longerTimeout,
|
||||
Wait: true,
|
||||
Now: time.Now,
|
||||
ProcessLog: slog.New(slog.DiscardHandler),
|
||||
Metrics: m,
|
||||
Alerts: alerts.New(alerts.Params{}),
|
||||
})
|
||||
g.SetTransport(&standIn{answers: hanging})
|
||||
|
||||
var (
|
||||
request sync.WaitGroup
|
||||
waited time.Duration
|
||||
)
|
||||
|
||||
request.Go(func() {
|
||||
began := time.Now()
|
||||
|
||||
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
|
||||
|
||||
waited = time.Since(began)
|
||||
})
|
||||
|
||||
// A moment before the timeout runs out, GeoJS is still being asked:
|
||||
// the request to it has not failed.
|
||||
time.Sleep(longerTimeout - time.Millisecond)
|
||||
synctest.Wait()
|
||||
wantFailures(t, m, 0)
|
||||
|
||||
// As it runs out, the client's request goes on, and the request to
|
||||
// GeoJS is abandoned, which counts as a failure.
|
||||
request.Wait()
|
||||
synctest.Wait()
|
||||
|
||||
if waited != longerTimeout {
|
||||
t.Errorf("waited %s for the answer, want %s", waited, longerTimeout)
|
||||
}
|
||||
|
||||
wantFailures(t, m, 1)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -767,6 +819,16 @@ func wantUnanswered(t *testing.T, m *metrics.Metrics, want float64) {
|
||||
}
|
||||
}
|
||||
|
||||
// wantFailures checks how many requests to GeoJS m counts as failed.
|
||||
func wantFailures(t *testing.T, m *metrics.Metrics, want float64) {
|
||||
t.Helper()
|
||||
|
||||
got := testutil.ToFloat64(m.GeoJSFailures)
|
||||
if got != want {
|
||||
t.Errorf("%v requests to GeoJS failed, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// waitForRequests waits until g has done all it can before time passes,
|
||||
// checks that GeoJS has had count requests, and returns the addresses each
|
||||
// asked about.
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
// Package lookuptest writes lookup databases, IPinfo Lite files in their
|
||||
// .mmdb form, for the tests of the packages that read them.
|
||||
package lookuptest
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"net"
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
"github.com/maxmind/mmdbwriter"
|
||||
"github.com/maxmind/mmdbwriter/mmdbtype"
|
||||
)
|
||||
|
||||
// fileMode is the mode of the files written: read and written by their
|
||||
// owner alone.
|
||||
const fileMode = 0o600
|
||||
|
||||
// Network is what a lookup database holds about a netblock, of the fields
|
||||
// smallwebwaf reads: its AS number, such as AS64496, the AS's name, and
|
||||
// its country, such as DE.
|
||||
type Network struct {
|
||||
ASN string
|
||||
ASName string
|
||||
Country string
|
||||
}
|
||||
|
||||
// Write writes a lookup database at path that places each netblock in
|
||||
// networks, such as 203.0.113.0/24, as its Network says, and no other
|
||||
// address.
|
||||
func Write(tb testing.TB, path string, networks map[string]Network) {
|
||||
tb.Helper()
|
||||
|
||||
records := make(map[string]mmdbtype.Map, len(networks))
|
||||
for netblock, network := range networks {
|
||||
records[netblock] = mmdbtype.Map{
|
||||
"asn": mmdbtype.String(network.ASN),
|
||||
"as_name": mmdbtype.String(network.ASName),
|
||||
"country_code": mmdbtype.String(network.Country),
|
||||
}
|
||||
}
|
||||
|
||||
WriteRecords(tb, path, records)
|
||||
}
|
||||
|
||||
// WriteRecords writes a lookup database at path that holds each record in
|
||||
// records for its netblock, and nothing for any other address.
|
||||
func WriteRecords(tb testing.TB, path string, records map[string]mmdbtype.Map) {
|
||||
tb.Helper()
|
||||
|
||||
tree, err := mmdbwriter.New(mmdbwriter.Options{
|
||||
DatabaseType: "ipinfo_lite",
|
||||
// The tests' clients are in the netblocks kept for documentation.
|
||||
IncludeReservedNetworks: true,
|
||||
})
|
||||
if err != nil {
|
||||
tb.Fatalf("new lookup database: %v", err)
|
||||
}
|
||||
|
||||
for netblock, record := range records {
|
||||
_, network, err := net.ParseCIDR(netblock)
|
||||
if err != nil {
|
||||
tb.Fatalf("netblock %q: %v", netblock, err)
|
||||
}
|
||||
|
||||
err = tree.Insert(network, record)
|
||||
if err != nil {
|
||||
tb.Fatalf("insert %s: %v", netblock, err)
|
||||
}
|
||||
}
|
||||
|
||||
var database bytes.Buffer
|
||||
|
||||
_, err = tree.WriteTo(&database)
|
||||
if err != nil {
|
||||
tb.Fatalf("write the lookup database: %v", err)
|
||||
}
|
||||
|
||||
err = os.WriteFile(path, database.Bytes(), fileMode)
|
||||
if err != nil {
|
||||
tb.Fatalf("write %s: %v", path, err)
|
||||
}
|
||||
}
|
||||
+140
-4
@@ -6,6 +6,7 @@ package metrics
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
@@ -13,8 +14,10 @@ import (
|
||||
"github.com/prometheus/client_golang/prometheus/promhttp"
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||
"sneak.berlin/go/smallwebwaf/internal/reputation"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||
)
|
||||
@@ -34,8 +37,10 @@ type Metrics struct {
|
||||
rateLimitHits *prometheus.CounterVec
|
||||
sizeAndTimeLimitHits *prometheus.CounterVec
|
||||
offences *prometheus.CounterVec
|
||||
// ruleMatches are made by AddRules.
|
||||
// ruleMatches are made by AddRules, and reputationHits by
|
||||
// AddReputation.
|
||||
ruleMatches *prometheus.CounterVec
|
||||
reputationHits *prometheus.CounterVec
|
||||
countries *busiest
|
||||
asns *busiest
|
||||
|
||||
@@ -89,8 +94,9 @@ func New(topN int, instanceName string) *Metrics {
|
||||
Help: "How long requests passed to the app took, from then to their end.",
|
||||
}),
|
||||
rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total",
|
||||
"Requests that broke a rate limit, by its window.",
|
||||
[]string{"window"}),
|
||||
"Requests that broke a rate limit or a byte limit, by its window and "+
|
||||
"its kind, requests or bytes.",
|
||||
[]string{"window", "kind"}),
|
||||
sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total",
|
||||
"Requests that passed a size or time limit, by its setting.",
|
||||
[]string{"limit"}),
|
||||
@@ -229,6 +235,102 @@ func (m *Metrics) AddRemoteLog(remote *remotelog.Sender) {
|
||||
)
|
||||
}
|
||||
|
||||
// AddLookupFile adds the metrics of the lookup database, read as the
|
||||
// metrics are asked for: when the file in use was read, which lastRead
|
||||
// returns, and the replacements of it that could not be read, which
|
||||
// readFailures returns. The lookup package's File, which has both, cannot
|
||||
// be named here: that package counts GeoJS's requests in these metrics.
|
||||
func (m *Metrics) AddLookupFile(lastRead func() time.Time, readFailures func() int) {
|
||||
m.registry.MustRegister(
|
||||
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
||||
Name: "smallwebwaf_lookup_database_last_read_timestamp_seconds",
|
||||
Help: "When the lookup database in use was read, in seconds since 1970.",
|
||||
}, func() float64 {
|
||||
return float64(lastRead().Unix())
|
||||
}),
|
||||
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||
Name: "smallwebwaf_lookup_database_read_failures_total",
|
||||
Help: "Replacements of the lookup database that could not be read.",
|
||||
}, func() float64 {
|
||||
return float64(readFailures())
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
// sourceLabel is the label of the reputation metrics: a list's URL, a
|
||||
// DNSBL zone, its key masked, or abuseipdb.
|
||||
const sourceLabel = "source"
|
||||
|
||||
// AddReputation adds the metrics of the lists fetched from URLs and of the
|
||||
// DNSBL zones, by source, each list's URL or each zone, its key masked as
|
||||
// config.MaskZoneKey masks it: the requests whose client a blocklist, a
|
||||
// zone's verdict or AbuseIPDB's score lists, which ReputationHit counts,
|
||||
// and, read from lists and dnsbl as the metrics are asked for, for a list,
|
||||
// the fetches that failed and when the copy in use was fetched, and for a
|
||||
// zone, the queries made and those that failed. It is called once, before
|
||||
// ReputationHit.
|
||||
func (m *Metrics) AddReputation(lists *reputation.Lists, dnsbl *reputation.DNSBL) {
|
||||
m.reputationHits = counterVec("smallwebwaf_reputation_hits_total",
|
||||
"Requests whose client a blocklist, a DNSBL zone or AbuseIPDB lists, by "+
|
||||
"the blocklist's URL, the zone, or abuseipdb.",
|
||||
[]string{sourceLabel})
|
||||
m.registry.MustRegister(m.reputationHits)
|
||||
|
||||
for _, zone := range dnsbl.Zones() {
|
||||
source := prometheus.Labels{sourceLabel: config.MaskZoneKey(zone)}
|
||||
|
||||
m.addReputationQueries(source, func() int { return dnsbl.Queries(zone) })
|
||||
m.addReputationFailures(source, func() int { return dnsbl.Failures(zone) })
|
||||
}
|
||||
|
||||
for _, listURL := range lists.URLs() {
|
||||
source := prometheus.Labels{sourceLabel: listURL}
|
||||
|
||||
m.addReputationFailures(source, func() int { return lists.Failures(listURL) })
|
||||
m.registry.MustRegister(
|
||||
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
||||
Name: "smallwebwaf_reputation_last_fetch_timestamp_seconds",
|
||||
Help: "When the copy of the list in use was fetched, in seconds since " +
|
||||
"1970, or 0 while there is none.",
|
||||
ConstLabels: source,
|
||||
}, func() float64 {
|
||||
fetched := lists.Fetched(listURL)
|
||||
if fetched.IsZero() {
|
||||
return 0
|
||||
}
|
||||
|
||||
return float64(fetched.Unix())
|
||||
}),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// AddAbuseIPDB adds the metrics of AbuseIPDB, with the source abuseipdb,
|
||||
// read from abuseIPDB as the metrics are asked for: the checks made, those
|
||||
// that failed, and how many checks the day's budget has left. It is
|
||||
// called once, after AddReputation, while SWWAF_ABUSEIPDB_KEY is set.
|
||||
func (m *Metrics) AddAbuseIPDB(abuseIPDB *reputation.AbuseIPDB) {
|
||||
source := prometheus.Labels{sourceLabel: reputation.AbuseIPDBSource}
|
||||
|
||||
m.addReputationQueries(source, abuseIPDB.Checked)
|
||||
m.addReputationFailures(source, abuseIPDB.Failures)
|
||||
m.registry.MustRegister(
|
||||
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
||||
Name: "smallwebwaf_reputation_daily_budget_remaining",
|
||||
Help: "Checks of the day's SWWAF_ABUSEIPDB_DAILY_BUDGET not yet spent.",
|
||||
ConstLabels: source,
|
||||
}, func() float64 {
|
||||
return float64(abuseIPDB.BudgetLeft())
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
// ReputationHit counts a request whose client source lists: a blocklist,
|
||||
// by its URL, a DNSBL zone, its key masked, or AbuseIPDB, abuseipdb.
|
||||
func (m *Metrics) ReputationHit(source string) {
|
||||
m.reputationHits.WithLabelValues(source).Inc()
|
||||
}
|
||||
|
||||
// AddAlerts adds the metrics of the alerts sent to each destination set,
|
||||
// read from queue as the metrics are asked for, by destination: the
|
||||
// alerts sent, the requests to the destination that failed, the alerts
|
||||
@@ -303,7 +405,15 @@ func (m *Metrics) RequestEnded(
|
||||
}
|
||||
|
||||
if line.LimitHit != "" {
|
||||
m.rateLimitHits.WithLabelValues(line.LimitHit).Inc()
|
||||
// The log line names a byte limit's window with _bytes after it.
|
||||
window, isBytes := strings.CutSuffix(line.LimitHit, "_bytes")
|
||||
|
||||
kind := ratelimit.KindRequests
|
||||
if isBytes {
|
||||
kind = ratelimit.KindBytes
|
||||
}
|
||||
|
||||
m.rateLimitHits.WithLabelValues(window, kind).Inc()
|
||||
}
|
||||
|
||||
if limit != "" {
|
||||
@@ -360,6 +470,32 @@ func (m *Metrics) StateFileEditSetAside(name string) {
|
||||
m.stateFileEditsSetAside.WithLabelValues(name).Inc()
|
||||
}
|
||||
|
||||
// addReputationQueries adds the counter of the queries to source, a DNSBL
|
||||
// zone, or of the checks of clients with AbuseIPDB, which count tells.
|
||||
func (m *Metrics) addReputationQueries(source prometheus.Labels, count func() int) {
|
||||
m.registry.MustRegister(prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||
Name: "smallwebwaf_reputation_queries_total",
|
||||
Help: "Queries to the DNSBL zone, or checks of clients with AbuseIPDB.",
|
||||
ConstLabels: source,
|
||||
}, func() float64 {
|
||||
return float64(count())
|
||||
}))
|
||||
}
|
||||
|
||||
// addReputationFailures adds the counter of the fetches of source, a
|
||||
// list, the queries to it, a DNSBL zone, or the checks with it, AbuseIPDB,
|
||||
// that failed, which count tells.
|
||||
func (m *Metrics) addReputationFailures(source prometheus.Labels, count func() int) {
|
||||
m.registry.MustRegister(prometheus.NewCounterFunc(prometheus.CounterOpts{
|
||||
Name: "smallwebwaf_reputation_failures_total",
|
||||
Help: "Fetches of the list, queries to the DNSBL zone, or checks with " +
|
||||
"AbuseIPDB, that failed.",
|
||||
ConstLabels: source,
|
||||
}, func() float64 {
|
||||
return float64(count())
|
||||
}))
|
||||
}
|
||||
|
||||
// statusClass returns the class of status, such as 2xx, or none when no
|
||||
// status was sent.
|
||||
func statusClass(status int) string {
|
||||
|
||||
@@ -282,7 +282,7 @@ func (rq *request) showClient() {
|
||||
|
||||
answer := clientAnswer{Bans: state.BanEntries(rq.h.ledger.Covering(addr))}
|
||||
|
||||
client, seen := rq.h.limiter.Client(clientGroup(addr))
|
||||
client, seen := rq.h.limiter.Client(rq.h.clientGroup(addr))
|
||||
if seen {
|
||||
answer.Client = &client
|
||||
}
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -48,7 +48,7 @@ func TestAdminEndpointsAreOffWhileTheTokenIsUnset(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
if after := server.Ledger.Snapshot(); !slices.Equal(after, before) {
|
||||
if after := server.Ledger.Snapshot(); !reflect.DeepEqual(after, before) {
|
||||
t.Errorf("the bans are now\n%+v\nwant them unchanged\n%+v", after, before)
|
||||
}
|
||||
}
|
||||
@@ -79,7 +79,7 @@ func TestAdminEndpointsNeedTheAdminToken(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
if after := server.Ledger.Snapshot(); !slices.Equal(after, before) {
|
||||
if after := server.Ledger.Snapshot(); !reflect.DeepEqual(after, before) {
|
||||
t.Errorf("%s %s without the token changed the bans to\n%+v\nfrom\n%+v",
|
||||
e.method, e.path, after, before)
|
||||
}
|
||||
|
||||
@@ -118,7 +118,8 @@ func TestObserveModeRaisesTheBanAlertsItWouldHave(t *testing.T) {
|
||||
line := s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
|
||||
|
||||
// No ban is made, and none made permanent.
|
||||
if held := server.Ledger.Snapshot(); len(held) != 1 || held[0] != attackBan ||
|
||||
held := server.Ledger.Snapshot()
|
||||
if len(held) != 1 || !reflect.DeepEqual(held[0], attackBan) ||
|
||||
line.BanExpires != requestlog.FormatTime(attackBan.Expires) {
|
||||
t.Errorf("the ledger holds %+v, and the log line gives %s, want the ban "+
|
||||
"for the attack alone, as it was", held, line.BanExpires)
|
||||
@@ -203,7 +204,16 @@ func startWithAlerts(
|
||||
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
|
||||
t.Helper()
|
||||
|
||||
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
|
||||
return startAppWithAlerts(t, func(http.ResponseWriter, *http.Request) {}, env)
|
||||
}
|
||||
|
||||
// startAppWithAlerts is startWithAlerts in front of the app handler.
|
||||
func startAppWithAlerts(
|
||||
t *testing.T, handler http.HandlerFunc, env map[string]string,
|
||||
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
|
||||
t.Helper()
|
||||
|
||||
app := startApp(t, handler)
|
||||
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
|
||||
settings := map[string]string{
|
||||
trustedProxies: trustLocalhost,
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
package proxy
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
)
|
||||
|
||||
func TestWithEveryAnomalyThresholdOffARequestIsNotCounted(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// A request from a client looked up through GeoJS, with every anomaly
|
||||
// threshold off. Its handler has neither GeoJS's answers nor the
|
||||
// anomaly counters, nor a clock, and the request no response: reading
|
||||
// any of them to count the request panics.
|
||||
rq := &request{
|
||||
h: &handler{config: &config.Config{LookupSource: "geojs"}},
|
||||
client: netip.MustParseAddr("203.0.113.9"),
|
||||
lookedUp: true,
|
||||
}
|
||||
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
t.Errorf("counting the request did work, with every threshold off: %v", r)
|
||||
}
|
||||
}()
|
||||
|
||||
rq.countAnomalies()
|
||||
}
|
||||
@@ -0,0 +1,377 @@
|
||||
package proxy_test
|
||||
|
||||
import (
|
||||
"maps"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"slices"
|
||||
"strconv"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/anomaly"
|
||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
)
|
||||
|
||||
// The anomaly thresholds: the prefix of a scope followed by the end of a
|
||||
// count.
|
||||
const (
|
||||
anomalyClient = "SWWAF_ANOMALY_CLIENT_"
|
||||
anomalyNet = "SWWAF_ANOMALY_NET_"
|
||||
anomalyASN = "SWWAF_ANOMALY_ASN_"
|
||||
anomalyTotal = "SWWAF_ANOMALY_TOTAL_"
|
||||
anomalyWatch = "SWWAF_WATCH_"
|
||||
|
||||
requestsPerMinute = "REQUESTS_PER_MINUTE"
|
||||
requestsPerHour = "REQUESTS_PER_HOUR"
|
||||
bytesPerMinute = "BYTES_PER_MINUTE"
|
||||
bytesPerHour = "BYTES_PER_HOUR"
|
||||
)
|
||||
|
||||
// The other anomaly settings.
|
||||
const (
|
||||
anomalyNetV4Prefix = "SWWAF_ANOMALY_NET_V4_PREFIX"
|
||||
anomalyNetV6Prefix = "SWWAF_ANOMALY_NET_V6_PREFIX"
|
||||
watchNets = "SWWAF_WATCH_NETS"
|
||||
)
|
||||
|
||||
const (
|
||||
// clientsNet is the netblock around client at the default length, and
|
||||
// office a named netblock of the same.
|
||||
clientsNet = "203.0.113.0/24"
|
||||
office = "office=" + clientsNet
|
||||
// aLot is a threshold no test reaches.
|
||||
aLot = "1000"
|
||||
// hour is the window an alert names for a threshold per hour.
|
||||
hour = "hour"
|
||||
)
|
||||
|
||||
func TestEachScopeAndWindowOverItsThresholdAlertsOncePerCooldown(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, scope := range []struct {
|
||||
prefix, scope string
|
||||
// netblock is the alert's, and counted what its reason names. extra
|
||||
// is what its detail gives besides what every anomaly alert's does.
|
||||
netblock netip.Prefix
|
||||
counted string
|
||||
extra map[string]any
|
||||
}{
|
||||
{
|
||||
anomalyClient, anomaly.ScopeClient, netip.MustParsePrefix(client + "/32"),
|
||||
"the client " + client + "/32", nil,
|
||||
},
|
||||
{
|
||||
anomalyNet, anomaly.ScopeNet, netip.MustParsePrefix(clientsNet),
|
||||
"the netblock " + clientsNet, nil,
|
||||
},
|
||||
{anomalyASN, anomaly.ScopeASN, netip.Prefix{}, asnDE, map[string]any{"asn": asnDE}},
|
||||
{anomalyTotal, anomaly.ScopeTotal, netip.Prefix{}, "the whole service", nil},
|
||||
{
|
||||
anomalyWatch, anomaly.ScopeWatch, netip.MustParsePrefix(clientsNet),
|
||||
"the named netblock office, " + clientsNet, map[string]any{"name": "office"},
|
||||
},
|
||||
} {
|
||||
for _, threshold := range []struct {
|
||||
end, kind, window string
|
||||
// value is the threshold, which the third upload of 100 bytes
|
||||
// takes the count over, to count.
|
||||
value int64
|
||||
count float64
|
||||
}{
|
||||
{requestsPerMinute, ratelimit.KindRequests, minute, 2, 3},
|
||||
{requestsPerHour, ratelimit.KindRequests, hour, 2, 3},
|
||||
{bytesPerMinute, ratelimit.KindBytes, minute, 250, 300},
|
||||
{bytesPerHour, ratelimit.KindBytes, hour, 250, 300},
|
||||
} {
|
||||
setting := scope.prefix + threshold.end
|
||||
value := strconv.FormatInt(threshold.value, 10)
|
||||
|
||||
t.Run(setting, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, clk, server, queue := startWithLookupsAndClock(t, map[string]string{
|
||||
setting: value, watchNets: office,
|
||||
})
|
||||
start := clk.Now()
|
||||
|
||||
// The third upload takes the count over the threshold, and the
|
||||
// fourth, within the cooldown, is held back. Each is passed to
|
||||
// the app.
|
||||
for range 4 {
|
||||
s.uploadFrom(client)
|
||||
}
|
||||
|
||||
detail := map[string]any{
|
||||
"scope": scope.scope, "window": threshold.window, "kind": threshold.kind,
|
||||
"count": threshold.count, "threshold": threshold.value,
|
||||
}
|
||||
maps.Copy(detail, scope.extra)
|
||||
|
||||
wantAlerts(t, queue, alerts.Alert{
|
||||
Instance: alertInstance,
|
||||
Time: start,
|
||||
Event: alerts.EventAnomaly,
|
||||
Client: netip.MustParseAddr(client),
|
||||
Netblock: scope.netblock,
|
||||
ASN: asnDE,
|
||||
ASName: asNameDE,
|
||||
Country: "DE",
|
||||
Reason: threshold.kind + " per " + threshold.window + " of " +
|
||||
scope.counted + " over the threshold of " + value,
|
||||
Detail: detail,
|
||||
})
|
||||
|
||||
wantAlertedAgainOnceTheCooldownHasRunOut(t, s, clk, queue)
|
||||
|
||||
if held := server.Ledger.Snapshot(); len(held) != 0 {
|
||||
t.Errorf("the ledger holds %+v, want no ban", held)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// wantAlertedAgainOnceTheCooldownHasRunOut checks that, once the cooldown
|
||||
// has run out after a first alert, which held back one repeat, the next
|
||||
// count over the threshold, at the latest three uploads from client on,
|
||||
// raises another alert, giving that repeat.
|
||||
func wantAlertedAgainOnceTheCooldownHasRunOut(
|
||||
t *testing.T, s *sender, clk *clock, queue *alerts.Queue,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
clk.advance(15 * time.Minute)
|
||||
|
||||
for range 3 {
|
||||
s.uploadFrom(client)
|
||||
}
|
||||
|
||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||
if len(waiting) != 2 || !waiting[1].Time.Equal(clk.Now()) ||
|
||||
waiting[1].SuppressedRepeats != 1 {
|
||||
t.Errorf("alerts wait %+v, want the first and another, with 1 repeat", waiting)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEveryRequestIsCountedWhateverIsDoneWithIt(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const (
|
||||
allowed = "192.0.2.7" // in SWWAF_ALLOW_NETS
|
||||
exempt = "192.0.2.10" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
|
||||
denied = "192.0.2.20" // in SWWAF_DENY_NETS
|
||||
)
|
||||
|
||||
s, _, _, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{
|
||||
anomalyClient + requestsPerMinute: "2",
|
||||
allowNets: allowed,
|
||||
rateLimitExemptNets: exempt,
|
||||
rateLimitExemptPaths: "/static/",
|
||||
denyNets: denied,
|
||||
})
|
||||
|
||||
// The third request of each takes its client's count over the threshold
|
||||
// of 2.
|
||||
for _, sent := range []struct {
|
||||
from, path string
|
||||
status int
|
||||
action string
|
||||
}{
|
||||
{allowed, "/", http.StatusOK, requestlog.ActionForward},
|
||||
{exempt, "/", http.StatusOK, requestlog.ActionForward},
|
||||
{client, "/static/app.js", http.StatusOK, requestlog.ActionForward},
|
||||
{denied, "/", http.StatusForbidden, requestlog.ActionDenied},
|
||||
} {
|
||||
for range 3 {
|
||||
s.request(sent.from, sent.path, sent.status, sent.action)
|
||||
}
|
||||
}
|
||||
|
||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||
|
||||
got := make([]string, 0, len(waiting))
|
||||
for _, alert := range waiting {
|
||||
got = append(got, alert.Client.String())
|
||||
}
|
||||
|
||||
if want := []string{allowed, exempt, client, denied}; !slices.Equal(got, want) {
|
||||
t.Errorf("alerts for the clients %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestThresholdsOffCountNothingAndAlertNothing(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// With every threshold off, nothing is counted.
|
||||
s, server, queue := startWithLookups(t, map[string]string{watchNets: office})
|
||||
|
||||
for range 5 {
|
||||
s.uploadFrom(client)
|
||||
}
|
||||
|
||||
if counters := server.Anomalies.Snapshot(); len(counters) != 0 {
|
||||
t.Errorf("counters %+v, want none", counters)
|
||||
}
|
||||
|
||||
wantAlerts(t, queue)
|
||||
|
||||
// With one set, its count alone is counted, in its scope alone.
|
||||
s, clk, server, queue := startWithLookupsAndClock(t, map[string]string{
|
||||
anomalyNet + requestsPerMinute: aLot, watchNets: office,
|
||||
})
|
||||
|
||||
for range 5 {
|
||||
s.uploadFrom(client)
|
||||
}
|
||||
|
||||
want := []anomaly.Counter{{
|
||||
Scope: anomaly.ScopeNet,
|
||||
Netblock: netip.MustParsePrefix(clientsNet),
|
||||
Minute: ratelimit.Buckets{Start: clk.Now(), Current: 5},
|
||||
}}
|
||||
if got := server.Anomalies.Snapshot(); !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("counters\n%+v\nwant\n%+v", got, want)
|
||||
}
|
||||
|
||||
wantAlerts(t, queue)
|
||||
}
|
||||
|
||||
func TestNetblockAroundAClientIsAsLongAsTheSettingsSay(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
env map[string]string
|
||||
// Each client of sent sends one request, and want gives the
|
||||
// netblocks they are counted in, each with its requests.
|
||||
sent []string
|
||||
want map[string]int64
|
||||
}{
|
||||
{
|
||||
"by default", nil,
|
||||
[]string{client, "203.0.113.200", "192.0.2.7", ipv6Client, "2001:db8:0:ffff::1"},
|
||||
map[string]int64{clientsNet: 2, "192.0.2.0/24": 1, "2001:db8::/48": 2},
|
||||
},
|
||||
{
|
||||
"as set", map[string]string{anomalyNetV4Prefix: "16", anomalyNetV6Prefix: "32"},
|
||||
[]string{client, "203.0.200.1", ipv6Client, "2001:db8:ffff::1"},
|
||||
map[string]int64{"203.0.0.0/16": 2, "2001:db8::/32": 2},
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := map[string]string{anomalyNet + requestsPerMinute: aLot}
|
||||
maps.Copy(env, tc.env)
|
||||
|
||||
s, _, server, _ := startAppWithAlerts(t, readAndAnswer, env)
|
||||
|
||||
for _, from := range tc.sent {
|
||||
s.get(from, http.StatusOK, requestlog.ActionForward)
|
||||
}
|
||||
|
||||
got := map[string]int64{}
|
||||
for _, counter := range server.Anomalies.Snapshot() {
|
||||
got[counter.Netblock.String()] = counter.Minute.Current
|
||||
}
|
||||
|
||||
if !maps.Equal(got, tc.want) {
|
||||
t.Errorf("requests by netblock %v, want %v", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientIsCountedForItsASNumberOnceTheLookupGivesOne(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, server, _ := startWithLookups(t, map[string]string{
|
||||
anomalyASN + requestsPerMinute: aLot,
|
||||
})
|
||||
|
||||
// The lookup database does not hold unplaced.
|
||||
for _, from := range []string{fromDE, fromDE, fromKP, noCountry, unplaced} {
|
||||
s.uploadFrom(from)
|
||||
}
|
||||
|
||||
got := map[string]int64{}
|
||||
for _, counter := range server.Anomalies.Snapshot() {
|
||||
got[counter.ASN] = counter.Minute.Current
|
||||
}
|
||||
|
||||
if want := map[string]int64{asnDE: 2, asnKP: 1, "AS64500": 1}; !maps.Equal(got, want) {
|
||||
t.Errorf("requests by AS number %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRequestCountsForTheASNumberGeoJSGivesBeforeItEnds(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// The stand-in for GeoJS answers only once released, which the app
|
||||
// does as it answers the request, and then waits until the answer is
|
||||
// kept.
|
||||
geojsURL, _, release := startHeldGeoJS(t)
|
||||
|
||||
var server atomic.Pointer[proxy.Server]
|
||||
|
||||
app := startApp(t, func(http.ResponseWriter, *http.Request) {
|
||||
release()
|
||||
waitUntil(func() bool {
|
||||
_, kept := server.Load().GeoJS.Kept(netip.MustParsePrefix(fromDE + "/32"))
|
||||
|
||||
return kept
|
||||
})
|
||||
})
|
||||
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
|
||||
addr, out, started := startProxyWithClock(t, app.URL, geojsURL, clk.Now,
|
||||
map[string]string{
|
||||
trustedProxies: trustLocalhost,
|
||||
lookupTimeout: "1h",
|
||||
anomalyASN + requestsPerMinute: aLot,
|
||||
})
|
||||
server.Store(started)
|
||||
|
||||
// The request went on without the answer, and is counted for the AS
|
||||
// number it gives.
|
||||
s := &sender{t: t, addr: addr, out: out}
|
||||
if line := s.get(fromDE, http.StatusOK, requestlog.ActionForward); line.ASN != "" {
|
||||
t.Errorf("log line has AS number %q, want none: the request waited", line.ASN)
|
||||
}
|
||||
|
||||
want := []anomaly.Counter{{
|
||||
Scope: anomaly.ScopeASN, ASN: asnDE,
|
||||
Minute: ratelimit.Buckets{Start: clk.Now(), Current: 1},
|
||||
}}
|
||||
if got := started.Anomalies.Snapshot(); !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("counters\n%+v\nwant\n%+v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEachNamedNetblockCountsTheClientsInIt(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _, server, _ := startAppWithAlerts(t, readAndAnswer, map[string]string{
|
||||
anomalyWatch + requestsPerMinute: aLot,
|
||||
watchNets: office + ",wide=203.0.0.0/16,other=198.51.100.0/25",
|
||||
})
|
||||
|
||||
// client is in office and in wide.
|
||||
for _, from := range []string{client, "203.0.200.1", "192.0.2.7"} {
|
||||
s.get(from, http.StatusOK, requestlog.ActionForward)
|
||||
}
|
||||
|
||||
got := map[string]int64{}
|
||||
for _, counter := range server.Anomalies.Snapshot() {
|
||||
got[counter.Name] = counter.Minute.Current
|
||||
}
|
||||
|
||||
if want := map[string]int64{"office": 1, "wide": 2}; !maps.Equal(got, want) {
|
||||
t.Errorf("requests by named netblock %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
+94
-23
@@ -6,6 +6,7 @@ import (
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||
)
|
||||
@@ -41,57 +42,126 @@ func (rq *request) banned(now time.Time) bool {
|
||||
|
||||
// limitBroken counts the request for the rate limits at now, notes the
|
||||
// client's counts for the log line, and reports whether the request takes
|
||||
// the client over a limit. In enforce mode such a request bans the
|
||||
// client's netblock, and sets the client's counters back to zero; in
|
||||
// observe mode it does neither, and raises the alert for the ban it would
|
||||
// have made, if that alert would be sent.
|
||||
// the client over a rate limit, as its limit percentage lowers it, which
|
||||
// breaks it.
|
||||
func (rq *request) limitBroken(now time.Time) bool {
|
||||
group := clientGroup(rq.client)
|
||||
|
||||
counts, hit, over := rq.h.limiter.Count(group, now)
|
||||
counts, hit, over := rq.h.limiter.Count(rq.h.clientGroup(rq.client), now,
|
||||
rq.limitPercent.percent)
|
||||
rq.line.Counts = counts
|
||||
|
||||
if !over {
|
||||
return false
|
||||
if over {
|
||||
rq.banForLimit(now, hit, rq.h.config.BanResponse)
|
||||
}
|
||||
|
||||
return over
|
||||
}
|
||||
|
||||
// countBytes counts the request's bytes, as countedBytes gives them, for
|
||||
// the byte limits, once its response has ended, and notes the client's
|
||||
// byte totals for the log line; its requests stay there as the rate limits
|
||||
// counted them. Only a request passed to the app has them counted, and
|
||||
// only one the rate limits counted; in observe mode, not one that enforce
|
||||
// mode would have refused. Bytes that take the client over a byte limit,
|
||||
// as its limit percentage for the byte limits lowers it, break it; the
|
||||
// response was passed on whole.
|
||||
func (rq *request) countBytes() {
|
||||
if !rq.counted || rq.line.WouldAction != "" {
|
||||
return
|
||||
}
|
||||
|
||||
now := rq.h.now()
|
||||
|
||||
counts, hit, over := rq.h.limiter.CountBytes(rq.h.clientGroup(rq.client), now,
|
||||
rq.countedBytes(), rq.bytesPercent.percent)
|
||||
rq.line.Counts.MinuteBytes = counts.MinuteBytes
|
||||
rq.line.Counts.HourBytes = counts.HourBytes
|
||||
rq.line.Counts.DayBytes = counts.DayBytes
|
||||
|
||||
if over {
|
||||
rq.banForLimit(now, hit, rq.out.status)
|
||||
}
|
||||
}
|
||||
|
||||
// countedBytes returns the request's bytes, once it has ended, as the
|
||||
// byte limits and the anomaly thresholds count them: the response's body
|
||||
// bytes, the request's, or both, as SWWAF_BYTES_COUNT says. For an
|
||||
// upgraded connection, such as a WebSocket, which has closed by then, what
|
||||
// it carried from the app counts with the response's and what it carried
|
||||
// from the client with the request's.
|
||||
func (rq *request) countedBytes() int64 {
|
||||
response, request := rq.out.bytes, rq.requestBytes()
|
||||
if rq.upgraded != nil {
|
||||
response += rq.upgraded.fromApp.Load()
|
||||
request += rq.upgraded.toApp.Load()
|
||||
}
|
||||
|
||||
switch rq.h.config.BytesCount {
|
||||
case "response":
|
||||
return response
|
||||
case "request":
|
||||
return request
|
||||
default: // both
|
||||
return response + request
|
||||
}
|
||||
}
|
||||
|
||||
// banForLimit bans the client's netblock at now for a broken limit, the
|
||||
// one hit names, and notes the offence for the log line. status is what
|
||||
// the client was sent, or is sent: SWWAF_BAN_RESPONSE for a request over
|
||||
// a rate limit, the app's answer for one whose bytes broke a byte limit.
|
||||
// The ban's notes give the client's limit percentage for that kind of
|
||||
// limit. The ban sets the client's counters back to zero. In observe mode
|
||||
// it makes no ban and sets nothing back, and raises the alert for the ban
|
||||
// it would have made, if that alert would be sent.
|
||||
func (rq *request) banForLimit(now time.Time, hit ratelimit.Hit, status int) {
|
||||
rq.line.LimitHit = hit.Window
|
||||
if hit.Kind == ratelimit.KindBytes {
|
||||
rq.line.LimitHit += "_bytes" // as counts names the byte totals
|
||||
}
|
||||
|
||||
rq.line.Offence = requestlog.OffenceLimit
|
||||
|
||||
netblock := rq.h.netblock(rq.client)
|
||||
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseLimit) {
|
||||
return true
|
||||
return
|
||||
}
|
||||
|
||||
notes := bans.Notes{
|
||||
ASN: rq.line.ASN,
|
||||
ASName: rq.line.ASName,
|
||||
Country: rq.line.Country,
|
||||
Kind: hit.Kind,
|
||||
Limit: hit.Limit,
|
||||
Window: hit.Window,
|
||||
Count: hit.Requests,
|
||||
Request: rq.noted(now),
|
||||
Count: hit.Count,
|
||||
Reputation: rq.reputation,
|
||||
Request: rq.noted(now, status),
|
||||
Requests: rq.netblockRequests(netblock),
|
||||
}
|
||||
|
||||
percent := rq.limitPercent
|
||||
if hit.Kind == ratelimit.KindBytes {
|
||||
percent = rq.bytesPercent
|
||||
}
|
||||
|
||||
notes.LimitPercent, notes.LimitPercentSetting = percent.logged()
|
||||
|
||||
if rq.h.config.Observe {
|
||||
ban, wouldBan := rq.h.ledger.WouldBanForLimit(netblock, now, notes)
|
||||
if wouldBan {
|
||||
rq.alertBan(ban)
|
||||
}
|
||||
|
||||
return true
|
||||
return
|
||||
}
|
||||
|
||||
ban, made := rq.h.ledger.BanForLimit(netblock, now, notes)
|
||||
rq.h.limiter.Reset(group)
|
||||
rq.h.limiter.Reset(rq.h.clientGroup(rq.client))
|
||||
rq.line.BanExpires = banExpires(ban)
|
||||
|
||||
if made {
|
||||
rq.alertBan(ban)
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// banForAttack bans the client's netblock at now for a clear sign of
|
||||
@@ -110,7 +180,8 @@ func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
|
||||
Country: rq.line.Country,
|
||||
RuleID: rule.ID,
|
||||
Target: rule.Target,
|
||||
Request: rq.noted(now),
|
||||
Reputation: rq.reputation,
|
||||
Request: rq.noted(now, rq.h.config.BanResponse),
|
||||
Requests: rq.netblockRequests(netblock),
|
||||
}
|
||||
|
||||
@@ -177,16 +248,16 @@ func (rq *request) alertBan(ban bans.Ban) {
|
||||
})
|
||||
}
|
||||
|
||||
// noted is the request, refused at now with SWWAF_BAN_RESPONSE, or in
|
||||
// observe mode as it would have been, as the notes of the ban it makes
|
||||
// keep it.
|
||||
func (rq *request) noted(now time.Time) bans.Request {
|
||||
// noted is the request, at now, with status, what the client was sent, or
|
||||
// in observe mode would have been, as the notes of the ban it makes keep
|
||||
// it.
|
||||
func (rq *request) noted(now time.Time, status int) bans.Request {
|
||||
return bans.Request{
|
||||
Time: now,
|
||||
Method: rq.in.Method,
|
||||
Host: rq.in.Host,
|
||||
Path: rq.in.URL.RequestURI(),
|
||||
Status: rq.h.config.BanResponse,
|
||||
Status: status,
|
||||
UserAgent: rq.in.UserAgent(),
|
||||
}
|
||||
}
|
||||
@@ -207,7 +278,7 @@ func (h *handler) netblock(client netip.Addr) netip.Prefix {
|
||||
return netip.PrefixFrom(addr, h.config.BanScopeV4Prefix).Masked()
|
||||
}
|
||||
|
||||
return clientGroup(addr)
|
||||
return h.clientGroup(addr)
|
||||
}
|
||||
|
||||
// banExpires is when ban ends, as the log line gives it: a time, or
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"maps"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"slices"
|
||||
"sync"
|
||||
"testing"
|
||||
@@ -165,9 +166,14 @@ func TestBanCoversTheClientsNetblock(t *testing.T) {
|
||||
[]string{otherClient, exempt}, []string{"203.0.112.9", allowed},
|
||||
},
|
||||
{
|
||||
"an IPv6 /64", nil, "2001:db8:5::1",
|
||||
"an IPv6 /64, by default", nil, "2001:db8:5::1",
|
||||
[]string{"2001:db8:5::ffff:1"}, []string{"2001:db8:5:1::1"},
|
||||
},
|
||||
{
|
||||
"the IPv6 netblock SWWAF_IPV6_GROUP_PREFIX sets",
|
||||
map[string]string{ipv6GroupPrefix: "48"}, "2001:db8:7::1",
|
||||
[]string{"2001:db8:7:ffff::1"}, []string{"2001:db8:8::1"},
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
@@ -284,6 +290,7 @@ func TestBanNotes(t *testing.T) {
|
||||
ASN: asnDE,
|
||||
ASName: asNameDE,
|
||||
Country: "DE",
|
||||
Kind: "requests",
|
||||
Limit: 1,
|
||||
Window: minute,
|
||||
Count: 2,
|
||||
@@ -306,7 +313,7 @@ func TestBanNotes(t *testing.T) {
|
||||
ledger := server.Ledger
|
||||
|
||||
got := ledger.Bans(netblock)
|
||||
if len(got) != 1 || got[0] != want {
|
||||
if len(got) != 1 || !reflect.DeepEqual(got[0], want) {
|
||||
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
package proxy
|
||||
|
||||
import (
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
)
|
||||
|
||||
// whole is the percentage of each limit a client gets when no biased
|
||||
// threshold lowers its limits.
|
||||
const whole = 100
|
||||
|
||||
// percentage is a client's limit percentage for the rate limits or for
|
||||
// the byte limits, as the biased thresholds give it, and the setting that
|
||||
// gave it: "" with whole when none lowers that kind of limit.
|
||||
type percentage struct {
|
||||
percent int64
|
||||
setting string
|
||||
}
|
||||
|
||||
// biasedThresholdsSet reports whether a biased threshold can lower a
|
||||
// client's limits: one of its lists is not empty,
|
||||
// SWWAF_UNKNOWN_LIMIT_PERCENT is below 100, or SWWAF_ASN_LIMIT_PERCENT_URL
|
||||
// is set. The client's lookup is then needed before its request goes on.
|
||||
func biasedThresholdsSet(cfg *config.Config) bool {
|
||||
return len(cfg.ASNLimitPercent) > 0 || len(cfg.CountryLimitPercent) > 0 ||
|
||||
len(cfg.ASNBytesPercent) > 0 || len(cfg.CountryBytesPercent) > 0 ||
|
||||
cfg.UnknownLimitPercent < whole || cfg.ASNLimitPercentURL != ""
|
||||
}
|
||||
|
||||
// limitPercentages returns the client's limit percentages, for the rate
|
||||
// limits and for the byte limits, by its AS number and country as looked
|
||||
// up, each "" when unknown, and the blocklists, DNSBL zones and AbuseIPDB
|
||||
// that list it. Each is the lowest of those the settings give it, the
|
||||
// first of them in the order below when several are lowest: the
|
||||
// percentage SWWAF_ASN_LIMIT_PERCENT gives its AS number, the one the file
|
||||
// SWWAF_ASN_LIMIT_PERCENT_URL names gives it, the one
|
||||
// SWWAF_COUNTRY_LIMIT_PERCENT gives its country, for a client without a
|
||||
// country, SWWAF_UNKNOWN_LIMIT_PERCENT, for a client a blocklist lists,
|
||||
// the percentage of SWWAF_BLOCKLIST_ACTION while it is limit, and for a
|
||||
// client a DNSBL zone's verdict lists, or whose AbuseIPDB score is a hit,
|
||||
// the percentage of SWWAF_REPUTATION_ACTION while it is limit. For the
|
||||
// byte limits, SWWAF_ASN_BYTES_PERCENT and SWWAF_COUNTRY_BYTES_PERCENT
|
||||
// take the place of the first three for an AS number or a country they
|
||||
// list.
|
||||
func (rq *request) limitPercentages() (percentage, percentage) {
|
||||
cfg := rq.h.config
|
||||
asn, country := rq.line.ASN, rq.line.Country
|
||||
|
||||
unknown := percentage{percent: whole}
|
||||
if country == "" {
|
||||
unknown = percentage{cfg.UnknownLimitPercent, "SWWAF_UNKNOWN_LIMIT_PERCENT"}
|
||||
}
|
||||
|
||||
fetched := percentage{percent: whole}
|
||||
if percent, listed := rq.h.lists.ASNLimitPercent(asn); listed {
|
||||
fetched = percentage{percent, "SWWAF_ASN_LIMIT_PERCENT_URL"}
|
||||
}
|
||||
|
||||
blocklisted := percentage{percent: whole}
|
||||
if rq.blocklisted && cfg.BlocklistAction == "limit" {
|
||||
blocklisted = percentage{cfg.BlocklistLimitPercent, "SWWAF_BLOCKLIST_ACTION"}
|
||||
}
|
||||
|
||||
reputationListed := percentage{percent: whole}
|
||||
if (rq.dnsblListed || rq.abuseIPDBHit) && cfg.ReputationAction == "limit" {
|
||||
reputationListed = percentage{cfg.ReputationLimitPercent, "SWWAF_REPUTATION_ACTION"}
|
||||
}
|
||||
|
||||
asnRequests := lowest(given(cfg.ASNLimitPercent, asn, "SWWAF_ASN_LIMIT_PERCENT"),
|
||||
fetched)
|
||||
countryRequests := given(cfg.CountryLimitPercent, country,
|
||||
"SWWAF_COUNTRY_LIMIT_PERCENT")
|
||||
|
||||
asnBytes, countryBytes := asnRequests, countryRequests
|
||||
if _, listed := cfg.ASNBytesPercent[asn]; listed {
|
||||
asnBytes = given(cfg.ASNBytesPercent, asn, "SWWAF_ASN_BYTES_PERCENT")
|
||||
}
|
||||
|
||||
if _, listed := cfg.CountryBytesPercent[country]; listed {
|
||||
countryBytes = given(cfg.CountryBytesPercent, country, "SWWAF_COUNTRY_BYTES_PERCENT")
|
||||
}
|
||||
|
||||
return lowest(asnRequests, countryRequests, unknown, blocklisted, reputationListed),
|
||||
lowest(asnBytes, countryBytes, unknown, blocklisted, reputationListed)
|
||||
}
|
||||
|
||||
// given returns the percentage percents, the setting named setting, gives
|
||||
// code, an AS number or a country, or whole when it does not list code.
|
||||
func given(percents map[string]int64, code, setting string) percentage {
|
||||
percent, listed := percents[code]
|
||||
if !listed {
|
||||
return percentage{percent: whole}
|
||||
}
|
||||
|
||||
return percentage{percent, setting}
|
||||
}
|
||||
|
||||
// lowest returns the lowest of percentages below whole, the first of them
|
||||
// when several are lowest, or whole when none is below it.
|
||||
func lowest(percentages ...percentage) percentage {
|
||||
low := percentage{percent: whole}
|
||||
|
||||
for _, p := range percentages {
|
||||
if p.percent < low.percent {
|
||||
low = p
|
||||
}
|
||||
}
|
||||
|
||||
return low
|
||||
}
|
||||
|
||||
// logged returns p as the log line and the notes of a ban give it: its
|
||||
// percent and setting, or nil and "" for whole, which they leave out.
|
||||
func (p percentage) logged() (*int64, string) {
|
||||
if p.percent == whole {
|
||||
return nil, ""
|
||||
}
|
||||
|
||||
return &p.percent, p.setting
|
||||
}
|
||||
@@ -0,0 +1,504 @@
|
||||
package proxy_test
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"maps"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"testing/synctest"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
|
||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
)
|
||||
|
||||
// The biased thresholds.
|
||||
const (
|
||||
asnLimitPercent = "SWWAF_ASN_LIMIT_PERCENT"
|
||||
countryLimitPercent = "SWWAF_COUNTRY_LIMIT_PERCENT"
|
||||
asnBytesPercent = "SWWAF_ASN_BYTES_PERCENT"
|
||||
countryBytesPercent = "SWWAF_COUNTRY_BYTES_PERCENT"
|
||||
unknownLimitPercent = "SWWAF_UNKNOWN_LIMIT_PERCENT"
|
||||
)
|
||||
|
||||
const (
|
||||
// asnDEHalf and countryDEHalf give fromDE's AS number and its country
|
||||
// half of every limit, and asnDEQuarter gives its AS number a quarter.
|
||||
asnDEHalf = asnDE + ":50"
|
||||
asnDEQuarter = asnDE + ":25"
|
||||
countryDEHalf = "de:50"
|
||||
// noCountry is in an AS of its own, AS64500, and in no country.
|
||||
noCountry = "192.0.2.80"
|
||||
// fourAMinute is the rate limit these tests set: half of it is 2
|
||||
// requests a minute, a quarter of it 1.
|
||||
fourAMinute = "4"
|
||||
// twoUploads is the byte limit these tests set: 199 bytes, which an
|
||||
// upload, a request with a body and its answer, 100 bytes, is within,
|
||||
// and half of which, 99 bytes, it is over.
|
||||
twoUploads = "199"
|
||||
// none is how percentText gives a percentage left out.
|
||||
none = "none"
|
||||
)
|
||||
|
||||
func TestEachBiasedThresholdLowersTheRateLimits(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct {
|
||||
setting, value, from string
|
||||
}{
|
||||
{asnLimitPercent, asnDEHalf, fromDE},
|
||||
{countryLimitPercent, countryDEHalf, fromDE},
|
||||
{unknownLimitPercent, "50", unplaced},
|
||||
} {
|
||||
t.Run(tc.setting, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _, _ := startWithLookups(t, map[string]string{
|
||||
rateLimitPerMinute: fourAMinute, tc.setting: tc.value,
|
||||
})
|
||||
|
||||
// Half of 4 requests a minute: the third breaks the limit.
|
||||
for _, sent := range []struct {
|
||||
status int
|
||||
action string
|
||||
}{
|
||||
{http.StatusOK, requestlog.ActionForward},
|
||||
{http.StatusOK, requestlog.ActionForward},
|
||||
{http.StatusForbidden, requestlog.ActionRateLimited},
|
||||
} {
|
||||
line := s.get(tc.from, sent.status, sent.action)
|
||||
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
|
||||
"50 from "+tc.setting)
|
||||
}
|
||||
|
||||
// fromKP, which no setting lists, has the whole limit.
|
||||
for range 3 {
|
||||
line := s.get(fromKP, http.StatusOK, requestlog.ActionForward)
|
||||
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
|
||||
none)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEachBiasedThresholdLowersTheByteLimits(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// The AS numbers and countries are given in either case.
|
||||
for _, tc := range []struct {
|
||||
setting, value, from string
|
||||
}{
|
||||
{asnLimitPercent, asnDEHalf, fromDE},
|
||||
{countryLimitPercent, "DE:50", fromDE},
|
||||
{unknownLimitPercent, "50", unplaced},
|
||||
{asnBytesPercent, "as64496:50", fromDE},
|
||||
{countryBytesPercent, countryDEHalf, fromDE},
|
||||
} {
|
||||
t.Run(tc.setting, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _, _ := startWithLookups(t, map[string]string{
|
||||
bytesLimitPerMinute: twoUploads, tc.setting: tc.value,
|
||||
})
|
||||
|
||||
// The upload's 100 bytes are over half of 199, 99.
|
||||
line := s.uploadFrom(tc.from)
|
||||
if line.LimitHit != minuteBytes {
|
||||
t.Errorf("log line has limit_hit %q, want %s", line.LimitHit, minuteBytes)
|
||||
}
|
||||
|
||||
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
|
||||
"50 from "+tc.setting)
|
||||
|
||||
// fromKP, which no setting lists, has the whole limit.
|
||||
line = s.uploadFrom(fromKP)
|
||||
if line.LimitHit != "" {
|
||||
t.Errorf("log line for %s has limit_hit %q, want none", fromKP, line.LimitHit)
|
||||
}
|
||||
|
||||
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting, none)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBytesPercentSettingsTakeThePlaceOfTheOthersForByteLimits(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
env map[string]string
|
||||
// limitPercent and bytesPercent are the log line's, as percentText
|
||||
// gives them, and limitHit is its limit_hit.
|
||||
limitPercent, bytesPercent, limitHit string
|
||||
}{
|
||||
{
|
||||
"lowering the byte limits alone",
|
||||
map[string]string{asnBytesPercent: asnDEHalf},
|
||||
none, "50 from " + asnBytesPercent, minuteBytes,
|
||||
},
|
||||
{
|
||||
"lowering the byte limits alone, by country",
|
||||
map[string]string{countryBytesPercent: countryDEHalf},
|
||||
none, "50 from " + countryBytesPercent, minuteBytes,
|
||||
},
|
||||
{
|
||||
"raising the byte limits back",
|
||||
map[string]string{asnLimitPercent: asnDEHalf, asnBytesPercent: asnDE + ":100"},
|
||||
"50 from " + asnLimitPercent, none, "",
|
||||
},
|
||||
{
|
||||
"raising the byte limits back, by country",
|
||||
map[string]string{countryLimitPercent: countryDEHalf, countryBytesPercent: "de:100"},
|
||||
"50 from " + countryLimitPercent, none, "",
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := map[string]string{bytesLimitPerMinute: twoUploads}
|
||||
maps.Copy(env, tc.env)
|
||||
s, _, _ := startWithLookups(t, env)
|
||||
|
||||
// The upload's 100 bytes are over 99, half of 199, and within 199.
|
||||
line := s.uploadFrom(fromDE)
|
||||
if line.LimitHit != tc.limitHit {
|
||||
t.Errorf("log line has limit_hit %q, want %q", line.LimitHit, tc.limitHit)
|
||||
}
|
||||
|
||||
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
|
||||
tc.limitPercent)
|
||||
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
|
||||
tc.bytesPercent)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestZeroPercentIsAZeroAllowance(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _, _ := startWithLookups(t, map[string]string{asnLimitPercent: asnDE + ":0"})
|
||||
|
||||
// The first request breaks the limit, and bans the client; the log line
|
||||
// gives the 0.
|
||||
line := s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||
if line.fields["limit_percent"] != float64(0) ||
|
||||
line.fields["limit_percent_setting"] != asnLimitPercent {
|
||||
t.Errorf("log line has limit_percent %v from %v, want 0 from %s",
|
||||
line.fields["limit_percent"], line.fields["limit_percent_setting"],
|
||||
asnLimitPercent)
|
||||
}
|
||||
|
||||
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
|
||||
}
|
||||
|
||||
func TestLowestPercentageApplies(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
env map[string]string
|
||||
from string
|
||||
// want is the log line's limit_percent, as percentText gives it.
|
||||
want string
|
||||
}{
|
||||
{
|
||||
"the country's",
|
||||
map[string]string{asnLimitPercent: asnDEHalf, countryLimitPercent: "de:25"},
|
||||
fromDE, "25 from " + countryLimitPercent,
|
||||
},
|
||||
{
|
||||
"the AS number's",
|
||||
map[string]string{asnLimitPercent: asnDEQuarter, countryLimitPercent: countryDEHalf},
|
||||
fromDE, "25 from " + asnLimitPercent,
|
||||
},
|
||||
{
|
||||
"the AS number's, the first of two alike",
|
||||
map[string]string{asnLimitPercent: asnDEQuarter, countryLimitPercent: "de:25"},
|
||||
fromDE, "25 from " + asnLimitPercent,
|
||||
},
|
||||
{
|
||||
"that for a client without a country",
|
||||
map[string]string{asnLimitPercent: "AS64500:50", unknownLimitPercent: "25"},
|
||||
noCountry, "25 from " + unknownLimitPercent,
|
||||
},
|
||||
{
|
||||
// SWWAF_UNKNOWN_LIMIT_PERCENT is left at its default, 100.
|
||||
"the AS number's, for a client without a country",
|
||||
map[string]string{asnLimitPercent: "AS64500:25"},
|
||||
noCountry, "25 from " + asnLimitPercent,
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
env := map[string]string{rateLimitPerMinute: fourAMinute}
|
||||
maps.Copy(env, tc.env)
|
||||
s, _, _ := startWithLookups(t, env)
|
||||
|
||||
// A quarter of 4 requests a minute: the second breaks the limit.
|
||||
s.get(tc.from, http.StatusOK, requestlog.ActionForward)
|
||||
|
||||
line := s.get(tc.from, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
|
||||
tc.want)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnknownLimitPercentGivesEveryClientWithoutACountryItsPercentage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _, _ := startWithLookups(t, map[string]string{
|
||||
rateLimitPerMinute: fourAMinute, unknownLimitPercent: "50",
|
||||
})
|
||||
|
||||
// One the lookup database does not hold, and one on a private address,
|
||||
// which is never looked up: the third request of each breaks half of 4.
|
||||
for _, from := range []string{unplaced, "10.0.0.8"} {
|
||||
s.get(from, http.StatusOK, requestlog.ActionForward)
|
||||
s.get(from, http.StatusOK, requestlog.ActionForward)
|
||||
s.get(from, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||
}
|
||||
|
||||
// One in a country has the whole limit.
|
||||
for range 3 {
|
||||
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientWithoutAnAnswerInTimeHasTheUnknownLimitPercent(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// In a synctest bubble, as TestRequestWaitsAsLongAsTheLookupTimeoutSays
|
||||
// says, with a GeoJS that never answers.
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
|
||||
time.Now, map[string]string{unknownLimitPercent: "0"})
|
||||
|
||||
// Once the second the request waits for its answer is up, the client
|
||||
// counts as without a country, and its zero allowance refuses the
|
||||
// request before it reaches the app.
|
||||
serveFromDE(t, server, http.MethodGet, http.NoBody)
|
||||
|
||||
line := out.requestLine(t)
|
||||
wantLine(t, line, http.StatusForbidden, requestlog.ActionRateLimited)
|
||||
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
|
||||
"0 from "+unknownLimitPercent)
|
||||
})
|
||||
}
|
||||
|
||||
func TestRequestWaitsForItsLookupWhileABiasedThresholdIsSet(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const timeout = 3 * time.Second
|
||||
|
||||
for _, tc := range []struct {
|
||||
setting, value string
|
||||
waits bool
|
||||
}{
|
||||
{asnLimitPercent, asnDEHalf, true},
|
||||
{countryLimitPercent, countryDEHalf, true},
|
||||
{asnBytesPercent, asnDEHalf, true},
|
||||
{countryBytesPercent, countryDEHalf, true},
|
||||
{unknownLimitPercent, "99", true},
|
||||
{asnLimitPercentURL, asnURL, true},
|
||||
// At 100, its default, it lowers no limit.
|
||||
{unknownLimitPercent, "100", false},
|
||||
} {
|
||||
t.Run(tc.setting+"="+tc.value, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// In a synctest bubble, as TestRequestWaitsAsLongAsTheLookupTimeoutSays
|
||||
// says, with a GeoJS that never answers.
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
// The request's body is over SWWAF_REQUEST_MAX_BYTES, so that it
|
||||
// is refused after the checks, and never reaches the app.
|
||||
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
|
||||
time.Now, map[string]string{
|
||||
lookupTimeout: timeout.String(), requestMaxBytes: "1",
|
||||
tc.setting: tc.value,
|
||||
})
|
||||
began := time.Now()
|
||||
|
||||
serveFromDE(t, server, http.MethodPost, strings.NewReader("ab"))
|
||||
|
||||
want := time.Duration(0)
|
||||
if tc.waits {
|
||||
want = timeout
|
||||
}
|
||||
|
||||
if waited := time.Since(began); waited != want {
|
||||
t.Errorf("the request waited %s for its answer, want %s", waited, want)
|
||||
}
|
||||
|
||||
wantLine(t, out.requestLine(t), http.StatusRequestEntityTooLarge,
|
||||
requestlog.ActionTooLarge)
|
||||
|
||||
// The bubble's clock stops once this function returns, so the
|
||||
// request to GeoJS, which a request that did not wait leaves
|
||||
// under way, has to be abandoned before then.
|
||||
time.Sleep(timeout)
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBanForALoweredLimitGivesThePercentageInItsNotesAndItsAlert(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
env map[string]string
|
||||
// before is how many uploads come before the one that breaks a
|
||||
// limit, which is answered with status and logged with action.
|
||||
before int
|
||||
status int
|
||||
action string
|
||||
// reason and want are the ban's reason, and its notes' limit
|
||||
// percentage, as percentText gives it.
|
||||
reason, want string
|
||||
}{
|
||||
{
|
||||
// A quarter of 12 requests a minute is 3: the fourth breaks it.
|
||||
"a rate limit",
|
||||
map[string]string{rateLimitPerMinute: "12", asnLimitPercent: asnDEQuarter},
|
||||
3, http.StatusForbidden, requestlog.ActionRateLimited,
|
||||
"requests per minute over the limit of 3", "25 from " + asnLimitPercent,
|
||||
},
|
||||
{
|
||||
// The byte limits' percentage, not the rate limits'.
|
||||
"a byte limit",
|
||||
map[string]string{
|
||||
bytesLimitPerMinute: twoUploads, asnLimitPercent: asnDEQuarter,
|
||||
asnBytesPercent: asnDEHalf,
|
||||
},
|
||||
0, http.StatusOK, requestlog.ActionForward,
|
||||
"bytes per minute over the limit of 99", "50 from " + asnBytesPercent,
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, server, queue := startWithLookups(t, tc.env)
|
||||
|
||||
for range tc.before {
|
||||
s.uploadFrom(fromDE)
|
||||
}
|
||||
|
||||
s.requestWithBody(http.MethodPost, fromDE, "/", uploadHeader, uploadBody,
|
||||
tc.status, tc.action)
|
||||
|
||||
held := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32"))
|
||||
if len(held) != 1 {
|
||||
t.Fatalf("bans %+v, want one", held)
|
||||
}
|
||||
|
||||
notes := held[0].Notes
|
||||
if held[0].Reason != tc.reason {
|
||||
t.Errorf("the ban's reason is %q, want %q", held[0].Reason, tc.reason)
|
||||
}
|
||||
|
||||
wantPercent(t, "the notes' limit_percent", notes.LimitPercent,
|
||||
notes.LimitPercentSetting, tc.want)
|
||||
|
||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||
if len(waiting) != 1 {
|
||||
t.Fatalf("%d alerts wait, want the ban's alone: %+v", len(waiting), waiting)
|
||||
}
|
||||
|
||||
alerted, _ := waiting[0].Detail["notes"].(bans.Notes)
|
||||
wantPercent(t, "the alert's notes' limit_percent", alerted.LimitPercent,
|
||||
alerted.LimitPercentSetting, tc.want)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// startWithLookups is startWithLookupsAndClock for a test that needs no
|
||||
// clock.
|
||||
func startWithLookups(
|
||||
t *testing.T, env map[string]string,
|
||||
) (*sender, *proxy.Server, *alerts.Queue) {
|
||||
t.Helper()
|
||||
|
||||
s, _, server, queue := startWithLookupsAndClock(t, env)
|
||||
|
||||
return s, server, queue
|
||||
}
|
||||
|
||||
// startWithLookupsAndClock is startAppWithAlerts in front of
|
||||
// readAndAnswer, with the settings in env on top of clients looked up in a
|
||||
// lookup database, which places fromDE and fromKP in the AS numbers and
|
||||
// countries the stand-in for GeoJS gives them, noCountry in AS64500 and no
|
||||
// country, and no other address.
|
||||
func startWithLookupsAndClock(
|
||||
t *testing.T, env map[string]string,
|
||||
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
|
||||
t.Helper()
|
||||
|
||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
||||
fromDE + "/32": {ASN: asnDE, ASName: asNameDE, Country: "DE"},
|
||||
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
|
||||
noCountry + "/32": {ASN: "AS64500", ASName: "Nowhere Net"},
|
||||
})
|
||||
|
||||
settings := map[string]string{lookupSource: fileSource, lookupDBPath: path}
|
||||
maps.Copy(settings, env)
|
||||
|
||||
return startAppWithAlerts(t, readAndAnswer, settings)
|
||||
}
|
||||
|
||||
// uploadFrom is upload from the client at from.
|
||||
func (s *sender) uploadFrom(from string) logLine {
|
||||
s.t.Helper()
|
||||
|
||||
line, _ := s.requestWithBody(http.MethodPost, from, "/", uploadHeader, uploadBody,
|
||||
http.StatusOK, requestlog.ActionForward)
|
||||
|
||||
return line
|
||||
}
|
||||
|
||||
// serveFromDE hands a request from fromDE with method and body straight to
|
||||
// server's handler, without the network, and returns once it is answered.
|
||||
func serveFromDE(t *testing.T, server *proxy.Server, method string, body io.Reader) {
|
||||
t.Helper()
|
||||
|
||||
req := httptest.NewRequestWithContext(t.Context(), method, "/", body)
|
||||
req.RemoteAddr = net.JoinHostPort(fromDE, "1234")
|
||||
|
||||
server.Handler.ServeHTTP(httptest.NewRecorder(), req)
|
||||
}
|
||||
|
||||
// wantPercent checks a limit percentage that a log line or a ban's notes
|
||||
// give, what, and the setting that gave it, against want, as percentText
|
||||
// gives them.
|
||||
func wantPercent(t *testing.T, what string, percent *int64, setting, want string) {
|
||||
t.Helper()
|
||||
|
||||
if got := percentText(percent, setting); got != want {
|
||||
t.Errorf("%s is %s, want %s", what, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// percentText gives a limit percentage and the setting that gave it as
|
||||
// text, such as "50 from SWWAF_ASN_LIMIT_PERCENT", or none when both are
|
||||
// left out.
|
||||
func percentText(percent *int64, setting string) string {
|
||||
switch {
|
||||
case percent == nil && setting == "":
|
||||
return none
|
||||
case percent == nil:
|
||||
return "none from " + setting
|
||||
default:
|
||||
return fmt.Sprintf("%d from %s", *percent, setting)
|
||||
}
|
||||
}
|
||||
@@ -103,6 +103,48 @@ func (b *responseBody) Close() error {
|
||||
return b.body.Close()
|
||||
}
|
||||
|
||||
// upgradedConn is the connection to the app once the app has switched
|
||||
// protocols, as for a WebSocket. ReverseProxy writes to it what the client
|
||||
// sends and reads from it what the app sends, on goroutines of its own,
|
||||
// until the connection closes; it counts the bytes each way, for the byte
|
||||
// limits.
|
||||
type upgradedConn struct {
|
||||
io.ReadWriteCloser
|
||||
|
||||
// fromApp is how many bytes the app has sent, and toApp how many the
|
||||
// client has.
|
||||
fromApp atomic.Int64
|
||||
toApp atomic.Int64
|
||||
}
|
||||
|
||||
// Read reads what the app sends.
|
||||
func (c *upgradedConn) Read(p []byte) (int, error) {
|
||||
n, err := c.ReadWriteCloser.Read(p)
|
||||
c.fromApp.Add(int64(n))
|
||||
|
||||
return n, err
|
||||
}
|
||||
|
||||
// Write sends the app what the client sent.
|
||||
func (c *upgradedConn) Write(p []byte) (int, error) {
|
||||
n, err := c.ReadWriteCloser.Write(p)
|
||||
c.toApp.Add(int64(n))
|
||||
|
||||
return n, err
|
||||
}
|
||||
|
||||
// CloseWrite tells the app that the client sends no more, while what the
|
||||
// app sends still passes. ReverseProxy calls it once the client has
|
||||
// stopped sending, and closes the connection there if it is not supported.
|
||||
func (c *upgradedConn) CloseWrite() error {
|
||||
conn, ok := c.ReadWriteCloser.(interface{ CloseWrite() error })
|
||||
if !ok {
|
||||
return http.ErrNotSupported
|
||||
}
|
||||
|
||||
return conn.CloseWrite()
|
||||
}
|
||||
|
||||
// limitBody returns body, cut off with an *http.MaxBytesError after
|
||||
// maxBytes, or unchanged if maxBytes is zero, which is off.
|
||||
func limitBody(body io.ReadCloser, maxBytes int64) io.ReadCloser {
|
||||
|
||||
@@ -0,0 +1,584 @@
|
||||
package proxy_test
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
)
|
||||
|
||||
// The byte limit settings.
|
||||
const (
|
||||
bytesLimitPerMinute = "SWWAF_BYTES_LIMIT_PER_MINUTE"
|
||||
bytesLimitPerHour = "SWWAF_BYTES_LIMIT_PER_HOUR"
|
||||
bytesLimitPerDay = "SWWAF_BYTES_LIMIT_PER_DAY"
|
||||
bytesCount = "SWWAF_BYTES_COUNT"
|
||||
)
|
||||
|
||||
// The values of SWWAF_BYTES_COUNT.
|
||||
const (
|
||||
countResponse = "response"
|
||||
countRequest = "request"
|
||||
countBoth = "both"
|
||||
)
|
||||
|
||||
const (
|
||||
// bodyBytes is the size of the body of each request these tests send
|
||||
// with one, and answerBytes that of each answer of the app.
|
||||
bodyBytes = 30
|
||||
answerBytes = 70
|
||||
// byteLimit is the byte limit these tests set, as a setting: a request
|
||||
// with a body and its answer, 100 bytes, go over it.
|
||||
byteLimit = "99"
|
||||
// minuteBytes is limit_hit for SWWAF_BYTES_LIMIT_PER_MINUTE.
|
||||
minuteBytes = "minute_bytes"
|
||||
)
|
||||
|
||||
func TestEachByteLimitBansOnceTheResponseHasEnded(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const scraper = "192.0.2.200"
|
||||
|
||||
for _, tc := range []struct {
|
||||
setting, window string
|
||||
// apart is the time between the two requests, which the window
|
||||
// still covers.
|
||||
apart time.Duration
|
||||
}{
|
||||
{bytesLimitPerMinute, minute, 0},
|
||||
{bytesLimitPerHour, "hour", 2 * time.Minute},
|
||||
{bytesLimitPerDay, "day", 2 * time.Hour},
|
||||
} {
|
||||
t.Run(tc.setting, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, clk := startWithAnswers(t, map[string]string{
|
||||
tc.setting: byteLimit, metricsToken: token,
|
||||
})
|
||||
|
||||
// 70 bytes are within the limit of 99.
|
||||
line, _ := s.download()
|
||||
if line.LimitHit != "" || line.Offence != "" {
|
||||
t.Errorf("log line has limit_hit %q and offence %q, want neither",
|
||||
line.LimitHit, line.Offence)
|
||||
}
|
||||
|
||||
// 140 bytes are over it. The response is passed on whole, and
|
||||
// then bans the client for an hour.
|
||||
clk.advance(tc.apart)
|
||||
expires := requestlog.FormatTime(clk.Now().Add(time.Hour))
|
||||
|
||||
line, got := s.download()
|
||||
if got.err != nil || len(got.body) != answerBytes ||
|
||||
line.ResponseBytes != answerBytes {
|
||||
t.Errorf("got %d bytes (%v), and the log line has response_bytes %d, "+
|
||||
"want %d", len(got.body), got.err, line.ResponseBytes, answerBytes)
|
||||
}
|
||||
|
||||
if line.LimitHit != tc.window+"_bytes" || line.Offence != requestlog.OffenceLimit ||
|
||||
line.BanExpires != expires {
|
||||
t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+
|
||||
"want %s_bytes, limit and %s", line.LimitHit, line.Offence,
|
||||
line.BanExpires, tc.window, expires)
|
||||
}
|
||||
|
||||
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||
|
||||
wantMetric(t, s.scrape(scraper), `smallwebwaf_rate_limit_hits_total{`+
|
||||
`instance="`+alertInstance+`",kind="bytes",window="`+tc.window+`"}`, 1)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResponseOverAByteLimitByItselfIsPassedOnWhole(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _ := startWithAnswers(t, map[string]string{bytesLimitPerMinute: "50"})
|
||||
|
||||
// The answer's 70 bytes are over the limit of 50 on their own.
|
||||
line, got := s.download()
|
||||
if got.err != nil || len(got.body) != answerBytes || line.LimitHit != minuteBytes {
|
||||
t.Errorf("got %d bytes (%v), and the log line has limit_hit %q, want %d and %s",
|
||||
len(got.body), got.err, line.LimitHit, answerBytes, minuteBytes)
|
||||
}
|
||||
|
||||
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||
}
|
||||
|
||||
func TestBytesOfAnAnswerThatBreaksOffAreCounted(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, clk, _, _ := startAppWithAlerts(t, breakOff, map[string]string{
|
||||
bytesLimitPerMinute: "50",
|
||||
})
|
||||
expires := requestlog.FormatTime(clk.Now().Add(time.Hour))
|
||||
|
||||
// The 70 bytes passed on before the app broke off are over the limit of
|
||||
// 50, and ban the client for an hour.
|
||||
line, got := s.requestWithBody(http.MethodGet, client, "/", "", "", http.StatusOK,
|
||||
requestlog.ActionUpstreamError)
|
||||
if len(got.body) != answerBytes || line.LimitHit != minuteBytes ||
|
||||
line.BanExpires != expires {
|
||||
t.Errorf("got %d bytes, and the log line has limit_hit %q and ban_expires %q, "+
|
||||
"want %d, %s and %s", len(got.body), line.LimitHit, line.BanExpires,
|
||||
answerBytes, minuteBytes, expires)
|
||||
}
|
||||
|
||||
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||
}
|
||||
|
||||
func TestWebSocketBytesAreCountedOnceItCloses(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct {
|
||||
setting string
|
||||
counted float64
|
||||
}{
|
||||
{countResponse, answerBytes},
|
||||
{countRequest, bodyBytes},
|
||||
{countBoth, bodyBytes + answerBytes},
|
||||
} {
|
||||
t.Run(tc.setting, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _, _, _ := startAppWithAlerts(t, answerAfterUpgrade, map[string]string{
|
||||
bytesLimitPerMinute: "29", bytesCount: tc.setting,
|
||||
})
|
||||
|
||||
// The client sends 30 bytes and the app 70, each over the limit
|
||||
// of 29, which bans the client once the WebSocket has closed.
|
||||
line := s.webSocket()
|
||||
if line.LimitHit != minuteBytes || line.Counts.MinuteBytes != tc.counted {
|
||||
t.Errorf("log line has limit_hit %q and minute_bytes %v, want %s and %v",
|
||||
line.LimitHit, line.Counts.MinuteBytes, minuteBytes, tc.counted)
|
||||
}
|
||||
|
||||
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebSocketPassesTheAnswerAfterTheClientStopsSending(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
app := startApp(t, echoOnceTheClientStops)
|
||||
addr, out := startProxy(t, app.URL,
|
||||
map[string]string{trustedProxies: trustLocalhost})
|
||||
s := &sender{t: t, addr: addr, out: out}
|
||||
|
||||
conn, reader := s.openWebSocket()
|
||||
send(t, conn, uploadBody)
|
||||
|
||||
// The client closes its sending side and waits for the answer, which the
|
||||
// app sends only once it has seen the client stop. smallwebwaf passes the
|
||||
// close on to the app through CloseWrite on upgradedConn; without that,
|
||||
// it closes both connections, and the answer is lost.
|
||||
tcp, ok := conn.(*net.TCPConn)
|
||||
if !ok {
|
||||
t.Fatalf("connection is a %T, want a *net.TCPConn", conn)
|
||||
}
|
||||
|
||||
err := tcp.CloseWrite()
|
||||
if err != nil {
|
||||
t.Fatalf("close the sending side: %v", err)
|
||||
}
|
||||
|
||||
got, err := io.ReadAll(reader)
|
||||
if err != nil || string(got) != uploadBody {
|
||||
t.Errorf("got %q (%v), want %q", got, err, uploadBody)
|
||||
}
|
||||
|
||||
s.closeWebSocket(conn)
|
||||
}
|
||||
|
||||
func TestBytesCountSaysWhichBytesCount(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct {
|
||||
setting string
|
||||
// each is the bytes each request counts, and breaking the request
|
||||
// that goes over the limit of 99.
|
||||
each float64
|
||||
breaking int
|
||||
}{
|
||||
{countResponse, answerBytes, 2},
|
||||
{countRequest, bodyBytes, 4},
|
||||
{countBoth, bodyBytes + answerBytes, 1},
|
||||
} {
|
||||
t.Run(tc.setting, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _ := startWithAnswers(t, map[string]string{
|
||||
bytesLimitPerMinute: byteLimit, bytesCount: tc.setting,
|
||||
})
|
||||
|
||||
for i := 1; i <= tc.breaking; i++ {
|
||||
line := s.upload()
|
||||
|
||||
want := ""
|
||||
if i == tc.breaking {
|
||||
want = minuteBytes
|
||||
}
|
||||
|
||||
counted := float64(i) * tc.each
|
||||
if line.LimitHit != want || line.Counts.MinuteBytes != counted {
|
||||
t.Errorf("request %d: log line has limit_hit %q and minute_bytes %v, "+
|
||||
"want %q and %v", i, line.LimitHit, line.Counts.MinuteBytes,
|
||||
want, counted)
|
||||
}
|
||||
}
|
||||
|
||||
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestByteLimitsLeaveOutWhatTheRateLimitsLeaveOut(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const (
|
||||
allowed = "192.0.2.7" // in SWWAF_ALLOW_NETS
|
||||
exempt = "192.0.2.10" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
|
||||
)
|
||||
|
||||
s, _ := startWithAnswers(t, map[string]string{
|
||||
bytesLimitPerMinute: byteLimit,
|
||||
allowNets: allowed,
|
||||
rateLimitExemptNets: exempt,
|
||||
rateLimitExemptPaths: "/assets/",
|
||||
})
|
||||
|
||||
// Each sends 200 bytes, none of which is counted.
|
||||
for _, sent := range []struct{ from, path string }{
|
||||
{allowed, "/"}, {exempt, "/"}, {client, "/assets/app.js"},
|
||||
} {
|
||||
for range 2 {
|
||||
line, _ := s.requestWithBody(http.MethodPost, sent.from, sent.path,
|
||||
uploadHeader, uploadBody, http.StatusOK, requestlog.ActionForward)
|
||||
if _, counted := line.fields["counts"]; counted || line.LimitHit != "" {
|
||||
t.Errorf("%s %s: log line has counts %v and limit_hit %q, want neither",
|
||||
sent.from, sent.path, line.fields["counts"], line.LimitHit)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A path that is not exempt is counted, and breaks the limit.
|
||||
line := s.upload()
|
||||
if line.LimitHit != minuteBytes {
|
||||
t.Errorf("log line has limit_hit %q, want %s", line.LimitHit, minuteBytes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestByteLimitsOffCountTheBytesAndBanNoOne(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const off = "off"
|
||||
|
||||
s, _ := startWithAnswers(t, map[string]string{
|
||||
bytesLimitPerMinute: off, bytesLimitPerHour: off, bytesLimitPerDay: off,
|
||||
})
|
||||
|
||||
for i := 1; i <= 3; i++ {
|
||||
line := s.upload()
|
||||
|
||||
counted := float64(i * (bodyBytes + answerBytes))
|
||||
if line.LimitHit != "" || line.Counts.MinuteBytes != counted ||
|
||||
line.Counts.HourBytes != counted || line.Counts.DayBytes != counted {
|
||||
t.Errorf("request %d: log line has limit_hit %q and counts %+v, "+
|
||||
"want none and %v bytes in each window", i, line.LimitHit,
|
||||
line.Counts, counted)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBanForABrokenByteLimitHasItsNotesAndItsAlert(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, clk, server, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{
|
||||
bytesLimitPerMinute: byteLimit,
|
||||
})
|
||||
start := clk.Now()
|
||||
|
||||
s.requestWithBody(http.MethodPost, client, "/upload?part=1", uploadHeader,
|
||||
uploadBody, http.StatusOK, requestlog.ActionForward)
|
||||
|
||||
netblock := netip.MustParsePrefix(client + "/32")
|
||||
want := bans.Ban{
|
||||
Netblock: netblock,
|
||||
Start: start,
|
||||
Expires: start.Add(time.Hour),
|
||||
Cause: bans.CauseLimit,
|
||||
Reason: "bytes per minute over the limit of " + byteLimit,
|
||||
Notes: bans.Notes{
|
||||
Kind: "bytes",
|
||||
Limit: 99,
|
||||
Window: minute,
|
||||
Count: bodyBytes + answerBytes,
|
||||
// The request as it was answered, by the app.
|
||||
Request: bans.Request{
|
||||
Time: start,
|
||||
Method: http.MethodPost,
|
||||
Host: appHost,
|
||||
Path: "/upload?part=1",
|
||||
Status: http.StatusOK,
|
||||
UserAgent: userAgent,
|
||||
},
|
||||
Requests: 1,
|
||||
},
|
||||
}
|
||||
|
||||
got := server.Ledger.Bans(netblock)
|
||||
if len(got) != 1 || !reflect.DeepEqual(got[0], want) {
|
||||
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
|
||||
}
|
||||
|
||||
wantAlerts(t, queue, banAlert(alerts.EventBan, start, client, want,
|
||||
requestlog.FormatTime(want.Expires)))
|
||||
|
||||
if offences := historyOf(t, server, client).Offences.Limit; offences != 1 {
|
||||
t.Errorf("history counts %d offences for a limit, want 1", offences)
|
||||
}
|
||||
}
|
||||
|
||||
func TestObserveModeLogsAndAlertsAByteLimitAndBansNoOne(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, clk, server, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{
|
||||
mode: observe,
|
||||
bytesLimitPerMinute: byteLimit,
|
||||
})
|
||||
start := clk.Now()
|
||||
|
||||
// No ban sets the client's counters back to zero, so each request
|
||||
// breaks the limit again. The answer is the app's either way, and the
|
||||
// alert for the ban is not sent twice within the cooldown.
|
||||
for range 2 {
|
||||
line := s.upload()
|
||||
wantWouldAction(t, line, "")
|
||||
|
||||
if line.LimitHit != minuteBytes || line.Offence != requestlog.OffenceLimit ||
|
||||
line.BanExpires != "" {
|
||||
t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+
|
||||
"want %s, limit and none", line.LimitHit, line.Offence, line.BanExpires,
|
||||
minuteBytes)
|
||||
}
|
||||
}
|
||||
|
||||
if held := server.Ledger.Snapshot(); len(held) != 0 {
|
||||
t.Errorf("the ledger holds %+v, want no ban", held)
|
||||
}
|
||||
|
||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||
if len(waiting) != 1 {
|
||||
t.Fatalf("%d alerts wait, want 1: %+v", len(waiting), waiting)
|
||||
}
|
||||
|
||||
notes, _ := waiting[0].Detail["notes"].(bans.Notes)
|
||||
alert := banAlert(alerts.EventBan, start, client, bans.Ban{
|
||||
Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit,
|
||||
Reason: "bytes per minute over the limit of " + byteLimit, Notes: notes,
|
||||
}, requestlog.FormatTime(start.Add(time.Hour)))
|
||||
alert.Detail["mode"] = observe
|
||||
wantAlerts(t, queue, alert)
|
||||
}
|
||||
|
||||
func TestObserveModeLeavesOutTheBytesOfARequestEnforceModeRefuses(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _ := startWithAnswers(t, map[string]string{
|
||||
mode: observe,
|
||||
rateLimitPerMinute: "1",
|
||||
bytesLimitPerMinute: "150",
|
||||
})
|
||||
|
||||
s.upload()
|
||||
|
||||
// The second request breaks the rate limit, which in enforce mode would
|
||||
// refuse it before the app sent anything, so its 100 bytes are not
|
||||
// counted, and the byte limit is not broken. Its line gives the bytes
|
||||
// counted before it.
|
||||
line := s.upload()
|
||||
wantWouldAction(t, line, requestlog.ActionRateLimited)
|
||||
|
||||
if line.LimitHit != minute || line.Counts.MinuteBytes != bodyBytes+answerBytes {
|
||||
t.Errorf("log line has limit_hit %q and minute_bytes %v, want minute and %d",
|
||||
line.LimitHit, line.Counts.MinuteBytes, bodyBytes+answerBytes)
|
||||
}
|
||||
}
|
||||
|
||||
// uploadHeader and uploadBody are the header and the body of a request
|
||||
// with a body of bodyBytes.
|
||||
//
|
||||
//nolint:gochecknoglobals // a constant cannot call strings.Repeat
|
||||
var (
|
||||
uploadHeader = "Content-Length: " + strconv.Itoa(bodyBytes)
|
||||
uploadBody = strings.Repeat("u", bodyBytes)
|
||||
)
|
||||
|
||||
// readAndAnswer is the app of these tests: it reads each request's whole
|
||||
// body and answers with answerBytes bytes.
|
||||
func readAndAnswer(w http.ResponseWriter, r *http.Request) {
|
||||
_, _ = io.Copy(io.Discard, r.Body)
|
||||
_, _ = io.WriteString(w, strings.Repeat("a", answerBytes))
|
||||
}
|
||||
|
||||
// breakOff is an app that announces an answer of twice answerBytes, and
|
||||
// breaks off after answerBytes.
|
||||
func breakOff(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Length", strconv.Itoa(2*answerBytes))
|
||||
_, _ = io.WriteString(w, strings.Repeat("a", answerBytes))
|
||||
}
|
||||
|
||||
// answerAfterUpgrade is an app that switches protocols, as for a
|
||||
// WebSocket, and then answers each line it receives with a line of
|
||||
// answerBytes.
|
||||
func answerAfterUpgrade(w http.ResponseWriter, _ *http.Request) {
|
||||
conn, buffered, err := http.NewResponseController(w).Hijack()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
defer func() {
|
||||
_ = conn.Close()
|
||||
}()
|
||||
|
||||
_, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" +
|
||||
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
|
||||
_ = buffered.Flush()
|
||||
|
||||
for {
|
||||
_, err := buffered.ReadString('\n')
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
_, _ = buffered.WriteString(strings.Repeat("a", answerBytes-1) + "\n")
|
||||
_ = buffered.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
// echoOnceTheClientStops is an app that switches protocols, as for a
|
||||
// WebSocket, reads what the client sends until the client stops sending,
|
||||
// and then sends it all back.
|
||||
func echoOnceTheClientStops(w http.ResponseWriter, _ *http.Request) {
|
||||
conn, buffered, err := http.NewResponseController(w).Hijack()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
defer func() {
|
||||
_ = conn.Close()
|
||||
}()
|
||||
|
||||
_, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" +
|
||||
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
|
||||
_ = buffered.Flush()
|
||||
|
||||
received, _ := io.ReadAll(buffered)
|
||||
_, _ = buffered.Write(received)
|
||||
_ = buffered.Flush()
|
||||
}
|
||||
|
||||
// webSocket opens a WebSocket from client to answerAfterUpgrade, sends a
|
||||
// line of bodyBytes on it, reads the answer, and closes it. It checks the
|
||||
// answer, and the log line as request does, and returns the log line.
|
||||
func (s *sender) webSocket() logLine {
|
||||
s.t.Helper()
|
||||
|
||||
conn, reader := s.openWebSocket()
|
||||
send(s.t, conn, strings.Repeat("u", bodyBytes-1)+"\n")
|
||||
|
||||
got, err := reader.ReadString('\n')
|
||||
if err != nil || len(got) != answerBytes {
|
||||
s.t.Errorf("got %d bytes (%v), want %d", len(got), err, answerBytes)
|
||||
}
|
||||
|
||||
return s.closeWebSocket(conn)
|
||||
}
|
||||
|
||||
// openWebSocket sends a request from client to switch protocols, as for a
|
||||
// WebSocket, and checks that the app switches. It returns the connection,
|
||||
// on which reading fails once waitLimit has passed, and a reader of what
|
||||
// the app sends on it.
|
||||
func (s *sender) openWebSocket() (net.Conn, *bufio.Reader) {
|
||||
s.t.Helper()
|
||||
|
||||
conn := dial(s.t, s.addr)
|
||||
send(s.t, conn, "GET /socket HTTP/1.1\r\nHost: "+appHost+"\r\n"+forwardedFor+
|
||||
": "+client+"\r\nConnection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
|
||||
|
||||
err := conn.SetReadDeadline(time.Now().Add(waitLimit))
|
||||
if err != nil {
|
||||
s.t.Fatalf("set read deadline: %v", err)
|
||||
}
|
||||
|
||||
reader := bufio.NewReader(conn)
|
||||
|
||||
res, err := http.ReadResponse(reader, nil)
|
||||
if err != nil {
|
||||
s.t.Fatalf("read the answer to the upgrade: %v", err)
|
||||
}
|
||||
|
||||
_ = res.Body.Close()
|
||||
|
||||
if res.StatusCode != http.StatusSwitchingProtocols {
|
||||
s.t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols)
|
||||
}
|
||||
|
||||
return conn, reader
|
||||
}
|
||||
|
||||
// closeWebSocket closes conn, a WebSocket openWebSocket opened, checks its
|
||||
// log line as request does, and returns it.
|
||||
func (s *sender) closeWebSocket(conn net.Conn) logLine {
|
||||
s.t.Helper()
|
||||
|
||||
_ = conn.Close()
|
||||
|
||||
line := s.out.requestLines(s.t, s.sent+1)[s.sent]
|
||||
s.sent++
|
||||
wantLine(s.t, line, http.StatusSwitchingProtocols, requestlog.ActionForward)
|
||||
|
||||
return line
|
||||
}
|
||||
|
||||
// startWithAnswers is startAppWithAlerts in front of readAndAnswer, for a
|
||||
// test that looks at neither the server nor the alerts.
|
||||
func startWithAnswers(t *testing.T, env map[string]string) (*sender, *clock) {
|
||||
t.Helper()
|
||||
|
||||
s, clk, _, _ := startAppWithAlerts(t, readAndAnswer, env)
|
||||
|
||||
return s, clk
|
||||
}
|
||||
|
||||
// download sends a GET request for / from client, and checks that the
|
||||
// app's answer is passed on, as request does. It returns the log line and
|
||||
// the answer.
|
||||
func (s *sender) download() (logLine, answer) {
|
||||
s.t.Helper()
|
||||
|
||||
return s.requestWithBody(http.MethodGet, client, "/", "", "", http.StatusOK,
|
||||
requestlog.ActionForward)
|
||||
}
|
||||
|
||||
// upload is download for a POST request with a body of bodyBytes, and
|
||||
// returns the log line.
|
||||
func (s *sender) upload() logLine {
|
||||
s.t.Helper()
|
||||
|
||||
line, _ := s.requestWithBody(http.MethodPost, client, "/", uploadHeader,
|
||||
uploadBody, http.StatusOK, requestlog.ActionForward)
|
||||
|
||||
return line
|
||||
}
|
||||
@@ -76,16 +76,14 @@ func scheme(r *http.Request, peerTrusted bool) string {
|
||||
return proto
|
||||
}
|
||||
|
||||
// ipv6GroupPrefix is the length of the IPv6 netblock that is one client.
|
||||
const ipv6GroupPrefix = 64
|
||||
|
||||
// clientGroup is the client a request is counted toward: its IPv4
|
||||
// address, or the /64 its IPv6 address is in, since one abuser usually
|
||||
// holds a whole /64. An IPv4 address in IPv6 form counts as IPv4.
|
||||
func clientGroup(addr netip.Addr) netip.Prefix {
|
||||
// address, or its IPv6 group, the netblock its IPv6 address is in of the
|
||||
// length SWWAF_IPV6_GROUP_PREFIX sets, a /64 by default, since one abuser
|
||||
// usually holds a whole /64. An IPv4 address in IPv6 form counts as IPv4.
|
||||
func (h *handler) clientGroup(addr netip.Addr) netip.Prefix {
|
||||
addr = addr.Unmap()
|
||||
if addr.Is6() {
|
||||
return netip.PrefixFrom(addr, ipv6GroupPrefix).Masked()
|
||||
return netip.PrefixFrom(addr, h.config.IPv6GroupPrefix).Masked()
|
||||
}
|
||||
|
||||
return netip.PrefixFrom(addr, addr.BitLen())
|
||||
|
||||
@@ -56,6 +56,24 @@ func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestTableOfClientsHoldsAtMostMaxTrackedClients(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _, server := startWithClock(t, "", map[string]string{maxTrackedClients: "2"})
|
||||
|
||||
// The third client drops the least recently seen, the first, with its
|
||||
// history.
|
||||
for _, from := range []string{"192.0.2.1", "192.0.2.2", "192.0.2.3"} {
|
||||
s.get(from, http.StatusOK, requestlog.ActionForward)
|
||||
}
|
||||
|
||||
_, held := server.Limiter.Client(netip.MustParsePrefix("192.0.2.1/32"))
|
||||
if server.Limiter.Len() != 2 || held {
|
||||
t.Errorf("the table holds %d clients, the first among them: %t; want 2, "+
|
||||
"without it", server.Limiter.Len(), held)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHistoryCountsTheBodiesEachWay(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
+24
-20
@@ -9,34 +9,42 @@ import (
|
||||
)
|
||||
|
||||
// The headers in which the app is passed the client's AS number and
|
||||
// country while SWWAF_ADD_LOOKUP_HEADERS is set. Go sends a header name in
|
||||
// this form, so X-Client-ASN arrives as X-Client-Asn; header names are not
|
||||
// case-sensitive.
|
||||
// country while SWWAF_ADD_LOOKUP_HEADERS is set. Go writes every header
|
||||
// name in this form, as it sends it and as it receives it, so X-Client-ASN
|
||||
// arrives as X-Client-Asn, and Del removes a client's own whatever their
|
||||
// case; header names are not case-sensitive.
|
||||
const (
|
||||
asnHeader = "X-Client-Asn"
|
||||
countryHeader = "X-Client-Country"
|
||||
)
|
||||
|
||||
// lookUp looks up the client's AS number and country, and notes them for
|
||||
// the log line, unless SWWAF_LOOKUP_SOURCE is off or the client is on a
|
||||
// private, loopback or link-local address, which no lookup can place.
|
||||
// While a setting needs the answer, a new client's request waits for it.
|
||||
// lookUp looks up the client's AS number and country, in the lookup
|
||||
// database or through GeoJS, and notes them for the log line, unless
|
||||
// SWWAF_LOOKUP_SOURCE is off or the client is on a private, loopback or
|
||||
// link-local address, which no lookup can place. The lookup database
|
||||
// answers at once. With GeoJS, while a setting needs the answer, such as a
|
||||
// country list or a biased threshold, a new client's request waits for it.
|
||||
// ctx is the request's own context.
|
||||
func (rq *request) lookUp(ctx context.Context) {
|
||||
if rq.h.config.LookupSource == "off" || !canBePlaced(rq.client) {
|
||||
return
|
||||
}
|
||||
|
||||
answer := rq.h.geojs.LookUp(ctx, clientGroup(rq.client))
|
||||
rq.lookedUp = true
|
||||
rq.line.ASN = answer.ASN
|
||||
rq.line.ASName = answer.ASName
|
||||
rq.line.Country = answer.Country
|
||||
if rq.h.config.LookupSource == "file" {
|
||||
rq.lookupAnswer = rq.h.lookupFile.LookUp(rq.h.clientGroup(rq.client))
|
||||
} else {
|
||||
rq.lookupAnswer = rq.h.geojs.LookUp(ctx, rq.h.clientGroup(rq.client))
|
||||
}
|
||||
|
||||
// addLookup adds answer, GeoJS's answer about a client, to the client's
|
||||
// history, and to the notes of the bans on its netblock that have no AS
|
||||
// number, AS name or country yet.
|
||||
rq.lookedUp = true
|
||||
rq.line.ASN = rq.lookupAnswer.ASN
|
||||
rq.line.ASName = rq.lookupAnswer.ASName
|
||||
rq.line.Country = rq.lookupAnswer.Country
|
||||
}
|
||||
|
||||
// addLookup adds answer, an answer about a client from the lookup
|
||||
// database or GeoJS, to the client's history, and to the notes of the bans
|
||||
// on its netblock that have no AS number, AS name or country yet.
|
||||
func (h *handler) addLookup(answer lookup.Answer) {
|
||||
h.limiter.AddLookup(answer.Client, answer.Answered,
|
||||
answer.ASN, answer.ASName, answer.Country)
|
||||
@@ -45,12 +53,8 @@ func (h *handler) addLookup(answer lookup.Answer) {
|
||||
}
|
||||
|
||||
// setLookupHeaders sets the headers in which the app is passed the
|
||||
// client's AS number and country, leaving out one that is unknown. Any
|
||||
// the client sent are removed, so that the app can believe them.
|
||||
// client's AS number and country, leaving out one that is unknown.
|
||||
func setLookupHeaders(header http.Header, asn, country string) {
|
||||
header.Del(asnHeader)
|
||||
header.Del(countryHeader)
|
||||
|
||||
if asn != "" {
|
||||
header.Set(asnHeader, asn)
|
||||
}
|
||||
|
||||
@@ -1,15 +1,20 @@
|
||||
package proxy_test
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/synctest"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
)
|
||||
|
||||
@@ -17,6 +22,10 @@ import (
|
||||
// and country.
|
||||
type asnAndCountry struct{ asn, asName, country string }
|
||||
|
||||
// fileSource is the SWWAF_LOOKUP_SOURCE that looks clients up in the
|
||||
// lookup database.
|
||||
const fileSource = "file"
|
||||
|
||||
func TestEveryClientIsLookedUpWithoutWaitingWhileNoSettingNeedsIt(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -163,6 +172,64 @@ func TestLookupSourceOffLooksNoClientUp(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientsAreLookedUpInTheLookupDatabaseAndGeoJSIsNotAsked(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
geojsURL, asked := startGeoJS(t)
|
||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
||||
fromDE + "/32": {ASN: asnDE, ASName: asNameDE, Country: "DE"},
|
||||
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
|
||||
})
|
||||
s, clk, server := startWithClock(t, geojsURL, map[string]string{
|
||||
lookupSource: fileSource,
|
||||
lookupDBPath: path,
|
||||
allowedCountries: "DE",
|
||||
rateLimitPerMinute: "1",
|
||||
})
|
||||
|
||||
// fromDE's second request breaks the limit and bans it. The list
|
||||
// refuses fromKP, and unplaced, which the file does not hold.
|
||||
de := asnAndCountry{asnDE, asNameDE, "DE"}
|
||||
|
||||
for _, tc := range []struct {
|
||||
line logLine
|
||||
want asnAndCountry
|
||||
}{
|
||||
{s.get(fromDE, http.StatusOK, requestlog.ActionForward), de},
|
||||
{s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited), de},
|
||||
{
|
||||
s.get(fromKP, http.StatusForbidden, requestlog.ActionCountryDenied),
|
||||
asnAndCountry{asnKP, asNameKP, "KP"},
|
||||
},
|
||||
{
|
||||
s.get(unplaced, http.StatusForbidden, requestlog.ActionCountryDenied),
|
||||
asnAndCountry{},
|
||||
},
|
||||
} {
|
||||
got := asnAndCountry{tc.line.ASN, tc.line.ASName, tc.line.Country}
|
||||
if got != tc.want {
|
||||
t.Errorf("log line has %+v, want %+v", got, tc.want)
|
||||
}
|
||||
}
|
||||
|
||||
h := historyOf(t, server, fromDE)
|
||||
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != de ||
|
||||
!h.LookedUp.Equal(clk.Now()) {
|
||||
t.Errorf("history has %+v, looked up at %s; want %+v, at %s",
|
||||
got, h.LookedUp, de, clk.Now())
|
||||
}
|
||||
|
||||
notes := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32"))[0].Notes
|
||||
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != de {
|
||||
t.Errorf("the ban's notes have %+v, want %+v", got, de)
|
||||
}
|
||||
|
||||
if len(asked()) != 0 {
|
||||
t.Errorf("GeoJS was asked about %v, want nothing", asked())
|
||||
}
|
||||
}
|
||||
|
||||
func TestLookupHeadersArePassedToTheAppAndTheClientsOwnRemoved(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -189,9 +256,9 @@ func TestLookupHeadersArePassedToTheAppAndTheClientsOwnRemoved(t *testing.T) {
|
||||
// Each client sends headers of its own. fromDE's first request waits
|
||||
// for its answer, which the app is passed; unplaced has none to pass,
|
||||
// and a client on a private address is not looked up.
|
||||
own := "X-Client-ASN: AS1\r\nX-Client-Country: KP"
|
||||
for _, from := range []string{fromDE, unplaced, "10.0.0.8"} {
|
||||
s.requestWithHeader(from, "/", own, http.StatusOK, requestlog.ActionForward)
|
||||
s.requestWithHeader(from, "/", clientsOwnLookupHeaders,
|
||||
http.StatusOK, requestlog.ActionForward)
|
||||
}
|
||||
|
||||
mu.Lock()
|
||||
@@ -205,6 +272,103 @@ func TestLookupHeadersArePassedToTheAppAndTheClientsOwnRemoved(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientsOwnLookupHeadersAreRemovedWhileTheSettingIsOff(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var (
|
||||
mu sync.Mutex
|
||||
asn, country []string
|
||||
)
|
||||
|
||||
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
|
||||
asn, country = r.Header.Values("X-Client-Asn"), r.Header.Values("X-Client-Country")
|
||||
})
|
||||
geojsURL, _ := startGeoJS(t)
|
||||
addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{
|
||||
trustedProxies: trustLocalhost,
|
||||
})
|
||||
s := &sender{t: t, addr: addr, out: out}
|
||||
|
||||
s.requestWithHeader(fromDE, "/", clientsOwnLookupHeaders,
|
||||
http.StatusOK, requestlog.ActionForward)
|
||||
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
|
||||
if asn != nil || country != nil {
|
||||
t.Errorf("the app was passed X-Client-ASN %v and X-Client-Country %v, want neither",
|
||||
asn, country)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRequestWaitsAsLongAsTheLookupTimeoutSays(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// The test runs in a synctest bubble, where the time package runs on a
|
||||
// clock of the test's own: the wait lasts exactly as long as it should,
|
||||
// however slowly the test process runs. Nothing in it may wait on the
|
||||
// network, which would keep that clock from moving on: the request is
|
||||
// handed to the proxy's handler, and GeoJS is one that never answers.
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
// Not the default second. The exclusive list needs the answer, and
|
||||
// the app is never reached.
|
||||
const timeout = 3 * time.Second
|
||||
|
||||
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
|
||||
time.Now, map[string]string{
|
||||
lookupTimeout: timeout.String(),
|
||||
allowedCountries: "DE",
|
||||
})
|
||||
|
||||
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/",
|
||||
http.NoBody)
|
||||
req.RemoteAddr = net.JoinHostPort(fromDE, "1234")
|
||||
began := time.Now()
|
||||
|
||||
server.Handler.ServeHTTP(httptest.NewRecorder(), req)
|
||||
|
||||
if waited := time.Since(began); waited != timeout {
|
||||
t.Errorf("the request waited %s for its answer, want %s", waited, timeout)
|
||||
}
|
||||
|
||||
// Without an answer, the client is in no country the list allows.
|
||||
wantLine(t, out.requestLine(t), http.StatusForbidden,
|
||||
requestlog.ActionCountryDenied)
|
||||
})
|
||||
}
|
||||
|
||||
// unansweredGeoJSURL is where a GeoJS that never answers is asked: a
|
||||
// request to it waits, without the network, until it is abandoned.
|
||||
// TestMain registers it with Go's default transport, through which GeoJS
|
||||
// is asked.
|
||||
const unansweredGeoJSURL = "unanswered://geojs/v1/ip/geo.json"
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
transport, _ := http.DefaultTransport.(*http.Transport)
|
||||
transport.RegisterProtocol("unanswered", unansweredGeoJS{})
|
||||
transport.RegisterProtocol("abuseipdb", abuseIPDBStandIn{})
|
||||
|
||||
m.Run()
|
||||
}
|
||||
|
||||
// unansweredGeoJS is the GeoJS at unansweredGeoJSURL.
|
||||
type unansweredGeoJS struct{}
|
||||
|
||||
// RoundTrip waits until req is abandoned.
|
||||
func (unansweredGeoJS) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
<-req.Context().Done()
|
||||
|
||||
return nil, req.Context().Err()
|
||||
}
|
||||
|
||||
// clientsOwnLookupHeaders are the X-Client-ASN and X-Client-Country a
|
||||
// client sends of its own, each twice, in two cases.
|
||||
const clientsOwnLookupHeaders = "X-Client-ASN: AS1\r\nx-client-asn: AS2\r\n" +
|
||||
"X-CLIENT-COUNTRY: KP\r\nx-client-country: CN"
|
||||
|
||||
// waitUntil waits until done reports true, for at most waitLimit.
|
||||
func waitUntil(done func() bool) {
|
||||
deadline := time.Now().Add(waitLimit)
|
||||
|
||||
@@ -222,8 +222,8 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
|
||||
metrics := s.scrape(scraper)
|
||||
wantMetric(t, metrics,
|
||||
`smallwebwaf_requests_total{action="denied",instance="app",status_class="none"}`, 1)
|
||||
wantMetric(t, metrics,
|
||||
`smallwebwaf_rate_limit_hits_total{instance="app",window="minute"}`, 1)
|
||||
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+
|
||||
`kind="requests",window="minute"}`, 1)
|
||||
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 1)
|
||||
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 1)
|
||||
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1)
|
||||
@@ -238,8 +238,8 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
|
||||
s.get(client, 0, requestlog.ActionRateLimited)
|
||||
|
||||
metrics = s.scrape(scraper)
|
||||
wantMetric(t, metrics,
|
||||
`smallwebwaf_rate_limit_hits_total{instance="app",window="minute"}`, 2)
|
||||
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+
|
||||
`kind="requests",window="minute"}`, 2)
|
||||
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 2)
|
||||
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 2)
|
||||
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1)
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -126,7 +127,7 @@ func TestObserveModeMakesNoBanAndKeepsTheBansItHas(t *testing.T) {
|
||||
}
|
||||
|
||||
got := server.Ledger.Snapshot()
|
||||
if len(got) != 1 || got[0] != kept {
|
||||
if len(got) != 1 || !reflect.DeepEqual(got[0], kept) {
|
||||
t.Errorf("bans\n%+v\nwant only\n%+v", got, kept)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -122,6 +122,8 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
|
||||
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
|
||||
t.Helper()
|
||||
|
||||
bytes := float64(sent + received)
|
||||
|
||||
want := withTimings(line, requestlog.Line{
|
||||
Type: requestType, Time: line.Time, Instance: "app",
|
||||
ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
|
||||
@@ -131,7 +133,10 @@ func wantRequestFields(t *testing.T, line logLine, host string, sent, received i
|
||||
RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
|
||||
ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
|
||||
UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
|
||||
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
|
||||
Counts: ratelimit.Counts{
|
||||
Minute: 1, Hour: 1, Day: 1,
|
||||
MinuteBytes: bytes, HourBytes: bytes, DayBytes: bytes,
|
||||
},
|
||||
})
|
||||
if !reflect.DeepEqual(line.Line, want) {
|
||||
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
||||
@@ -315,7 +320,7 @@ func TestServerHasTheDefaultLimits(t *testing.T) {
|
||||
server := proxy.New(proxy.Params{
|
||||
Config: cfg,
|
||||
RequestLog: io.Discard,
|
||||
ProcessLog: requestlog.NewProcessLogger(io.Discard, cfg.InstanceName),
|
||||
ProcessLog: requestlog.NewProcessLogger(io.Discard, cfg.InstanceName, cfg.LogLevel),
|
||||
})
|
||||
|
||||
if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 ||
|
||||
@@ -394,6 +399,25 @@ func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogLevelHoldsBackNoRequestLine(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// At error the warning that the request to the app failed is held back,
|
||||
// and is written before the answer is.
|
||||
addr, out := startProxy(t, "http://"+localhost+":1", map[string]string{
|
||||
"SWWAF_LOG_LEVEL": "error",
|
||||
})
|
||||
|
||||
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
|
||||
wantLine(t, out.requestLine(t), http.StatusBadGateway, requestlog.ActionUpstreamError)
|
||||
|
||||
for _, line := range out.lines(t) {
|
||||
if line["type"] == "process" {
|
||||
t.Errorf("process line %v, want none at error", line)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogsAnAnswerThatBrokeOff(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
+96
-8
@@ -12,11 +12,13 @@ import (
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/anomaly"
|
||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
"sneak.berlin/go/smallwebwaf/internal/reputation"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||
)
|
||||
@@ -54,9 +56,15 @@ type Params struct {
|
||||
RequestLog io.Writer
|
||||
// ProcessLog receives the process's own messages.
|
||||
ProcessLog *slog.Logger
|
||||
// GeoJSURL is where clients' AS numbers and countries are looked up,
|
||||
// normally lookup.URL, unless SWWAF_LOOKUP_SOURCE is off.
|
||||
// GeoJSURL is where clients' AS numbers and countries are looked up
|
||||
// while SWWAF_LOOKUP_SOURCE is geojs, normally lookup.URL.
|
||||
GeoJSURL string
|
||||
// AbuseIPDBURL is where clients are checked with AbuseIPDB while
|
||||
// SWWAF_ABUSEIPDB_KEY is set, normally reputation.AbuseIPDBURL.
|
||||
AbuseIPDBURL string
|
||||
// LookupFile is the lookup database they are looked up in while
|
||||
// SWWAF_LOOKUP_SOURCE is file, and nil otherwise.
|
||||
LookupFile *lookup.File
|
||||
// Now tells the time by which requests are counted for the rate
|
||||
// limits, bans are made and run out, and GeoJS's answers are kept,
|
||||
// normally time.Now in UTC, the time the state files give.
|
||||
@@ -65,18 +73,29 @@ type Params struct {
|
||||
// against.
|
||||
Rules *rules.Files
|
||||
// Alerts receive the alert for each ban the proxy makes or makes
|
||||
// permanent, and for GeoJS failing.
|
||||
// permanent, for each count over an anomaly threshold, for each request
|
||||
// whose client a blocklist, a DNSBL zone or AbuseIPDB lists, and for
|
||||
// GeoJS failing, a fetch of a list failing, a query to a DNSBL zone or
|
||||
// a check with AbuseIPDB failing, or the day's AbuseIPDB checks used up.
|
||||
Alerts *alerts.Queue
|
||||
}
|
||||
|
||||
// Server is the server smallwebwaf runs, with the parts of the proxy
|
||||
// whose state the state files keep, and the metrics.
|
||||
// whose state the state files keep, the lookup database, nil unless
|
||||
// SWWAF_LOOKUP_SOURCE is file, the lists fetched from URLs, which its Run
|
||||
// fetches, the DNSBL zones' verdicts, AbuseIPDB's scores and checks
|
||||
// spent, and the metrics.
|
||||
type Server struct {
|
||||
*http.Server
|
||||
|
||||
Ledger *bans.Ledger
|
||||
Limiter *ratelimit.Limiter
|
||||
GeoJS *lookup.GeoJS
|
||||
Anomalies *anomaly.Counters
|
||||
LookupFile *lookup.File
|
||||
Lists *reputation.Lists
|
||||
DNSBL *reputation.DNSBL
|
||||
AbuseIPDB *reputation.AbuseIPDB
|
||||
Metrics *metrics.Metrics
|
||||
}
|
||||
|
||||
@@ -89,6 +108,7 @@ type Server struct {
|
||||
func New(params Params) *Server {
|
||||
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
|
||||
m := metrics.New(params.Config.MetricsTopN, params.Config.InstanceName)
|
||||
lists, dnsbl, abuseIPDB := newReputation(params, m)
|
||||
h := &handler{
|
||||
config: params.Config,
|
||||
requestLog: params.RequestLog,
|
||||
@@ -101,7 +121,10 @@ func New(params Params) *Server {
|
||||
PerMinute: params.Config.RateLimitPerMinute,
|
||||
PerHour: params.Config.RateLimitPerHour,
|
||||
PerDay: params.Config.RateLimitPerDay,
|
||||
}),
|
||||
BytesPerMinute: params.Config.BytesLimitPerMinute,
|
||||
BytesPerHour: params.Config.BytesLimitPerHour,
|
||||
BytesPerDay: params.Config.BytesLimitPerDay,
|
||||
}, params.Config.MaxTrackedClients),
|
||||
ledger: bans.New(bans.Rules{
|
||||
LimitBanDuration: params.Config.LimitBanDuration,
|
||||
LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow,
|
||||
@@ -109,17 +132,32 @@ func New(params Params) *Server {
|
||||
AttackBanDuration: params.Config.AttackBanDuration,
|
||||
MaxBans: params.Config.MaxBans,
|
||||
}),
|
||||
anomalies: anomaly.New(anomaly.Params{
|
||||
Client: params.Config.AnomalyClient,
|
||||
Net: params.Config.AnomalyNet,
|
||||
ASN: params.Config.AnomalyASN,
|
||||
Total: params.Config.AnomalyTotal,
|
||||
Watch: params.Config.AnomalyWatch,
|
||||
NetV4Prefix: params.Config.AnomalyNetV4Prefix,
|
||||
NetV6Prefix: params.Config.AnomalyNetV6Prefix,
|
||||
NamedNetblocks: params.Config.WatchNets,
|
||||
Alerts: params.Alerts,
|
||||
}),
|
||||
lookupFile: params.LookupFile,
|
||||
lists: lists,
|
||||
dnsbl: dnsbl,
|
||||
abuseIPDB: abuseIPDB,
|
||||
rules: params.Rules,
|
||||
alerts: params.Alerts,
|
||||
}
|
||||
h.geojs = lookup.New(lookup.Params{
|
||||
URL: params.GeoJSURL,
|
||||
Timeout: params.Config.LookupTimeout,
|
||||
// The country lists and the headers act on the answer before the
|
||||
// request goes on.
|
||||
// The country lists, the headers and the biased thresholds act on
|
||||
// the answer before the request goes on.
|
||||
Wait: len(params.Config.DeniedCountries) > 0 ||
|
||||
len(params.Config.ExclusivelyAllowedCountries) > 0 ||
|
||||
params.Config.AddLookupHeaders,
|
||||
params.Config.AddLookupHeaders || biasedThresholdsSet(params.Config),
|
||||
Answered: h.addLookup,
|
||||
Now: params.Now,
|
||||
ProcessLog: params.ProcessLog,
|
||||
@@ -145,10 +183,49 @@ func New(params Params) *Server {
|
||||
Ledger: h.ledger,
|
||||
Limiter: h.limiter,
|
||||
GeoJS: h.geojs,
|
||||
Anomalies: h.anomalies,
|
||||
LookupFile: h.lookupFile,
|
||||
Lists: h.lists,
|
||||
DNSBL: h.dnsbl,
|
||||
AbuseIPDB: h.abuseIPDB,
|
||||
Metrics: m,
|
||||
}
|
||||
}
|
||||
|
||||
// newReputation returns the lists fetched from URLs, the DNSBL zones'
|
||||
// verdicts and AbuseIPDB's scores, as the settings in params name them,
|
||||
// with none fetched, asked for or checked yet, and adds their metrics to
|
||||
// m, AbuseIPDB's while SWWAF_ABUSEIPDB_KEY is set.
|
||||
func newReputation(
|
||||
params Params, m *metrics.Metrics,
|
||||
) (*reputation.Lists, *reputation.DNSBL, *reputation.AbuseIPDB) {
|
||||
cfg := params.Config
|
||||
lists := reputation.New(reputation.Params{
|
||||
BlocklistURLs: cfg.BlocklistURLs, Refresh: cfg.BlocklistRefresh,
|
||||
ASNLimitPercentURL: cfg.ASNLimitPercentURL, Now: params.Now,
|
||||
ProcessLog: params.ProcessLog, Alerts: params.Alerts,
|
||||
})
|
||||
dnsbl := reputation.NewDNSBL(reputation.DNSBLParams{
|
||||
Zones: cfg.DNSBLZones, Resolver: cfg.DNSBLResolver, CacheTTL: cfg.ReputationCacheTTL,
|
||||
Timeout: cfg.ReputationTimeout, Now: params.Now, ProcessLog: params.ProcessLog,
|
||||
Alerts: params.Alerts,
|
||||
})
|
||||
abuseIPDB := reputation.NewAbuseIPDB(reputation.AbuseIPDBParams{
|
||||
URL: params.AbuseIPDBURL, Key: cfg.AbuseIPDBKey, MinScore: cfg.AbuseIPDBMinScore,
|
||||
DailyBudget: cfg.AbuseIPDBDailyBudget, CacheTTL: cfg.ReputationCacheTTL,
|
||||
Timeout: cfg.ReputationTimeout, Now: params.Now, ProcessLog: params.ProcessLog,
|
||||
Alerts: params.Alerts,
|
||||
})
|
||||
|
||||
m.AddReputation(lists, dnsbl)
|
||||
|
||||
if cfg.AbuseIPDBKey != "" {
|
||||
m.AddAbuseIPDB(abuseIPDB)
|
||||
}
|
||||
|
||||
return lists, dnsbl, abuseIPDB
|
||||
}
|
||||
|
||||
// handler is the proxy. It holds what every request shares; what belongs
|
||||
// to one request is in a request.
|
||||
type handler struct {
|
||||
@@ -162,6 +239,11 @@ type handler struct {
|
||||
limiter *ratelimit.Limiter
|
||||
ledger *bans.Ledger
|
||||
geojs *lookup.GeoJS
|
||||
anomalies *anomaly.Counters
|
||||
lookupFile *lookup.File
|
||||
lists *reputation.Lists
|
||||
dnsbl *reputation.DNSBL
|
||||
abuseIPDB *reputation.AbuseIPDB
|
||||
rules *rules.Files
|
||||
alerts *alerts.Queue
|
||||
}
|
||||
@@ -200,6 +282,7 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
// Once the request has ended, before its log line is written.
|
||||
defer rq.addToHistory()
|
||||
defer rq.countAnomalies()
|
||||
|
||||
refused := rq.check(r.Context())
|
||||
rq.checked = time.Now()
|
||||
@@ -218,5 +301,10 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// Once the response has ended, before the request is added to its
|
||||
// client's history. Deferred, since ReverseProxy panics to end a
|
||||
// response it cannot finish.
|
||||
defer rq.countBytes()
|
||||
|
||||
rq.forward(r.Context())
|
||||
}
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||
@@ -60,6 +61,8 @@ const (
|
||||
requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES"
|
||||
responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES"
|
||||
trustedProxies = "SWWAF_TRUSTED_PROXIES"
|
||||
ipv6GroupPrefix = "SWWAF_IPV6_GROUP_PREFIX"
|
||||
maxTrackedClients = "SWWAF_MAX_TRACKED_CLIENTS"
|
||||
allowNets = "SWWAF_ALLOW_NETS"
|
||||
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
|
||||
denyNets = "SWWAF_DENY_NETS"
|
||||
@@ -67,6 +70,7 @@ const (
|
||||
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
||||
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
|
||||
lookupSource = "SWWAF_LOOKUP_SOURCE"
|
||||
lookupDBPath = "SWWAF_LOOKUP_DB_PATH"
|
||||
lookupTimeout = "SWWAF_LOOKUP_TIMEOUT"
|
||||
addLookupHeaders = "SWWAF_ADD_LOOKUP_HEADERS"
|
||||
deniedCountries = "SWWAF_DENIED_COUNTRIES"
|
||||
@@ -235,16 +239,45 @@ func startProxyWithClock(
|
||||
}
|
||||
|
||||
// startProxyWithAlerts is startProxyWithClock, and returns the queue of
|
||||
// the alerts the proxy raises as well, as the settings in env make it. No
|
||||
// alert is sent from it: they wait in it, for the test to look at. With
|
||||
// no geojsURL, there is no stand-in for GeoJS to look clients up at, and
|
||||
// SWWAF_LOOKUP_SOURCE is off unless env sets it.
|
||||
// the alerts the proxy raises as well, as newProxy makes them.
|
||||
func startProxyWithAlerts(
|
||||
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
||||
env map[string]string,
|
||||
) (string, *output, *proxy.Server, *alerts.Queue) {
|
||||
t.Helper()
|
||||
|
||||
server, out, alertQueue := newProxy(t, appURL, geojsURL, now, env)
|
||||
|
||||
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
|
||||
go func() {
|
||||
_ = server.Serve(listener)
|
||||
}()
|
||||
|
||||
t.Cleanup(func() {
|
||||
_ = server.Close()
|
||||
})
|
||||
|
||||
return listener.Addr().String(), out, server, alertQueue
|
||||
}
|
||||
|
||||
// newProxy makes the server startProxyWithClock starts, without starting
|
||||
// it, and returns it, what it writes, and the queue of the alerts the
|
||||
// proxy raises, as the settings in env make it. No alert is sent from the
|
||||
// queue: they wait in it, for the test to look at. With no geojsURL, there
|
||||
// is no stand-in for GeoJS to look clients up at, and SWWAF_LOOKUP_SOURCE
|
||||
// is off unless env sets it. While it is file, the lookup database
|
||||
// SWWAF_LOOKUP_DB_PATH names is read. Clients are checked with AbuseIPDB
|
||||
// at abuseIPDBURL while env sets SWWAF_ABUSEIPDB_KEY.
|
||||
func newProxy(
|
||||
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
||||
env map[string]string,
|
||||
) (*proxy.Server, *output, *alerts.Queue) {
|
||||
t.Helper()
|
||||
|
||||
settings := map[string]string{
|
||||
"SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir(), instanceName: "app",
|
||||
}
|
||||
@@ -264,7 +297,7 @@ func startProxyWithAlerts(
|
||||
}
|
||||
|
||||
out := &output{}
|
||||
processLog := requestlog.NewProcessLogger(out, cfg.InstanceName)
|
||||
processLog := requestlog.NewProcessLogger(out, cfg.InstanceName, cfg.LogLevel)
|
||||
|
||||
ruleFiles, err := rules.Load(rules.Params{
|
||||
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
|
||||
@@ -283,30 +316,30 @@ func startProxyWithAlerts(
|
||||
ProcessLog: processLog,
|
||||
})
|
||||
|
||||
var lookupFile *lookup.File
|
||||
|
||||
if cfg.LookupSource == fileSource {
|
||||
lookupFile, err = lookup.OpenFile(lookup.FileParams{
|
||||
Path: cfg.LookupDBPath, Now: now, ProcessLog: processLog, Alerts: alertQueue,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("lookup database: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
server := proxy.New(proxy.Params{
|
||||
Config: cfg,
|
||||
RequestLog: out,
|
||||
ProcessLog: processLog,
|
||||
GeoJSURL: geojsURL,
|
||||
AbuseIPDBURL: abuseIPDBURL,
|
||||
LookupFile: lookupFile,
|
||||
Now: now,
|
||||
Rules: ruleFiles,
|
||||
Alerts: alertQueue,
|
||||
})
|
||||
|
||||
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
|
||||
go func() {
|
||||
_ = server.Serve(listener)
|
||||
}()
|
||||
|
||||
t.Cleanup(func() {
|
||||
_ = server.Close()
|
||||
})
|
||||
|
||||
return listener.Addr().String(), out, server, alertQueue
|
||||
return server, out, alertQueue
|
||||
}
|
||||
|
||||
// newClient returns an HTTP client that sends requests as they are made,
|
||||
|
||||
@@ -72,6 +72,49 @@ func TestRateLimitRefusesBeforeTheApp(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestIPv6GroupPrefixSetsTheClientTheLimitsCount(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// With SWWAF_IPV6_GROUP_PREFIX at 48, the first two addresses, in two
|
||||
// /64s of one /48, are one client, and the second's request breaks the
|
||||
// limit; the third, in the next /48, is another client.
|
||||
const (
|
||||
first = "2001:db8:9::1"
|
||||
second = "2001:db8:9:1::1"
|
||||
other = "2001:db8:a::1"
|
||||
)
|
||||
|
||||
for _, tc := range []struct {
|
||||
setting, value string
|
||||
// status and action are those of the request that breaks the
|
||||
// limit: a rate limit refuses it, a byte limit passes it on.
|
||||
status int
|
||||
action string
|
||||
}{
|
||||
{rateLimitPerMinute, "1", http.StatusForbidden, requestlog.ActionRateLimited},
|
||||
{bytesLimitPerMinute, byteLimit, http.StatusOK, requestlog.ActionForward},
|
||||
} {
|
||||
t.Run(tc.setting, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s, _ := startWithAnswers(t, map[string]string{
|
||||
ipv6GroupPrefix: "48", tc.setting: tc.value,
|
||||
})
|
||||
|
||||
s.get(first, http.StatusOK, requestlog.ActionForward)
|
||||
|
||||
line := s.get(second, tc.status, tc.action)
|
||||
if line.ClientGroup != "2001:db8:9::/48" ||
|
||||
line.Offence != requestlog.OffenceLimit {
|
||||
t.Errorf("log line has client_group %q and offence %q, "+
|
||||
"want 2001:db8:9::/48 and limit", line.ClientGroup, line.Offence)
|
||||
}
|
||||
|
||||
s.get(other, http.StatusOK, requestlog.ActionForward)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
package proxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
"sneak.berlin/go/smallwebwaf/internal/reputation"
|
||||
)
|
||||
|
||||
// deny is the SWWAF_BLOCKLIST_ACTION and the SWWAF_REPUTATION_ACTION that
|
||||
// refuses the requests of a client a source lists.
|
||||
const deny = "deny"
|
||||
|
||||
// blocklistDenied notes the blocklists that list the client, as
|
||||
// noteListed does, and reports whether SWWAF_BLOCKLIST_ACTION, being deny,
|
||||
// refuses the request. Being limit, it lowers the client's limits instead
|
||||
// (see limitPercentages), and being log, it does nothing more.
|
||||
func (rq *request) blocklistDenied() bool {
|
||||
listedBy := rq.h.lists.ListedBy(rq.client)
|
||||
rq.blocklisted = len(listedBy) > 0
|
||||
rq.noteListed(listedBy, "listed by a blocklist")
|
||||
|
||||
return rq.blocklisted && rq.h.config.BlocklistAction == deny
|
||||
}
|
||||
|
||||
// dnsblDenied notes the DNSBL zones whose verdict lists the client, as
|
||||
// noteListed does, and reports whether SWWAF_REPUTATION_ACTION, being
|
||||
// deny, refuses the request. Being limit, it lowers the client's limits
|
||||
// instead (see limitPercentages), and being log, it does nothing more. A
|
||||
// zone without a verdict on the client is asked about it in the
|
||||
// background, and the request does not wait for the answer. ctx is the
|
||||
// request's own context.
|
||||
func (rq *request) dnsblDenied(ctx context.Context) bool {
|
||||
listedBy := rq.h.dnsbl.ListedBy(ctx, rq.client)
|
||||
rq.dnsblListed = len(listedBy) > 0
|
||||
rq.noteListed(listedBy, "listed by a DNSBL zone")
|
||||
|
||||
return rq.dnsblListed && rq.h.config.ReputationAction == deny
|
||||
}
|
||||
|
||||
// abuseIPDBDenied notes AbuseIPDB, as noteHit does, with the score, when
|
||||
// its score of the client is a hit, and reports whether
|
||||
// SWWAF_REPUTATION_ACTION, being deny, refuses the request, as dnsblDenied
|
||||
// does for a zone. While SWWAF_ABUSEIPDB_KEY is unset it does nothing. A
|
||||
// client without a score is checked in the background, by the request's
|
||||
// address, if its history counts an offence, and the request does not
|
||||
// wait for the answer. The score is then used for each address of the
|
||||
// client. ctx is the request's own context.
|
||||
func (rq *request) abuseIPDBDenied(ctx context.Context) bool {
|
||||
if rq.h.config.AbuseIPDBKey == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
client := rq.h.clientGroup(rq.client)
|
||||
held, _ := rq.h.limiter.Client(client)
|
||||
offender := held.History.Offences != ratelimit.Offences{}
|
||||
|
||||
score, hit := rq.h.abuseIPDB.Hit(ctx, client, rq.client, offender)
|
||||
if !hit {
|
||||
return false
|
||||
}
|
||||
|
||||
rq.abuseIPDBHit = true
|
||||
rq.noteHit(bans.ReputationHit{Source: reputation.AbuseIPDBSource, Score: &score},
|
||||
"scored by AbuseIPDB at or over SWWAF_ABUSEIPDB_MIN_SCORE")
|
||||
|
||||
return rq.h.config.ReputationAction == deny
|
||||
}
|
||||
|
||||
// noteListed notes each of sources, the URLs of the blocklists or the
|
||||
// DNSBL zones, their keys masked, that list the client, as noteHit does,
|
||||
// with reason.
|
||||
func (rq *request) noteListed(sources []string, reason string) {
|
||||
for _, source := range sources {
|
||||
rq.noteHit(bans.ReputationHit{Source: source}, reason)
|
||||
}
|
||||
}
|
||||
|
||||
// noteHit adds hit's source, which lists the client, to the log line's
|
||||
// reputation, and hit to the notes of a ban the request makes, counts the
|
||||
// source in the metrics, and raises a reputation_hit alert with reason,
|
||||
// whose detail gives hit's source and score.
|
||||
func (rq *request) noteHit(hit bans.ReputationHit, reason string) {
|
||||
detail := map[string]any{"source": hit.Source}
|
||||
if hit.Score != nil {
|
||||
detail["score"] = *hit.Score
|
||||
}
|
||||
|
||||
rq.line.Reputation = append(rq.line.Reputation, hit.Source)
|
||||
rq.reputation = append(rq.reputation, hit)
|
||||
rq.h.metrics.ReputationHit(hit.Source)
|
||||
rq.h.alerts.Raise(alerts.Alert{
|
||||
Event: alerts.EventReputationHit,
|
||||
Client: rq.client,
|
||||
Netblock: rq.h.clientGroup(rq.client),
|
||||
ASN: rq.line.ASN,
|
||||
ASName: rq.line.ASName,
|
||||
Country: rq.line.Country,
|
||||
Reason: reason,
|
||||
Detail: detail,
|
||||
})
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
+133
-33
@@ -3,6 +3,7 @@ package proxy
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptrace"
|
||||
"net/http/httputil"
|
||||
@@ -15,6 +16,10 @@ import (
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/anomaly"
|
||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
)
|
||||
@@ -49,8 +54,27 @@ type request struct {
|
||||
peer netip.Addr
|
||||
peerTrusted bool
|
||||
// lookedUp is true once the client's AS number and country have been
|
||||
// looked up, whether or not an answer was there.
|
||||
// looked up, whether or not an answer was there, and lookupAnswer is
|
||||
// what the lookup gave then, the zero Answer while GeoJS had given none.
|
||||
lookedUp bool
|
||||
lookupAnswer lookup.Answer
|
||||
// counted is true for a request the rate limits counted, whose bytes
|
||||
// the byte limits count once it has ended. limitPercent and
|
||||
// bytesPercent are then its client's limit percentages for the rate
|
||||
// limits and for the byte limits.
|
||||
counted bool
|
||||
limitPercent, bytesPercent percentage
|
||||
// attack is true for a request that matched a ban rule, and
|
||||
// ruleBlocked for one a block rule refused, each an offence its
|
||||
// client's history counts.
|
||||
attack, ruleBlocked bool
|
||||
// blocklisted is true once a blocklist is found to list the client,
|
||||
// dnsblListed once a DNSBL zone's verdict is, and abuseIPDBHit once
|
||||
// AbuseIPDB's score of it is a hit.
|
||||
blocklisted, dnsblListed, abuseIPDBHit bool
|
||||
// reputation is the reputation sources that list the client, for the
|
||||
// notes of a ban the request makes.
|
||||
reputation []bans.ReputationHit
|
||||
start time.Time
|
||||
// checked is when the checks were done, and upstreamStart when the
|
||||
// request was handed to the app.
|
||||
@@ -62,6 +86,9 @@ type request struct {
|
||||
refused atomic.Pointer[refusal]
|
||||
// complete is true once the app's whole answer has been passed on.
|
||||
complete bool
|
||||
// upgraded is the connection to the app once the app has switched
|
||||
// protocols, as for a WebSocket, and nil otherwise.
|
||||
upgraded *upgradedConn
|
||||
|
||||
// mu guards what follows. The timeouts run on goroutines of their
|
||||
// own, and the transport starts and stops them, and notes the times
|
||||
@@ -117,7 +144,7 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
|
||||
RequestID: requestID(r, peerTrusted),
|
||||
PeerIP: peer.String(),
|
||||
ForwardedFor: strings.Join(forwardedFor, ", "),
|
||||
ClientGroup: clientGroup(client).String(),
|
||||
ClientGroup: h.clientGroup(client).String(),
|
||||
ContentType: r.Header.Get("Content-Type"),
|
||||
RequestHeaders: requestHeaders(r, h.config.LogRequestHeaders),
|
||||
HasAuthorization: len(r.Header.Values("Authorization")) > 0,
|
||||
@@ -198,12 +225,15 @@ func (rq *request) check(ctx context.Context) *refusal {
|
||||
// client in SWWAF_ALLOW_NETS skips them, and is not looked up. For any
|
||||
// other client, SWWAF_DENY_NETS comes first, then a ban on its netblock,
|
||||
// so that a client either refuses is not looked up, then the lookup of
|
||||
// its AS number and country, 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 or the
|
||||
// its AS number and country, then the country lists, then the blocklists,
|
||||
// then the DNSBL zones' verdicts, and then AbuseIPDB's score; 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 or the
|
||||
// request's path is exempt under SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that
|
||||
// every other request is counted, and last the rule files. ctx is the
|
||||
// request's own context.
|
||||
// every other request is counted, each of them by the client's limit
|
||||
// percentages, and last the rule files. A request exempt from the rate
|
||||
// limits is exempt from the byte limits too. ctx is the request's own
|
||||
// context.
|
||||
func (rq *request) checkClient(ctx context.Context) string {
|
||||
cfg := rq.h.config
|
||||
if isInside(rq.client, cfg.AllowNets) {
|
||||
@@ -226,9 +256,23 @@ func (rq *request) checkClient(ctx context.Context) string {
|
||||
return requestlog.ActionCountryDenied
|
||||
}
|
||||
|
||||
exempt := isInside(rq.client, cfg.RateLimitExemptNets) ||
|
||||
pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
|
||||
if !exempt && rq.limitBroken(now) {
|
||||
if rq.blocklistDenied() {
|
||||
return requestlog.ActionDenied
|
||||
}
|
||||
|
||||
if rq.dnsblDenied(ctx) || rq.abuseIPDBDenied(ctx) {
|
||||
return requestlog.ActionDenied
|
||||
}
|
||||
|
||||
rq.counted = !isInside(rq.client, cfg.RateLimitExemptNets) &&
|
||||
!pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
|
||||
if rq.counted {
|
||||
rq.limitPercent, rq.bytesPercent = rq.limitPercentages()
|
||||
rq.line.LimitPercent, rq.line.LimitPercentSetting = rq.limitPercent.logged()
|
||||
rq.line.BytesPercent, rq.line.BytesPercentSetting = rq.bytesPercent.logged()
|
||||
}
|
||||
|
||||
if rq.counted && rq.limitBroken(now) {
|
||||
return requestlog.ActionRateLimited
|
||||
}
|
||||
|
||||
@@ -293,8 +337,9 @@ func (rq *request) forward(ctx context.Context) {
|
||||
|
||||
// rewrite makes the request the app receives: the client's request,
|
||||
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
|
||||
// the request's id set, and, while SWWAF_ADD_LOOKUP_HEADERS is set, the
|
||||
// client's AS number and country.
|
||||
// the request's id set, without any X-Client-ASN or X-Client-Country the
|
||||
// client sent, whatever SWWAF_ADD_LOOKUP_HEADERS says, and, while it is
|
||||
// set, with the client's AS number and country in them.
|
||||
func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
||||
upstream := rq.h.config.UpstreamURL
|
||||
pr.Out.URL.Scheme = upstream.Scheme
|
||||
@@ -304,6 +349,8 @@ func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
||||
pr.Out.URL.RawQuery = pr.In.URL.RawQuery
|
||||
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
|
||||
pr.Out.Header.Set(requestIDHeader, rq.line.RequestID)
|
||||
pr.Out.Header.Del(asnHeader)
|
||||
pr.Out.Header.Del(countryHeader)
|
||||
|
||||
if rq.h.config.AddLookupHeaders {
|
||||
setLookupHeaders(pr.Out.Header, rq.line.ASN, rq.line.Country)
|
||||
@@ -318,11 +365,18 @@ func (rq *request) modifyResponse(res *http.Response) error {
|
||||
if res.StatusCode == http.StatusSwitchingProtocols {
|
||||
// An upgraded connection, such as a WebSocket, is not cut by the
|
||||
// timeouts. ReverseProxy writes this answer straight to the
|
||||
// connection it takes over, not through rq.out.
|
||||
// connection it takes over, not through rq.out, and then copies
|
||||
// what passes each way through res.Body, the connection to the app.
|
||||
rq.stopTimers()
|
||||
rq.out.status = res.StatusCode
|
||||
rq.line.Websocket = true
|
||||
|
||||
conn, ok := res.Body.(io.ReadWriteCloser)
|
||||
if ok {
|
||||
rq.upgraded = &upgradedConn{ReadWriteCloser: conn}
|
||||
res.Body = rq.upgraded
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -425,10 +479,7 @@ func (rq *request) finish() {
|
||||
line.ResponseContentType = header.Get("Content-Type")
|
||||
line.CacheControl = header.Get("Cache-Control")
|
||||
line.Location = header.Get("Location")
|
||||
|
||||
if rq.body != nil {
|
||||
line.RequestBytes = rq.body.bytes.Load()
|
||||
}
|
||||
line.RequestBytes = rq.requestBytes()
|
||||
|
||||
// limit is the setting whose size or time limit the request passed.
|
||||
var limit string
|
||||
@@ -485,36 +536,85 @@ func timing(start, end time.Time) *float64 {
|
||||
}
|
||||
|
||||
// addToHistory adds the request, which has ended, to its client's
|
||||
// history, and then, for a client that was looked up, the answer kept
|
||||
// about it to that history and to the notes of the bans on its netblock:
|
||||
// an answer that came during the request may have come before either was
|
||||
// there, and one that comes later is added when it comes.
|
||||
// history, and then the lookup's answer about the client, as
|
||||
// answerAtTheEnd gives it, to that history and to the notes of the bans
|
||||
// on its netblock: an answer may have come before either was there, and
|
||||
// one from GeoJS that comes later is added when it comes.
|
||||
func (rq *request) addToHistory() {
|
||||
var requestBytes int64
|
||||
if rq.body != nil {
|
||||
requestBytes = rq.body.bytes.Load()
|
||||
}
|
||||
|
||||
forwarded := !rq.upstreamStart.IsZero()
|
||||
group := clientGroup(rq.client)
|
||||
|
||||
rq.h.limiter.AddToHistory(group, rq.h.now(), ratelimit.Request{
|
||||
rq.h.limiter.AddToHistory(rq.h.clientGroup(rq.client), rq.h.now(), ratelimit.Request{
|
||||
Forwarded: forwarded,
|
||||
Refused: !forwarded && rq.refused.Load() != nil,
|
||||
Status: rq.out.status,
|
||||
RequestBytes: requestBytes,
|
||||
RequestBytes: rq.requestBytes(),
|
||||
ResponseBytes: rq.out.bytes,
|
||||
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
|
||||
Attack: rq.attack,
|
||||
RuleBlocked: rq.ruleBlocked,
|
||||
})
|
||||
|
||||
if !rq.lookedUp {
|
||||
answer, found := rq.answerAtTheEnd()
|
||||
if found {
|
||||
rq.h.addLookup(answer)
|
||||
}
|
||||
}
|
||||
|
||||
// countAnomalies counts the request, which has ended, and its bytes, as
|
||||
// countedBytes gives them, for the anomaly thresholds, whatever was done
|
||||
// with it: a request refused, one from a client in SWWAF_ALLOW_NETS or
|
||||
// SWWAF_RATE_LIMIT_EXEMPT_NETS, and one for a path in
|
||||
// SWWAF_RATE_LIMIT_EXEMPT_PATHS are counted too. It is counted for its
|
||||
// client's AS number when answerAtTheEnd gives one. With every anomaly
|
||||
// threshold off, the default, it does nothing.
|
||||
func (rq *request) countAnomalies() {
|
||||
if !anomalyThresholdsSet(rq.h.config) {
|
||||
return
|
||||
}
|
||||
|
||||
answer, kept := rq.h.geojs.Kept(group)
|
||||
if kept {
|
||||
rq.h.addLookup(answer)
|
||||
answer, _ := rq.answerAtTheEnd()
|
||||
|
||||
rq.h.anomalies.Count(rq.h.now(), anomaly.Request{
|
||||
Client: rq.client,
|
||||
ClientGroup: rq.h.clientGroup(rq.client),
|
||||
ASN: answer.ASN,
|
||||
ASName: answer.ASName,
|
||||
Country: answer.Country,
|
||||
Bytes: rq.countedBytes(),
|
||||
})
|
||||
}
|
||||
|
||||
// anomalyThresholdsSet reports whether any anomaly threshold is set.
|
||||
func anomalyThresholdsSet(cfg *config.Config) bool {
|
||||
off := anomaly.Thresholds{}
|
||||
|
||||
return cfg.AnomalyClient != off || cfg.AnomalyNet != off || cfg.AnomalyASN != off ||
|
||||
cfg.AnomalyTotal != off || cfg.AnomalyWatch != off
|
||||
}
|
||||
|
||||
// answerAtTheEnd returns, for a client that was looked up, the lookup's
|
||||
// answer about it as the request ends, and whether there is one: the
|
||||
// lookup database's, which was there at once, or the one GeoJS has given
|
||||
// by then, which a request does not wait for unless a setting needs it.
|
||||
func (rq *request) answerAtTheEnd() (lookup.Answer, bool) {
|
||||
if !rq.lookedUp {
|
||||
return lookup.Answer{}, false
|
||||
}
|
||||
|
||||
if rq.h.config.LookupSource == "file" {
|
||||
return rq.lookupAnswer, true
|
||||
}
|
||||
|
||||
return rq.h.geojs.Kept(rq.h.clientGroup(rq.client))
|
||||
}
|
||||
|
||||
// requestBytes is how many bytes of the request's body have been read.
|
||||
func (rq *request) requestBytes() int64 {
|
||||
if rq.body == nil {
|
||||
return 0
|
||||
}
|
||||
|
||||
return rq.body.bytes.Load()
|
||||
}
|
||||
|
||||
// clientRequestDeadline is when the client must have sent its whole
|
||||
|
||||
@@ -119,7 +119,10 @@ func wantFullLine(t *testing.T, line logLine) {
|
||||
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
|
||||
CacheControl: "no-store", Location: "/elsewhere",
|
||||
Action: requestlog.ActionForward,
|
||||
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
|
||||
// Its 3 bytes in and 5 out, each way counted by default.
|
||||
Counts: ratelimit.Counts{
|
||||
Minute: 1, Hour: 1, Day: 1, MinuteBytes: 8, HourBytes: 8, DayBytes: 8,
|
||||
},
|
||||
})
|
||||
if !reflect.DeepEqual(line.Line, want) {
|
||||
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"net/netip"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"slices"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -73,7 +74,7 @@ func TestEachRuleAction(t *testing.T) {
|
||||
}
|
||||
|
||||
got := server.Ledger.Bans(netblock)
|
||||
if len(got) != 1 || got[0] != want {
|
||||
if len(got) != 1 || !reflect.DeepEqual(got[0], want) {
|
||||
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
|
||||
}
|
||||
|
||||
|
||||
@@ -12,7 +12,8 @@ import (
|
||||
// action of the rule that refuses it, ActionRuleBlocked for a block rule
|
||||
// and ActionBanned for a ban rule, or "" when none does. A ban rule bans
|
||||
// the client's netblock for a clear sign of attack, or in observe mode
|
||||
// raises the alert for the ban it would have made.
|
||||
// raises the alert for the ban it would have made. Either rule's match
|
||||
// is noted as an offence, for the client's history.
|
||||
func (rq *request) checkRules(now time.Time) string {
|
||||
matched := rq.h.rules.Match(rq.in)
|
||||
|
||||
@@ -28,8 +29,11 @@ func (rq *request) checkRules(now time.Time) string {
|
||||
// Only the last rule matched can refuse the request.
|
||||
switch last := matched[len(matched)-1]; last.Action {
|
||||
case rules.ActionBlock:
|
||||
rq.ruleBlocked = true
|
||||
|
||||
return requestlog.ActionRuleBlocked
|
||||
case rules.ActionBan:
|
||||
rq.attack = true
|
||||
rq.banForAttack(now, last)
|
||||
|
||||
return requestlog.ActionBanned
|
||||
|
||||
@@ -11,7 +11,7 @@ import (
|
||||
func TestHistoryKeepsEveryRequest(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{})
|
||||
limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
@@ -53,7 +53,7 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
|
||||
func TestLookupReachesTheHistoryOfAClientInTheTable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{})
|
||||
limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
other := netip.MustParsePrefix("198.51.100.7/32")
|
||||
start := midnight()
|
||||
@@ -90,7 +90,7 @@ func TestLookupReachesTheHistoryOfAClientInTheTable(t *testing.T) {
|
||||
func TestResetKeepsTheHistory(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit})
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
@@ -109,7 +109,7 @@ func TestResetKeepsTheHistory(t *testing.T) {
|
||||
func TestRequestsAddsUpTheClientsInsideTheNetblock(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{})
|
||||
limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
|
||||
|
||||
for client, requests := range map[string]int{
|
||||
"198.51.100.9/32": 2,
|
||||
|
||||
+190
-72
@@ -1,9 +1,10 @@
|
||||
// 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.
|
||||
// and bytes 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 or a byte limit, and each client's history
|
||||
// since it was first seen. At most SWWAF_MAX_TRACKED_CLIENTS clients are
|
||||
// kept, in memory, and written to clients.json and read from it by the
|
||||
// state package.
|
||||
package ratelimit
|
||||
|
||||
import (
|
||||
@@ -16,26 +17,32 @@ import (
|
||||
"github.com/hashicorp/golang-lru/v2/simplelru"
|
||||
)
|
||||
|
||||
// maxClients is how many clients are kept. Past it, the least recently
|
||||
// seen client is dropped, with its history, and starts afresh if it comes
|
||||
// back.
|
||||
const maxClients = 20000
|
||||
|
||||
const day = 24 * time.Hour
|
||||
|
||||
// The kinds of limits, as the metrics name them.
|
||||
const (
|
||||
// KindRequests is a rate limit, on a client's requests.
|
||||
KindRequests = "requests"
|
||||
// KindBytes is a byte limit, on a client's bytes.
|
||||
KindBytes = "bytes"
|
||||
)
|
||||
|
||||
// Limits are the most requests a client may make in a minute, an hour and
|
||||
// a day. Zero is no limit.
|
||||
// a day, and the most bytes. Zero is no limit.
|
||||
type Limits struct {
|
||||
PerMinute int64
|
||||
PerHour int64
|
||||
PerDay int64
|
||||
BytesPerMinute int64
|
||||
BytesPerHour int64
|
||||
BytesPerDay int64
|
||||
}
|
||||
|
||||
// Limiter counts each client's requests against the limits, and keeps
|
||||
// its history. It is safe for concurrent use.
|
||||
// Limiter counts each client's requests and bytes 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.
|
||||
// Client.buckets and Client.byteBuckets.
|
||||
windows [3]window
|
||||
|
||||
mu sync.Mutex
|
||||
@@ -43,17 +50,23 @@ type Limiter struct {
|
||||
}
|
||||
|
||||
// Client is a client in the table, as clients.json holds it: its buckets
|
||||
// in each window, and its history.
|
||||
// of requests and of bytes in each window, and its history.
|
||||
//
|
||||
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||
type Client struct {
|
||||
Client netip.Prefix `json:"client"`
|
||||
Minute Buckets `json:"minute"`
|
||||
Hour Buckets `json:"hour"`
|
||||
Day Buckets `json:"day"`
|
||||
MinuteBytes Buckets `json:"minute_bytes"`
|
||||
HourBytes Buckets `json:"hour_bytes"`
|
||||
DayBytes Buckets `json:"day_bytes"`
|
||||
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.
|
||||
// Buckets are a client's two buckets in one window: the requests, or the
|
||||
// bytes, 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"`
|
||||
@@ -68,8 +81,8 @@ type History struct {
|
||||
LastSeen time.Time `json:"last_seen"`
|
||||
// ASN, ASName and Country are the client's AS number, AS name and
|
||||
// country as last looked up, each empty when the lookup could not
|
||||
// find it, and LookedUp is when GeoJS gave that answer; all are empty
|
||||
// while the client never was looked up.
|
||||
// find it, and LookedUp is when the lookup gave that answer; all are
|
||||
// empty while the client never was looked up.
|
||||
ASN string `json:"asn,omitempty"`
|
||||
ASName string `json:"as_name,omitempty"`
|
||||
Country string `json:"country,omitempty"`
|
||||
@@ -100,9 +113,15 @@ type Responses struct {
|
||||
}
|
||||
|
||||
// Offences are a client's offences, by kind.
|
||||
//
|
||||
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||
type Offences struct {
|
||||
// Limit is its requests that broke a rate limit.
|
||||
// Limit is its requests that broke a rate limit or a byte limit,
|
||||
// Attack those that matched a ban rule, a clear sign of attack, and
|
||||
// RuleBlocked those a block rule refused.
|
||||
Limit int64 `json:"limit"`
|
||||
Attack int64 `json:"attack"`
|
||||
RuleBlocked int64 `json:"rule_blocked"`
|
||||
}
|
||||
|
||||
// Request is what a client's history keeps of one of its requests.
|
||||
@@ -119,12 +138,19 @@ type Request struct {
|
||||
// and of its response.
|
||||
RequestBytes int64
|
||||
ResponseBytes int64
|
||||
// BrokeLimit is true for a request that broke a rate limit.
|
||||
// BrokeLimit is true for a request that broke a rate limit or a byte
|
||||
// limit, Attack for one that matched a ban rule, and RuleBlocked for
|
||||
// one a block rule refused.
|
||||
BrokeLimit bool
|
||||
Attack bool
|
||||
RuleBlocked bool
|
||||
}
|
||||
|
||||
// New returns a Limiter for limits, with no client counted yet.
|
||||
func New(limits Limits) *Limiter {
|
||||
// New returns a Limiter for limits, with no client counted yet, whose
|
||||
// table holds at most maxClients clients (SWWAF_MAX_TRACKED_CLIENTS). Past
|
||||
// it, the least recently seen client is dropped, with its history, and
|
||||
// starts afresh if it comes back.
|
||||
func New(limits Limits, maxClients int) *Limiter {
|
||||
clients, err := simplelru.NewLRU[netip.Prefix, *Client](maxClients, nil)
|
||||
if err != nil {
|
||||
panic(err) // NewLRU fails only for a size below one
|
||||
@@ -132,62 +158,74 @@ func New(limits Limits) *Limiter {
|
||||
|
||||
return &Limiter{
|
||||
windows: [3]window{
|
||||
{name: "minute", length: time.Minute, limit: limits.PerMinute},
|
||||
{name: "hour", length: time.Hour, limit: limits.PerHour},
|
||||
{name: "day", length: day, limit: limits.PerDay},
|
||||
{
|
||||
name: "minute", length: time.Minute,
|
||||
limit: limits.PerMinute, byteLimit: limits.BytesPerMinute,
|
||||
},
|
||||
{
|
||||
name: "hour", length: time.Hour,
|
||||
limit: limits.PerHour, byteLimit: limits.BytesPerHour,
|
||||
},
|
||||
{
|
||||
name: "day", length: day,
|
||||
limit: limits.PerDay, byteLimit: limits.BytesPerDay,
|
||||
},
|
||||
},
|
||||
clients: clients,
|
||||
}
|
||||
}
|
||||
|
||||
// Hit is a request that takes a client over a rate limit.
|
||||
// Hit is a request that takes a client over a rate limit, or whose bytes
|
||||
// take it over a byte limit.
|
||||
type Hit struct {
|
||||
// Kind is KindRequests for a rate limit, KindBytes for a byte limit.
|
||||
Kind string
|
||||
// Window is "minute", "hour" or "day".
|
||||
Window string
|
||||
// Limit is the window's limit.
|
||||
// Limit is the window's limit, as the client's percentage of it.
|
||||
Limit int64
|
||||
// Requests is the client's requests counted in the window, this one
|
||||
// included.
|
||||
Requests float64
|
||||
// Count is the client's requests, or bytes, counted in the window,
|
||||
// this request's included.
|
||||
Count float64
|
||||
}
|
||||
|
||||
// Counts are a client's requests in the minute, the hour and the day that
|
||||
// end at a request, that request included.
|
||||
// Counts are a client's requests and bytes in the minute, the hour and
|
||||
// the day that end at a request, that request's included.
|
||||
//
|
||||
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
|
||||
type Counts struct {
|
||||
Minute float64 `json:"minute"`
|
||||
Hour float64 `json:"hour"`
|
||||
Day float64 `json:"day"`
|
||||
MinuteBytes float64 `json:"minute_bytes"`
|
||||
HourBytes float64 `json:"hour_bytes"`
|
||||
DayBytes float64 `json:"day_bytes"`
|
||||
}
|
||||
|
||||
// Count counts a request from client at now, in every window, whether or
|
||||
// not it is refused, and returns the client's requests in each window. 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) (Counts, Hit, bool) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
var (
|
||||
requests [3]float64
|
||||
hit Hit
|
||||
)
|
||||
|
||||
for i, b := range l.get(client).buckets() {
|
||||
w := l.windows[i]
|
||||
|
||||
requests[i] = b.add(now, w.length)
|
||||
if hit.Window == "" && w.limit > 0 && requests[i] > float64(w.limit) {
|
||||
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests[i]}
|
||||
}
|
||||
// not it is refused, and returns the client's counts in each window. It
|
||||
// reports whether the request takes the client over a rate limit, of
|
||||
// which the client gets the percentage percent, rounded down, and the hit:
|
||||
// the window whose limit it goes over, the shortest if it is over
|
||||
// several. A limit that is off stays off.
|
||||
func (l *Limiter) Count(
|
||||
client netip.Prefix, now time.Time, percent int64,
|
||||
) (Counts, Hit, bool) {
|
||||
return l.count(client, now, 1, 0, percent)
|
||||
}
|
||||
|
||||
counts := Counts{Minute: requests[0], Hour: requests[1], Day: requests[2]}
|
||||
|
||||
return counts, hit, hit.Window != ""
|
||||
// CountBytes counts bytes, those of a request from client that has ended,
|
||||
// at now, in every window, and returns the client's counts in each window.
|
||||
// It reports whether the bytes take the client over a byte limit, of which
|
||||
// the client gets the percentage percent, and the hit, as Count does.
|
||||
func (l *Limiter) CountBytes(
|
||||
client netip.Prefix, now time.Time, bytes, percent int64,
|
||||
) (Counts, Hit, bool) {
|
||||
return l.count(client, now, 0, bytes, percent)
|
||||
}
|
||||
|
||||
// Reset sets client's counts in every window back to zero. Its history
|
||||
// keeps its totals.
|
||||
// Reset sets client's counts of requests and of bytes 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()
|
||||
@@ -195,6 +233,7 @@ func (l *Limiter) Reset(client netip.Prefix) {
|
||||
c, seen := l.clients.Peek(client)
|
||||
if seen {
|
||||
c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{}
|
||||
c.MinuteBytes, c.HourBytes, c.DayBytes = Buckets{}, Buckets{}, Buckets{}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -227,10 +266,19 @@ func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
|
||||
if r.BrokeLimit {
|
||||
h.Offences.Limit++
|
||||
}
|
||||
|
||||
if r.Attack {
|
||||
h.Offences.Attack++
|
||||
}
|
||||
|
||||
if r.RuleBlocked {
|
||||
h.Offences.RuleBlocked++
|
||||
}
|
||||
}
|
||||
|
||||
// AddLookup gives client's history its AS number, AS name and country, as
|
||||
// GeoJS gave them at lookedUp, if the table of clients holds the client.
|
||||
// the lookup gave them at lookedUp, if the table of clients holds the
|
||||
// client.
|
||||
// It does not make the client the most recently seen.
|
||||
func (l *Limiter) AddLookup(
|
||||
client netip.Prefix, lookedUp time.Time, asn, asName, country string,
|
||||
@@ -327,19 +375,63 @@ func (l *Limiter) Load(clients []Client, now time.Time) {
|
||||
l.clients.Purge()
|
||||
|
||||
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)) {
|
||||
for i, w := range l.windows {
|
||||
for _, b := range []*Buckets{c.buckets()[i], c.byteBuckets()[i]} {
|
||||
if b.Passed(now, w.length) {
|
||||
*b = Buckets{}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
l.clients.Add(c.Client, &c)
|
||||
}
|
||||
}
|
||||
|
||||
// count adds requests and bytes from client at now to its buckets in
|
||||
// every window, and returns its counts. A limit is broken only by what is
|
||||
// added to it, so that a request whose bytes are counted after another of
|
||||
// the client's requests broke a rate limit does not break it too. The
|
||||
// client gets the percentage percent of each limit.
|
||||
func (l *Limiter) count(
|
||||
client netip.Prefix, now time.Time, requests, bytes, percent int64,
|
||||
) (Counts, Hit, bool) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
c := l.get(client)
|
||||
requestBuckets, byteBuckets := c.buckets(), c.byteBuckets()
|
||||
|
||||
var (
|
||||
requestCounts, byteCounts [3]float64
|
||||
hit Hit
|
||||
)
|
||||
|
||||
for i, w := range l.windows {
|
||||
requestCounts[i] = requestBuckets[i].Add(now, w.length, requests)
|
||||
byteCounts[i] = byteBuckets[i].Add(now, w.length, bytes)
|
||||
limit, byteLimit := percentOf(w.limit, percent), percentOf(w.byteLimit, percent)
|
||||
|
||||
switch {
|
||||
case hit.Window != "":
|
||||
case requests > 0 && w.limit > 0 && requestCounts[i] > float64(limit):
|
||||
hit = Hit{
|
||||
Kind: KindRequests, Window: w.name, Limit: limit, Count: requestCounts[i],
|
||||
}
|
||||
case bytes > 0 && w.byteLimit > 0 && byteCounts[i] > float64(byteLimit):
|
||||
hit = Hit{
|
||||
Kind: KindBytes, Window: w.name, Limit: byteLimit, Count: byteCounts[i],
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
counts := Counts{
|
||||
Minute: requestCounts[0], Hour: requestCounts[1], Day: requestCounts[2],
|
||||
MinuteBytes: byteCounts[0], HourBytes: byteCounts[1], DayBytes: byteCounts[2],
|
||||
}
|
||||
|
||||
return counts, hit, hit.Window != ""
|
||||
}
|
||||
|
||||
// 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 {
|
||||
@@ -352,30 +444,48 @@ func (l *Limiter) get(client netip.Prefix) *Client {
|
||||
return c
|
||||
}
|
||||
|
||||
// buckets returns c's buckets in the minute, the hour and the day.
|
||||
// buckets returns c's buckets of requests 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
|
||||
// most requests a client may make in it.
|
||||
// byteBuckets returns c's buckets of bytes in the minute, the hour and the
|
||||
// day.
|
||||
func (c *Client) byteBuckets() [3]*Buckets {
|
||||
return [3]*Buckets{&c.MinuteBytes, &c.HourBytes, &c.DayBytes}
|
||||
}
|
||||
|
||||
// window is a length of time over which requests and bytes are counted,
|
||||
// and the most requests and the most bytes a client may have in it.
|
||||
type window struct {
|
||||
name string
|
||||
length time.Duration
|
||||
limit int64
|
||||
byteLimit 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
|
||||
// that bucket the window still covers.
|
||||
// percentOf returns the percentage percent of limit, rounded down. It is
|
||||
// written as limit's hundreds times percent, plus the rest's share, since
|
||||
// limit*percent can overflow for a byte limit.
|
||||
func percentOf(limit, percent int64) int64 {
|
||||
const hundred = 100
|
||||
|
||||
return limit/hundred*percent + limit%hundred*percent/hundred
|
||||
}
|
||||
|
||||
// Add counts n requests, or n bytes, at now in a window of length, and
|
||||
// returns the count in the window that ends at now: what is in the bucket
|
||||
// under way, and what is in the bucket before it weighted by how much of
|
||||
// that bucket the window still covers. With n zero it counts nothing, and
|
||||
// returns the count. The anomaly counters count in Buckets too.
|
||||
//
|
||||
// Concurrent requests can be counted out of order, so now can be a moment
|
||||
// before the bucket under way began; such a request is counted in that
|
||||
// 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 {
|
||||
func (b *Buckets) Add(now time.Time, length time.Duration, n int64) float64 {
|
||||
if now.Before(b.Start.Add(-time.Second)) {
|
||||
*b = Buckets{}
|
||||
}
|
||||
@@ -392,7 +502,7 @@ func (b *Buckets) add(now time.Time, length time.Duration) float64 {
|
||||
b.Current = 0
|
||||
}
|
||||
|
||||
b.Current++
|
||||
b.Current += n
|
||||
|
||||
elapsed := max(now.Sub(b.Start), 0)
|
||||
covered := 1 - float64(elapsed)/float64(length)
|
||||
@@ -400,6 +510,14 @@ func (b *Buckets) add(now time.Time, length time.Duration) float64 {
|
||||
return float64(b.Previous)*covered + float64(b.Current)
|
||||
}
|
||||
|
||||
// Passed reports whether the window of length that ends at now covers
|
||||
// neither of b's buckets: it begins after the bucket under way has ended.
|
||||
// What they hold then counts no more, and a state file read at now drops
|
||||
// it.
|
||||
func (b *Buckets) Passed(now time.Time, length time.Duration) bool {
|
||||
return !now.Add(-length).Before(b.Start.Add(length))
|
||||
}
|
||||
|
||||
// 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) {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package ratelimit_test
|
||||
|
||||
import (
|
||||
"math"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -11,6 +12,14 @@ import (
|
||||
// limit is the limit the tests set.
|
||||
const limit = 3
|
||||
|
||||
// tableSize is the most clients the tests' tables hold, the default of
|
||||
// SWWAF_MAX_TRACKED_CLIENTS.
|
||||
const tableSize = 20000
|
||||
|
||||
// whole is the percentage of each limit a client gets when nothing lowers
|
||||
// its limits.
|
||||
const whole = 100
|
||||
|
||||
// The windows, as Count names them.
|
||||
const (
|
||||
minute = "minute"
|
||||
@@ -32,7 +41,7 @@ func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) {
|
||||
t.Run(tc.window, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(tc.limits)
|
||||
limiter := ratelimit.New(tc.limits, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
quarter := tc.length / 4
|
||||
@@ -57,43 +66,204 @@ func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) {
|
||||
func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: limit})
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: limit}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
for range limit {
|
||||
_, _, over := limiter.Count(client, start)
|
||||
_, _, over := limiter.Count(client, start, whole)
|
||||
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)
|
||||
_, hit, over := limiter.Count(client, start, whole)
|
||||
|
||||
want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1}
|
||||
want := ratelimit.Hit{
|
||||
Kind: ratelimit.KindRequests, Window: minute, Limit: limit, Count: limit + 1,
|
||||
}
|
||||
if !over || hit != want {
|
||||
t.Errorf("request over the limit gives %+v and %t, want %+v and true",
|
||||
hit, over, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientGetsItsPercentageOfEachLimitRoundedDown(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 5, BytesPerDay: math.MaxInt64},
|
||||
tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
// Half of 5 requests is 2.5, rounded down to 2: the third is over.
|
||||
for range 2 {
|
||||
_, _, over := limiter.Count(client, start, 50)
|
||||
if over {
|
||||
t.Fatal("a request within half the limit is over it")
|
||||
}
|
||||
}
|
||||
|
||||
_, hit, over := limiter.Count(client, start, 50)
|
||||
|
||||
want := ratelimit.Hit{Kind: ratelimit.KindRequests, Window: minute, Limit: 2, Count: 3}
|
||||
if !over || hit != want {
|
||||
t.Errorf("the third request gives %+v and %t, want %+v and true", hit, over, want)
|
||||
}
|
||||
|
||||
// Half of the largest byte limit is still far above a TiB: working it
|
||||
// out does not overflow.
|
||||
_, hit, over = limiter.CountBytes(client, start, 1<<40, 50)
|
||||
if over {
|
||||
t.Errorf("a TiB is over half the largest byte limit: %+v", hit)
|
||||
}
|
||||
}
|
||||
|
||||
func TestZeroPercentIsAZeroAllowanceAndALimitOffStaysOff(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Only the hour has limits: the minute's and the day's are off.
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit, BytesPerHour: 1000},
|
||||
tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
// At 0 percent, the first request and the first byte are over the
|
||||
// hour's limits, which are 0; the minute's, which are off, stay off.
|
||||
_, hit, _ := limiter.Count(client, start, 0)
|
||||
|
||||
want := ratelimit.Hit{Kind: ratelimit.KindRequests, Window: hour, Limit: 0, Count: 1}
|
||||
if hit != want {
|
||||
t.Errorf("the first request gives %+v, want %+v", hit, want)
|
||||
}
|
||||
|
||||
_, hit, _ = limiter.CountBytes(client, start, 1, 0)
|
||||
|
||||
want = ratelimit.Hit{Kind: ratelimit.KindBytes, Window: hour, Limit: 0, Count: 1}
|
||||
if hit != want {
|
||||
t.Errorf("the first byte gives %+v, want %+v", hit, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEachByteLimitIsBrokenByTheBytesCounted(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const byteLimit = 1000
|
||||
|
||||
for _, tc := range []struct {
|
||||
window string
|
||||
limits ratelimit.Limits
|
||||
}{
|
||||
{minute, ratelimit.Limits{BytesPerMinute: byteLimit}},
|
||||
{hour, ratelimit.Limits{BytesPerHour: byteLimit}},
|
||||
{"day", ratelimit.Limits{BytesPerDay: byteLimit}},
|
||||
} {
|
||||
t.Run(tc.window, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(tc.limits, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
|
||||
// 600 bytes are within the limit, 600 more over it.
|
||||
_, _, over := limiter.CountBytes(client, midnight(), 600, whole)
|
||||
if over {
|
||||
t.Fatal("600 bytes are over the limit of 1000")
|
||||
}
|
||||
|
||||
_, hit, over := limiter.CountBytes(client, midnight(), 600, whole)
|
||||
|
||||
want := ratelimit.Hit{
|
||||
Kind: ratelimit.KindBytes, Window: tc.window, Limit: byteLimit, Count: 1200,
|
||||
}
|
||||
if !over || hit != want {
|
||||
t.Errorf("1200 bytes give %+v and %t, want %+v and true", hit, over, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestALimitIsBrokenOnlyByWhatIsAddedToIt(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 2, BytesPerMinute: 1000},
|
||||
tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
other := netip.MustParsePrefix("203.0.113.10/32")
|
||||
start := midnight()
|
||||
|
||||
// The third request breaks the rate limit. The bytes of a request
|
||||
// counted after it, within the byte limit, do not break it again.
|
||||
for range 2 {
|
||||
wantCount(t, limiter, client, start, "")
|
||||
}
|
||||
|
||||
wantCount(t, limiter, client, start, minute)
|
||||
wantBytesCount(t, limiter, client, start, 500, "")
|
||||
wantBytesCount(t, limiter, client, start, 600, ratelimit.KindBytes)
|
||||
|
||||
// Bytes over the byte limit do not have the next request break it, nor
|
||||
// the rate limit, which that request is within.
|
||||
wantBytesCount(t, limiter, other, start, 1200, ratelimit.KindBytes)
|
||||
wantCount(t, limiter, other, start, "")
|
||||
}
|
||||
|
||||
func TestCountGivesTheBytesInEachWindow(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
limiter.CountBytes(client, start, 300, whole)
|
||||
|
||||
// A quarter into the next hour, the minute has only these 100 bytes.
|
||||
// The hour still covers three quarters of the bucket before, whose 300
|
||||
// bytes count 225, and these: 325. The day covers all 400.
|
||||
later := start.Add(time.Hour + time.Hour/4)
|
||||
limiter.CountBytes(client, later, 100, whole)
|
||||
|
||||
// A request's counts give the bytes counted so far too.
|
||||
counts, _, _ := limiter.Count(client, later, whole)
|
||||
|
||||
want := ratelimit.Counts{
|
||||
Minute: 1, Hour: 1, Day: 1, MinuteBytes: 100, HourBytes: 325, DayBytes: 400,
|
||||
}
|
||||
if counts != want {
|
||||
t.Errorf("counts %+v, want %+v", counts, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResetSetsTheBytesBackToZero(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{BytesPerDay: 1000}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
wantBytesCount(t, limiter, client, start, 1200, ratelimit.KindBytes)
|
||||
limiter.Reset(client)
|
||||
|
||||
// The client has its whole allowance of bytes again.
|
||||
wantBytesCount(t, limiter, client, start, 1000, "")
|
||||
}
|
||||
|
||||
func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{})
|
||||
limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
for range 3 {
|
||||
limiter.Count(client, start)
|
||||
limiter.Count(client, start, whole)
|
||||
}
|
||||
|
||||
// A quarter into the next hour, the minute has only this request. The
|
||||
// hour still covers three quarters of the bucket before, with its three
|
||||
// requests, which count 2.25, and this one: 3.25. The day covers all
|
||||
// four.
|
||||
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4))
|
||||
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4), whole)
|
||||
|
||||
want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4}
|
||||
if counts != want {
|
||||
@@ -104,7 +274,7 @@ func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
|
||||
func TestResetSetsTheCountsBackToZero(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerDay: limit})
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerDay: limit}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
@@ -126,7 +296,7 @@ func TestResetSetsTheCountsBackToZero(t *testing.T) {
|
||||
func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit})
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
@@ -145,7 +315,8 @@ func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) {
|
||||
func TestRefusedRequestsCount(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: 2 * limit})
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: 2 * limit},
|
||||
tableSize)
|
||||
refused := netip.MustParsePrefix("203.0.113.9/32")
|
||||
within := netip.MustParsePrefix("203.0.113.10/32")
|
||||
start := midnight()
|
||||
@@ -178,7 +349,7 @@ func TestRefusedRequestsCount(t *testing.T) {
|
||||
func TestRequestCountedLateGoesInTheBucketUnderWay(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit})
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
@@ -194,7 +365,7 @@ func TestRequestCountedLateGoesInTheBucketUnderWay(t *testing.T) {
|
||||
func TestClockSetBackStartsTheBucketsAfresh(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit})
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize)
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
@@ -217,12 +388,12 @@ func TestClockSetBackStartsTheBucketsAfresh(t *testing.T) {
|
||||
wantCount(t, limiter, client, setBack, hour)
|
||||
}
|
||||
|
||||
func TestKeepsAtMost20000ClientsDroppingTheLeastRecentlySeen(t *testing.T) {
|
||||
func TestKeepsAtMostMaxClientsDroppingTheLeastRecentlySeen(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const maxClients = 20000
|
||||
const maxClients = 3
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 1})
|
||||
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 1}, maxClients)
|
||||
now := midnight()
|
||||
|
||||
clients := make([]netip.Prefix, maxClients+1)
|
||||
@@ -244,6 +415,11 @@ func TestKeepsAtMost20000ClientsDroppingTheLeastRecentlySeen(t *testing.T) {
|
||||
// One client more drops the least recently seen, the second, which
|
||||
// starts afresh, while the first is kept.
|
||||
wantCount(t, limiter, clients[maxClients], now, "")
|
||||
|
||||
if limiter.Len() != maxClients {
|
||||
t.Errorf("the table holds %d clients, want %d", limiter.Len(), maxClients)
|
||||
}
|
||||
|
||||
wantCount(t, limiter, clients[1], now, "")
|
||||
wantCount(t, limiter, clients[0], now, minute)
|
||||
}
|
||||
@@ -261,9 +437,24 @@ func wantCount(
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
_, hit, _ := limiter.Count(client, now)
|
||||
_, hit, _ := limiter.Count(client, now, whole)
|
||||
if hit.Window != want {
|
||||
t.Errorf("request from %s at %s is over %q, want %q",
|
||||
client, now.Format(time.RFC3339), hit.Window, want)
|
||||
}
|
||||
}
|
||||
|
||||
// wantBytesCount counts bytes from client at now, and checks the kind of
|
||||
// the limit they break, "" for none.
|
||||
func wantBytesCount(
|
||||
t *testing.T, limiter *ratelimit.Limiter, client netip.Prefix, now time.Time,
|
||||
bytes int64, want string,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
_, hit, _ := limiter.CountBytes(client, now, bytes, whole)
|
||||
if hit.Kind != want {
|
||||
t.Errorf("%d bytes from %s at %s break a limit on %q, want %q",
|
||||
bytes, client, now.Format(time.RFC3339), hit.Kind, want)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,9 +14,9 @@ func TestSnapshotListsTheClientsByAddress(t *testing.T) {
|
||||
|
||||
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{})
|
||||
limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
|
||||
for _, i := range []int{2, 3, 0, 1} {
|
||||
limiter.Count(netip.MustParsePrefix(want[i]), midnight())
|
||||
limiter.Count(netip.MustParsePrefix(want[i]), midnight(), whole)
|
||||
}
|
||||
|
||||
snapshot := limiter.Snapshot()
|
||||
@@ -43,7 +43,7 @@ func TestLoadedCountsCarryOn(t *testing.T) {
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
before := ratelimit.New(ratelimit.Limits{PerHour: limit})
|
||||
before := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize)
|
||||
for range limit {
|
||||
wantCount(t, before, client, start, "")
|
||||
}
|
||||
@@ -51,7 +51,7 @@ func TestLoadedCountsCarryOn(t *testing.T) {
|
||||
// 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 := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize)
|
||||
after.Load(before.Snapshot(), later)
|
||||
wantCount(t, after, client, later, hour)
|
||||
}
|
||||
@@ -62,40 +62,47 @@ func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
start := midnight()
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{})
|
||||
limiter.Count(client, start)
|
||||
limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
|
||||
limiter.Count(client, start, whole)
|
||||
limiter.CountBytes(client, start, 5, whole)
|
||||
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
|
||||
|
||||
loaded := func(now time.Time) ratelimit.Client {
|
||||
t.Helper()
|
||||
|
||||
after := ratelimit.New(ratelimit.Limits{})
|
||||
after := ratelimit.New(ratelimit.Limits{}, tableSize)
|
||||
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.
|
||||
// minute's buckets, of requests and of bytes, 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)
|
||||
}
|
||||
|
||||
if got.MinuteBytes != (ratelimit.Buckets{}) || got.HourBytes.Current != 5 ||
|
||||
got.DayBytes.Current != 5 {
|
||||
t.Errorf("loaded two minutes on with buckets of bytes %+v, %+v and %+v",
|
||||
got.MinuteBytes, got.HourBytes, got.DayBytes)
|
||||
}
|
||||
|
||||
// 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)
|
||||
if got.Minute.Current != 1 || got.MinuteBytes.Current != 5 {
|
||||
t.Errorf("loaded just under two minutes on with minute buckets %+v and %+v",
|
||||
got.Minute, got.MinuteBytes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const maxClients = 20000
|
||||
const maxClients = 3
|
||||
|
||||
// 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
|
||||
@@ -109,7 +116,7 @@ func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) {
|
||||
addr = addr.Next()
|
||||
}
|
||||
|
||||
limiter := ratelimit.New(ratelimit.Limits{})
|
||||
limiter := ratelimit.New(ratelimit.Limits{}, maxClients)
|
||||
limiter.Load(clients, midnight())
|
||||
|
||||
got := limiter.Snapshot()
|
||||
|
||||
@@ -0,0 +1,328 @@
|
||||
package reputation
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"slices"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/hashicorp/golang-lru/v2/simplelru"
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
)
|
||||
|
||||
const (
|
||||
// AbuseIPDBURL is where clients are checked: the check endpoint of
|
||||
// AbuseIPDB's API.
|
||||
AbuseIPDBURL = "https://api.abuseipdb.com/api/v2/check"
|
||||
// AbuseIPDBSource is how the request log, the alerts and the metrics
|
||||
// name AbuseIPDB.
|
||||
AbuseIPDBSource = "abuseipdb"
|
||||
// maxAnswerBytes is the most of an answer of AbuseIPDB that is read.
|
||||
maxAnswerBytes = 64 << 10
|
||||
// day is the length of the day the checks are counted in, in UTC.
|
||||
day = 24 * time.Hour
|
||||
)
|
||||
|
||||
var (
|
||||
errNoScore = errors.New("the answer gives no abuseConfidenceScore")
|
||||
errBudgetUsedUp = errors.New(
|
||||
"checks spent; none is made until the day ends at 00:00 UTC")
|
||||
)
|
||||
|
||||
// Score is what AbuseIPDB said about a client, as reputation.json holds
|
||||
// it: the client, an IPv4 address or an IPv6 group, its abuse confidence
|
||||
// score, from 0 to 100, and when AbuseIPDB answered.
|
||||
type Score struct {
|
||||
Client netip.Prefix `json:"client"`
|
||||
Score int64 `json:"score"`
|
||||
Fetched time.Time `json:"fetched"`
|
||||
}
|
||||
|
||||
// Checks are what reputation.json keeps of the checks of clients with
|
||||
// AbuseIPDB: the day, in UTC, of the checks Spent counts, zero before the
|
||||
// first, and the scores still in use.
|
||||
type Checks struct {
|
||||
Day time.Time `json:"day,omitzero"`
|
||||
Spent int `json:"spent"`
|
||||
Scores []Score `json:"scores"`
|
||||
}
|
||||
|
||||
// AbuseIPDBParams are what NewAbuseIPDB needs.
|
||||
type AbuseIPDBParams struct {
|
||||
// URL is where clients are checked, normally AbuseIPDBURL, with Key,
|
||||
// the account's key (SWWAF_ABUSEIPDB_KEY).
|
||||
URL string
|
||||
Key string
|
||||
// MinScore is the least score that is a hit (SWWAF_ABUSEIPDB_MIN_SCORE),
|
||||
// and DailyBudget the most checks made in a day, in UTC
|
||||
// (SWWAF_ABUSEIPDB_DAILY_BUDGET).
|
||||
MinScore int64
|
||||
DailyBudget int
|
||||
// CacheTTL is how long a score is used after it was fetched
|
||||
// (SWWAF_REPUTATION_CACHE_TTL), and Timeout how long a check may take
|
||||
// (SWWAF_REPUTATION_TIMEOUT).
|
||||
CacheTTL time.Duration
|
||||
Timeout time.Duration
|
||||
// Now tells the time, normally time.Now in UTC.
|
||||
Now func() time.Time
|
||||
// ProcessLog receives each check that fails, and why, and the day's
|
||||
// budget used up.
|
||||
ProcessLog *slog.Logger
|
||||
// Alerts receive a source_failure alert for each.
|
||||
Alerts *alerts.Queue
|
||||
}
|
||||
|
||||
// AbuseIPDB checks clients with AbuseIPDB, in the background, and keeps
|
||||
// their scores. It is safe for concurrent use.
|
||||
type AbuseIPDB struct {
|
||||
params AbuseIPDBParams
|
||||
httpClient *http.Client
|
||||
|
||||
mu sync.Mutex
|
||||
// scores are by client. Each is added as it is fetched and never moved
|
||||
// up, so that the one fetched longest ago is the first dropped.
|
||||
scores *simplelru.LRU[netip.Prefix, Score]
|
||||
// checking are the clients whose check is under way.
|
||||
checking map[netip.Prefix]bool
|
||||
// day is the day, in UTC, of the checks spent counts.
|
||||
day time.Time
|
||||
spent int
|
||||
// checks and failures count the checks made and those that failed,
|
||||
// and retryAt is when a client may be checked again after the last
|
||||
// check failed.
|
||||
checks int
|
||||
failures int
|
||||
retryAt time.Time
|
||||
}
|
||||
|
||||
// NewAbuseIPDB returns an AbuseIPDB with no score yet, and no check spent.
|
||||
func NewAbuseIPDB(params AbuseIPDBParams) *AbuseIPDB {
|
||||
scores, err := simplelru.NewLRU[netip.Prefix, Score](maxVerdicts, nil)
|
||||
if err != nil {
|
||||
panic(err) // NewLRU fails only for a size below one
|
||||
}
|
||||
|
||||
return &AbuseIPDB{
|
||||
params: params,
|
||||
httpClient: &http.Client{},
|
||||
scores: scores,
|
||||
checking: map[netip.Prefix]bool{},
|
||||
}
|
||||
}
|
||||
|
||||
// Hit returns AbuseIPDB's score of client, an IPv4 address or an IPv6
|
||||
// group, and whether it is a hit: MinScore or more. A score is used until
|
||||
// CacheTTL has passed since it was fetched, whichever of the client's
|
||||
// addresses its request comes from. A client without one is checked in
|
||||
// the background, by addr, the address its request came from, if
|
||||
// offender, if it has committed an offence, unless its check is under
|
||||
// way, a check failed less than failureDelay ago, or the day's checks
|
||||
// have used up DailyBudget; Hit never waits for a check. The check that
|
||||
// uses the budget up is logged and raised as a source_failure alert. ctx
|
||||
// is the context of the client's request, and a check goes on after the
|
||||
// request ends.
|
||||
func (a *AbuseIPDB) Hit(
|
||||
ctx context.Context, client netip.Prefix, addr netip.Addr, offender bool,
|
||||
) (int64, bool) {
|
||||
a.mu.Lock()
|
||||
|
||||
now := a.params.Now()
|
||||
|
||||
kept, found := a.scores.Peek(client)
|
||||
if found && now.Sub(kept.Fetched) < a.params.CacheTTL {
|
||||
a.mu.Unlock()
|
||||
|
||||
return kept.Score, kept.Score >= a.params.MinScore
|
||||
}
|
||||
|
||||
if today := now.Truncate(day); !a.day.Equal(today) {
|
||||
a.day, a.spent = today, 0
|
||||
}
|
||||
|
||||
check := offender && !a.checking[client] && !now.Before(a.retryAt) &&
|
||||
a.spent < a.params.DailyBudget
|
||||
if check {
|
||||
a.checking[client] = true
|
||||
a.checks++
|
||||
a.spent++
|
||||
|
||||
go a.check(context.WithoutCancel(ctx), client, addr)
|
||||
}
|
||||
|
||||
usedUp := check && a.spent == a.params.DailyBudget
|
||||
|
||||
a.mu.Unlock()
|
||||
|
||||
if usedUp {
|
||||
a.alert("the daily budget of AbuseIPDB checks is used up",
|
||||
fmt.Errorf("%d %w", a.params.DailyBudget, errBudgetUsedUp))
|
||||
}
|
||||
|
||||
return 0, false
|
||||
}
|
||||
|
||||
// Checked returns how many checks were made.
|
||||
func (a *AbuseIPDB) Checked() int {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
return a.checks
|
||||
}
|
||||
|
||||
// Failures returns how many checks failed.
|
||||
func (a *AbuseIPDB) Failures() int {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
return a.failures
|
||||
}
|
||||
|
||||
// BudgetLeft returns how many checks the day's budget has left.
|
||||
func (a *AbuseIPDB) BudgetLeft() int {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
if !a.day.Equal(a.params.Now().Truncate(day)) {
|
||||
return a.params.DailyBudget
|
||||
}
|
||||
|
||||
return max(a.params.DailyBudget-a.spent, 0)
|
||||
}
|
||||
|
||||
// Snapshot returns the checks spent and every score still in use, sorted
|
||||
// by client, as reputation.json keeps them.
|
||||
func (a *AbuseIPDB) Snapshot() Checks {
|
||||
a.mu.Lock()
|
||||
|
||||
now := a.params.Now()
|
||||
checks := Checks{Day: a.day, Spent: a.spent, Scores: make([]Score, 0, a.scores.Len())}
|
||||
|
||||
for _, kept := range a.scores.Values() {
|
||||
if now.Sub(kept.Fetched) < a.params.CacheTTL {
|
||||
checks.Scores = append(checks.Scores, kept)
|
||||
}
|
||||
}
|
||||
|
||||
a.mu.Unlock()
|
||||
|
||||
slices.SortFunc(checks.Scores, func(x, y Score) int {
|
||||
return x.Client.Compare(y.Client)
|
||||
})
|
||||
|
||||
return checks
|
||||
}
|
||||
|
||||
// Load keeps checks, read from reputation.json, in place of those it
|
||||
// keeps, but for the scores past maxVerdicts, those fetched longest ago.
|
||||
// One fetched CacheTTL ago or more is neither used nor written, as for any
|
||||
// score.
|
||||
func (a *AbuseIPDB) Load(checks Checks) {
|
||||
scores := slices.Clone(checks.Scores)
|
||||
slices.SortStableFunc(scores, func(x, y Score) int {
|
||||
return x.Fetched.Compare(y.Fetched)
|
||||
})
|
||||
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
a.day, a.spent = checks.Day, checks.Spent
|
||||
a.scores.Purge()
|
||||
|
||||
for _, kept := range scores {
|
||||
a.scores.Add(kept.Client, kept)
|
||||
}
|
||||
}
|
||||
|
||||
// check checks client with AbuseIPDB by addr, one of its addresses, keeps
|
||||
// the score as client's, and notes the check as no longer under way. A
|
||||
// check that fails gives no score: it is counted, logged and raised as a
|
||||
// source_failure alert, and no client is checked for failureDelay.
|
||||
func (a *AbuseIPDB) check(ctx context.Context, client netip.Prefix, addr netip.Addr) {
|
||||
score, err := a.ask(ctx, addr)
|
||||
now := a.params.Now()
|
||||
|
||||
a.mu.Lock()
|
||||
|
||||
delete(a.checking, client)
|
||||
|
||||
if err == nil {
|
||||
a.scores.Add(client, Score{Client: client, Score: score, Fetched: now})
|
||||
} else {
|
||||
a.failures++
|
||||
a.retryAt = now.Add(failureDelay)
|
||||
}
|
||||
|
||||
a.mu.Unlock()
|
||||
|
||||
if err != nil {
|
||||
a.alert("checking a client with AbuseIPDB failed", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ask asks AbuseIPDB for addr's abuse confidence score, sending the key
|
||||
// in the header Key. An answer other than 200, one that gives no score,
|
||||
// and none within Timeout, fail.
|
||||
func (a *AbuseIPDB) ask(ctx context.Context, addr netip.Addr) (int64, error) {
|
||||
ctx, cancel := context.WithTimeout(ctx, a.params.Timeout)
|
||||
defer cancel()
|
||||
|
||||
query := url.Values{"ipAddress": {addr.String()}}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet,
|
||||
a.params.URL+"?"+query.Encode(), http.NoBody)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("make the request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Key", a.params.Key)
|
||||
req.Header.Set("Accept", "application/json")
|
||||
|
||||
res, err := a.httpClient.Do(req)
|
||||
if err != nil {
|
||||
// Do's error names the URL, which holds the client's address, which
|
||||
// is not to be logged: only what went wrong is kept.
|
||||
return 0, fmt.Errorf("check the client: %w", errors.Unwrap(err))
|
||||
}
|
||||
|
||||
defer func() {
|
||||
_ = res.Body.Close()
|
||||
}()
|
||||
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return 0, fmt.Errorf("%w %s", errStatus, res.Status)
|
||||
}
|
||||
|
||||
var answer struct {
|
||||
Data struct {
|
||||
AbuseConfidenceScore *int64 `json:"abuseConfidenceScore"`
|
||||
} `json:"data"`
|
||||
}
|
||||
|
||||
err = json.NewDecoder(io.LimitReader(res.Body, maxAnswerBytes)).Decode(&answer)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("read the answer: %w", err)
|
||||
}
|
||||
|
||||
if answer.Data.AbuseConfidenceScore == nil {
|
||||
return 0, errNoScore
|
||||
}
|
||||
|
||||
return *answer.Data.AbuseConfidenceScore, nil
|
||||
}
|
||||
|
||||
// alert raises a source_failure alert from AbuseIPDB with reason and err,
|
||||
// and logs them.
|
||||
func (a *AbuseIPDB) alert(reason string, err error) {
|
||||
// Raised before it is logged, so that the alert is there once the log
|
||||
// line is.
|
||||
raiseFailure(a.params.Alerts, reason, AbuseIPDBSource, err)
|
||||
a.params.ProcessLog.Warn(reason, "source", AbuseIPDBSource, "error", err.Error())
|
||||
}
|
||||
@@ -0,0 +1,663 @@
|
||||
package reputation_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/synctest"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||
"sneak.berlin/go/smallwebwaf/internal/reputation"
|
||||
)
|
||||
|
||||
// The tests of AbuseIPDB run in synctest bubbles, as those of the lists
|
||||
// do, and AbuseIPDB is a stand-in reached without the network, for the
|
||||
// same reason. A bubble's clock starts at midnight UTC, as a day the
|
||||
// checks are counted in starts.
|
||||
|
||||
const (
|
||||
// key is the account's key the tests give, the only one the stand-in
|
||||
// takes.
|
||||
key = "abuseipdb-key-0123456789abcdef"
|
||||
// suspect and other are clients that have committed an offence.
|
||||
suspect = "203.0.113.9"
|
||||
other = "2001:db8::9"
|
||||
)
|
||||
|
||||
func TestOnlyAnOffenderWithoutAScoreIsChecked(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
abuseIPDB := &abuseIPDBStandIn{scores: map[string]int64{suspect: 100}}
|
||||
checker := newAbuseIPDB(abuseIPDB, abuseIPDBParams())
|
||||
|
||||
// A client that has committed no offence is not checked.
|
||||
wantScore(t, checker, suspect, false, 0, false)
|
||||
synctest.Wait()
|
||||
wantChecked(t, abuseIPDB)
|
||||
|
||||
// An offender is, and from then on its score is used, whether or not
|
||||
// it is an offender.
|
||||
wantScore(t, checker, suspect, true, 0, false)
|
||||
synctest.Wait()
|
||||
wantScore(t, checker, suspect, true, 100, true)
|
||||
wantScore(t, checker, suspect, false, 100, true)
|
||||
synctest.Wait()
|
||||
wantChecked(t, abuseIPDB, suspect)
|
||||
})
|
||||
}
|
||||
|
||||
func TestIPv6ClientIsCheckedOnceAndItsScoreUsedForEachOfItsAddresses(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
// 15 addresses of 2001:db8:1:2::/64, one client, each in a part of
|
||||
// it of its own.
|
||||
var addresses []string
|
||||
for i := 1; i < 16; i++ {
|
||||
addresses = append(addresses, fmt.Sprintf("2001:db8:1:2:%x::9", i<<12))
|
||||
}
|
||||
|
||||
abuseIPDB := &abuseIPDBStandIn{scores: map[string]int64{addresses[0]: 100}}
|
||||
checker := newAbuseIPDB(abuseIPDB, abuseIPDBParams())
|
||||
|
||||
// A request from each has the client checked once, by the first.
|
||||
for _, address := range addresses {
|
||||
hitFrom(t, checker, address, true)
|
||||
}
|
||||
|
||||
synctest.Wait()
|
||||
wantChecked(t, abuseIPDB, addresses[0])
|
||||
|
||||
// Its score is the whole client's.
|
||||
for _, address := range addresses {
|
||||
wantScore(t, checker, address, true, 100, true)
|
||||
}
|
||||
|
||||
synctest.Wait()
|
||||
wantChecked(t, abuseIPDB, addresses[0])
|
||||
})
|
||||
}
|
||||
|
||||
func TestScoreAtOrOverTheMinimumIsAHit(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
scores := map[string]int64{"192.0.2.74": 74, "192.0.2.75": 75, "192.0.2.100": 100}
|
||||
p := abuseIPDBParams()
|
||||
p.MinScore = 75
|
||||
checker := newAbuseIPDB(&abuseIPDBStandIn{scores: scores}, p)
|
||||
|
||||
for client := range scores {
|
||||
hitFrom(t, checker, client, true)
|
||||
}
|
||||
|
||||
synctest.Wait()
|
||||
|
||||
for client, score := range scores {
|
||||
wantScore(t, checker, client, true, score, score >= 75)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestScoreUsedUntilTheCacheTTLHasPassedSinceItWasFetched(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
abuseIPDB := &abuseIPDBStandIn{scores: map[string]int64{suspect: 100}}
|
||||
checker := newAbuseIPDB(abuseIPDB, abuseIPDBParams())
|
||||
|
||||
hitFrom(t, checker, suspect, true)
|
||||
synctest.Wait()
|
||||
|
||||
// AbuseIPDB gives another score from now on, but the one kept is
|
||||
// used, and the client is not checked again, until the TTL has
|
||||
// passed.
|
||||
abuseIPDB.setScore(suspect, 80)
|
||||
time.Sleep(cacheTTL - time.Nanosecond)
|
||||
wantScore(t, checker, suspect, true, 100, true)
|
||||
synctest.Wait()
|
||||
wantChecked(t, abuseIPDB, suspect)
|
||||
|
||||
// Then it is not used, and the client is checked again.
|
||||
time.Sleep(time.Nanosecond)
|
||||
wantScore(t, checker, suspect, true, 0, false)
|
||||
synctest.Wait()
|
||||
wantScore(t, checker, suspect, true, 80, true)
|
||||
wantChecked(t, abuseIPDB, suspect, suspect)
|
||||
})
|
||||
}
|
||||
|
||||
func TestDailyBudgetKeptAcrossARestartAndWholeAgainAsTheDayEnds(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
var log bytes.Buffer
|
||||
|
||||
queue := newQueue()
|
||||
p := abuseIPDBParams()
|
||||
p.DailyBudget = 3
|
||||
p.Alerts = queue
|
||||
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
|
||||
abuseIPDB := &abuseIPDBStandIn{scores: map[string]int64{suspect: 100}}
|
||||
checker := newAbuseIPDB(abuseIPDB, p)
|
||||
|
||||
// At noon, the first three offenders spend the budget, and the
|
||||
// fourth, unchecked, is not.
|
||||
time.Sleep(12 * time.Hour)
|
||||
|
||||
const unchecked = "192.0.2.4"
|
||||
|
||||
clients := []string{suspect, "192.0.2.2", "192.0.2.3", unchecked}
|
||||
for _, client := range clients {
|
||||
hitFrom(t, checker, client, true)
|
||||
}
|
||||
|
||||
synctest.Wait()
|
||||
wantChecked(t, abuseIPDB, clients[:3]...)
|
||||
wantBudgetLeft(t, checker, 0)
|
||||
|
||||
// The check that used the budget up raised the alert, and logged it.
|
||||
const usedUp = "the daily budget of AbuseIPDB checks is used up"
|
||||
|
||||
wantFailureAlert(t, queue, time.Now(), usedUp,
|
||||
"3 checks spent; none is made until the day ends at 00:00 UTC", 0)
|
||||
|
||||
if !strings.Contains(log.String(), `"msg":"`+usedUp+`"`) {
|
||||
t.Errorf("logged\n%s\nwant the budget used up", log.String())
|
||||
}
|
||||
|
||||
// Restarted with what reputation.json keeps, it uses the scores, and
|
||||
// checks no client until the day ends.
|
||||
restarted := &abuseIPDBStandIn{}
|
||||
again := newAbuseIPDB(restarted, p)
|
||||
again.Load(checker.Snapshot())
|
||||
|
||||
wantScore(t, again, suspect, true, 100, true)
|
||||
wantBudgetLeft(t, again, 0)
|
||||
time.Sleep(12*time.Hour - time.Nanosecond)
|
||||
wantScore(t, again, unchecked, true, 0, false)
|
||||
synctest.Wait()
|
||||
wantChecked(t, restarted)
|
||||
|
||||
// At midnight the budget is whole again.
|
||||
time.Sleep(time.Nanosecond)
|
||||
wantBudgetLeft(t, again, 3)
|
||||
wantScore(t, again, unchecked, true, 0, false)
|
||||
synctest.Wait()
|
||||
wantChecked(t, restarted, unchecked)
|
||||
wantBudgetLeft(t, again, 2)
|
||||
})
|
||||
}
|
||||
|
||||
func TestFailedCheckGivesNoScoreAndNoClientIsCheckedForAMinute(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
// key is the key sent, status and body what AbuseIPDB answers with,
|
||||
// and error the failure.
|
||||
key, body string
|
||||
status int
|
||||
error string
|
||||
}{
|
||||
{
|
||||
"a refusal, past AbuseIPDB's own limit", key,
|
||||
`{"errors":[{"detail":"Daily rate limit of 1000 requests exceeded"}]}`,
|
||||
http.StatusTooManyRequests, "the server answered 429 Too Many Requests",
|
||||
},
|
||||
{
|
||||
"a refusal of a wrong key", "wrong-key-0123456789abcdef", "", 0,
|
||||
"the server answered 401 Unauthorized",
|
||||
},
|
||||
{
|
||||
"a server failure", key, "", http.StatusInternalServerError,
|
||||
"the server answered 500 Internal Server Error",
|
||||
},
|
||||
{
|
||||
"an answer without a score", key, `{"data":{"ipAddress":"` + suspect + `"}}`,
|
||||
http.StatusOK, "the answer gives no abuseConfidenceScore",
|
||||
},
|
||||
{
|
||||
"an answer that is not JSON", key, "<html>", http.StatusOK,
|
||||
"read the answer: invalid character '<' looking for beginning of value",
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
var log bytes.Buffer
|
||||
|
||||
queue := newQueue()
|
||||
p := abuseIPDBParams()
|
||||
p.Key = tc.key
|
||||
p.Alerts = queue
|
||||
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
|
||||
abuseIPDB := &abuseIPDBStandIn{status: tc.status, body: tc.body}
|
||||
checker := newAbuseIPDB(abuseIPDB, p)
|
||||
|
||||
// The failure gives no score, and no client is checked within a
|
||||
// minute of it.
|
||||
wantScore(t, checker, suspect, true, 0, false)
|
||||
synctest.Wait()
|
||||
time.Sleep(time.Minute - time.Nanosecond)
|
||||
wantScore(t, checker, other, true, 0, false)
|
||||
synctest.Wait()
|
||||
wantChecked(t, abuseIPDB, suspect)
|
||||
wantFailures(t, checker, 1)
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
wantScore(t, checker, other, true, 0, false)
|
||||
synctest.Wait()
|
||||
wantChecked(t, abuseIPDB, suspect, other)
|
||||
wantFailures(t, checker, 2)
|
||||
|
||||
if scores := checker.Snapshot().Scores; len(scores) != 0 {
|
||||
t.Errorf("scores %+v, want none", scores)
|
||||
}
|
||||
|
||||
// One alert for the first failure; the cooldown holds back the
|
||||
// second.
|
||||
wantFailureAlert(t, queue, time.Now().Add(-time.Minute),
|
||||
"checking a client with AbuseIPDB failed", tc.error, 1)
|
||||
|
||||
if !strings.Contains(log.String(), `"msg":"checking a client with `+
|
||||
`AbuseIPDB failed","source":"abuseipdb","error":"`+tc.error) {
|
||||
t.Errorf("logged\n%s\nwant the failures", log.String())
|
||||
}
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckNotAnsweredWithinTheTimeoutFailsAndHitNeverWaits(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
queue := newQueue()
|
||||
p := abuseIPDBParams()
|
||||
p.Alerts = queue
|
||||
abuseIPDB := &abuseIPDBStandIn{hanging: true}
|
||||
checker := newAbuseIPDB(abuseIPDB, p)
|
||||
began := time.Now()
|
||||
|
||||
// The second, while the first's check is under way, starts none.
|
||||
wantScore(t, checker, suspect, true, 0, false)
|
||||
wantScore(t, checker, suspect, true, 0, false)
|
||||
|
||||
if waited := time.Since(began); waited != 0 {
|
||||
t.Errorf("waited %s for the check, want no wait", waited)
|
||||
}
|
||||
|
||||
time.Sleep(timeout - time.Nanosecond)
|
||||
synctest.Wait()
|
||||
wantChecked(t, abuseIPDB, suspect)
|
||||
wantFailures(t, checker, 0)
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
synctest.Wait()
|
||||
wantFailures(t, checker, 1)
|
||||
wantFailureAlert(t, queue, time.Now(), "checking a client with AbuseIPDB failed",
|
||||
"check the client: context deadline exceeded", 0)
|
||||
})
|
||||
}
|
||||
|
||||
func TestKeyIsSentInTheKeyHeaderAndNeverShown(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
var log bytes.Buffer
|
||||
|
||||
queue := newQueue()
|
||||
p := abuseIPDBParams()
|
||||
p.Alerts = queue
|
||||
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
|
||||
abuseIPDB := &abuseIPDBStandIn{scores: map[string]int64{suspect: 100}}
|
||||
checker := newAbuseIPDB(abuseIPDB, p)
|
||||
m := metrics.New(1, "app")
|
||||
m.AddReputation(reputation.New(params()), reputation.NewDNSBL(dnsblParams()))
|
||||
m.AddAbuseIPDB(checker)
|
||||
|
||||
// One check that AbuseIPDB answers, and one it refuses with an answer
|
||||
// that names the key.
|
||||
wantScore(t, checker, suspect, true, 0, false)
|
||||
synctest.Wait()
|
||||
wantScore(t, checker, suspect, true, 100, true)
|
||||
|
||||
abuseIPDB.answerWith(http.StatusUnauthorized, `{"errors":[{"detail":"`+key+`"}]}`)
|
||||
wantScore(t, checker, other, true, 0, false)
|
||||
synctest.Wait()
|
||||
wantFailures(t, checker, 1)
|
||||
|
||||
abuseIPDB.mu.Lock()
|
||||
sent := slices.Clone(abuseIPDB.keys)
|
||||
abuseIPDB.mu.Unlock()
|
||||
|
||||
if !slices.Equal(sent, []string{key, key}) {
|
||||
t.Errorf("checks sent the keys %v, want %s twice", sent, key)
|
||||
}
|
||||
|
||||
alerted, err := json.Marshal(waiting(queue))
|
||||
if err != nil {
|
||||
t.Fatalf("encode the alerts: %v", err)
|
||||
}
|
||||
|
||||
kept, err := json.Marshal(checker.Snapshot())
|
||||
if err != nil {
|
||||
t.Fatalf("encode the checks: %v", err)
|
||||
}
|
||||
|
||||
for name, shown := range map[string]string{
|
||||
"the log": log.String(), "the alerts": string(alerted),
|
||||
"the metrics": scrapeMetrics(t, m), "reputation.json": string(kept),
|
||||
} {
|
||||
if strings.Contains(shown, key) {
|
||||
t.Errorf("%s shows the key:\n%s", name, shown)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestMetricsCountTheChecksTheFailuresAndTheBudgetLeft(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
abuseIPDB := &abuseIPDBStandIn{}
|
||||
p := abuseIPDBParams()
|
||||
p.DailyBudget = 5
|
||||
checker := newAbuseIPDB(abuseIPDB, p)
|
||||
m := metrics.New(1, "app")
|
||||
m.AddReputation(reputation.New(params()), reputation.NewDNSBL(dnsblParams()))
|
||||
m.AddAbuseIPDB(checker)
|
||||
|
||||
// One check that AbuseIPDB answers, and one that fails.
|
||||
hitFrom(t, checker, suspect, true)
|
||||
synctest.Wait()
|
||||
abuseIPDB.answerWith(http.StatusInternalServerError, "")
|
||||
hitFrom(t, checker, other, true)
|
||||
synctest.Wait()
|
||||
|
||||
scraped := scrapeMetrics(t, m)
|
||||
|
||||
for series, want := range map[string]string{
|
||||
"queries_total": "2",
|
||||
"failures_total": "1",
|
||||
"daily_budget_remaining": "3",
|
||||
} {
|
||||
line := "\nsmallwebwaf_reputation_" + series +
|
||||
`{instance="app",source="abuseipdb"} ` + want + "\n"
|
||||
if !strings.Contains(scraped, line) {
|
||||
t.Errorf("metrics\n%s\nwant%s", scraped, line)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestScoreFetchedATTLAgoIsNeitherUsedNorKept(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
|
||||
p := abuseIPDBParams()
|
||||
p.Now = func() time.Time { return now }
|
||||
checker := reputation.NewAbuseIPDB(p)
|
||||
|
||||
// The last score still in use, and one, of other's /64, fetched a TTL
|
||||
// ago.
|
||||
inUse := reputation.Score{
|
||||
Client: netip.MustParsePrefix(suspect + "/32"), Score: 100,
|
||||
Fetched: now.Add(-cacheTTL + time.Nanosecond),
|
||||
}
|
||||
stale := reputation.Score{
|
||||
Client: netip.MustParsePrefix("2001:db8::/64"), Score: 100,
|
||||
Fetched: now.Add(-cacheTTL),
|
||||
}
|
||||
|
||||
checker.Load(reputation.Checks{Scores: []reputation.Score{stale, inUse}})
|
||||
|
||||
wantScore(t, checker, suspect, false, 100, true)
|
||||
wantScore(t, checker, other, false, 0, false)
|
||||
|
||||
got := checker.Snapshot().Scores
|
||||
if !reflect.DeepEqual(got, []reputation.Score{inUse}) {
|
||||
t.Errorf("scores %+v, want only %+v", got, inUse)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAtMost100000ScoresKeptTheOneFetchedLongestAgoDroppedFirst(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
|
||||
p := abuseIPDBParams()
|
||||
p.Now = func() time.Time { return now }
|
||||
checker := reputation.NewAbuseIPDB(p)
|
||||
|
||||
// 100,001 scores, listed by client, as reputation.json lists them, each
|
||||
// fetched a millisecond before the one before it: the last is one too
|
||||
// many.
|
||||
const count = 100001
|
||||
|
||||
scores := make([]reputation.Score, 0, count)
|
||||
|
||||
addr := netip.MustParseAddr("198.18.0.0")
|
||||
for i := range count {
|
||||
scores = append(scores, reputation.Score{
|
||||
Client: netip.PrefixFrom(addr, 32),
|
||||
Fetched: now.Add(-time.Duration(i) * time.Millisecond),
|
||||
})
|
||||
addr = addr.Next()
|
||||
}
|
||||
|
||||
checker.Load(reputation.Checks{Scores: scores})
|
||||
|
||||
got := checker.Snapshot().Scores
|
||||
if len(got) != count-1 || !slices.Contains(got, scores[0]) ||
|
||||
slices.Contains(got, scores[count-1]) {
|
||||
t.Errorf("%d scores kept, want all but the one fetched longest ago", len(got))
|
||||
}
|
||||
}
|
||||
|
||||
// abuseIPDBStandIn is a stand-in for AbuseIPDB. It answers a check sent
|
||||
// with key by the client's score, as scores gives it, 0 for a client it
|
||||
// does not give; a check sent with another key with 401; and, while
|
||||
// status is not 0, every check with status and body; and while hanging,
|
||||
// none at all. It notes each client checked, and the key sent.
|
||||
type abuseIPDBStandIn struct {
|
||||
mu sync.Mutex
|
||||
scores map[string]int64
|
||||
status int
|
||||
body string
|
||||
hanging bool
|
||||
checked []string
|
||||
keys []string
|
||||
}
|
||||
|
||||
// RoundTrip has the stand-in answer req, in place of the network. A check
|
||||
// abandoned before the stand-in answers fails, as over the network.
|
||||
func (s *abuseIPDBStandIn) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
client := req.URL.Query().Get("ipAddress")
|
||||
sent := req.Header.Get("Key")
|
||||
|
||||
s.mu.Lock()
|
||||
s.checked = append(s.checked, client)
|
||||
s.keys = append(s.keys, sent)
|
||||
score := s.scores[client]
|
||||
status, body, hanging := s.status, s.body, s.hanging
|
||||
s.mu.Unlock()
|
||||
|
||||
switch {
|
||||
case hanging:
|
||||
<-req.Context().Done()
|
||||
|
||||
return nil, req.Context().Err()
|
||||
case sent != key:
|
||||
status = http.StatusUnauthorized
|
||||
case status == 0:
|
||||
status = http.StatusOK
|
||||
body = fmt.Sprintf(`{"data":{"ipAddress":%q,"abuseConfidenceScore":%d}}`, client,
|
||||
score)
|
||||
}
|
||||
|
||||
return &http.Response{
|
||||
StatusCode: status,
|
||||
Status: fmt.Sprintf("%d %s", status, http.StatusText(status)),
|
||||
Header: http.Header{},
|
||||
Body: io.NopCloser(strings.NewReader(body)),
|
||||
Request: req,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// setScore has the stand-in give client score.
|
||||
func (s *abuseIPDBStandIn) setScore(client string, score int64) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
s.scores[client] = score
|
||||
}
|
||||
|
||||
// answerWith has the stand-in answer every check with status and body.
|
||||
func (s *abuseIPDBStandIn) answerWith(status int, body string) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
s.status, s.body = status, body
|
||||
}
|
||||
|
||||
// abuseIPDBParams returns the AbuseIPDBParams of the tests: key, a minimum
|
||||
// score of 75, a daily budget of 900, and the cache TTL and timeout of the
|
||||
// DNSBL tests, by the bubble's clock, with alerts to a queue that sends
|
||||
// none.
|
||||
func abuseIPDBParams() reputation.AbuseIPDBParams {
|
||||
return reputation.AbuseIPDBParams{
|
||||
URL: "https://abuseipdb.example/api/v2/check",
|
||||
Key: key,
|
||||
MinScore: 75,
|
||||
DailyBudget: 900,
|
||||
CacheTTL: cacheTTL,
|
||||
Timeout: timeout,
|
||||
Now: time.Now,
|
||||
ProcessLog: slog.New(slog.DiscardHandler),
|
||||
Alerts: newQueue(),
|
||||
}
|
||||
}
|
||||
|
||||
// newAbuseIPDB returns the AbuseIPDB of p, checking clients with
|
||||
// abuseIPDB.
|
||||
func newAbuseIPDB(
|
||||
abuseIPDB *abuseIPDBStandIn, p reputation.AbuseIPDBParams,
|
||||
) *reputation.AbuseIPDB {
|
||||
checker := reputation.NewAbuseIPDB(p)
|
||||
checker.SetTransport(abuseIPDB)
|
||||
|
||||
return checker
|
||||
}
|
||||
|
||||
// wantScore checks the score checker gives client, and whether it is a
|
||||
// hit, as a request from client finds them, offender or not.
|
||||
func wantScore(
|
||||
t *testing.T, checker *reputation.AbuseIPDB, client string, offender bool,
|
||||
score int64, hit bool,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
gotScore, gotHit := hitFrom(t, checker, client, offender)
|
||||
if gotScore != score || gotHit != hit {
|
||||
t.Errorf("%s has the score %d, a hit %t, want %d, %t", client, gotScore, gotHit,
|
||||
score, hit)
|
||||
}
|
||||
}
|
||||
|
||||
// hitFrom is checker's Hit for a request from address, offender or not.
|
||||
// Its client is address for an IPv4 address, and its /64 for an IPv6 one,
|
||||
// as smallwebwaf counts clients.
|
||||
func hitFrom(
|
||||
t *testing.T, checker *reputation.AbuseIPDB, address string, offender bool,
|
||||
) (int64, bool) {
|
||||
t.Helper()
|
||||
|
||||
addr := netip.MustParseAddr(address)
|
||||
|
||||
client := netip.PrefixFrom(addr, addr.BitLen())
|
||||
if addr.Is6() {
|
||||
client = netip.PrefixFrom(addr, 64).Masked()
|
||||
}
|
||||
|
||||
return checker.Hit(t.Context(), client, addr, offender)
|
||||
}
|
||||
|
||||
// wantChecked checks the clients the stand-in was asked about, in any
|
||||
// order.
|
||||
func wantChecked(t *testing.T, abuseIPDB *abuseIPDBStandIn, want ...string) {
|
||||
t.Helper()
|
||||
|
||||
abuseIPDB.mu.Lock()
|
||||
got := slices.Sorted(slices.Values(abuseIPDB.checked))
|
||||
abuseIPDB.mu.Unlock()
|
||||
|
||||
want = slices.Sorted(slices.Values(want))
|
||||
|
||||
if !slices.Equal(got, want) {
|
||||
t.Errorf("checked %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// wantFailures checks how many checks failed.
|
||||
func wantFailures(t *testing.T, checker *reputation.AbuseIPDB, want int) {
|
||||
t.Helper()
|
||||
|
||||
if got := checker.Failures(); got != want {
|
||||
t.Errorf("%d checks failed, want %d", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// wantFailureAlert checks that the one alert waiting in queue is a
|
||||
// source_failure alert from AbuseIPDB, raised at raised, with reason and
|
||||
// the error failure, and that the cooldown has held back held repeats of
|
||||
// it.
|
||||
func wantFailureAlert(
|
||||
t *testing.T, queue *alerts.Queue, raised time.Time, reason, failure string,
|
||||
held int64,
|
||||
) {
|
||||
t.Helper()
|
||||
|
||||
got := waiting(queue)
|
||||
if len(got) != 1 || !got[0].Time.Equal(raised) ||
|
||||
got[0].Event != alerts.EventSourceFailure || got[0].Reason != reason ||
|
||||
got[0].Detail["source"] != reputation.AbuseIPDBSource ||
|
||||
got[0].Detail["error"] != failure || queue.Suppressed() != held {
|
||||
t.Errorf("alerts waiting %+v, %d held back, want only AbuseIPDB's %q with %q, "+
|
||||
"and %d", got, queue.Suppressed(), reason, failure, held)
|
||||
}
|
||||
}
|
||||
|
||||
// wantBudgetLeft checks how many checks the day's budget has left.
|
||||
func wantBudgetLeft(t *testing.T, checker *reputation.AbuseIPDB, want int) {
|
||||
t.Helper()
|
||||
|
||||
if got := checker.BudgetLeft(); got != want {
|
||||
t.Errorf("%d checks left, want %d", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// scrapeMetrics returns the metrics m serves.
|
||||
func scrapeMetrics(t *testing.T, m *metrics.Metrics) string {
|
||||
t.Helper()
|
||||
|
||||
scraped := httptest.NewRecorder()
|
||||
m.ServeHTTP(scraped, httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/",
|
||||
http.NoBody))
|
||||
|
||||
return scraped.Body.String()
|
||||
}
|
||||
@@ -0,0 +1,332 @@
|
||||
package reputation
|
||||
|
||||
import (
|
||||
"cmp"
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/hashicorp/golang-lru/v2/simplelru"
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
)
|
||||
|
||||
const (
|
||||
// maxVerdicts is how many verdicts of the DNSBL zones are kept, and how
|
||||
// many scores of AbuseIPDB. Past it, the one fetched longest ago is
|
||||
// dropped.
|
||||
maxVerdicts = 100000
|
||||
// maxQueries is how many queries may be under way at once. Past it, a
|
||||
// zone is not asked about a client until the client's next request, so
|
||||
// that a swarm of new addresses cannot fill the memory.
|
||||
maxQueries = 1000
|
||||
// failureDelay is how long a zone is not asked again after a query to
|
||||
// it fails, and no client is checked with AbuseIPDB after a check
|
||||
// fails, so that a source refusing them is not asked on every request.
|
||||
failureDelay = time.Minute
|
||||
)
|
||||
|
||||
var (
|
||||
errAsk = errors.New("ask the zone")
|
||||
errRefused = errors.New("the zone refused the query")
|
||||
errNotListing = errors.New("the answer is outside 127.0.0.0/8")
|
||||
)
|
||||
|
||||
// Verdict is what a zone said about a client, as reputation.json holds
|
||||
// it: the zone, the client's address, whether the zone lists it, and when
|
||||
// the zone answered.
|
||||
type Verdict struct {
|
||||
Zone string `json:"zone"`
|
||||
Client netip.Addr `json:"client"`
|
||||
Listed bool `json:"listed"`
|
||||
Fetched time.Time `json:"fetched"`
|
||||
}
|
||||
|
||||
// DNSBLParams are what NewDNSBL needs.
|
||||
type DNSBLParams struct {
|
||||
// Zones are the DNSBL zones clients are asked about in
|
||||
// (SWWAF_DNSBL_ZONES).
|
||||
Zones []string
|
||||
// Resolver is the resolver they are asked through
|
||||
// (SWWAF_DNSBL_RESOLVER), or, while it is the zero AddrPort, the
|
||||
// host's, as /etc/resolv.conf names it.
|
||||
Resolver netip.AddrPort
|
||||
// CacheTTL is how long a verdict is used after it was fetched
|
||||
// (SWWAF_REPUTATION_CACHE_TTL), and Timeout how long a query may take
|
||||
// (SWWAF_REPUTATION_TIMEOUT).
|
||||
CacheTTL time.Duration
|
||||
Timeout time.Duration
|
||||
// Now tells the time, normally time.Now in UTC.
|
||||
Now func() time.Time
|
||||
// ProcessLog receives each query that fails, and why.
|
||||
ProcessLog *slog.Logger
|
||||
// Alerts receive a source_failure alert for each query that fails.
|
||||
Alerts *alerts.Queue
|
||||
}
|
||||
|
||||
// DNSBL asks the DNSBL zones about clients, in the background, and keeps
|
||||
// their verdicts. It is safe for concurrent use.
|
||||
type DNSBL struct {
|
||||
params DNSBLParams
|
||||
resolver *net.Resolver
|
||||
|
||||
mu sync.Mutex
|
||||
// verdicts are by query. Each is added as it is fetched and never moved
|
||||
// up, so that the one fetched longest ago is the first dropped.
|
||||
verdicts *simplelru.LRU[query, Verdict]
|
||||
// asking are the queries under way.
|
||||
asking map[query]bool
|
||||
// queries and failures count, by zone, the queries made and those that
|
||||
// failed, and retryAt is when a zone whose last query failed may be
|
||||
// asked again.
|
||||
queries map[string]int
|
||||
failures map[string]int
|
||||
retryAt map[string]time.Time
|
||||
}
|
||||
|
||||
// query is a client's address, to ask a zone about.
|
||||
type query struct {
|
||||
zone string
|
||||
client netip.Addr
|
||||
}
|
||||
|
||||
// NewDNSBL returns a DNSBL with no verdict yet.
|
||||
func NewDNSBL(params DNSBLParams) *DNSBL {
|
||||
verdicts, err := simplelru.NewLRU[query, Verdict](maxVerdicts, nil)
|
||||
if err != nil {
|
||||
panic(err) // NewLRU fails only for a size below one
|
||||
}
|
||||
|
||||
resolver := &net.Resolver{}
|
||||
if params.Resolver.IsValid() {
|
||||
// Dial is used by Go's own resolver alone.
|
||||
resolver.PreferGo = true
|
||||
resolver.Dial = func(ctx context.Context, network, _ string) (net.Conn, error) {
|
||||
var dialer net.Dialer
|
||||
|
||||
return dialer.DialContext(ctx, network, params.Resolver.String())
|
||||
}
|
||||
}
|
||||
|
||||
return &DNSBL{
|
||||
params: params,
|
||||
resolver: resolver,
|
||||
verdicts: verdicts,
|
||||
asking: map[query]bool{},
|
||||
queries: map[string]int{},
|
||||
failures: map[string]int{},
|
||||
retryAt: map[string]time.Time{},
|
||||
}
|
||||
}
|
||||
|
||||
// Zones returns the zones, in the order SWWAF_DNSBL_ZONES names them.
|
||||
func (d *DNSBL) Zones() []string {
|
||||
return slices.Clone(d.params.Zones)
|
||||
}
|
||||
|
||||
// ListedBy returns the zones whose verdict on addr, a client's address,
|
||||
// lists it, in the order SWWAF_DNSBL_ZONES names them, each with its key
|
||||
// masked, as config.MaskZoneKey masks it, since they go to the request
|
||||
// log, the alerts and the metrics. A verdict is used until CacheTTL has
|
||||
// passed since it was fetched. Each zone without one is asked about addr
|
||||
// in the background, unless a query about addr to it is under way, the
|
||||
// zone is left alone after a failure, or maxQueries are under way;
|
||||
// ListedBy never waits for a query. ctx is the context of the client's
|
||||
// request, and a query goes on after the request ends.
|
||||
func (d *DNSBL) ListedBy(ctx context.Context, addr netip.Addr) []string {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
|
||||
now := d.params.Now()
|
||||
|
||||
var listedBy []string
|
||||
|
||||
for _, zone := range d.params.Zones {
|
||||
q := query{zone: zone, client: addr}
|
||||
|
||||
kept, found := d.verdicts.Peek(q)
|
||||
|
||||
switch {
|
||||
case found && now.Sub(kept.Fetched) < d.params.CacheTTL:
|
||||
if kept.Listed {
|
||||
listedBy = append(listedBy, config.MaskZoneKey(zone))
|
||||
}
|
||||
case !d.asking[q] && !now.Before(d.retryAt[zone]) && len(d.asking) < maxQueries:
|
||||
d.asking[q] = true
|
||||
d.queries[zone]++
|
||||
|
||||
go d.ask(context.WithoutCancel(ctx), q)
|
||||
}
|
||||
}
|
||||
|
||||
return listedBy
|
||||
}
|
||||
|
||||
// Queries returns how many queries were made to zone.
|
||||
func (d *DNSBL) Queries(zone string) int {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
|
||||
return d.queries[zone]
|
||||
}
|
||||
|
||||
// Failures returns how many queries to zone failed.
|
||||
func (d *DNSBL) Failures(zone string) int {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
|
||||
return d.failures[zone]
|
||||
}
|
||||
|
||||
// Snapshot returns every verdict still in use, sorted by client, then by
|
||||
// zone, as reputation.json lists them.
|
||||
func (d *DNSBL) Snapshot() []Verdict {
|
||||
d.mu.Lock()
|
||||
|
||||
now := d.params.Now()
|
||||
verdicts := make([]Verdict, 0, d.verdicts.Len())
|
||||
|
||||
for _, kept := range d.verdicts.Values() {
|
||||
if now.Sub(kept.Fetched) < d.params.CacheTTL {
|
||||
verdicts = append(verdicts, kept)
|
||||
}
|
||||
}
|
||||
|
||||
d.mu.Unlock()
|
||||
|
||||
slices.SortFunc(verdicts, func(a, b Verdict) int {
|
||||
return cmp.Or(a.Client.Compare(b.Client), strings.Compare(a.Zone, b.Zone))
|
||||
})
|
||||
|
||||
return verdicts
|
||||
}
|
||||
|
||||
// Load keeps verdicts, read from reputation.json, in place of those it
|
||||
// keeps, but for those of a zone SWWAF_DNSBL_ZONES does not name, and,
|
||||
// past maxVerdicts, those fetched longest ago. One fetched CacheTTL ago or
|
||||
// more is neither used nor written, as for any verdict.
|
||||
func (d *DNSBL) Load(verdicts []Verdict) {
|
||||
verdicts = slices.Clone(verdicts)
|
||||
slices.SortStableFunc(verdicts, func(a, b Verdict) int {
|
||||
return a.Fetched.Compare(b.Fetched)
|
||||
})
|
||||
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
|
||||
d.verdicts.Purge()
|
||||
|
||||
for _, kept := range verdicts {
|
||||
if slices.Contains(d.params.Zones, kept.Zone) {
|
||||
d.verdicts.Add(query{zone: kept.Zone, client: kept.Client}, kept)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ask asks q's zone about q's client, keeps the verdict, and notes the
|
||||
// query as no longer under way. A query that fails gives no verdict: it
|
||||
// is counted, logged and raised as a source_failure alert, which show the
|
||||
// zone with its key masked, and the zone is not asked again for
|
||||
// failureDelay.
|
||||
func (d *DNSBL) ask(ctx context.Context, q query) {
|
||||
listed, err := d.lookUp(ctx, q)
|
||||
now := d.params.Now()
|
||||
|
||||
d.mu.Lock()
|
||||
|
||||
delete(d.asking, q)
|
||||
|
||||
if err == nil {
|
||||
d.verdicts.Add(q, Verdict{
|
||||
Zone: q.zone, Client: q.client, Listed: listed, Fetched: now,
|
||||
})
|
||||
} else {
|
||||
d.failures[q.zone]++
|
||||
d.retryAt[q.zone] = now.Add(failureDelay)
|
||||
}
|
||||
|
||||
d.mu.Unlock()
|
||||
|
||||
if err != nil {
|
||||
const failed = "asking a DNSBL zone failed"
|
||||
|
||||
shown := config.MaskZoneKey(q.zone)
|
||||
|
||||
// Raised before it is logged, so that the alert is there once the
|
||||
// log line is.
|
||||
raiseFailure(d.params.Alerts, failed, shown, err)
|
||||
d.params.ProcessLog.Warn(failed, "zone", shown, "error", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
// lookUp asks q's zone about q's client through the resolver, and returns
|
||||
// whether the zone lists it, as readAnswer reads the answer. No such name
|
||||
// is a client the zone does not list. A query not answered within Timeout
|
||||
// fails.
|
||||
func (d *DNSBL) lookUp(ctx context.Context, q query) (bool, error) {
|
||||
ctx, cancel := context.WithTimeout(ctx, d.params.Timeout)
|
||||
defer cancel()
|
||||
|
||||
answer, err := d.resolver.LookupNetIP(ctx, "ip4", queryName(q.zone, q.client))
|
||||
|
||||
var dnsErr *net.DNSError
|
||||
|
||||
switch {
|
||||
case err == nil:
|
||||
return readAnswer(answer)
|
||||
case errors.As(err, &dnsErr) && dnsErr.IsNotFound:
|
||||
return false, nil
|
||||
case errors.As(err, &dnsErr):
|
||||
// The error names the name asked about, which holds the client's
|
||||
// address, which is not to be logged: only what went wrong is kept.
|
||||
return false, fmt.Errorf("%w: %s", errAsk, dnsErr.Err)
|
||||
default:
|
||||
return false, fmt.Errorf("%w: %w", errAsk, err)
|
||||
}
|
||||
}
|
||||
|
||||
// queryName returns the name a zone is asked about addr by, as RFC 5782
|
||||
// builds it: the four numbers of an IPv4 address, or the 32 hex digits of
|
||||
// an IPv6 address, in reverse order, each followed by a dot, then the zone
|
||||
// and a dot, which makes it a full name, to which the resolver adds no
|
||||
// search domain of /etc/resolv.conf.
|
||||
func queryName(zone string, addr netip.Addr) string {
|
||||
parts := strings.Split(addr.String(), ".")
|
||||
if addr.Is6() {
|
||||
parts = strings.Split(hex.EncodeToString(addr.AsSlice()), "")
|
||||
}
|
||||
|
||||
slices.Reverse(parts)
|
||||
|
||||
return strings.Join(parts, ".") + "." + zone + "."
|
||||
}
|
||||
|
||||
// readAnswer reads the addresses a zone answered with. An address in
|
||||
// 127.0.0.0/8 lists the client, as RFC 5782 has zones answer, but one in
|
||||
// 127.255.255.0/24 is how Spamhaus refuses a query, such as one sent
|
||||
// through a public resolver or one past its limit, and is a failure. So is
|
||||
// an address outside 127.0.0.0/8, such as a resolver gives that answers
|
||||
// even for names that do not exist.
|
||||
func readAnswer(answer []netip.Addr) (bool, error) {
|
||||
listing := netip.MustParsePrefix("127.0.0.0/8")
|
||||
refusal := netip.MustParsePrefix("127.255.255.0/24")
|
||||
|
||||
for _, addr := range answer {
|
||||
switch {
|
||||
case refusal.Contains(addr):
|
||||
return false, fmt.Errorf("%w: %s", errRefused, addr)
|
||||
case !listing.Contains(addr):
|
||||
return false, fmt.Errorf("%w: %s", errNotListing, addr)
|
||||
}
|
||||
}
|
||||
|
||||
return len(answer) > 0, nil
|
||||
}
|
||||
@@ -0,0 +1,737 @@
|
||||
package reputation_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/synctest"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||
"sneak.berlin/go/smallwebwaf/internal/reputation"
|
||||
)
|
||||
|
||||
// The tests of the DNSBL zones run in synctest bubbles, as those of the
|
||||
// lists do, and the resolver the zones are asked through is a stand-in
|
||||
// reached through an in-memory connection, net.Pipe's, for the same
|
||||
// reason. They run one at a time, none in parallel with another test of
|
||||
// this package: Go's resolver counts the queries under way in one
|
||||
// sync.WaitGroup for the whole process, and the process fails when
|
||||
// queries from two bubbles, or from a bubble and from outside one, are
|
||||
// under way at once. TestMain has the resolver make its configuration,
|
||||
// which it makes on its first query, outside every bubble, since the
|
||||
// configuration holds a channel, which the bubble it was made in would
|
||||
// keep to itself.
|
||||
|
||||
const (
|
||||
// zone and otherZone are the DNSBL zones the tests name.
|
||||
zone = "dnsbl.example"
|
||||
otherZone = "other.example"
|
||||
// cacheTTL is the tests' SWWAF_REPUTATION_CACHE_TTL, and timeout their
|
||||
// SWWAF_REPUTATION_TIMEOUT: a second, the least time /etc/resolv.conf
|
||||
// can have Go's resolver wait for one server, so that it is the
|
||||
// DNSBL's own timeout that ends a query, whatever that file says.
|
||||
cacheTTL = 24 * time.Hour
|
||||
timeout = time.Second
|
||||
// listed and unlisted are clients zone is asked about by the names
|
||||
// listedName and unlistedName, and most tests have zone list the first
|
||||
// alone, by answering with listing.
|
||||
listed = "192.0.2.99"
|
||||
unlisted = "192.0.2.100"
|
||||
listedName = "99.2.0.192." + zone + "."
|
||||
unlistedName = "100.2.0.192." + zone + "."
|
||||
listing = "127.0.0.2"
|
||||
)
|
||||
|
||||
// The DNS response codes the stand-in answers with, besides no error.
|
||||
const (
|
||||
serverFailure = 2
|
||||
noSuchName = 3
|
||||
refused = 5
|
||||
)
|
||||
|
||||
var errNoNetwork = errors.New("the test dials nothing")
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
// A query that fails at once, as nothing is dialled for it.
|
||||
resolver := &net.Resolver{
|
||||
PreferGo: true,
|
||||
Dial: func(context.Context, string, string) (net.Conn, error) {
|
||||
return nil, errNoNetwork
|
||||
},
|
||||
}
|
||||
_, _ = resolver.LookupNetIP(context.Background(), "ip4", "warm-up.invalid.")
|
||||
|
||||
m.Run()
|
||||
}
|
||||
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
func TestZonesListOrNotClientsByTheirIPv4AndIPv6Addresses(t *testing.T) {
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
// The addresses of the examples of RFC 5782, and the names it
|
||||
// gives for them.
|
||||
const (
|
||||
v4 = "192.0.2.99"
|
||||
v6 = "2001:db8:1:2:3:4:567:89ab"
|
||||
// v6Name is the hex digits of v6, in reverse order.
|
||||
v6Name = "b.a.9.8.7.6.5.0.4.0.0.0.3.0.0.0.2.0.0.0.1.0.0.0.8.b.d.0.1.0.0.2."
|
||||
)
|
||||
|
||||
resolver := &resolverStandIn{answers: map[string]answer{
|
||||
"99.2.0.192." + zone + ".": {addrs: []string{listing}},
|
||||
v6Name + otherZone + ".": {addrs: []string{"127.0.0.4", "127.0.0.10"}},
|
||||
"99.2.0.192." + otherZone + ".": {},
|
||||
}}
|
||||
dnsbl := newDNSBL(resolver, dnsblParams(zone, otherZone))
|
||||
|
||||
// Neither client has a verdict yet, so neither is listed, and each
|
||||
// zone is asked about each.
|
||||
wantZones(t, dnsbl, v4)
|
||||
wantZones(t, dnsbl, v6)
|
||||
synctest.Wait()
|
||||
|
||||
wantZones(t, dnsbl, v4, zone)
|
||||
wantZones(t, dnsbl, v6, otherZone)
|
||||
wantAsked(t, resolver,
|
||||
"99.2.0.192."+zone+".", "99.2.0.192."+otherZone+".",
|
||||
v6Name+zone+".", v6Name+otherZone+".")
|
||||
})
|
||||
}
|
||||
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
func TestListedByNeverWaitsForAQuery(t *testing.T) {
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
dnsbl := newDNSBL(&resolverStandIn{hanging: true}, dnsblParams(zone))
|
||||
began := time.Now()
|
||||
|
||||
// The second, while the first's query is under way, starts none.
|
||||
wantZones(t, dnsbl, listed)
|
||||
wantZones(t, dnsbl, listed)
|
||||
|
||||
if waited := time.Since(began); waited != 0 {
|
||||
t.Errorf("waited %s for the query, want no wait", waited)
|
||||
}
|
||||
|
||||
synctest.Wait()
|
||||
wantQueries(t, dnsbl, 1, 0)
|
||||
waitForTheResolver()
|
||||
})
|
||||
}
|
||||
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
func TestVerdictUsedUntilTheCacheTTLHasPassedSinceItWasFetched(t *testing.T) {
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
resolver := &resolverStandIn{answers: map[string]answer{
|
||||
listedName: {addrs: []string{listing}},
|
||||
}}
|
||||
dnsbl := newDNSBL(resolver, dnsblParams(zone))
|
||||
|
||||
wantZones(t, dnsbl, listed)
|
||||
wantZones(t, dnsbl, unlisted)
|
||||
synctest.Wait()
|
||||
|
||||
// The zone lists the other client from now on, but the verdicts
|
||||
// kept are used, and the zone is not asked again, until the TTL
|
||||
// has passed.
|
||||
resolver.set(listedName, answer{rcode: noSuchName})
|
||||
resolver.set(unlistedName, answer{addrs: []string{listing}})
|
||||
time.Sleep(cacheTTL - time.Nanosecond)
|
||||
|
||||
wantZones(t, dnsbl, listed, zone)
|
||||
wantZones(t, dnsbl, unlisted)
|
||||
synctest.Wait()
|
||||
wantQueries(t, dnsbl, 2, 0)
|
||||
|
||||
// Then neither verdict is used, and both clients are asked about
|
||||
// again.
|
||||
time.Sleep(time.Nanosecond)
|
||||
wantZones(t, dnsbl, listed)
|
||||
wantZones(t, dnsbl, unlisted)
|
||||
synctest.Wait()
|
||||
|
||||
wantQueries(t, dnsbl, 4, 0)
|
||||
wantZones(t, dnsbl, listed)
|
||||
wantZones(t, dnsbl, unlisted, zone)
|
||||
})
|
||||
}
|
||||
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
func TestQueryNotAnsweredWithinTheTimeoutFails(t *testing.T) {
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
queue := newQueue()
|
||||
p := dnsblParams(zone)
|
||||
p.Alerts = queue
|
||||
dnsbl := newDNSBL(&resolverStandIn{hanging: true}, p)
|
||||
|
||||
wantZones(t, dnsbl, listed)
|
||||
time.Sleep(timeout - time.Nanosecond)
|
||||
synctest.Wait()
|
||||
wantQueries(t, dnsbl, 1, 0)
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
synctest.Wait()
|
||||
wantQueries(t, dnsbl, 1, 1)
|
||||
|
||||
if got := waiting(queue); len(got) != 1 ||
|
||||
got[0].Detail["error"] != "ask the zone: i/o timeout" {
|
||||
t.Errorf("alerts waiting %+v, want the timeout's", got)
|
||||
}
|
||||
|
||||
if verdicts := dnsbl.Snapshot(); len(verdicts) != 0 {
|
||||
t.Errorf("verdicts %+v, want none", verdicts)
|
||||
}
|
||||
|
||||
waitForTheResolver()
|
||||
})
|
||||
}
|
||||
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
func TestZoneThatFailsOrRefusesGivesNoVerdictAndIsLeftAloneForAMinute(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
answer answer
|
||||
error string
|
||||
}{
|
||||
{
|
||||
"a server failure", answer{rcode: serverFailure},
|
||||
"ask the zone: server misbehaving",
|
||||
},
|
||||
{"a refusal", answer{rcode: refused}, "ask the zone: server misbehaving"},
|
||||
{
|
||||
"an answer in 127.255.255.0/24, with which Spamhaus refuses a query",
|
||||
answer{addrs: []string{"127.255.255.254"}},
|
||||
"the zone refused the query: 127.255.255.254",
|
||||
},
|
||||
{
|
||||
"an answer outside 127.0.0.0/8, as for a name that does not exist",
|
||||
answer{addrs: []string{"192.0.2.1"}},
|
||||
"the answer is outside 127.0.0.0/8: 192.0.2.1",
|
||||
},
|
||||
} {
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
var log bytes.Buffer
|
||||
|
||||
queue := newQueue()
|
||||
p := dnsblParams(zone)
|
||||
p.Alerts = queue
|
||||
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
|
||||
dnsbl := newDNSBL(&resolverStandIn{answers: map[string]answer{
|
||||
listedName: tc.answer,
|
||||
}}, p)
|
||||
|
||||
// The failure gives no verdict, and the zone is not asked
|
||||
// again within a minute of it.
|
||||
wantZones(t, dnsbl, listed)
|
||||
synctest.Wait()
|
||||
time.Sleep(time.Minute - time.Nanosecond)
|
||||
wantZones(t, dnsbl, listed)
|
||||
synctest.Wait()
|
||||
wantQueries(t, dnsbl, 1, 1)
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
wantZones(t, dnsbl, listed)
|
||||
synctest.Wait()
|
||||
wantQueries(t, dnsbl, 2, 2)
|
||||
|
||||
if verdicts := dnsbl.Snapshot(); len(verdicts) != 0 {
|
||||
t.Errorf("verdicts %+v, want none", verdicts)
|
||||
}
|
||||
|
||||
// One alert for the first failure; the cooldown holds back
|
||||
// the second.
|
||||
wantAlert(t, queue, alerts.Alert{
|
||||
Time: time.Now().Add(-time.Minute),
|
||||
Event: alerts.EventSourceFailure,
|
||||
Reason: "asking a DNSBL zone failed",
|
||||
Detail: map[string]any{"source": zone, "error": tc.error},
|
||||
})
|
||||
|
||||
if !strings.Contains(log.String(), `"msg":"asking a DNSBL zone failed",`+
|
||||
`"zone":"`+zone+`","error":"`+tc.error) {
|
||||
t.Errorf("logged\n%s\nwant the failures", log.String())
|
||||
}
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
func TestAtMost1000QueriesUnderWay(t *testing.T) {
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
dnsbl := newDNSBL(&resolverStandIn{hanging: true}, dnsblParams(zone))
|
||||
|
||||
client := netip.MustParseAddr("198.18.0.0")
|
||||
for range 1001 {
|
||||
dnsbl.ListedBy(t.Context(), client)
|
||||
client = client.Next()
|
||||
}
|
||||
|
||||
synctest.Wait()
|
||||
wantQueries(t, dnsbl, 1000, 0)
|
||||
waitForTheResolver()
|
||||
})
|
||||
}
|
||||
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
func TestMetricsCountEachZonesQueriesAndThoseThatFailed(t *testing.T) {
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
resolver := &resolverStandIn{answers: map[string]answer{
|
||||
listedName: {addrs: []string{listing}},
|
||||
"99.2.0.192." + otherZone + ".": {rcode: serverFailure},
|
||||
}}
|
||||
dnsbl := newDNSBL(resolver, dnsblParams(zone, otherZone))
|
||||
m := metrics.New(1, "app")
|
||||
m.AddReputation(reputation.New(params()), dnsbl)
|
||||
|
||||
wantZones(t, dnsbl, listed)
|
||||
synctest.Wait()
|
||||
|
||||
scraped := httptest.NewRecorder()
|
||||
m.ServeHTTP(scraped, httptest.NewRequestWithContext(t.Context(), http.MethodGet,
|
||||
"/", http.NoBody))
|
||||
|
||||
for series, want := range map[string]string{
|
||||
"queries_total" + `{instance="app",source="` + zone + `"}`: "1",
|
||||
"failures_total" + `{instance="app",source="` + zone + `"}`: "0",
|
||||
"queries_total" + `{instance="app",source="` + otherZone + `"}`: "1",
|
||||
"failures_total" + `{instance="app",source="` + otherZone + `"}`: "1",
|
||||
} {
|
||||
line := "\nsmallwebwaf_reputation_" + series + " " + want + "\n"
|
||||
if !strings.Contains(scraped.Body.String(), line) {
|
||||
t.Errorf("metrics\n%s\nwant%s", scraped.Body.String(), line)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
func TestZoneKeyIsMaskedInTheVerdictsTheFailuresAndTheMetrics(t *testing.T) {
|
||||
const (
|
||||
key = "abcdefghijklmnopqrstuvwxyz"
|
||||
keyed = key + ".xbl.dq.spamhaus.net"
|
||||
masked = "********.xbl.dq.spamhaus.net"
|
||||
)
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
var log bytes.Buffer
|
||||
|
||||
queue := newQueue()
|
||||
p := dnsblParams(keyed)
|
||||
p.Alerts = queue
|
||||
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
|
||||
dnsbl := newDNSBL(&resolverStandIn{answers: map[string]answer{
|
||||
"99.2.0.192." + keyed + ".": {addrs: []string{listing}},
|
||||
"100.2.0.192." + keyed + ".": {rcode: serverFailure},
|
||||
}}, p)
|
||||
m := metrics.New(1, "app")
|
||||
m.AddReputation(reputation.New(params()), dnsbl)
|
||||
|
||||
// Both clients are asked about before either answer comes, so that
|
||||
// the failure does not keep the zone from the other query.
|
||||
wantZones(t, dnsbl, listed)
|
||||
wantZones(t, dnsbl, unlisted)
|
||||
synctest.Wait()
|
||||
wantZones(t, dnsbl, listed, masked)
|
||||
|
||||
if got := waiting(queue); len(got) != 1 || got[0].Detail["source"] != masked {
|
||||
t.Errorf("alerts waiting %+v, want the failure's, from %s", got, masked)
|
||||
}
|
||||
|
||||
scraped := httptest.NewRecorder()
|
||||
m.ServeHTTP(scraped, httptest.NewRequestWithContext(t.Context(), http.MethodGet,
|
||||
"/", http.NoBody))
|
||||
|
||||
for name, shown := range map[string]string{
|
||||
"the log": log.String(), "the metrics": scraped.Body.String(),
|
||||
} {
|
||||
if strings.Contains(shown, key) || !strings.Contains(shown, masked) {
|
||||
t.Errorf("%s shows the key, or does not name the zone:\n%s", name, shown)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
func TestVerdictsKeptAcrossARestart(t *testing.T) {
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
resolver := &resolverStandIn{answers: map[string]answer{
|
||||
listedName: {addrs: []string{listing}},
|
||||
}}
|
||||
dnsbl := newDNSBL(resolver, dnsblParams(zone))
|
||||
fetched := time.Now()
|
||||
|
||||
wantZones(t, dnsbl, unlisted)
|
||||
wantZones(t, dnsbl, listed)
|
||||
synctest.Wait()
|
||||
|
||||
kept := dnsbl.Snapshot()
|
||||
|
||||
want := []reputation.Verdict{
|
||||
{Zone: zone, Client: netip.MustParseAddr(listed), Listed: true, Fetched: fetched},
|
||||
{Zone: zone, Client: netip.MustParseAddr(unlisted), Fetched: fetched},
|
||||
}
|
||||
if !reflect.DeepEqual(kept, want) {
|
||||
t.Errorf("verdicts %+v, want %+v", kept, want)
|
||||
}
|
||||
|
||||
// Restarted an hour later with what reputation.json keeps, it uses
|
||||
// the verdicts, and asks the zone nothing, until the TTL has passed
|
||||
// since they were fetched.
|
||||
time.Sleep(time.Hour)
|
||||
|
||||
restarted := &resolverStandIn{}
|
||||
again := newDNSBL(restarted, dnsblParams(zone))
|
||||
again.Load(kept)
|
||||
|
||||
wantZones(t, again, listed, zone)
|
||||
wantZones(t, again, unlisted)
|
||||
synctest.Wait()
|
||||
wantAsked(t, restarted)
|
||||
|
||||
time.Sleep(cacheTTL - time.Hour)
|
||||
wantZones(t, again, listed)
|
||||
synctest.Wait()
|
||||
wantAsked(t, restarted, listedName)
|
||||
})
|
||||
}
|
||||
|
||||
func TestNeitherAVerdictOfAZoneNotNamedNorOnePastItsTTLIsKept(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
|
||||
p := dnsblParams(zone)
|
||||
p.Now = func() time.Time { return now }
|
||||
dnsbl := reputation.NewDNSBL(p)
|
||||
|
||||
// The last verdict still in use, one fetched a TTL ago, and one of a
|
||||
// zone SWWAF_DNSBL_ZONES does not name.
|
||||
inUse := reputation.Verdict{
|
||||
Zone: zone, Client: netip.MustParseAddr(listed), Listed: true,
|
||||
Fetched: now.Add(-cacheTTL + time.Nanosecond),
|
||||
}
|
||||
stale := reputation.Verdict{
|
||||
Zone: zone, Client: netip.MustParseAddr(unlisted), Fetched: now.Add(-cacheTTL),
|
||||
}
|
||||
notNamed := reputation.Verdict{
|
||||
Zone: otherZone, Client: netip.MustParseAddr(listed), Listed: true, Fetched: now,
|
||||
}
|
||||
|
||||
dnsbl.Load([]reputation.Verdict{notNamed, stale, inUse})
|
||||
|
||||
if got := dnsbl.Snapshot(); !reflect.DeepEqual(got, []reputation.Verdict{inUse}) {
|
||||
t.Errorf("verdicts %+v, want only %+v", got, inUse)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAtMost100000VerdictsKeptTheOneFetchedLongestAgoDroppedFirst(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
|
||||
p := dnsblParams(zone)
|
||||
p.Now = func() time.Time { return now }
|
||||
dnsbl := reputation.NewDNSBL(p)
|
||||
|
||||
// 100,001 verdicts, listed by client, as reputation.json lists them,
|
||||
// each fetched a millisecond before the one before it: the last is one
|
||||
// too many.
|
||||
const count = 100001
|
||||
|
||||
verdicts := make([]reputation.Verdict, 0, count)
|
||||
|
||||
client := netip.MustParseAddr("198.18.0.0")
|
||||
for i := range count {
|
||||
verdicts = append(verdicts, reputation.Verdict{
|
||||
Zone: zone, Client: client, Fetched: now.Add(-time.Duration(i) * time.Millisecond),
|
||||
})
|
||||
client = client.Next()
|
||||
}
|
||||
|
||||
dnsbl.Load(verdicts)
|
||||
|
||||
got := dnsbl.Snapshot()
|
||||
if len(got) != count-1 || !slices.Contains(got, verdicts[0]) ||
|
||||
slices.Contains(got, verdicts[count-1]) {
|
||||
t.Errorf("%d verdicts kept, want all but the one fetched longest ago", len(got))
|
||||
}
|
||||
}
|
||||
|
||||
//nolint:paralleltest // one at a time, as the comment at the top of this file says
|
||||
func TestQueriesGoToTheResolverSWWAFDNSBLResolverNames(t *testing.T) {
|
||||
resolver := &resolverStandIn{answers: map[string]answer{
|
||||
listedName: {addrs: []string{listing}},
|
||||
}}
|
||||
|
||||
conn, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
|
||||
served := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
resolver.serveUDP(conn)
|
||||
close(served)
|
||||
}()
|
||||
|
||||
t.Cleanup(func() {
|
||||
_ = conn.Close()
|
||||
|
||||
<-served
|
||||
})
|
||||
|
||||
p := dnsblParams(zone)
|
||||
p.Resolver = netip.MustParseAddrPort(conn.LocalAddr().String())
|
||||
// On the real clock: the stand-in answers at once, so only a test
|
||||
// process held up for a whole minute would see the query fail.
|
||||
p.Timeout = time.Minute
|
||||
|
||||
isListed, err := reputation.NewDNSBL(p).LookUp(zone, netip.MustParseAddr(listed))
|
||||
if err != nil || !isListed {
|
||||
t.Errorf("listed %t (%v), want true", isListed, err)
|
||||
}
|
||||
|
||||
wantAsked(t, resolver, listedName)
|
||||
}
|
||||
|
||||
// resolverStandIn is a stand-in for the resolver the zones are asked
|
||||
// through. It answers each query by the name asked about, as answers
|
||||
// gives, with no such name for a name answers does not give, and not at
|
||||
// all while hanging. It notes each name asked about.
|
||||
type resolverStandIn struct {
|
||||
mu sync.Mutex
|
||||
answers map[string]answer
|
||||
hanging bool
|
||||
names []string
|
||||
}
|
||||
|
||||
// answer is how the stand-in answers a name: with an A record of each of
|
||||
// addrs, or with the response code rcode, unless it is 0, for no error.
|
||||
type answer struct {
|
||||
addrs []string
|
||||
rcode uint16
|
||||
}
|
||||
|
||||
// What the stand-in reads of a query, and writes in its reply.
|
||||
const (
|
||||
// headerLength is the length of a DNS message's header, which the
|
||||
// question follows: its id, its flags, and how many questions,
|
||||
// answers and other records it holds, two bytes each.
|
||||
headerLength = 12
|
||||
// typeAndClass is the length of the type and the class that end a
|
||||
// question, after its name.
|
||||
typeAndClass = 4
|
||||
// replyFlags mark a reply to a query that asked for recursion, which
|
||||
// is available, with no error. The response code goes in their last
|
||||
// four bits.
|
||||
replyFlags = 0x8180
|
||||
// maxMessage is the longest query read over UDP.
|
||||
maxMessage = 1232
|
||||
)
|
||||
|
||||
// set has the stand-in answer name with given.
|
||||
func (s *resolverStandIn) set(name string, given answer) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
s.answers[name] = given
|
||||
}
|
||||
|
||||
// dial connects Go's resolver to the stand-in through an in-memory
|
||||
// connection, on which it sends each query, and reads each reply, after
|
||||
// its length, as over TCP.
|
||||
func (s *resolverStandIn) dial(context.Context, string, string) (net.Conn, error) {
|
||||
client, server := net.Pipe()
|
||||
|
||||
go s.serve(server)
|
||||
|
||||
return client, nil
|
||||
}
|
||||
|
||||
// serve answers the queries that come on conn until the resolver closes
|
||||
// it.
|
||||
func (s *resolverStandIn) serve(conn net.Conn) {
|
||||
defer func() {
|
||||
_ = conn.Close()
|
||||
}()
|
||||
|
||||
for {
|
||||
var length [2]byte
|
||||
|
||||
_, err := io.ReadFull(conn, length[:])
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
message := make([]byte, binary.BigEndian.Uint16(length[:]))
|
||||
|
||||
_, err = io.ReadFull(conn, message)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
reply, answered := s.reply(message)
|
||||
if !answered {
|
||||
continue // the resolver gives up, and closes conn
|
||||
}
|
||||
|
||||
//nolint:gosec // a reply of a few dozen bytes
|
||||
_, err = conn.Write(append(binary.BigEndian.AppendUint16(nil, uint16(len(reply))),
|
||||
reply...))
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// serveUDP answers the queries that come on conn, each in a datagram, as
|
||||
// a resolver does, until conn is closed.
|
||||
func (s *resolverStandIn) serveUDP(conn net.PacketConn) {
|
||||
message := make([]byte, maxMessage)
|
||||
|
||||
for {
|
||||
n, from, err := conn.ReadFrom(message)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
reply, answered := s.reply(message[:n])
|
||||
if answered {
|
||||
_, _ = conn.WriteTo(reply, from)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// reply returns the stand-in's reply to message, a query, and false for
|
||||
// none, while it hangs. It notes the name asked about.
|
||||
func (s *resolverStandIn) reply(message []byte) ([]byte, bool) {
|
||||
// The name is labels, each after its length, ended by a length of 0.
|
||||
var labels []string
|
||||
|
||||
end := headerLength
|
||||
for message[end] != 0 {
|
||||
length := int(message[end])
|
||||
labels = append(labels, string(message[end+1:end+1+length]))
|
||||
end += 1 + length
|
||||
}
|
||||
|
||||
end += 1 + typeAndClass
|
||||
name := strings.Join(labels, ".") + "."
|
||||
|
||||
s.mu.Lock()
|
||||
s.names = append(s.names, name)
|
||||
given, found := s.answers[name]
|
||||
hanging := s.hanging
|
||||
s.mu.Unlock()
|
||||
|
||||
if hanging {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
if !found {
|
||||
given = answer{rcode: noSuchName}
|
||||
}
|
||||
|
||||
// The query's id, the flags, one question, the answers, and no other
|
||||
// records, then the question, as asked.
|
||||
reply := slices.Clone(message[:2])
|
||||
reply = binary.BigEndian.AppendUint16(reply, replyFlags|given.rcode)
|
||||
reply = binary.BigEndian.AppendUint16(reply, 1)
|
||||
//nolint:gosec // a handful of answers
|
||||
reply = binary.BigEndian.AppendUint16(reply, uint16(len(given.addrs)))
|
||||
reply = append(reply, 0, 0, 0, 0)
|
||||
reply = append(reply, message[headerLength:end]...)
|
||||
|
||||
// An A record starts with the name asked about, by a pointer to it in
|
||||
// the question, then its type, A, its class, IN, how long it may be
|
||||
// kept, 60 seconds, and the length of its address, 4 bytes.
|
||||
record := []byte{0xc0, headerLength, 0, 1, 0, 1, 0, 0, 0, 60, 0, 4}
|
||||
|
||||
for _, addr := range given.addrs {
|
||||
reply = append(reply, record...)
|
||||
reply = append(reply, netip.MustParseAddr(addr).AsSlice()...)
|
||||
}
|
||||
|
||||
return reply, true
|
||||
}
|
||||
|
||||
// dnsblParams returns the DNSBLParams of zones, with the tests' cache TTL
|
||||
// and timeout, by the bubble's clock, with alerts to a queue that sends
|
||||
// none.
|
||||
func dnsblParams(zones ...string) reputation.DNSBLParams {
|
||||
return reputation.DNSBLParams{
|
||||
Zones: zones,
|
||||
CacheTTL: cacheTTL,
|
||||
Timeout: timeout,
|
||||
Now: time.Now,
|
||||
ProcessLog: slog.New(slog.DiscardHandler),
|
||||
Alerts: newQueue(),
|
||||
}
|
||||
}
|
||||
|
||||
// waitForTheResolver waits, on the bubble's clock, an hour, until Go's
|
||||
// resolver has given up on every stand-in that does not answer: it waits
|
||||
// for a server as long as /etc/resolv.conf has it wait, a few seconds,
|
||||
// even after the query was given up, and a bubble cannot end before it.
|
||||
func waitForTheResolver() {
|
||||
time.Sleep(time.Hour)
|
||||
}
|
||||
|
||||
// newDNSBL returns the DNSBL of p, asking resolver.
|
||||
func newDNSBL(resolver *resolverStandIn, p reputation.DNSBLParams) *reputation.DNSBL {
|
||||
dnsbl := reputation.NewDNSBL(p)
|
||||
dnsbl.SetDial(resolver.dial)
|
||||
|
||||
return dnsbl
|
||||
}
|
||||
|
||||
// wantZones checks the zones whose verdict dnsbl says lists client, as a
|
||||
// request from client finds them.
|
||||
func wantZones(t *testing.T, dnsbl *reputation.DNSBL, client string, want ...string) {
|
||||
t.Helper()
|
||||
|
||||
got := dnsbl.ListedBy(t.Context(), netip.MustParseAddr(client))
|
||||
if !slices.Equal(got, want) {
|
||||
t.Errorf("%s is listed by %v, want %v", client, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// wantQueries checks how many queries dnsbl made to zone, and how many of
|
||||
// them failed.
|
||||
func wantQueries(t *testing.T, dnsbl *reputation.DNSBL, queries, failures int) {
|
||||
t.Helper()
|
||||
|
||||
if dnsbl.Queries(zone) != queries || dnsbl.Failures(zone) != failures {
|
||||
t.Errorf("%d queries and %d failures, want %d and %d", dnsbl.Queries(zone),
|
||||
dnsbl.Failures(zone), queries, failures)
|
||||
}
|
||||
}
|
||||
|
||||
// wantAsked checks the names the stand-in was asked about, in any order.
|
||||
func wantAsked(t *testing.T, resolver *resolverStandIn, want ...string) {
|
||||
t.Helper()
|
||||
|
||||
resolver.mu.Lock()
|
||||
got := slices.Sorted(slices.Values(resolver.names))
|
||||
resolver.mu.Unlock()
|
||||
|
||||
slices.Sort(want)
|
||||
|
||||
if !slices.Equal(got, want) {
|
||||
t.Errorf("asked about %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
package reputation
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
)
|
||||
|
||||
// SetTransport has l's fetches go through transport instead of the
|
||||
// network.
|
||||
func (l *Lists) SetTransport(transport http.RoundTripper) {
|
||||
l.httpClient.Transport = transport
|
||||
}
|
||||
|
||||
// SetTransport has a's checks go through transport instead of the
|
||||
// network.
|
||||
func (a *AbuseIPDB) SetTransport(transport http.RoundTripper) {
|
||||
a.httpClient.Transport = transport
|
||||
}
|
||||
|
||||
// SetDial has d's queries go through dial instead of the network.
|
||||
func (d *DNSBL) SetDial(
|
||||
dial func(ctx context.Context, network, address string) (net.Conn, error),
|
||||
) {
|
||||
d.resolver = &net.Resolver{PreferGo: true, Dial: dial}
|
||||
}
|
||||
|
||||
// LookUp asks zone about addr at once, as a query in the background does,
|
||||
// and returns whether zone lists addr.
|
||||
func (d *DNSBL) LookUp(zone string, addr netip.Addr) (bool, error) {
|
||||
return d.lookUp(context.Background(), query{zone: zone, client: addr})
|
||||
}
|
||||
@@ -0,0 +1,509 @@
|
||||
// Package reputation fetches the lists the settings name by URL: the
|
||||
// blocklists of SWWAF_BLOCKLIST_URLS, and the file of AS:percent lines
|
||||
// SWWAF_ASN_LIMIT_PERCENT_URL names. It keeps the last good copy of each,
|
||||
// whole, comment lines included, which is used while a fetch fails, and
|
||||
// when each was last tried. It also asks the DNSBL zones of
|
||||
// SWWAF_DNSBL_ZONES about clients, and keeps their verdicts, and checks
|
||||
// clients with AbuseIPDB, and keeps their scores and the checks spent
|
||||
// today. The state package writes all of these to reputation.json and
|
||||
// reads them from it, so that a restart keeps them too.
|
||||
package reputation
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
)
|
||||
|
||||
const (
|
||||
// maxListBytes is the most of a list that is read. A longer one is a
|
||||
// failure, so that a wrong URL cannot fill the memory.
|
||||
maxListBytes = 16 << 20
|
||||
// fetchTimeout bounds one fetch of a list.
|
||||
fetchTimeout = time.Minute
|
||||
// mappedBits is the length of ::ffff:0.0.0.0/96, the netblock of every
|
||||
// IPv4-mapped address.
|
||||
mappedBits = 96
|
||||
)
|
||||
|
||||
var (
|
||||
errStatus = errors.New("the server answered")
|
||||
errTooLong = errors.New("the list is longer than 16 MiB")
|
||||
errNotNetblock = errors.New("is not an address or a netblock, such as 192.0.2.0/24")
|
||||
errNotASNPercent = errors.New(
|
||||
"is not an AS number, : and a percentage, such as AS64496:50")
|
||||
)
|
||||
|
||||
// List is a list as reputation.json holds it: the URL it is fetched from,
|
||||
// when it was last tried, the fetch failed or not, and its last good copy:
|
||||
// when that was fetched, and its lines, as fetched, comment lines
|
||||
// included, both left out while no fetch of it has succeeded.
|
||||
type List struct {
|
||||
URL string `json:"url"`
|
||||
Tried time.Time `json:"tried"`
|
||||
Fetched time.Time `json:"fetched,omitzero"`
|
||||
Lines []string `json:"lines,omitzero"`
|
||||
}
|
||||
|
||||
// Params are what New needs.
|
||||
type Params struct {
|
||||
// BlocklistURLs are the blocklists (SWWAF_BLOCKLIST_URLS), and
|
||||
// ASNLimitPercentURL the file of AS:percent lines
|
||||
// (SWWAF_ASN_LIMIT_PERCENT_URL), "" while it is unset.
|
||||
BlocklistURLs []string
|
||||
ASNLimitPercentURL string
|
||||
// Refresh is how long after a list was last fetched or tried it is
|
||||
// fetched again (SWWAF_BLOCKLIST_REFRESH).
|
||||
Refresh time.Duration
|
||||
// Now tells the time, normally time.Now in UTC.
|
||||
Now func() time.Time
|
||||
// ProcessLog receives each fetch of a list, and why one failed.
|
||||
ProcessLog *slog.Logger
|
||||
// Alerts receive a source_failure alert for each fetch that fails.
|
||||
Alerts *alerts.Queue
|
||||
}
|
||||
|
||||
// Lists are the lists Params names, each with its last good copy. They
|
||||
// are safe for concurrent use.
|
||||
type Lists struct {
|
||||
params Params
|
||||
httpClient *http.Client
|
||||
|
||||
mu sync.Mutex
|
||||
// lists are by URL, one for each URL Params names.
|
||||
lists map[string]*list
|
||||
}
|
||||
|
||||
// list is one list: what reputation.json keeps of it, its last try, zero
|
||||
// before the first, and its last good copy, what that copy says, and how
|
||||
// many fetches of it failed.
|
||||
type list struct {
|
||||
kept List
|
||||
entries entries
|
||||
failures int
|
||||
}
|
||||
|
||||
// entries are what the lines of a copy say: for a blocklist, the netblocks
|
||||
// it names, with the lengths among them, and for the file of AS:percent
|
||||
// lines, the percentage it gives each AS number.
|
||||
type entries struct {
|
||||
netblocks map[netip.Prefix]bool
|
||||
lengths []int
|
||||
percents map[string]int64
|
||||
}
|
||||
|
||||
// New returns the lists, without a copy of any yet.
|
||||
func New(params Params) *Lists {
|
||||
l := &Lists{params: params, httpClient: &http.Client{}, lists: map[string]*list{}}
|
||||
|
||||
for _, listURL := range l.URLs() {
|
||||
l.lists[listURL] = &list{kept: List{URL: listURL}}
|
||||
}
|
||||
|
||||
return l
|
||||
}
|
||||
|
||||
// URLs returns the URL of every list: the blocklists' in the order
|
||||
// SWWAF_BLOCKLIST_URLS names them, then SWWAF_ASN_LIMIT_PERCENT_URL.
|
||||
func (l *Lists) URLs() []string {
|
||||
urls := slices.Clone(l.params.BlocklistURLs)
|
||||
if l.params.ASNLimitPercentURL != "" {
|
||||
urls = append(urls, l.params.ASNLimitPercentURL)
|
||||
}
|
||||
|
||||
return urls
|
||||
}
|
||||
|
||||
// ListedBy returns the URLs of the blocklists whose copy lists addr, in
|
||||
// the order SWWAF_BLOCKLIST_URLS names them.
|
||||
func (l *Lists) ListedBy(addr netip.Addr) []string {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
var listedBy []string
|
||||
|
||||
for _, listURL := range l.params.BlocklistURLs {
|
||||
if l.lists[listURL].entries.contain(addr) {
|
||||
listedBy = append(listedBy, listURL)
|
||||
}
|
||||
}
|
||||
|
||||
return listedBy
|
||||
}
|
||||
|
||||
// ASNLimitPercent returns the percentage the copy of the file of
|
||||
// AS:percent lines gives asn, and whether it lists asn.
|
||||
func (l *Lists) ASNLimitPercent(asn string) (int64, bool) {
|
||||
if l.params.ASNLimitPercentURL == "" {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
percent, listed := l.lists[l.params.ASNLimitPercentURL].entries.percents[asn]
|
||||
|
||||
return percent, listed
|
||||
}
|
||||
|
||||
// Fetched returns when the copy in use of the list at listURL was
|
||||
// fetched, or zero while there is none.
|
||||
func (l *Lists) Fetched(listURL string) time.Time {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
return l.lists[listURL].kept.Fetched
|
||||
}
|
||||
|
||||
// Failures returns how many fetches of the list at listURL failed.
|
||||
func (l *Lists) Failures(listURL string) int {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
return l.lists[listURL].failures
|
||||
}
|
||||
|
||||
// Run fetches each list once Refresh has passed since it was last fetched
|
||||
// or tried, the later of the two, until ctx is done. A list never tried is
|
||||
// fetched at once, and so is one whose last try or copy, read from
|
||||
// reputation.json, is that old.
|
||||
func (l *Lists) Run(ctx context.Context) {
|
||||
if len(l.lists) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
for ctx.Err() == nil {
|
||||
next := l.fetchDue(ctx)
|
||||
timer := time.NewTimer(next.Sub(l.params.Now()))
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
case <-timer.C:
|
||||
}
|
||||
|
||||
timer.Stop()
|
||||
}
|
||||
}
|
||||
|
||||
// Snapshot returns each list that has been tried, with its copy, if it
|
||||
// has one, sorted by URL, as reputation.json lists them.
|
||||
func (l *Lists) Snapshot() []List {
|
||||
l.mu.Lock()
|
||||
|
||||
tried := make([]List, 0, len(l.lists))
|
||||
|
||||
for _, held := range l.lists {
|
||||
if !held.kept.Tried.IsZero() {
|
||||
tried = append(tried, held.kept)
|
||||
}
|
||||
}
|
||||
|
||||
l.mu.Unlock()
|
||||
|
||||
slices.SortFunc(tried, func(a, b List) int {
|
||||
return strings.Compare(a.URL, b.URL)
|
||||
})
|
||||
|
||||
return tried
|
||||
}
|
||||
|
||||
// Load puts lists, read from reputation.json, in place of the last tries
|
||||
// and copies held. A list Params does not name is dropped. A copy with a
|
||||
// line that parse refuses is an error, and then nothing changes.
|
||||
func (l *Lists) Load(lists []List) error {
|
||||
found := make(map[string]entries, len(lists))
|
||||
|
||||
for _, kept := range lists {
|
||||
if _, named := l.lists[kept.URL]; !named {
|
||||
continue
|
||||
}
|
||||
|
||||
read, err := l.parse(kept.URL, kept.Lines)
|
||||
if err != nil {
|
||||
return fmt.Errorf("the copy of %s: %w", kept.URL, err)
|
||||
}
|
||||
|
||||
found[kept.URL] = read
|
||||
}
|
||||
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
for listURL, held := range l.lists {
|
||||
held.kept, held.entries = List{URL: listURL}, entries{}
|
||||
}
|
||||
|
||||
for _, kept := range lists {
|
||||
read, named := found[kept.URL]
|
||||
if named {
|
||||
l.lists[kept.URL].kept, l.lists[kept.URL].entries = kept, read
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// fetchDue fetches each list that is due, one after another, and returns
|
||||
// when the next is due. Once ctx has ended, it starts none, since a fetch
|
||||
// cut off is noted as a try.
|
||||
func (l *Lists) fetchDue(ctx context.Context) time.Time {
|
||||
var next time.Time
|
||||
|
||||
for _, listURL := range l.URLs() {
|
||||
due := l.due(listURL)
|
||||
if ctx.Err() == nil && !l.params.Now().Before(due) {
|
||||
l.fetch(ctx, listURL)
|
||||
due = l.due(listURL)
|
||||
}
|
||||
|
||||
if next.IsZero() || due.Before(next) {
|
||||
next = due
|
||||
}
|
||||
}
|
||||
|
||||
return next
|
||||
}
|
||||
|
||||
// due returns when the list at listURL is to be fetched: Refresh after it
|
||||
// was last fetched or tried, the later of the two.
|
||||
func (l *Lists) due(listURL string) time.Time {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
held := l.lists[listURL]
|
||||
|
||||
last := held.kept.Fetched
|
||||
if held.kept.Tried.After(last) {
|
||||
last = held.kept.Tried
|
||||
}
|
||||
|
||||
return last.Add(l.params.Refresh)
|
||||
}
|
||||
|
||||
// fetch fetches the list at listURL, and notes the try. A good copy takes
|
||||
// the place of the one held. A failure leaves that in use, and is counted,
|
||||
// logged and raised as a source_failure alert. A fetch cut off as ctx
|
||||
// ends, as smallwebwaf stops, is no failure, but is still noted as a try,
|
||||
// so that a restart waits for it: the server may have had its request.
|
||||
func (l *Lists) fetch(ctx context.Context, listURL string) {
|
||||
lines, err := l.get(ctx, listURL)
|
||||
|
||||
var found entries
|
||||
if err == nil {
|
||||
found, err = l.parse(listURL, lines)
|
||||
}
|
||||
|
||||
cutOff := err != nil && ctx.Err() != nil
|
||||
now := l.params.Now()
|
||||
|
||||
l.mu.Lock()
|
||||
|
||||
held := l.lists[listURL]
|
||||
held.kept.Tried = now
|
||||
|
||||
if err == nil {
|
||||
held.kept.Fetched, held.kept.Lines = now, lines
|
||||
held.entries = found
|
||||
} else if !cutOff {
|
||||
held.failures++
|
||||
}
|
||||
|
||||
l.mu.Unlock()
|
||||
|
||||
if cutOff {
|
||||
return
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
const failed = "fetching a list failed"
|
||||
|
||||
// Raised before it is logged, so that the alert is there once the
|
||||
// log line is.
|
||||
raiseFailure(l.params.Alerts, failed, listURL, err)
|
||||
l.params.ProcessLog.Warn(failed, "url", listURL, "error", err.Error())
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
l.params.ProcessLog.Info("fetched a list", "url", listURL, "lines", len(lines))
|
||||
}
|
||||
|
||||
// raiseFailure raises a source_failure alert into queue, with reason, and
|
||||
// in its detail the source that failed, a list's URL, a zone with its key
|
||||
// masked or abuseipdb, and err.
|
||||
func raiseFailure(queue *alerts.Queue, reason, source string, err error) {
|
||||
queue.Raise(alerts.Alert{
|
||||
Event: alerts.EventSourceFailure,
|
||||
Reason: reason,
|
||||
Detail: map[string]any{"source": source, "error": err.Error()},
|
||||
})
|
||||
}
|
||||
|
||||
// get fetches the list at listURL, and returns its lines. An answer other
|
||||
// than 200, or a list longer than maxListBytes, is a failure.
|
||||
func (l *Lists) get(ctx context.Context, listURL string) ([]string, error) {
|
||||
ctx, cancel := context.WithTimeout(ctx, fetchTimeout)
|
||||
defer cancel()
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, listURL, http.NoBody)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("make the request: %w", err)
|
||||
}
|
||||
|
||||
res, err := l.httpClient.Do(req)
|
||||
if err != nil {
|
||||
// Do's error names the URL, which the log line and the alert name
|
||||
// already: only what went wrong is kept.
|
||||
return nil, fmt.Errorf("fetch the list: %w", errors.Unwrap(err))
|
||||
}
|
||||
|
||||
defer func() {
|
||||
_ = res.Body.Close()
|
||||
}()
|
||||
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("%w %s", errStatus, res.Status)
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(res.Body, maxListBytes+1))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read the list: %w", err)
|
||||
}
|
||||
|
||||
if len(body) > maxListBytes {
|
||||
return nil, errTooLong
|
||||
}
|
||||
|
||||
lines := []string{}
|
||||
for line := range strings.Lines(string(body)) {
|
||||
lines = append(lines, strings.TrimSuffix(line, "\n"))
|
||||
}
|
||||
|
||||
return lines, nil
|
||||
}
|
||||
|
||||
// parse reads the lines of the list at listURL: those of a blocklist, or
|
||||
// of the file of AS:percent lines. Anything after a ; or a # on a line is
|
||||
// left out, and so is a line left blank. Any other line that does not read
|
||||
// is an error naming it by its number.
|
||||
func (l *Lists) parse(listURL string, lines []string) (entries, error) {
|
||||
if listURL == l.params.ASNLimitPercentURL {
|
||||
return parsePercents(lines)
|
||||
}
|
||||
|
||||
return parseNetblocks(lines)
|
||||
}
|
||||
|
||||
// parseNetblocks reads a blocklist's lines, each an address or a netblock
|
||||
// as the settings take them.
|
||||
func parseNetblocks(lines []string) (entries, error) {
|
||||
found := entries{netblocks: map[netip.Prefix]bool{}}
|
||||
|
||||
for i, line := range lines {
|
||||
text := withoutComment(line)
|
||||
if text == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
netblock, ok := parseNetblock(text)
|
||||
if !ok {
|
||||
return entries{}, fmt.Errorf("line %d %w", i+1, errNotNetblock)
|
||||
}
|
||||
|
||||
found.netblocks[netblock] = true
|
||||
|
||||
if !slices.Contains(found.lengths, netblock.Bits()) {
|
||||
found.lengths = append(found.lengths, netblock.Bits())
|
||||
}
|
||||
}
|
||||
|
||||
return found, nil
|
||||
}
|
||||
|
||||
// parseNetblock reads text, a line of a blocklist, and reports whether it
|
||||
// is an address or a netblock as the settings take them. A client's IPv4
|
||||
// address is checked as IPv4, never IPv4-mapped, so an IPv4-mapped line,
|
||||
// such as ::ffff:192.0.2.0/120, is read as the IPv4 address or netblock it
|
||||
// stands for, 192.0.2.0/24, and a mapped netblock shorter than /96, which
|
||||
// stands for none, is refused.
|
||||
func parseNetblock(text string) (netip.Prefix, bool) {
|
||||
netblock, err := config.ParseNetblock(text)
|
||||
if err != nil {
|
||||
return netip.Prefix{}, false
|
||||
}
|
||||
|
||||
// The address as written: ParseNetblock's has the bits past the
|
||||
// netblock's length cleared, the ::ffff among them below /96.
|
||||
written, _, _ := strings.Cut(text, "/")
|
||||
if addr, _ := netip.ParseAddr(written); !addr.Is4In6() {
|
||||
return netblock, true
|
||||
}
|
||||
|
||||
if netblock.Bits() < mappedBits {
|
||||
return netip.Prefix{}, false
|
||||
}
|
||||
|
||||
return netip.PrefixFrom(netblock.Addr().Unmap(), netblock.Bits()-mappedBits), true
|
||||
}
|
||||
|
||||
// parsePercents reads the lines of the file of AS:percent lines, each an
|
||||
// AS number, : and a percentage, as SWWAF_ASN_LIMIT_PERCENT takes them. An
|
||||
// AS number listed more than once gets the lowest of its percentages.
|
||||
func parsePercents(lines []string) (entries, error) {
|
||||
found := entries{percents: map[string]int64{}}
|
||||
|
||||
for i, line := range lines {
|
||||
text := withoutComment(line)
|
||||
if text == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
asnText, percentText, _ := strings.Cut(text, ":")
|
||||
asn, asnErr := config.ParseASN(asnText)
|
||||
|
||||
percent, percentErr := config.ParsePercent(percentText)
|
||||
if asnErr != nil || percentErr != nil {
|
||||
return entries{}, fmt.Errorf("line %d %w", i+1, errNotASNPercent)
|
||||
}
|
||||
|
||||
earlier, listed := found.percents[asn]
|
||||
if !listed || percent < earlier {
|
||||
found.percents[asn] = percent
|
||||
}
|
||||
}
|
||||
|
||||
return found, nil
|
||||
}
|
||||
|
||||
// withoutComment returns line without anything after a ; or a #, and
|
||||
// without the spaces around what is left.
|
||||
func withoutComment(line string) string {
|
||||
text, _, _ := strings.Cut(line, ";")
|
||||
text, _, _ = strings.Cut(text, "#")
|
||||
|
||||
return strings.TrimSpace(text)
|
||||
}
|
||||
|
||||
// contain reports whether the netblocks of a blocklist's copy hold addr:
|
||||
// whether addr, cut to one of their lengths, is one of them.
|
||||
func (e entries) contain(addr netip.Addr) bool {
|
||||
for _, length := range e.lengths {
|
||||
netblock, err := addr.Prefix(length)
|
||||
if err == nil && e.netblocks[netblock] {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,610 @@
|
||||
package reputation_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"reflect"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"testing/synctest"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/reputation"
|
||||
)
|
||||
|
||||
// The tests run in a synctest bubble, where the time package runs on a
|
||||
// clock of the test's own: time.Sleep moves it on at once, and
|
||||
// synctest.Wait returns once Run waits for the next list to be due, so
|
||||
// that every fetch due by then has been made. The stand-in for the
|
||||
// servers the lists are fetched from answers without the network, since a
|
||||
// fetch waiting on the network would keep that clock from moving on.
|
||||
|
||||
const (
|
||||
// dropURL and torURL are the blocklists, and asnURL the file of
|
||||
// AS:percent lines.
|
||||
dropURL = "https://lists.example/drop.txt"
|
||||
torURL = "https://lists.example/tor.txt"
|
||||
asnURL = "https://lists.example/asn.txt"
|
||||
// refresh is the tests' SWWAF_BLOCKLIST_REFRESH, and cooldown their
|
||||
// SWWAF_ALERT_COOLDOWN, longer than it.
|
||||
refresh = 24 * time.Hour
|
||||
cooldown = 48 * time.Hour
|
||||
// drop is a blocklist as the Spamhaus DROP list is written, with an
|
||||
// address and a netblock in each of its comments, which list nothing.
|
||||
drop = "; Spamhaus DROP List 2026/10/07 - (c) 2026 The Spamhaus Project SLL\n" +
|
||||
"; Last-Modified: Wed, 07 Oct 2026 00:00:00 GMT ; 192.0.2.1\n" +
|
||||
"# 198.51.100.0/24\n" +
|
||||
"\n" +
|
||||
"203.0.113.0/24 ; SBL1\n" +
|
||||
" 192.0.2.9 # one address\n" +
|
||||
"2001:db8:1::/48 ; SBL2\n"
|
||||
)
|
||||
|
||||
func TestListedAddressesAndNetblocksWithTheCommentsLeftOut(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
servers := &standIn{lists: map[string]string{dropURL: drop}}
|
||||
lists := start(t, servers, params(dropURL))
|
||||
|
||||
for addr, want := range map[string][]string{
|
||||
"203.0.113.0": {dropURL},
|
||||
"203.0.113.255": {dropURL},
|
||||
"192.0.2.9": {dropURL},
|
||||
"2001:db8:1::7": {dropURL},
|
||||
"203.0.114.0": nil,
|
||||
"192.0.2.8": nil,
|
||||
"192.0.2.1": nil,
|
||||
"198.51.100.7": nil,
|
||||
"2001:db8:2::7": nil,
|
||||
} {
|
||||
wantListedBy(t, lists, addr, want...)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestIPv4MappedLineListsTheIPv4AddressOrNetblockItStandsFor(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
now := time.Now()
|
||||
lists := reputation.New(params(dropURL))
|
||||
|
||||
err := lists.Load([]reputation.List{{
|
||||
URL: dropURL, Tried: now, Fetched: now,
|
||||
Lines: []string{"::ffff:192.0.2.9", "::ffff:203.0.113.0/120"},
|
||||
}})
|
||||
if err != nil {
|
||||
t.Fatalf("load: %v", err)
|
||||
}
|
||||
|
||||
for addr, want := range map[string][]string{
|
||||
"192.0.2.9": {dropURL},
|
||||
"203.0.113.255": {dropURL},
|
||||
"192.0.2.8": nil,
|
||||
"203.0.114.0": nil,
|
||||
} {
|
||||
wantListedBy(t, lists, addr, want...)
|
||||
}
|
||||
|
||||
// A mapped netblock shorter than /96 stands for no IPv4 one.
|
||||
err = lists.Load([]reputation.List{{
|
||||
URL: dropURL, Tried: now, Fetched: now, Lines: []string{"::ffff:198.51.100.0/88"},
|
||||
}})
|
||||
|
||||
const want = "the copy of " + dropURL +
|
||||
": line 1 is not an address or a netblock, such as 192.0.2.0/24"
|
||||
if err == nil || err.Error() != want {
|
||||
t.Errorf("error %v, want %s", err, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientIsListedByEachBlocklistThatListsIt(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
servers := &standIn{lists: map[string]string{
|
||||
dropURL: "203.0.113.0/24\n", torURL: "203.0.113.9\n",
|
||||
}}
|
||||
lists := start(t, servers, params(torURL, dropURL))
|
||||
|
||||
// In the order SWWAF_BLOCKLIST_URLS names them.
|
||||
wantListedBy(t, lists, "203.0.113.9", torURL, dropURL)
|
||||
wantListedBy(t, lists, "203.0.113.8", dropURL)
|
||||
})
|
||||
}
|
||||
|
||||
func TestListFetchedAgainOnceRefreshHasPassed(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
servers := &standIn{lists: map[string]string{dropURL: "203.0.113.9\n"}}
|
||||
lists := start(t, servers, params(dropURL))
|
||||
began := time.Now()
|
||||
|
||||
wantFetches(t, servers, 1)
|
||||
|
||||
servers.set(dropURL, "203.0.113.10\n")
|
||||
time.Sleep(refresh - time.Nanosecond)
|
||||
wantFetches(t, servers, 1)
|
||||
wantListedBy(t, lists, "203.0.113.9", dropURL)
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
wantFetches(t, servers, 2)
|
||||
wantListedBy(t, lists, "203.0.113.9")
|
||||
wantListedBy(t, lists, "203.0.113.10", dropURL)
|
||||
|
||||
if fetched := lists.Fetched(dropURL); !fetched.Equal(began.Add(refresh)) {
|
||||
t.Errorf("the copy in use was fetched at %s, want %s", fetched,
|
||||
began.Add(refresh))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestFailedFetchKeepsTheLastGoodCopyAndAlertsOncePerCooldown(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
// fail has the stand-in answer the fetches after the first so that
|
||||
// they fail with error.
|
||||
fail func(servers *standIn)
|
||||
error string
|
||||
}{
|
||||
{
|
||||
"an answer other than 200",
|
||||
func(servers *standIn) { servers.set(dropURL, "") },
|
||||
"the server answered 503 Service Unavailable",
|
||||
},
|
||||
{
|
||||
"a line that does not read",
|
||||
func(servers *standIn) { servers.set(dropURL, "203.0.113.10\n<html>\n") },
|
||||
"line 2 is not an address or a netblock, such as 192.0.2.0/24",
|
||||
},
|
||||
{
|
||||
"a list longer than 16 MiB",
|
||||
func(servers *standIn) {
|
||||
servers.set(dropURL, strings.Repeat("#\n", 8<<20+1))
|
||||
},
|
||||
"the list is longer than 16 MiB",
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
var log bytes.Buffer
|
||||
|
||||
servers := &standIn{lists: map[string]string{dropURL: drop}}
|
||||
queue := newQueue()
|
||||
p := params(dropURL)
|
||||
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
|
||||
p.Alerts = queue
|
||||
lists := start(t, servers, p)
|
||||
kept := lists.Snapshot()
|
||||
|
||||
tc.fail(servers)
|
||||
|
||||
// Each failure is tried again once refresh has passed since it.
|
||||
for range 2 {
|
||||
time.Sleep(refresh)
|
||||
synctest.Wait()
|
||||
}
|
||||
|
||||
wantFetches(t, servers, 3)
|
||||
wantListedBy(t, lists, "203.0.113.9", dropURL)
|
||||
|
||||
want := kept[0]
|
||||
want.Tried = time.Now()
|
||||
|
||||
if got := lists.Snapshot(); !reflect.DeepEqual(got, []reputation.List{want}) {
|
||||
t.Errorf("lists %+v, want the first copy, last tried now, %+v", got, want)
|
||||
}
|
||||
|
||||
if lists.Failures(dropURL) != 2 {
|
||||
t.Errorf("%d failures, want 2", lists.Failures(dropURL))
|
||||
}
|
||||
|
||||
// One alert for the first failure; the cooldown holds back the
|
||||
// second.
|
||||
wantAlert(t, queue, alerts.Alert{
|
||||
Time: time.Now().Add(-refresh),
|
||||
Event: alerts.EventSourceFailure,
|
||||
Reason: "fetching a list failed",
|
||||
Detail: map[string]any{"source": dropURL, "error": tc.error},
|
||||
})
|
||||
|
||||
if !strings.Contains(log.String(), `"msg":"fetching a list failed",`+
|
||||
`"url":"`+dropURL+`","error":"`+tc.error) {
|
||||
t.Errorf("logged\n%s\nwant the failures", log.String())
|
||||
}
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchNotDoneWithinAMinuteFails(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
servers := &standIn{lists: map[string]string{}, hanging: true}
|
||||
lists := start(t, servers, params(dropURL))
|
||||
|
||||
time.Sleep(time.Minute - time.Nanosecond)
|
||||
synctest.Wait()
|
||||
|
||||
if lists.Failures(dropURL) != 0 {
|
||||
t.Errorf("%d failures before a minute, want none", lists.Failures(dropURL))
|
||||
}
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
synctest.Wait()
|
||||
|
||||
if lists.Failures(dropURL) != 1 {
|
||||
t.Errorf("%d failures after a minute, want 1", lists.Failures(dropURL))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestKeptCopyIsFetchedAgainOnceRefreshHasPassedSinceItWasFetched(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
servers := &standIn{lists: map[string]string{
|
||||
dropURL: "203.0.113.10\n", torURL: "198.51.100.10\n",
|
||||
}}
|
||||
lists := reputation.New(params(dropURL, torURL))
|
||||
lists.SetTransport(servers)
|
||||
|
||||
// drop.txt was fetched an hour ago, and tor.txt a refresh ago, as
|
||||
// reputation.json says at start.
|
||||
err := lists.Load([]reputation.List{
|
||||
{URL: dropURL, Fetched: time.Now().Add(-time.Hour), Lines: []string{"203.0.113.9"}},
|
||||
{URL: torURL, Fetched: time.Now().Add(-refresh), Lines: []string{"198.51.100.9"}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("load: %v", err)
|
||||
}
|
||||
|
||||
run(t, lists)
|
||||
|
||||
wantFetches(t, servers, 1)
|
||||
wantListedBy(t, lists, "203.0.113.9", dropURL)
|
||||
wantListedBy(t, lists, "198.51.100.10", torURL)
|
||||
|
||||
time.Sleep(refresh - time.Hour - time.Nanosecond)
|
||||
wantFetches(t, servers, 1)
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
wantFetches(t, servers, 2)
|
||||
wantListedBy(t, lists, "203.0.113.10", dropURL)
|
||||
})
|
||||
}
|
||||
|
||||
func TestRestartWaitsRefreshAfterTheLastTryEvenOneThatFailed(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
// drop.txt is fetched, and a refresh later the fetch downloads it
|
||||
// whole but fails on a line that does not read.
|
||||
servers := &standIn{lists: map[string]string{dropURL: "198.51.100.1\n"}}
|
||||
lists := start(t, servers, params(dropURL))
|
||||
|
||||
servers.set(dropURL, "198.51.100.2\n<html>\n")
|
||||
time.Sleep(refresh)
|
||||
wantFetches(t, servers, 2)
|
||||
|
||||
// Restarted with what reputation.json keeps, it waits a refresh
|
||||
// after the failed try, as it does while it runs.
|
||||
restarted := &standIn{lists: map[string]string{dropURL: "198.51.100.2\n"}}
|
||||
again := reputation.New(params(dropURL))
|
||||
again.SetTransport(restarted)
|
||||
|
||||
err := again.Load(lists.Snapshot())
|
||||
if err != nil {
|
||||
t.Fatalf("load: %v", err)
|
||||
}
|
||||
|
||||
run(t, again)
|
||||
wantFetches(t, restarted, 0)
|
||||
wantListedBy(t, again, "198.51.100.1", dropURL)
|
||||
|
||||
time.Sleep(refresh - time.Nanosecond)
|
||||
wantFetches(t, restarted, 0)
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
wantFetches(t, restarted, 1)
|
||||
wantListedBy(t, again, "198.51.100.2", dropURL)
|
||||
})
|
||||
}
|
||||
|
||||
func TestFetchCutOffAsItStopsIsNoFailureButARestartWaitsForIt(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
// Stopped 30 seconds into the fetch of drop.txt, before tor.txt's.
|
||||
servers := &standIn{lists: map[string]string{}, hanging: true}
|
||||
queue := newQueue()
|
||||
p := params(dropURL, torURL)
|
||||
p.Alerts = queue
|
||||
|
||||
lists := reputation.New(p)
|
||||
lists.SetTransport(servers)
|
||||
|
||||
ctx, stop := context.WithCancel(t.Context())
|
||||
stopped := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
lists.Run(ctx)
|
||||
close(stopped)
|
||||
}()
|
||||
|
||||
time.Sleep(30 * time.Second)
|
||||
stop()
|
||||
<-stopped
|
||||
|
||||
wantFetches(t, servers, 1)
|
||||
|
||||
if lists.Failures(dropURL) != 0 || len(waiting(queue)) != 0 {
|
||||
t.Errorf("%d failures and alerts %+v, want none", lists.Failures(dropURL),
|
||||
waiting(queue))
|
||||
}
|
||||
|
||||
// Restarted an hour later with what reputation.json keeps, it fetches
|
||||
// tor.txt, never tried, at once, and drop.txt a refresh after its
|
||||
// cut-off try.
|
||||
time.Sleep(time.Hour)
|
||||
|
||||
restarted := &standIn{lists: map[string]string{
|
||||
dropURL: "203.0.113.7\n", torURL: "198.51.100.7\n",
|
||||
}}
|
||||
again := reputation.New(params(dropURL, torURL))
|
||||
again.SetTransport(restarted)
|
||||
|
||||
err := again.Load(lists.Snapshot())
|
||||
if err != nil {
|
||||
t.Fatalf("load: %v", err)
|
||||
}
|
||||
|
||||
run(t, again)
|
||||
wantFetches(t, restarted, 1)
|
||||
wantListedBy(t, again, "198.51.100.7", torURL)
|
||||
|
||||
time.Sleep(refresh - time.Hour - time.Nanosecond)
|
||||
wantFetches(t, restarted, 1)
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
wantFetches(t, restarted, 2)
|
||||
wantListedBy(t, again, "203.0.113.7", dropURL)
|
||||
})
|
||||
}
|
||||
|
||||
func TestASNLimitPercentFileGivesEachASNumberItsLowestPercentage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
synctest.Test(t, func(t *testing.T) {
|
||||
servers := &standIn{lists: map[string]string{
|
||||
asnURL: "# hosting networks\nAS14061:50 ; DigitalOcean\nas16276:25\n\n" +
|
||||
"AS14061:10\nAS14061:30\n",
|
||||
}}
|
||||
p := params()
|
||||
p.ASNLimitPercentURL = asnURL
|
||||
lists := start(t, servers, p)
|
||||
|
||||
for asn, want := range map[string]int64{"AS14061": 10, "AS16276": 25} {
|
||||
percent, listed := lists.ASNLimitPercent(asn)
|
||||
if !listed || percent != want {
|
||||
t.Errorf("%s has %d (listed %t), want %d", asn, percent, listed, want)
|
||||
}
|
||||
}
|
||||
|
||||
if _, listed := lists.ASNLimitPercent("AS64496"); listed {
|
||||
t.Error("AS64496 is listed")
|
||||
}
|
||||
|
||||
// A line that does not read fails the fetch.
|
||||
servers.set(asnURL, "AS14061:50\nAS16276\n")
|
||||
time.Sleep(refresh)
|
||||
synctest.Wait()
|
||||
|
||||
if lists.Failures(asnURL) != 1 {
|
||||
t.Errorf("%d failures, want 1", lists.Failures(asnURL))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestLoadDropsCopiesOfListsNotNamedAndRefusesOnesThatDoNotRead(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
fetched := time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
|
||||
kept := reputation.List{
|
||||
URL: dropURL, Tried: fetched, Fetched: fetched, Lines: []string{"203.0.113.9"},
|
||||
}
|
||||
lists := reputation.New(params(dropURL))
|
||||
|
||||
err := lists.Load([]reputation.List{kept, {
|
||||
URL: torURL, Tried: fetched, Fetched: fetched, Lines: []string{"198.51.100.9"},
|
||||
}})
|
||||
if err != nil {
|
||||
t.Fatalf("load: %v", err)
|
||||
}
|
||||
|
||||
if got := lists.Snapshot(); !reflect.DeepEqual(got, []reputation.List{kept}) {
|
||||
t.Errorf("copies %+v, want only %+v", got, kept)
|
||||
}
|
||||
|
||||
err = lists.Load([]reputation.List{{URL: dropURL, Fetched: fetched, Lines: []string{
|
||||
"; DROP", "203.0.113.300",
|
||||
}}})
|
||||
|
||||
const want = "the copy of " + dropURL +
|
||||
": line 2 is not an address or a netblock, such as 192.0.2.0/24"
|
||||
if err == nil || err.Error() != want {
|
||||
t.Errorf("error %v, want %s", err, want)
|
||||
}
|
||||
|
||||
if got := lists.Snapshot(); !reflect.DeepEqual(got, []reputation.List{kept}) {
|
||||
t.Errorf("copies %+v after the error, want %+v still", got, kept)
|
||||
}
|
||||
}
|
||||
|
||||
// standIn is a stand-in for the servers the lists are fetched from. It
|
||||
// notes the URL of each fetch.
|
||||
type standIn struct {
|
||||
mu sync.Mutex
|
||||
// lists are what it answers with, by URL; it answers a URL it has no
|
||||
// list for with 503, and none at all while hanging.
|
||||
lists map[string]string
|
||||
hanging bool
|
||||
fetches []string
|
||||
}
|
||||
|
||||
// RoundTrip has the stand-in answer req, in place of the network. A fetch
|
||||
// abandoned before the stand-in answers fails, as over the network.
|
||||
func (s *standIn) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
s.mu.Lock()
|
||||
s.fetches = append(s.fetches, req.URL.String())
|
||||
list, found := s.lists[req.URL.String()]
|
||||
hanging := s.hanging
|
||||
s.mu.Unlock()
|
||||
|
||||
if hanging {
|
||||
<-req.Context().Done()
|
||||
|
||||
return nil, req.Context().Err()
|
||||
}
|
||||
|
||||
status := http.StatusOK
|
||||
if !found {
|
||||
status = http.StatusServiceUnavailable
|
||||
}
|
||||
|
||||
return &http.Response{
|
||||
StatusCode: status,
|
||||
Status: fmt.Sprintf("%d %s", status, http.StatusText(status)),
|
||||
Header: http.Header{},
|
||||
Body: io.NopCloser(strings.NewReader(list)),
|
||||
Request: req,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// set has the stand-in answer listURL with list, or with 503 for "".
|
||||
func (s *standIn) set(listURL, list string) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
if list == "" {
|
||||
delete(s.lists, listURL)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
s.lists[listURL] = list
|
||||
}
|
||||
|
||||
// params returns the Params of the blocklists at urls, refreshed every
|
||||
// refresh, by the bubble's clock, with alerts to a queue that sends none.
|
||||
func params(urls ...string) reputation.Params {
|
||||
return reputation.Params{
|
||||
BlocklistURLs: urls,
|
||||
Refresh: refresh,
|
||||
Now: time.Now,
|
||||
ProcessLog: slog.New(slog.DiscardHandler),
|
||||
Alerts: newQueue(),
|
||||
}
|
||||
}
|
||||
|
||||
// newQueue returns a queue of alerts to a webhook that is never sent
|
||||
// them, with a cooldown of cooldown.
|
||||
func newQueue() *alerts.Queue {
|
||||
return alerts.New(alerts.Params{
|
||||
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
||||
Events: alerts.Events(),
|
||||
Cooldown: cooldown,
|
||||
Now: time.Now,
|
||||
})
|
||||
}
|
||||
|
||||
// start returns the lists of p, fetched through servers by Run, which runs
|
||||
// until the test ends, once Run has fetched those due at start.
|
||||
func start(t *testing.T, servers *standIn, p reputation.Params) *reputation.Lists {
|
||||
t.Helper()
|
||||
|
||||
lists := reputation.New(p)
|
||||
lists.SetTransport(servers)
|
||||
run(t, lists)
|
||||
|
||||
return lists
|
||||
}
|
||||
|
||||
// run runs lists' Run until the test ends, and waits until it has fetched
|
||||
// the lists due.
|
||||
func run(t *testing.T, lists *reputation.Lists) {
|
||||
t.Helper()
|
||||
|
||||
ctx, stop := context.WithCancel(t.Context())
|
||||
stopped := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
lists.Run(ctx)
|
||||
close(stopped)
|
||||
}()
|
||||
|
||||
t.Cleanup(func() {
|
||||
stop()
|
||||
<-stopped
|
||||
})
|
||||
|
||||
synctest.Wait()
|
||||
}
|
||||
|
||||
// wantFetches waits until Run has made the fetches due, and checks how
|
||||
// many the servers have had.
|
||||
func wantFetches(t *testing.T, servers *standIn, want int) {
|
||||
t.Helper()
|
||||
|
||||
synctest.Wait()
|
||||
|
||||
servers.mu.Lock()
|
||||
got := len(servers.fetches)
|
||||
servers.mu.Unlock()
|
||||
|
||||
if got != want {
|
||||
t.Errorf("%d fetches, want %d", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// wantListedBy checks the URLs of the blocklists lists says list addr.
|
||||
func wantListedBy(t *testing.T, lists *reputation.Lists, addr string, want ...string) {
|
||||
t.Helper()
|
||||
|
||||
got := lists.ListedBy(netip.MustParseAddr(addr))
|
||||
if !slices.Equal(got, want) {
|
||||
t.Errorf("%s is listed by %v, want %v", addr, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// waiting returns the alerts waiting in queue.
|
||||
func waiting(queue *alerts.Queue) []alerts.Alert {
|
||||
return queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||
}
|
||||
|
||||
// wantAlert checks that want is the one alert waiting in queue, and that
|
||||
// the cooldown has held back one repeat of it.
|
||||
func wantAlert(t *testing.T, queue *alerts.Queue, want alerts.Alert) {
|
||||
t.Helper()
|
||||
|
||||
got := waiting(queue)
|
||||
if len(got) != 1 || !reflect.DeepEqual(got[0], want) || queue.Suppressed() != 1 {
|
||||
t.Errorf("alerts waiting %+v, %d held back, want only %+v and 1", got,
|
||||
queue.Suppressed(), want)
|
||||
}
|
||||
}
|
||||
@@ -35,7 +35,9 @@ const (
|
||||
// rule.
|
||||
ActionRuleBlocked = "rule_blocked"
|
||||
// ActionDenied is a request refused because its client is in
|
||||
// SWWAF_DENY_NETS.
|
||||
// SWWAF_DENY_NETS, in a blocklist while SWWAF_BLOCKLIST_ACTION is deny,
|
||||
// or listed by a DNSBL zone, or scored a hit by AbuseIPDB, while
|
||||
// SWWAF_REPUTATION_ACTION is deny.
|
||||
ActionDenied = "denied"
|
||||
// ActionCountryDenied is a request refused for its client's country.
|
||||
ActionCountryDenied = "country_denied"
|
||||
@@ -45,7 +47,7 @@ const (
|
||||
)
|
||||
|
||||
// OffenceLimit is the offence a request line names for a request that
|
||||
// broke a rate limit.
|
||||
// broke a rate limit, or whose bytes broke a byte limit.
|
||||
const OffenceLimit = "limit"
|
||||
|
||||
// timeLayout is RFC 3339 with milliseconds.
|
||||
@@ -117,14 +119,30 @@ type Line struct {
|
||||
// ActionBanned, ActionCountryDenied, ActionRateLimited or
|
||||
// ActionRuleBlocked.
|
||||
WouldAction string `json:"would_action,omitempty"`
|
||||
// Counts are the client's requests as the rate limits counted them
|
||||
// with this one, for a request they counted.
|
||||
// LimitPercent and LimitPercentSetting are, for a request the rate
|
||||
// limits counted whose client a biased threshold gives a percentage of
|
||||
// the rate limits below 100, that percentage and the setting that gave
|
||||
// it. BytesPercent and BytesPercentSetting are the same for the byte
|
||||
// limits.
|
||||
LimitPercent *int64 `json:"limit_percent,omitempty"`
|
||||
LimitPercentSetting string `json:"limit_percent_setting,omitempty"`
|
||||
BytesPercent *int64 `json:"bytes_percent,omitempty"`
|
||||
BytesPercentSetting string `json:"bytes_percent_setting,omitempty"`
|
||||
// Counts are, for a request the rate limits counted, the client's
|
||||
// requests as they counted them with this one, and its bytes as the
|
||||
// byte limits counted them, with this request's once it has ended if
|
||||
// they count them.
|
||||
Counts ratelimit.Counts `json:"counts,omitzero"`
|
||||
// RuleIDs are the ids of the rule file rules the request matched.
|
||||
RuleIDs []string `json:"rule_ids,omitempty"`
|
||||
// LimitHit is the window whose rate limit the request went over:
|
||||
// minute, hour or day.
|
||||
// LimitHit is the window whose limit the request went over, named as
|
||||
// Counts names its count: minute, hour or day for a rate limit, and
|
||||
// minute_bytes, hour_bytes or day_bytes for a byte limit.
|
||||
LimitHit string `json:"limit_hit,omitempty"`
|
||||
// Reputation are the URLs of the blocklists that list the client, then
|
||||
// the DNSBL zones whose verdict lists it, their keys masked, then
|
||||
// abuseipdb when its score is a hit.
|
||||
Reputation []string `json:"reputation,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,
|
||||
@@ -175,8 +193,11 @@ func Milliseconds(d time.Duration) float64 {
|
||||
// NewProcessLogger returns the logger for the process's own messages:
|
||||
// JSON lines on w, marked "type":"process", with the time in the same form
|
||||
// as a request line's, and instanceName, SWWAF_INSTANCE_NAME, as instance.
|
||||
func NewProcessLogger(w io.Writer, instanceName string) *slog.Logger {
|
||||
// It writes only the messages at level, SWWAF_LOG_LEVEL, or more severe;
|
||||
// the request lines Write writes are never held back.
|
||||
func NewProcessLogger(w io.Writer, instanceName string, level slog.Level) *slog.Logger {
|
||||
handler := slog.NewJSONHandler(w, &slog.HandlerOptions{
|
||||
Level: level,
|
||||
ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr {
|
||||
if attr.Key == slog.TimeKey && len(groups) == 0 {
|
||||
return slog.String(slog.TimeKey, FormatTime(attr.Value.Time()))
|
||||
|
||||
@@ -3,6 +3,8 @@ package requestlog_test
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -70,7 +72,8 @@ func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
|
||||
|
||||
var out bytes.Buffer
|
||||
|
||||
requestlog.NewProcessLogger(&out, "fsn1app1/gitea").Info("starting", "version", "v1")
|
||||
requestlog.NewProcessLogger(&out, "fsn1app1/gitea", slog.LevelInfo).Info("starting",
|
||||
"version", "v1")
|
||||
|
||||
var fields map[string]any
|
||||
|
||||
@@ -94,3 +97,47 @@ func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
|
||||
t.Errorf("process line time %q, want now in UTC with milliseconds", timeText)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessLoggerWritesTheMessagesAtItsLevelOrMoreSevere(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
levels := []slog.Level{
|
||||
slog.LevelDebug, slog.LevelInfo, slog.LevelWarn, slog.LevelError,
|
||||
}
|
||||
|
||||
for i, level := range levels {
|
||||
t.Run(level.String(), func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var out bytes.Buffer
|
||||
|
||||
processLog := requestlog.NewProcessLogger(&out, "fsn1app1/gitea", level)
|
||||
for _, at := range levels {
|
||||
processLog.Log(t.Context(), at, "message")
|
||||
}
|
||||
|
||||
var got, want []string
|
||||
|
||||
for line := range strings.Lines(out.String()) {
|
||||
var fields struct {
|
||||
Level string `json:"level"`
|
||||
}
|
||||
|
||||
err := json.Unmarshal([]byte(line), &fields)
|
||||
if err != nil {
|
||||
t.Fatalf("decode %q: %v", line, err)
|
||||
}
|
||||
|
||||
got = append(got, fields.Level)
|
||||
}
|
||||
|
||||
for _, written := range levels[i:] {
|
||||
want = append(want, written.String())
|
||||
}
|
||||
|
||||
if !slices.Equal(got, want) {
|
||||
t.Errorf("lines at %v, want %v", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
|
||||
// the rule files and the state files, serves requests until it is told to
|
||||
// stop, and then stops in an orderly way, writing the state files.
|
||||
// the rule files, the lookup database 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 (
|
||||
@@ -20,6 +21,7 @@ import (
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
||||
"sneak.berlin/go/smallwebwaf/internal/reputation"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
||||
"sneak.berlin/go/smallwebwaf/internal/state"
|
||||
@@ -64,12 +66,14 @@ func Main(version string) int {
|
||||
})
|
||||
}
|
||||
|
||||
// Run reads the settings, the rule files and the state files, then serves
|
||||
// requests until ctx is done. It returns the process's exit status, 1
|
||||
// when smallwebwaf cannot start.
|
||||
// Run reads the settings, the rule files, the lookup database 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 {
|
||||
// Until the settings are read, the one message is an invalid setting's
|
||||
// error, which every SWWAF_LOG_LEVEL lets through.
|
||||
processLog := requestlog.NewProcessLogger(params.Stdout,
|
||||
config.InstanceName(params.LookupEnv))
|
||||
config.InstanceName(params.LookupEnv), slog.LevelError)
|
||||
|
||||
cfg, err := config.FromEnvironment(params.LookupEnv)
|
||||
if err != nil {
|
||||
@@ -87,8 +91,11 @@ func Run(ctx context.Context, params Params) int {
|
||||
if cfg.LogRemoteURL != nil {
|
||||
remote = newRemoteLogSender(cfg)
|
||||
stdout = io.MultiWriter(params.Stdout, remote)
|
||||
processLog = requestlog.NewProcessLogger(stdout, cfg.InstanceName)
|
||||
}
|
||||
|
||||
processLog = requestlog.NewProcessLogger(stdout, cfg.InstanceName, cfg.LogLevel)
|
||||
|
||||
if remote != nil {
|
||||
stopSending := startSending(ctx, remote, processLog)
|
||||
defer stopSending()
|
||||
}
|
||||
@@ -110,21 +117,17 @@ func Run(ctx context.Context, params Params) int {
|
||||
return 1
|
||||
}
|
||||
|
||||
server := proxy.New(proxy.Params{
|
||||
Config: cfg,
|
||||
RequestLog: stdout,
|
||||
ProcessLog: processLog,
|
||||
GeoJSURL: lookup.URL,
|
||||
Now: now,
|
||||
Rules: ruleFiles,
|
||||
Alerts: alertQueue,
|
||||
})
|
||||
server, err := newServer(cfg, stdout, processLog, now, ruleFiles, alertQueue)
|
||||
if err != nil {
|
||||
processLog.Error("cannot use the lookup database", "error", err.Error())
|
||||
|
||||
return 1
|
||||
}
|
||||
|
||||
if remote != nil {
|
||||
server.Metrics.AddRemoteLog(remote)
|
||||
}
|
||||
|
||||
server.Metrics.AddAlerts(alertQueue)
|
||||
|
||||
files, err := loadStateFiles(cfg, server, alertQueue, now, processLog)
|
||||
if err != nil {
|
||||
processLog.Error("cannot use the state files", "error", err.Error())
|
||||
@@ -145,7 +148,50 @@ func Run(ctx context.Context, params Params) int {
|
||||
"address", listener.Addr().String(),
|
||||
"settings", cfg)
|
||||
|
||||
return serve(ctx, server.Server, listener, files, ruleFiles, alertQueue, processLog)
|
||||
return serve(ctx, server, listener, files, ruleFiles, alertQueue, processLog)
|
||||
}
|
||||
|
||||
// newServer returns the server smallwebwaf runs, with the metrics of the
|
||||
// alerts, after reading the lookup database while SWWAF_LOOKUP_SOURCE is
|
||||
// file. A lookup database that cannot be read is an error.
|
||||
func newServer(
|
||||
cfg *config.Config, stdout io.Writer, processLog *slog.Logger,
|
||||
now func() time.Time, ruleFiles *rules.Files, alertQueue *alerts.Queue,
|
||||
) (*proxy.Server, error) {
|
||||
var lookupFile *lookup.File
|
||||
|
||||
if cfg.LookupSource == "file" {
|
||||
var err error
|
||||
|
||||
lookupFile, err = lookup.OpenFile(lookup.FileParams{
|
||||
Path: cfg.LookupDBPath,
|
||||
Now: now,
|
||||
ProcessLog: processLog,
|
||||
Alerts: alertQueue,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
server := proxy.New(proxy.Params{
|
||||
Config: cfg,
|
||||
RequestLog: stdout,
|
||||
ProcessLog: processLog,
|
||||
GeoJSURL: lookup.URL,
|
||||
AbuseIPDBURL: reputation.AbuseIPDBURL,
|
||||
LookupFile: lookupFile,
|
||||
Now: now,
|
||||
Rules: ruleFiles,
|
||||
Alerts: alertQueue,
|
||||
})
|
||||
server.Metrics.AddAlerts(alertQueue)
|
||||
|
||||
if lookupFile != nil {
|
||||
server.Metrics.AddLookupFile(lookupFile.LastRead, lookupFile.ReadFailures)
|
||||
}
|
||||
|
||||
return server, nil
|
||||
}
|
||||
|
||||
// loadStateFiles reads the state files into the parts of server and into
|
||||
@@ -161,7 +207,11 @@ func loadStateFiles(
|
||||
Ledger: server.Ledger,
|
||||
Limiter: server.Limiter,
|
||||
GeoJS: server.GeoJS,
|
||||
Lists: server.Lists,
|
||||
DNSBL: server.DNSBL,
|
||||
AbuseIPDB: server.AbuseIPDB,
|
||||
Alerts: alertQueue,
|
||||
Anomalies: server.Anomalies,
|
||||
Now: now,
|
||||
ProcessLog: processLog,
|
||||
Metrics: server.Metrics,
|
||||
@@ -227,11 +277,13 @@ func startSending(
|
||||
|
||||
// serve serves requests on listener, writes the state files as they are
|
||||
// due, takes in an admin's edits of them, reads the rule files again as
|
||||
// they change, and sends the alerts, until ctx is done. Then it gives the
|
||||
// requests in progress shutdownTimeout to finish, and writes every state
|
||||
// file, alerts.json with the alerts still waiting.
|
||||
// they change, and the lookup database when it is replaced, fetches the
|
||||
// lists the settings name by URL as they are due, and sends the alerts,
|
||||
// until ctx is done. Then it gives the requests in progress
|
||||
// shutdownTimeout to finish, and writes every state file, alerts.json with
|
||||
// the alerts still waiting.
|
||||
func serve(
|
||||
ctx context.Context, server *http.Server, listener net.Listener,
|
||||
ctx context.Context, server *proxy.Server, listener net.Listener,
|
||||
files *state.Files, ruleFiles *rules.Files, alertQueue *alerts.Queue,
|
||||
processLog *slog.Logger,
|
||||
) int {
|
||||
@@ -247,6 +299,12 @@ func serve(
|
||||
written := inBackground(func() { files.Run(writing) })
|
||||
watched := inBackground(func() { files.Watch(writing) })
|
||||
rulesWatched := inBackground(func() { ruleFiles.Watch(writing) })
|
||||
lookupFileWatched := inBackground(func() {
|
||||
if server.LookupFile != nil {
|
||||
server.LookupFile.Watch(writing)
|
||||
}
|
||||
})
|
||||
listsFetched := inBackground(func() { server.Lists.Run(writing) })
|
||||
alertsSent := inBackground(func() { alertQueue.Run(writing) })
|
||||
|
||||
select {
|
||||
@@ -289,6 +347,8 @@ func serve(
|
||||
<-written
|
||||
<-watched
|
||||
<-rulesWatched
|
||||
<-lookupFileWatched
|
||||
<-listsFetched
|
||||
<-alertsSent
|
||||
|
||||
err = files.WriteAll()
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
|
||||
"sneak.berlin/go/smallwebwaf/internal/smallwebwaf"
|
||||
)
|
||||
|
||||
@@ -37,7 +38,10 @@ const (
|
||||
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
|
||||
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
|
||||
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
||||
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
|
||||
rulesDir = "SWWAF_RULES_DIR"
|
||||
lookupSource = "SWWAF_LOOKUP_SOURCE"
|
||||
lookupDBPath = "SWWAF_LOOKUP_DB_PATH"
|
||||
instanceName = "SWWAF_INSTANCE_NAME"
|
||||
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
|
||||
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
|
||||
@@ -48,6 +52,8 @@ const (
|
||||
instance = "fsn1app1/gitea"
|
||||
// greeting is what the tests' app answers.
|
||||
greeting = "hello from the app"
|
||||
// placed is the client the tests' lookup databases place.
|
||||
placed = "203.0.113.9"
|
||||
)
|
||||
|
||||
// output collects what smallwebwaf writes on stdout.
|
||||
@@ -243,6 +249,54 @@ func TestServesUntilToldToStop(t *testing.T) {
|
||||
out.line(t, "msg", "stopped")
|
||||
}
|
||||
|
||||
func TestLogLevelHoldsBackTheLessSevereProcessLines(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// A list that cannot be fetched has a warning written once smallwebwaf
|
||||
// serves, after its starting line.
|
||||
lists := httptest.NewServer(http.HandlerFunc(
|
||||
func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
}))
|
||||
t.Cleanup(lists.Close)
|
||||
|
||||
ctx, stop := context.WithCancel(t.Context())
|
||||
out := &output{}
|
||||
exited := make(chan int, 1)
|
||||
|
||||
go func() {
|
||||
exited <- run(ctx, map[string]string{
|
||||
listenAddr: localhost + ":0",
|
||||
stateDir: t.TempDir(),
|
||||
rulesDir: t.TempDir(),
|
||||
"SWWAF_BLOCKLIST_URLS": lists.URL + "/tor.txt",
|
||||
"SWWAF_LOG_LEVEL": "warn",
|
||||
}, out)
|
||||
}()
|
||||
|
||||
out.line(t, "msg", "fetching a list failed")
|
||||
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")
|
||||
}
|
||||
|
||||
// Not one of the info lines from the start to the stop.
|
||||
for line := range strings.Lines(out.text()) {
|
||||
var fields map[string]any
|
||||
|
||||
err := json.Unmarshal([]byte(line), &fields)
|
||||
if err != nil || fields["level"] == "INFO" {
|
||||
t.Errorf("line %q (%v), want none at info", line, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestEveryLogLineAndMetricCarriesTheInstanceName(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -495,6 +549,212 @@ func TestRulesDirThatDoesNotExistStopsTheStart(t *testing.T) {
|
||||
": no such file or directory")
|
||||
}
|
||||
|
||||
func TestLookupDatabaseReplacedWhileRunningTakesEffect(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const token = "0123456789abcdef0123456789abcdef"
|
||||
|
||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||
writeLookupDatabase(t, path, "DE")
|
||||
|
||||
env := map[string]string{
|
||||
listenAddr: localhost + ":0",
|
||||
upstreamURL: startApp(t),
|
||||
stateDir: t.TempDir(),
|
||||
rulesDir: t.TempDir(),
|
||||
trustedProxies: localhost + "/32",
|
||||
metricsToken: token,
|
||||
instanceName: instance,
|
||||
lookupSource: "file",
|
||||
lookupDBPath: path,
|
||||
"SWWAF_DENIED_COUNTRIES": "kp",
|
||||
// The requests sent until a replacement takes effect, and those
|
||||
// for the metrics, must not break a rate limit, whose ban would
|
||||
// refuse them too.
|
||||
rateLimitExemptNets: placed + "," + localhost,
|
||||
}
|
||||
began := time.Now()
|
||||
// Each replacement is written beside the file and renamed over it, as
|
||||
// a refresh is.
|
||||
replacement := path + ".new"
|
||||
|
||||
var lastRead float64
|
||||
|
||||
runUntilStopped(t, env, func(url string) {
|
||||
wantStatus(t, url, placed, http.StatusOK)
|
||||
|
||||
// smallwebwaf reads the file again once it watches its directory,
|
||||
// which may be after the first replacement. Once that has taken
|
||||
// effect, only the watch can show the next. Each takes as long as
|
||||
// it takes, so that a slow test process cannot fail the test.
|
||||
writeLookupDatabase(t, replacement, "KP")
|
||||
rename(t, replacement, path)
|
||||
|
||||
for statusFrom(t, url, placed) != http.StatusForbidden {
|
||||
time.Sleep(pollInterval)
|
||||
}
|
||||
|
||||
writeLookupDatabase(t, replacement, "DE")
|
||||
rename(t, replacement, path)
|
||||
|
||||
for statusFrom(t, url, placed) != http.StatusOK {
|
||||
time.Sleep(pollInterval)
|
||||
}
|
||||
|
||||
// One that cannot be read leaves the file read before in use.
|
||||
err := os.WriteFile(replacement, []byte("not a lookup database\n"), 0o600)
|
||||
if err != nil {
|
||||
t.Fatalf("write %s: %v", replacement, err)
|
||||
}
|
||||
|
||||
rename(t, replacement, path)
|
||||
|
||||
metrics := metricsWith(t, url+"_smallwebwaf/metrics", token,
|
||||
`smallwebwaf_lookup_database_read_failures_total{instance="fsn1app1/gitea"} 1`)
|
||||
lastRead = seriesValue(t, metrics,
|
||||
`smallwebwaf_lookup_database_last_read_timestamp_seconds{instance="fsn1app1/gitea"}`)
|
||||
|
||||
wantStatus(t, url, placed, http.StatusOK)
|
||||
})
|
||||
|
||||
if lastRead < float64(began.Unix()) || lastRead > float64(time.Now().Unix()) {
|
||||
t.Errorf("the lookup database was last read at %v, want a time in the test", lastRead)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBlocklistTriesAndCopiesKeptInReputationJSONAcrossRestarts(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const (
|
||||
token = "0123456789abcdef0123456789abcdef"
|
||||
torPath = "/tor.txt"
|
||||
dropPath = "/drop.txt"
|
||||
// copyright is the DROP list's date and copyright line.
|
||||
copyright = "; Spamhaus DROP List 2026/10/07 - (c) 2026 The Spamhaus Project SLL"
|
||||
)
|
||||
|
||||
failing, torFetches := new(atomic.Bool), new(atomic.Int32)
|
||||
lists := map[string]string{
|
||||
torPath: "198.51.100.0/24\n",
|
||||
dropPath: copyright + "\n" + placed + " ; SBL1\n",
|
||||
}
|
||||
server := httptest.NewServer(http.HandlerFunc(
|
||||
func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == torPath {
|
||||
torFetches.Add(1)
|
||||
}
|
||||
|
||||
if failing.Load() {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
_, _ = io.WriteString(w, lists[r.URL.Path])
|
||||
}))
|
||||
t.Cleanup(server.Close)
|
||||
|
||||
torURL, dropURL := server.URL+torPath, server.URL+dropPath
|
||||
dir := t.TempDir()
|
||||
env := map[string]string{
|
||||
listenAddr: localhost + ":0",
|
||||
upstreamURL: startApp(t),
|
||||
stateDir: dir,
|
||||
rulesDir: t.TempDir(),
|
||||
trustedProxies: localhost + "/32",
|
||||
metricsToken: token,
|
||||
instanceName: instance,
|
||||
"SWWAF_BLOCKLIST_URLS": torURL,
|
||||
// The requests sent until the list takes effect, and those for the
|
||||
// metrics, must not break a rate limit, whose ban would refuse them
|
||||
// too.
|
||||
rateLimitExemptNets: placed + "," + localhost,
|
||||
}
|
||||
failures := `smallwebwaf_reputation_failures_total{instance="fsn1app1/gitea",` +
|
||||
`source="` + torURL + `"} `
|
||||
|
||||
// tor.txt cannot be fetched at first, which is counted.
|
||||
failing.Store(true)
|
||||
runUntilStopped(t, env, func(url string) {
|
||||
metricsWith(t, url+"_smallwebwaf/metrics", token, failures+"1")
|
||||
})
|
||||
|
||||
// Restarted with drop.txt named after it and the server answering, tor.txt
|
||||
// waits SWWAF_BLOCKLIST_REFRESH after its failed try, kept in reputation.json,
|
||||
// while drop.txt, never tried, is fetched at once. Lists are fetched in the
|
||||
// order named, so once drop.txt refuses the client, tor.txt has had its turn.
|
||||
failing.Store(false)
|
||||
|
||||
env["SWWAF_BLOCKLIST_URLS"] = torURL + "," + dropURL
|
||||
out := runUntilStopped(t, env, func(url string) {
|
||||
for statusFrom(t, url, placed) != http.StatusForbidden {
|
||||
time.Sleep(pollInterval)
|
||||
}
|
||||
})
|
||||
wantDeniedByList(t, out.line(t, "action", "denied"), dropURL)
|
||||
|
||||
if fetches := torFetches.Load(); fetches != 1 {
|
||||
t.Errorf("tor.txt fetched %d times, want once, before the restart", fetches)
|
||||
}
|
||||
|
||||
// After another restart, with the server failing, the copy of drop.txt
|
||||
// kept in reputation.json, its copyright line included, refuses the
|
||||
// client from the first request.
|
||||
failing.Store(true)
|
||||
|
||||
out = runUntilStopped(t, env, func(url string) {
|
||||
wantStatus(t, url, placed, http.StatusForbidden)
|
||||
})
|
||||
wantDeniedByList(t, out.line(t, "type", "request"), dropURL)
|
||||
|
||||
path := filepath.Join(dir, "reputation.json")
|
||||
|
||||
kept, err := os.ReadFile(path) //nolint:gosec // a file in the test's directory
|
||||
if err != nil || !strings.Contains(string(kept), `"`+copyright+`"`) {
|
||||
t.Errorf("reputation.json holds\n%s\nwant the copy with %q (%v)", kept, copyright,
|
||||
err)
|
||||
}
|
||||
}
|
||||
|
||||
// wantDeniedByList checks that the request log line is of a request the
|
||||
// blocklist at listURL refused.
|
||||
func wantDeniedByList(t *testing.T, line map[string]any, listURL string) {
|
||||
t.Helper()
|
||||
|
||||
reputation, _ := line["reputation"].([]any)
|
||||
if line["action"] != "denied" || len(reputation) != 1 || reputation[0] != listURL {
|
||||
t.Errorf("request log line %v, want one denied for %s", line, listURL)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLookupDatabaseThatCannotBeReadStopsTheStart(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
||||
|
||||
// If it starts instead, it is stopped after waitLimit.
|
||||
ctx, stop := context.WithTimeout(t.Context(), waitLimit)
|
||||
defer stop()
|
||||
|
||||
out := &output{}
|
||||
|
||||
status := run(ctx, map[string]string{
|
||||
listenAddr: localhost + ":0", stateDir: t.TempDir(), rulesDir: t.TempDir(),
|
||||
lookupSource: "file", lookupDBPath: path,
|
||||
}, out)
|
||||
if status != 1 {
|
||||
t.Fatalf("exit status %d, want 1", status)
|
||||
}
|
||||
|
||||
want := "SWWAF_LOOKUP_DB_PATH cannot be read: open " + path +
|
||||
": no such file or directory"
|
||||
|
||||
line := out.line(t, "msg", "cannot use the lookup database")
|
||||
if line["error"] != want {
|
||||
t.Errorf("start refused with %v, want the error %q", line, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEveryLineIsAlsoSentToTheRemoteLogEndpoint(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -908,7 +1168,7 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
|
||||
"SWWAF_REQUEST_MAX_BYTES": "100M",
|
||||
"SWWAF_RESPONSE_MAX_BYTES": "5G",
|
||||
"SWWAF_ALLOW_NETS": "",
|
||||
"SWWAF_RATE_LIMIT_EXEMPT_NETS": "",
|
||||
rateLimitExemptNets: "",
|
||||
"SWWAF_DENY_NETS": "",
|
||||
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
|
||||
"SWWAF_RATE_LIMIT_PER_HOUR": "10000",
|
||||
@@ -1046,6 +1306,50 @@ func metricsWith(t *testing.T, url, token string, series ...string) string {
|
||||
}
|
||||
}
|
||||
|
||||
// seriesValue returns the value of series, such as
|
||||
// name{instance="app"}, in metrics.
|
||||
func seriesValue(t *testing.T, metrics, series string) float64 {
|
||||
t.Helper()
|
||||
|
||||
for line := range strings.Lines(metrics) {
|
||||
value, found := strings.CutPrefix(strings.TrimSpace(line), series+" ")
|
||||
if !found {
|
||||
continue
|
||||
}
|
||||
|
||||
number, err := strconv.ParseFloat(value, 64)
|
||||
if err != nil {
|
||||
t.Fatalf("%s has the value %q: %v", series, value, err)
|
||||
}
|
||||
|
||||
return number
|
||||
}
|
||||
|
||||
t.Fatalf("no %s in the metrics:\n%s", series, metrics)
|
||||
|
||||
return 0
|
||||
}
|
||||
|
||||
// rename renames the file at from to, replacing any file there.
|
||||
func rename(t *testing.T, from, to string) {
|
||||
t.Helper()
|
||||
|
||||
err := os.Rename(from, to)
|
||||
if err != nil {
|
||||
t.Fatalf("rename %s to %s: %v", from, to, err)
|
||||
}
|
||||
}
|
||||
|
||||
// writeLookupDatabase writes a lookup database at path that places the
|
||||
// client placed in country, and no other address.
|
||||
func writeLookupDatabase(t *testing.T, path, country string) {
|
||||
t.Helper()
|
||||
|
||||
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
||||
placed + "/32": {ASN: "AS64496", ASName: "Example Net", Country: country},
|
||||
})
|
||||
}
|
||||
|
||||
// askAsAdmin sends a request with method to url, with body and
|
||||
// adminSecret, and checks that it is answered 200.
|
||||
func askAsAdmin(t *testing.T, method, url, body string) {
|
||||
|
||||
+234
-34
@@ -1,8 +1,11 @@
|
||||
// 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, lookups.json GeoJS's answers, and alerts.json the cooldowns,
|
||||
// the hour under way and the alerts waiting for each destination. Load
|
||||
// history, lookups.json GeoJS's answers, reputation.json the last try and
|
||||
// last good copy of each list fetched from a URL, the DNSBL zones'
|
||||
// verdicts, and AbuseIPDB's scores and checks spent, and alerts.json the
|
||||
// cooldowns, the hour under way, the alerts waiting for each destination
|
||||
// and the anomaly counters. Load
|
||||
// reads them at start, Watch takes in an admin's edit of one while
|
||||
// smallwebwaf runs, and Run and WriteAll write them. The disk is read and
|
||||
// written outside the parts' locks, which are held only to take a
|
||||
@@ -29,10 +32,12 @@ import (
|
||||
|
||||
"github.com/fsnotify/fsnotify"
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/anomaly"
|
||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
"sneak.berlin/go/smallwebwaf/internal/reputation"
|
||||
)
|
||||
|
||||
// version is the version of the files' format, the only one read.
|
||||
@@ -47,6 +52,7 @@ const (
|
||||
bansJSON = "bans.json"
|
||||
clientsJSON = "clients.json"
|
||||
lookupsJSON = "lookups.json"
|
||||
reputationJSON = "reputation.json"
|
||||
alertsJSON = "alerts.json"
|
||||
)
|
||||
|
||||
@@ -56,6 +62,7 @@ var (
|
||||
errMissing = errors.New("has no")
|
||||
errCause = errors.New("is not limit, attack or admin")
|
||||
errDestination = errors.New("is not webhook, slack or ntfy")
|
||||
errScope = errors.New("is not client, net, asn, total or watch")
|
||||
errWaitingList = errors.New(`waiting is a list, but now lists the alerts by ` +
|
||||
`destination: put the list under "webhook", as "waiting": {"webhook": [...]}, ` +
|
||||
`or remove the file`)
|
||||
@@ -70,13 +77,17 @@ type Params struct {
|
||||
// is (SWWAF_STATE_COUNTER_INTERVAL).
|
||||
WriteDelay time.Duration
|
||||
CounterInterval time.Duration
|
||||
// Ledger, Limiter, GeoJS and Alerts hold the state. Alerts also
|
||||
// receive a file_error alert for an edit set aside, and for a write
|
||||
// that fails while smallwebwaf runs.
|
||||
// Ledger, Limiter, GeoJS, Lists, DNSBL, AbuseIPDB, Alerts and Anomalies
|
||||
// hold the state. Alerts also receive a file_error alert for an edit set
|
||||
// aside, and for a write that fails while smallwebwaf runs.
|
||||
Ledger *bans.Ledger
|
||||
Limiter *ratelimit.Limiter
|
||||
GeoJS *lookup.GeoJS
|
||||
Lists *reputation.Lists
|
||||
DNSBL *reputation.DNSBL
|
||||
AbuseIPDB *reputation.AbuseIPDB
|
||||
Alerts *alerts.Queue
|
||||
Anomalies *anomaly.Counters
|
||||
// Now tells the time by which the counters' buckets run out, normally
|
||||
// time.Now in UTC.
|
||||
Now func() time.Time
|
||||
@@ -134,12 +145,24 @@ type lookupsFile struct {
|
||||
Lookups []lookup.Answer `json:"lookups"`
|
||||
}
|
||||
|
||||
// reputationFile is reputation.json, indented for an admin to read and
|
||||
// edit, so that each line of a list's copy is on a line of its own.
|
||||
type reputationFile struct {
|
||||
Version int `json:"version"`
|
||||
Lists []reputation.List `json:"lists"`
|
||||
Verdicts []reputation.Verdict `json:"verdicts"`
|
||||
AbuseIPDB reputation.Checks `json:"abuseipdb"`
|
||||
}
|
||||
|
||||
// alertsFile is alerts.json, indented for an admin to read and edit.
|
||||
//
|
||||
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
||||
type alertsFile struct {
|
||||
Version int `json:"version"`
|
||||
Cooldowns []alerts.Cooldown `json:"cooldowns"`
|
||||
Hour alerts.Hour `json:"hour"`
|
||||
Waiting map[string][]alerts.Alert `json:"waiting"`
|
||||
AnomalyCounters []anomaly.Counter `json:"anomaly_counters"`
|
||||
}
|
||||
|
||||
// stateFile is the struct of a state file. Once the file is decoded, its
|
||||
@@ -152,10 +175,10 @@ type stateFile interface {
|
||||
}
|
||||
|
||||
// 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
|
||||
// in it into the parts of Params that hold the state. 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)
|
||||
@@ -168,16 +191,17 @@ func Load(params Params) (*Files, error) {
|
||||
bansRead, bansErr := f.read(bansJSON)
|
||||
clientsRead, clientsErr := f.read(clientsJSON)
|
||||
lookupsRead, lookupsErr := f.read(lookupsJSON)
|
||||
reputationRead, reputationErr := f.read(reputationJSON)
|
||||
alertsRead, alertsErr := f.read(alertsJSON)
|
||||
|
||||
err = errors.Join(bansErr, clientsErr, lookupsErr, alertsErr)
|
||||
err = errors.Join(bansErr, clientsErr, lookupsErr, reputationErr, alertsErr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
params.ProcessLog.Info("read the state files", "directory", params.Dir,
|
||||
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead,
|
||||
"alerts_waiting", alertsRead)
|
||||
"lists", reputationRead, "alerts_waiting", alertsRead)
|
||||
|
||||
return f, nil
|
||||
}
|
||||
@@ -206,7 +230,9 @@ func (f *Files) Run(ctx context.Context) {
|
||||
|
||||
f.logFailure(bansJSON, f.writeFile(bansJSON))
|
||||
case <-interval.C:
|
||||
for _, name := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} {
|
||||
for _, name := range []string{
|
||||
bansJSON, clientsJSON, lookupsJSON, reputationJSON, alertsJSON,
|
||||
} {
|
||||
f.logFailure(name, f.writeFile(name))
|
||||
}
|
||||
}
|
||||
@@ -217,7 +243,7 @@ func (f *Files) Run(ctx context.Context) {
|
||||
// fails does not keep the others from being written.
|
||||
func (f *Files) WriteAll() error {
|
||||
return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON),
|
||||
f.writeFile(lookupsJSON), f.writeFile(alertsJSON))
|
||||
f.writeFile(lookupsJSON), f.writeFile(reputationJSON), f.writeFile(alertsJSON))
|
||||
}
|
||||
|
||||
// Watch watches Dir until ctx is done, and takes in an admin's edit of a
|
||||
@@ -252,7 +278,7 @@ func (f *Files) Watch(ctx context.Context) {
|
||||
return
|
||||
case event := <-watcher.Events:
|
||||
switch name := filepath.Base(event.Name); name {
|
||||
case bansJSON, clientsJSON, lookupsJSON, alertsJSON:
|
||||
case bansJSON, clientsJSON, lookupsJSON, reputationJSON, alertsJSON:
|
||||
f.fileChanged(name)
|
||||
}
|
||||
case err = <-watcher.Errors:
|
||||
@@ -398,7 +424,40 @@ func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
|
||||
|
||||
f.params.GeoJS.Load(file.Lookups)
|
||||
entries = len(file.Lookups)
|
||||
case reputationJSON:
|
||||
var file reputationFile
|
||||
|
||||
err := parse(path, data, &file)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
err = f.params.Lists.Load(file.Lists)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("%s: %w", path, err)
|
||||
}
|
||||
|
||||
f.params.DNSBL.Load(file.Verdicts)
|
||||
f.params.AbuseIPDB.Load(file.AbuseIPDB)
|
||||
entries = len(file.Lists)
|
||||
case alertsJSON:
|
||||
waiting, err := f.takeInAlerts(path, data)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
entries = waiting
|
||||
}
|
||||
|
||||
f.sums[name] = sha256.Sum256(data)
|
||||
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
// takeInAlerts parses data, what alerts.json, at path, holds, puts it
|
||||
// into the alerts and the anomaly counters, in place of what they held,
|
||||
// and returns how many alerts wait in it, as takeIn describes.
|
||||
func (f *Files) takeInAlerts(path string, data []byte) (int, error) {
|
||||
// waiting was a list, of the alerts waiting for the webhook, before
|
||||
// alerts went to Slack and ntfy too.
|
||||
var written struct {
|
||||
@@ -420,13 +479,12 @@ func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
|
||||
f.params.Alerts.Load(alerts.State{
|
||||
Cooldowns: file.Cooldowns, Hour: file.Hour, Waiting: file.Waiting,
|
||||
})
|
||||
f.params.Anomalies.Load(file.AnomalyCounters, f.params.Now())
|
||||
|
||||
entries := 0
|
||||
for _, waiting := range file.Waiting {
|
||||
entries += len(waiting)
|
||||
}
|
||||
}
|
||||
|
||||
f.sums[name] = sha256.Sum256(data)
|
||||
|
||||
return entries, nil
|
||||
}
|
||||
@@ -506,25 +564,31 @@ func (f *Files) setAside(name string, parseErr error) error {
|
||||
func (f *Files) encode(name string) ([]byte, error) {
|
||||
switch name {
|
||||
case bansJSON:
|
||||
file := bansFile{Version: version, Bans: BanEntries(f.params.Ledger.Snapshot())}
|
||||
|
||||
data, err := json.MarshalIndent(file, "", " ")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return append(data, '\n'), nil
|
||||
return encodeIndented(bansFile{
|
||||
Version: version, Bans: BanEntries(f.params.Ledger.Snapshot()),
|
||||
})
|
||||
case clientsJSON:
|
||||
return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
|
||||
case lookupsJSON:
|
||||
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
|
||||
case reputationJSON:
|
||||
return encodeIndented(reputationFile{
|
||||
Version: version, Lists: f.params.Lists.Snapshot(),
|
||||
Verdicts: f.params.DNSBL.Snapshot(), AbuseIPDB: f.params.AbuseIPDB.Snapshot(),
|
||||
})
|
||||
default: // alerts.json
|
||||
held := f.params.Alerts.Snapshot()
|
||||
file := alertsFile{
|
||||
|
||||
return encodeIndented(alertsFile{
|
||||
Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour,
|
||||
Waiting: held.Waiting,
|
||||
Waiting: held.Waiting, AnomalyCounters: f.params.Anomalies.Snapshot(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// encodeIndented encodes file, a state file's struct, indented for an
|
||||
// admin to read and edit.
|
||||
func encodeIndented(file any) ([]byte, error) {
|
||||
data, err := json.MarshalIndent(file, "", " ")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -532,7 +596,6 @@ func (f *Files) encode(name string) ([]byte, error) {
|
||||
|
||||
return append(data, '\n'), nil
|
||||
}
|
||||
}
|
||||
|
||||
// BanEntries returns held as bans.json lists them, an empty list for
|
||||
// none.
|
||||
@@ -615,8 +678,8 @@ func (f *bansFile) check(data []byte) error {
|
||||
}
|
||||
|
||||
// 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.
|
||||
// requests, or with requests or bytes 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 {
|
||||
@@ -628,6 +691,12 @@ func (f *clientsFile) check([]byte) error {
|
||||
return missing(i, "hour.start")
|
||||
case countsWithoutStart(client.Day):
|
||||
return missing(i, "day.start")
|
||||
case countsWithoutStart(client.MinuteBytes):
|
||||
return missing(i, "minute_bytes.start")
|
||||
case countsWithoutStart(client.HourBytes):
|
||||
return missing(i, "hour_bytes.start")
|
||||
case countsWithoutStart(client.DayBytes):
|
||||
return missing(i, "day_bytes.start")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -665,10 +734,92 @@ func (f *lookupsFile) check(data []byte) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// check refuses a list without its URL, which would name no list, or the
|
||||
// time it was last tried, which would have it fetched at once, and a copy
|
||||
// of it without the time it was fetched, or without its lines, which hold
|
||||
// the list. It refuses a verdict without its zone or its client, which
|
||||
// would be about no one, whether the zone lists the client, or the time
|
||||
// it was fetched, which would drop it, and so an AbuseIPDB score without
|
||||
// its client, the score, or the time it was fetched. A verdict's listed is
|
||||
// false for a client the zone does not list, and a score can be 0, which
|
||||
// the structs cannot tell from a missing one, so each is read again as
|
||||
// written.
|
||||
func (f *reputationFile) check(data []byte) error {
|
||||
for i, kept := range f.Lists {
|
||||
switch {
|
||||
case kept.URL == "":
|
||||
return missing(i, "url")
|
||||
case kept.Tried.IsZero():
|
||||
return missing(i, "tried")
|
||||
case kept.Fetched.IsZero() && kept.Lines != nil:
|
||||
return missing(i, "fetched")
|
||||
case kept.Lines == nil && !kept.Fetched.IsZero():
|
||||
return missing(i, "lines")
|
||||
}
|
||||
}
|
||||
|
||||
var written struct {
|
||||
Verdicts []struct {
|
||||
Listed *bool `json:"listed"`
|
||||
} `json:"verdicts"`
|
||||
}
|
||||
|
||||
err := json.Unmarshal(data, &written)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for i, verdict := range f.Verdicts {
|
||||
switch {
|
||||
case verdict.Zone == "":
|
||||
return fmt.Errorf("verdicts %w", missing(i, "zone"))
|
||||
case !verdict.Client.IsValid():
|
||||
return fmt.Errorf("verdicts %w", missing(i, "client"))
|
||||
case written.Verdicts[i].Listed == nil:
|
||||
return fmt.Errorf("verdicts %w", missing(i, "listed"))
|
||||
case verdict.Fetched.IsZero():
|
||||
return fmt.Errorf("verdicts %w", missing(i, "fetched"))
|
||||
}
|
||||
}
|
||||
|
||||
return checkScores(f.AbuseIPDB.Scores, data)
|
||||
}
|
||||
|
||||
// checkScores refuses an AbuseIPDB score, of scores, read from data, as
|
||||
// reputationFile's check describes.
|
||||
func checkScores(scores []reputation.Score, data []byte) error {
|
||||
var written struct {
|
||||
AbuseIPDB struct {
|
||||
Scores []struct {
|
||||
Score *int64 `json:"score"`
|
||||
} `json:"scores"`
|
||||
} `json:"abuseipdb"`
|
||||
}
|
||||
|
||||
err := json.Unmarshal(data, &written)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for i, kept := range scores {
|
||||
switch {
|
||||
case !kept.Client.IsValid():
|
||||
return fmt.Errorf("abuseipdb scores %w", missing(i, "client"))
|
||||
case written.AbuseIPDB.Scores[i].Score == nil:
|
||||
return fmt.Errorf("abuseipdb scores %w", missing(i, "score"))
|
||||
case kept.Fetched.IsZero():
|
||||
return fmt.Errorf("abuseipdb scores %w", missing(i, "fetched"))
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// check refuses a cooldown without its event or when its alert was sent,
|
||||
// which would hold back no repeat, alerts waiting for a destination with
|
||||
// another name than webhook, slack or ntfy, most likely misspelt, and an
|
||||
// alert waiting without its event or its time.
|
||||
// another name than webhook, slack or ntfy, most likely misspelt, an
|
||||
// alert waiting without its event or its time, and an anomaly counter as
|
||||
// checkAnomalyCounters does.
|
||||
func (f *alertsFile) check([]byte) error {
|
||||
for i, cooldown := range f.Cooldowns {
|
||||
switch {
|
||||
@@ -694,11 +845,60 @@ func (f *alertsFile) check([]byte) error {
|
||||
}
|
||||
}
|
||||
|
||||
return checkAnomalyCounters(f.AnomalyCounters)
|
||||
}
|
||||
|
||||
// checkAnomalyCounters refuses an anomaly counter whose scope is not
|
||||
// client, net, asn, total or watch, most likely misspelt, and one without
|
||||
// a field it needs, as missingFromCounter tells.
|
||||
func checkAnomalyCounters(counters []anomaly.Counter) error {
|
||||
for i, counter := range counters {
|
||||
if !slices.Contains(anomaly.Scopes(), counter.Scope) {
|
||||
return fmt.Errorf("anomaly_counters entry %d's scope %q %w", i+1,
|
||||
counter.Scope, errScope)
|
||||
}
|
||||
|
||||
field := missingFromCounter(counter)
|
||||
if field != "" {
|
||||
return fmt.Errorf("anomaly_counters %w", missing(i, field))
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// countsWithoutStart reports whether b holds requests but no start, which
|
||||
// places them in time.
|
||||
// missingFromCounter returns the first field counter, an anomaly counter,
|
||||
// needs and has not, or "" when it has them all: what tells it from the
|
||||
// others in its scope, without which it would never be counted again, the
|
||||
// netblock of a client, net or watch counter, the AS number of an asn one
|
||||
// and the name of a watch one; and the start of a window in which it has
|
||||
// requests or bytes, without which they would be dropped.
|
||||
func missingFromCounter(counter anomaly.Counter) string {
|
||||
scope := counter.Scope
|
||||
|
||||
switch {
|
||||
case scope != anomaly.ScopeASN && scope != anomaly.ScopeTotal &&
|
||||
!counter.Netblock.IsValid():
|
||||
return "netblock"
|
||||
case scope == anomaly.ScopeASN && counter.ASN == "":
|
||||
return "asn"
|
||||
case scope == anomaly.ScopeWatch && counter.Name == "":
|
||||
return "name"
|
||||
case countsWithoutStart(counter.Minute):
|
||||
return "minute.start"
|
||||
case countsWithoutStart(counter.Hour):
|
||||
return "hour.start"
|
||||
case countsWithoutStart(counter.MinuteBytes):
|
||||
return "minute_bytes.start"
|
||||
case countsWithoutStart(counter.HourBytes):
|
||||
return "hour_bytes.start"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
// countsWithoutStart reports whether b holds requests, or bytes, but no
|
||||
// start, which places them in time.
|
||||
func countsWithoutStart(b ratelimit.Buckets) bool {
|
||||
return b.Start.IsZero() && (b.Current != 0 || b.Previous != 0)
|
||||
}
|
||||
|
||||
+494
-27
@@ -22,10 +22,12 @@ import (
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
||||
"sneak.berlin/go/smallwebwaf/internal/anomaly"
|
||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
||||
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
||||
"sneak.berlin/go/smallwebwaf/internal/reputation"
|
||||
"sneak.berlin/go/smallwebwaf/internal/state"
|
||||
)
|
||||
|
||||
@@ -34,7 +36,13 @@ const (
|
||||
bansJSON = "bans.json"
|
||||
clientsJSON = "clients.json"
|
||||
lookupsJSON = "lookups.json"
|
||||
reputationJSON = "reputation.json"
|
||||
alertsJSON = "alerts.json"
|
||||
// blocklistURL and torURL are the blocklists the tests' lists name, and
|
||||
// dnsblZone the DNSBL zone of the tests' verdicts.
|
||||
blocklistURL = "https://lists.example/drop.txt"
|
||||
torURL = "https://lists.example/tor.txt"
|
||||
dnsblZone = "dnsbl.example"
|
||||
// The AS number and AS name the tests' clients are looked up in.
|
||||
asn = "AS64496"
|
||||
asName = "Example Net"
|
||||
@@ -45,6 +53,9 @@ const (
|
||||
// maxLogLines is how many lines of the process log wait for a test to
|
||||
// read them.
|
||||
maxLogLines = 64
|
||||
// whole is the percentage of each limit a client gets when nothing
|
||||
// lowers its limits.
|
||||
whole = 100
|
||||
)
|
||||
|
||||
// permanentBansJSON is bans.json holding permanentBan.
|
||||
@@ -64,6 +75,15 @@ const permanentBansJSON = `{
|
||||
"limit": 1000,
|
||||
"window": "minute",
|
||||
"count": 1000.5,
|
||||
"reputation": [
|
||||
{
|
||||
"source": "https://lists.example/drop.txt"
|
||||
},
|
||||
{
|
||||
"source": "abuseipdb",
|
||||
"score": 100
|
||||
}
|
||||
],
|
||||
"request": {
|
||||
"time": "2026-10-06T00:00:00Z",
|
||||
"method": "GET",
|
||||
@@ -153,6 +173,104 @@ const filledAlertsJSON = `{
|
||||
"suppressed_repeats": 0
|
||||
}
|
||||
]
|
||||
},
|
||||
"anomaly_counters": [
|
||||
{
|
||||
"scope": "asn",
|
||||
"asn": "AS64496",
|
||||
"hour_bytes": {
|
||||
"start": "2026-10-06T00:00:00Z",
|
||||
"current": 8,
|
||||
"previous": 0
|
||||
}
|
||||
},
|
||||
{
|
||||
"scope": "net",
|
||||
"netblock": "203.0.113.0/24",
|
||||
"minute": {
|
||||
"start": "2026-10-06T00:00:00Z",
|
||||
"current": 1,
|
||||
"previous": 0
|
||||
}
|
||||
},
|
||||
{
|
||||
"scope": "total",
|
||||
"minute": {
|
||||
"start": "2026-10-06T00:00:00Z",
|
||||
"current": 1,
|
||||
"previous": 0
|
||||
},
|
||||
"minute_bytes": {
|
||||
"start": "2026-10-06T00:00:00Z",
|
||||
"current": 8,
|
||||
"previous": 0
|
||||
}
|
||||
},
|
||||
{
|
||||
"scope": "watch",
|
||||
"netblock": "203.0.113.0/24",
|
||||
"name": "office",
|
||||
"hour": {
|
||||
"start": "2026-10-06T00:00:00Z",
|
||||
"current": 1,
|
||||
"previous": 0
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
`
|
||||
|
||||
// filledReputationJSON is reputation.json holding the blocklists' last
|
||||
// tries and the copy of one, with its comment line, two verdicts of a
|
||||
// DNSBL zone, and the AbuseIPDB checks spent today with two scores, as
|
||||
// fill puts them in.
|
||||
const filledReputationJSON = `{
|
||||
"version": 1,
|
||||
"lists": [
|
||||
{
|
||||
"url": "https://lists.example/drop.txt",
|
||||
"tried": "2026-10-06T00:00:00Z",
|
||||
"fetched": "2026-10-05T23:00:00Z",
|
||||
"lines": [
|
||||
"; Spamhaus DROP List 2026/10/05 - (c) 2026 The Spamhaus Project SLL",
|
||||
"203.0.113.0/24 ; SBL1",
|
||||
"2001:db8::/32 ; SBL2"
|
||||
]
|
||||
},
|
||||
{
|
||||
"url": "https://lists.example/tor.txt",
|
||||
"tried": "2026-10-06T00:00:00Z"
|
||||
}
|
||||
],
|
||||
"verdicts": [
|
||||
{
|
||||
"zone": "dnsbl.example",
|
||||
"client": "203.0.113.9",
|
||||
"listed": true,
|
||||
"fetched": "2026-10-05T23:00:00Z"
|
||||
},
|
||||
{
|
||||
"zone": "dnsbl.example",
|
||||
"client": "2001:db8::1",
|
||||
"listed": false,
|
||||
"fetched": "2026-10-05T22:00:00Z"
|
||||
}
|
||||
],
|
||||
"abuseipdb": {
|
||||
"day": "2026-10-06T00:00:00Z",
|
||||
"spent": 3,
|
||||
"scores": [
|
||||
{
|
||||
"client": "203.0.113.9/32",
|
||||
"score": 100,
|
||||
"fetched": "2026-10-05T23:00:00Z"
|
||||
},
|
||||
{
|
||||
"client": "2001:db8::/64",
|
||||
"score": 0,
|
||||
"fetched": "2026-10-05T22:00:00Z"
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
`
|
||||
@@ -183,18 +301,52 @@ func TestFilesWrittenAndReadBack(t *testing.T) {
|
||||
wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot())
|
||||
wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot())
|
||||
|
||||
if got, want := after.Lists.Snapshot(), before.Lists.Snapshot(); !reflect.DeepEqual(
|
||||
got, want) {
|
||||
t.Errorf("%s read back\n%+v\nwant\n%+v", reputationJSON, got, want)
|
||||
}
|
||||
|
||||
wantEqual(t, reputationJSON, after.DNSBL.Snapshot(), before.DNSBL.Snapshot())
|
||||
|
||||
checks, wantChecks := after.AbuseIPDB.Snapshot(), before.AbuseIPDB.Snapshot()
|
||||
if !reflect.DeepEqual(checks, wantChecks) {
|
||||
t.Errorf("%s read back\n%+v\nwant\n%+v", reputationJSON, checks, wantChecks)
|
||||
}
|
||||
|
||||
if got, want := after.Alerts.Snapshot(), before.Alerts.Snapshot(); !reflect.DeepEqual(
|
||||
got, want) {
|
||||
t.Errorf("%s read back\n%+v\nwant\n%+v", alertsJSON, got, want)
|
||||
}
|
||||
|
||||
wantEqual(t, alertsJSON, after.Anomalies.Snapshot(), before.Anomalies.Snapshot())
|
||||
|
||||
// Each one-per-line file lists its entries by client, and nothing
|
||||
// but the four files is left in the directory.
|
||||
// but the five 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, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
|
||||
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON)
|
||||
}
|
||||
|
||||
func TestReputationJSONKeepsEachCopyWholeOneLineToALine(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
params := newParams(dir)
|
||||
fill(params)
|
||||
|
||||
files := load(t, params)
|
||||
|
||||
err := files.WriteAll()
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
got := readFile(t, filepath.Join(dir, reputationJSON))
|
||||
if got != filledReputationJSON {
|
||||
t.Errorf("reputation.json\n%s\nwant\n%s", got, filledReputationJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAlertsJSONIsIndentedWithTheCooldownsTheHourAndTheAlertsWaiting(t *testing.T) {
|
||||
@@ -248,6 +400,49 @@ func TestSourceFailureCooldownKeptInAlertsJSONAcrossARestart(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnomalyCountersKeptInAlertsJSONAcrossARestart(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
request := anomaly.Request{
|
||||
Client: netip.MustParseAddr("203.0.113.9"),
|
||||
ClientGroup: netip.MustParsePrefix("203.0.113.9/32"),
|
||||
}
|
||||
|
||||
// The whole service may have two requests a minute.
|
||||
withThreshold := func() state.Params {
|
||||
params := newParams(dir)
|
||||
params.Anomalies = anomaly.New(anomaly.Params{
|
||||
Total: anomaly.Thresholds{RequestsPerMinute: 2}, Alerts: params.Alerts,
|
||||
})
|
||||
|
||||
return params
|
||||
}
|
||||
|
||||
before := withThreshold()
|
||||
files := load(t, before)
|
||||
|
||||
for range 2 {
|
||||
before.Anomalies.Count(midnight(), request)
|
||||
}
|
||||
|
||||
err := files.WriteAll()
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
// After the restart, the third request in the minute is over it.
|
||||
after := withThreshold()
|
||||
load(t, after)
|
||||
after.Anomalies.Count(midnight(), request)
|
||||
|
||||
waiting := after.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
|
||||
if len(waiting) != 1 || waiting[0].Event != alerts.EventAnomaly ||
|
||||
waiting[0].Detail["count"] != float64(3) {
|
||||
t.Errorf("alerts wait %+v, want one for 3 requests", waiting)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -275,9 +470,13 @@ func TestMissingFilesAreEmptyState(t *testing.T) {
|
||||
load(t, params)
|
||||
|
||||
held := params.Alerts.Snapshot()
|
||||
checks := params.AbuseIPDB.Snapshot()
|
||||
|
||||
if len(params.Ledger.Snapshot()) != 0 || len(params.Limiter.Snapshot()) != 0 ||
|
||||
len(params.GeoJS.Snapshot()) != 0 || len(held.Cooldowns) != 0 ||
|
||||
len(held.Waiting[alerts.DestinationWebhook]) != 0 || held.Hour.Sent != 0 {
|
||||
len(params.GeoJS.Snapshot()) != 0 || len(params.Lists.Snapshot()) != 0 ||
|
||||
len(params.DNSBL.Snapshot()) != 0 || len(checks.Scores) != 0 || checks.Spent != 0 ||
|
||||
len(held.Cooldowns) != 0 || len(held.Waiting[alerts.DestinationWebhook]) != 0 ||
|
||||
held.Hour.Sent != 0 {
|
||||
t.Error("state from no files")
|
||||
}
|
||||
}
|
||||
@@ -329,6 +528,21 @@ func TestFileThatDoesNotParseStopsTheStart(t *testing.T) {
|
||||
`{"version": 1, "waiting": {"webhook": [], "slak": []}}`,
|
||||
`: waiting "slak" is not webhook, slack or ntfy`,
|
||||
},
|
||||
{
|
||||
"an anomaly counter of an unknown scope", alertsJSON,
|
||||
`{"version": 1, "anomaly_counters": [{"scope": "total"}, ` +
|
||||
`{"scope": "nett", "netblock": "203.0.113.0/24"}]}`,
|
||||
`: anomaly_counters entry 2's scope "nett" is not client, net, asn, total ` +
|
||||
`or watch`,
|
||||
},
|
||||
{
|
||||
"a copy of a list with a line that does not read", reputationJSON,
|
||||
`{"version": 1, "lists": [{"url": "` + blocklistURL + `", ` +
|
||||
`"tried": "2026-10-06T00:00:00Z", "fetched": "2026-10-06T00:00:00Z", ` +
|
||||
`"lines": ["; DROP", "203.0.113.300"]}]}`,
|
||||
`: the copy of ` + blocklistURL + `: line 2 is not an address or a netblock, ` +
|
||||
`such as 192.0.2.0/24`,
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
@@ -392,6 +606,12 @@ func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
|
||||
`"hour": {"current": 3}}]}`,
|
||||
`: entry 1 has no "hour.start"`,
|
||||
},
|
||||
{
|
||||
"a client with bytes in a window without its start", clientsJSON,
|
||||
`{"version": 1, "clients": [{"client": "203.0.113.9/32", ` +
|
||||
`"minute_bytes": {"previous": 5120}}]}`,
|
||||
`: entry 1 has no "minute_bytes.start"`,
|
||||
},
|
||||
{
|
||||
"an answer without a client", lookupsJSON,
|
||||
`{"version": 1, "lookups": [{"country": "DE", ` + answer + `}]}`,
|
||||
@@ -418,6 +638,121 @@ func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestReputationJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
const (
|
||||
drop = `"url": "` + blocklistURL + `", `
|
||||
tried = `"tried": "2026-10-06T00:00:00Z", `
|
||||
fetched = `"fetched": "2026-10-06T00:00:00Z"`
|
||||
// verdictZone, verdictClient and listed start a verdict, which
|
||||
// fetched ends.
|
||||
verdictZone = `"zone": "` + dnsblZone + `", `
|
||||
verdictClient = `"client": "198.51.100.7", `
|
||||
listed = `"listed": false, `
|
||||
)
|
||||
|
||||
for _, tc := range []struct {
|
||||
name, content string
|
||||
// want is what the error says after the file's path.
|
||||
want string
|
||||
}{
|
||||
{
|
||||
"a list without its URL",
|
||||
`{"version": 1, "lists": [{` + tried + fetched + `, "lines": []}]}`,
|
||||
`: entry 1 has no "url"`,
|
||||
},
|
||||
{
|
||||
"a list without the time it was last tried",
|
||||
`{"version": 1, "lists": [{` + drop + fetched + `, "lines": []}]}`,
|
||||
`: entry 1 has no "tried"`,
|
||||
},
|
||||
{
|
||||
"a copy of a list without the time it was fetched",
|
||||
`{"version": 1, "lists": [{` + drop + tried + `"lines": []}]}`,
|
||||
`: entry 1 has no "fetched"`,
|
||||
},
|
||||
{
|
||||
// An empty list has no lines, which is not having none.
|
||||
"a copy of a list without its lines",
|
||||
`{"version": 1, "lists": [{` + drop + tried + fetched + `, "lines": []}, ` +
|
||||
`{"url": "` + torURL + `", ` + tried + fetched + `}]}`,
|
||||
`: entry 2 has no "lines"`,
|
||||
},
|
||||
{
|
||||
"a verdict without its zone",
|
||||
`{"version": 1, "verdicts": [{` + verdictClient + listed + fetched + `}]}`,
|
||||
`: verdicts entry 1 has no "zone"`,
|
||||
},
|
||||
{
|
||||
"a verdict without its client",
|
||||
`{"version": 1, "verdicts": [{` + verdictZone + listed + fetched + `}]}`,
|
||||
`: verdicts entry 1 has no "client"`,
|
||||
},
|
||||
{
|
||||
// A client the zone does not list has a listed of false, which is
|
||||
// not having none.
|
||||
"a verdict without whether the zone lists the client",
|
||||
`{"version": 1, "verdicts": [{` + verdictZone + verdictClient + listed +
|
||||
fetched + `}, {` + verdictZone + verdictClient + fetched + `}]}`,
|
||||
`: verdicts entry 2 has no "listed"`,
|
||||
},
|
||||
{
|
||||
"a verdict without the time it was fetched",
|
||||
`{"version": 1, "verdicts": [{` + verdictZone + verdictClient +
|
||||
`"listed": true}]}`,
|
||||
`: verdicts entry 1 has no "fetched"`,
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
wantRefused(t, reputationJSON, tc.content, tc.want)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReputationJSONScoreWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// scores opens the list of AbuseIPDB scores, and ends closes it; client,
|
||||
// score and fetched make a score.
|
||||
const (
|
||||
scores = `{"version": 1, "abuseipdb": {"scores": [`
|
||||
client = `"client": "198.51.100.7/32", `
|
||||
score = `"score": 0, `
|
||||
fetched = `"fetched": "2026-10-06T00:00:00Z"`
|
||||
ends = `}]}}`
|
||||
)
|
||||
|
||||
for _, tc := range []struct {
|
||||
name, content string
|
||||
// want is what the error says after the file's path.
|
||||
want string
|
||||
}{
|
||||
{
|
||||
"without its client", scores + `{` + score + fetched + ends,
|
||||
`: abuseipdb scores entry 1 has no "client"`,
|
||||
},
|
||||
{
|
||||
// A score of 0 is not having none.
|
||||
"without the score",
|
||||
scores + `{` + client + score + fetched + `}, {` + client + fetched + ends,
|
||||
`: abuseipdb scores entry 2 has no "score"`,
|
||||
},
|
||||
{
|
||||
"without the time it was fetched", scores + `{` + client + `"score": 100` + ends,
|
||||
`: abuseipdb scores entry 1 has no "fetched"`,
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
wantRefused(t, reputationJSON, tc.content, tc.want)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAlertsJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -446,6 +781,39 @@ func TestAlertsJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
|
||||
`{"version": 1, "waiting": {"slack": [{"event": "ban"}]}}`,
|
||||
`: waiting slack entry 1 has no "time"`,
|
||||
},
|
||||
{
|
||||
// The whole service's counter needs nothing to tell it apart.
|
||||
"an anomaly counter of a netblock without it",
|
||||
`{"version": 1, "anomaly_counters": [{"scope": "total"}, {"scope": "net"}]}`,
|
||||
`: anomaly_counters entry 2 has no "netblock"`,
|
||||
},
|
||||
{
|
||||
"an anomaly counter of a client without its netblock",
|
||||
`{"version": 1, "anomaly_counters": [{"scope": "client"}]}`,
|
||||
`: anomaly_counters entry 1 has no "netblock"`,
|
||||
},
|
||||
{
|
||||
"an anomaly counter of an AS number without it",
|
||||
`{"version": 1, "anomaly_counters": [{"scope": "asn"}]}`,
|
||||
`: anomaly_counters entry 1 has no "asn"`,
|
||||
},
|
||||
{
|
||||
"an anomaly counter of a named netblock without its name",
|
||||
`{"version": 1, "anomaly_counters": [` +
|
||||
`{"scope": "watch", "netblock": "203.0.113.0/24"}]}`,
|
||||
`: anomaly_counters entry 1 has no "name"`,
|
||||
},
|
||||
{
|
||||
"an anomaly counter of a named netblock without its netblock",
|
||||
`{"version": 1, "anomaly_counters": [{"scope": "watch", "name": "office"}]}`,
|
||||
`: anomaly_counters entry 1 has no "netblock"`,
|
||||
},
|
||||
{
|
||||
"an anomaly counter with bytes in a window without its start",
|
||||
`{"version": 1, "anomaly_counters": [` +
|
||||
`{"scope": "total", "hour_bytes": {"current": 5}}]}`,
|
||||
`: anomaly_counters entry 1 has no "hour_bytes.start"`,
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
@@ -488,7 +856,9 @@ func TestBanWithAnotherCauseStopsTheStart(t *testing.T) {
|
||||
func TestUnknownVersionStopsTheStart(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, file := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} {
|
||||
for _, file := range []string{
|
||||
bansJSON, clientsJSON, lookupsJSON, reputationJSON, alertsJSON,
|
||||
} {
|
||||
for _, content := range []string{`{"version": 2}`, `{}`} {
|
||||
t.Run(file+" "+content, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
@@ -558,7 +928,7 @@ func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) {
|
||||
load(t, read)
|
||||
|
||||
want := []bans.Ban{first, second}
|
||||
if got := read.Ledger.Snapshot(); !slices.Equal(got, want) {
|
||||
if got := read.Ledger.Snapshot(); !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("bans.json holds %+v, want %+v", got, want)
|
||||
}
|
||||
|
||||
@@ -590,8 +960,9 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
|
||||
|
||||
time.Sleep(time.Nanosecond)
|
||||
synctest.Wait()
|
||||
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
|
||||
removeFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
|
||||
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON)
|
||||
removeFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON,
|
||||
reputationJSON)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -707,7 +1078,7 @@ func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
|
||||
|
||||
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
|
||||
bans.Notes{})
|
||||
params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight())
|
||||
params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight(), whole)
|
||||
|
||||
err = files.WriteAll()
|
||||
if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") {
|
||||
@@ -810,7 +1181,7 @@ func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) {
|
||||
t.Errorf("bans.json is now %v (%v), want the socket", info, err)
|
||||
}
|
||||
|
||||
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
|
||||
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON)
|
||||
wantWriteFailed(t, params, bansJSON)
|
||||
}
|
||||
|
||||
@@ -883,6 +1254,32 @@ func TestEditOfEachFileTakenIn(t *testing.T) {
|
||||
wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(),
|
||||
[]lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}})
|
||||
|
||||
edit(t, dir, reputationJSON, `{"version": 1, "lists": [{"url": "`+blocklistURL+`", `+
|
||||
`"tried": "2026-10-06T00:00:00Z", "fetched": "2026-10-06T00:00:00Z", `+
|
||||
`"lines": ["198.51.100.7"]}], "verdicts": [{"zone": "`+dnsblZone+`", `+
|
||||
`"client": "198.51.100.7", "listed": true, "fetched": "2026-10-06T00:00:00Z"}], `+
|
||||
`"abuseipdb": {"day": "2026-10-06T00:00:00Z", "spent": 9, "scores": [`+
|
||||
`{"client": "198.51.100.7/32", "score": 80, "fetched": "2026-10-06T00:00:00Z"}]}}`)
|
||||
wantTakenIn(t, lines, dir, reputationJSON)
|
||||
|
||||
listedBy := params.Lists.ListedBy(client.Addr())
|
||||
if !slices.Equal(listedBy, []string{blocklistURL}) {
|
||||
t.Errorf("%s taken in lists %s on %v, want on the blocklist", reputationJSON,
|
||||
client.Addr(), listedBy)
|
||||
}
|
||||
|
||||
wantEqual(t, reputationJSON, params.DNSBL.Snapshot(), []reputation.Verdict{{
|
||||
Zone: dnsblZone, Client: client.Addr(), Listed: true, Fetched: midnight(),
|
||||
}})
|
||||
|
||||
checks := reputation.Checks{
|
||||
Day: midnight(), Spent: 9,
|
||||
Scores: []reputation.Score{{Client: client, Score: 80, Fetched: midnight()}},
|
||||
}
|
||||
if got := params.AbuseIPDB.Snapshot(); !reflect.DeepEqual(got, checks) {
|
||||
t.Errorf("%s taken in as\n%+v\nwant\n%+v", reputationJSON, got, checks)
|
||||
}
|
||||
|
||||
// A netblock with bits past its length is read as the netblock it is
|
||||
// in.
|
||||
edit(t, dir, alertsJSON, `{"version": 1, "cooldowns": [{"event": "ban", `+
|
||||
@@ -1110,7 +1507,7 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
|
||||
edit(t, dir, bansJSON, broken)
|
||||
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
|
||||
wantTakenIn(t, lines, dir, clientsJSON)
|
||||
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
|
||||
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON)
|
||||
|
||||
// The next write sets it aside, logged with where the error is, and
|
||||
// writes bans.json again from what smallwebwaf still holds.
|
||||
@@ -1134,7 +1531,8 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
|
||||
t.Errorf("alerts waiting %+v, want a file_error alert for %s", waiting, path+".bad")
|
||||
}
|
||||
|
||||
wantFiles(t, dir, alertsJSON, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON)
|
||||
wantFiles(t, dir, alertsJSON, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON,
|
||||
reputationJSON)
|
||||
|
||||
if got := readFile(t, path+".bad"); got != broken {
|
||||
t.Errorf("bans.json.bad holds\n%s\nwant the edit", got)
|
||||
@@ -1263,11 +1661,21 @@ func midnight() time.Time {
|
||||
}
|
||||
|
||||
// newParams returns Params for the state files in dir, with parts that
|
||||
// hold nothing yet. GeoJS is never asked, and the alerts, at most two an
|
||||
// hour, are never sent.
|
||||
// hold nothing yet. GeoJS is never asked, the lists, two blocklists, are
|
||||
// never fetched, and the alerts, at most two an hour, are
|
||||
// never sent. The anomaly counters count the scopes fill counts, with
|
||||
// thresholds fill does not reach.
|
||||
func newParams(dir string) state.Params {
|
||||
discard := slog.New(slog.DiscardHandler)
|
||||
m := metrics.New(1, "app")
|
||||
queue := alerts.New(alerts.Params{
|
||||
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
||||
Events: alerts.Events(),
|
||||
Cooldown: 15 * time.Minute,
|
||||
MaxPerHour: 2,
|
||||
Instance: "fsn1app1/gitea",
|
||||
Now: midnight,
|
||||
})
|
||||
|
||||
return state.Params{
|
||||
Dir: dir,
|
||||
@@ -1280,17 +1688,32 @@ func newParams(dir string) state.Params {
|
||||
AttackBanDuration: 7 * 24 * time.Hour,
|
||||
MaxBans: 5000,
|
||||
}),
|
||||
Limiter: ratelimit.New(ratelimit.Limits{}),
|
||||
Limiter: ratelimit.New(ratelimit.Limits{}, 20000),
|
||||
GeoJS: lookup.New(lookup.Params{
|
||||
Now: midnight, ProcessLog: discard, Metrics: m,
|
||||
}),
|
||||
Alerts: alerts.New(alerts.Params{
|
||||
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
||||
Events: alerts.Events(),
|
||||
Cooldown: 15 * time.Minute,
|
||||
MaxPerHour: 2,
|
||||
Instance: "fsn1app1/gitea",
|
||||
Now: midnight,
|
||||
Lists: reputation.New(reputation.Params{
|
||||
BlocklistURLs: []string{blocklistURL, torURL}, Refresh: 24 * time.Hour,
|
||||
Now: midnight, ProcessLog: discard, Alerts: queue,
|
||||
}),
|
||||
DNSBL: reputation.NewDNSBL(reputation.DNSBLParams{
|
||||
Zones: []string{dnsblZone}, CacheTTL: 24 * time.Hour, Timeout: time.Second,
|
||||
Now: midnight, ProcessLog: discard, Alerts: queue,
|
||||
}),
|
||||
AbuseIPDB: reputation.NewAbuseIPDB(reputation.AbuseIPDBParams{
|
||||
MinScore: 75, DailyBudget: 900, CacheTTL: 24 * time.Hour, Timeout: time.Second,
|
||||
Now: midnight, ProcessLog: discard, Alerts: queue,
|
||||
}),
|
||||
Alerts: queue,
|
||||
Anomalies: anomaly.New(anomaly.Params{
|
||||
Net: anomaly.Thresholds{RequestsPerMinute: 1000},
|
||||
ASN: anomaly.Thresholds{BytesPerHour: 1 << 30},
|
||||
Total: anomaly.Thresholds{RequestsPerMinute: 1000, BytesPerMinute: 1 << 30},
|
||||
Watch: anomaly.Thresholds{RequestsPerHour: 1000},
|
||||
NetV4Prefix: 24,
|
||||
NetV6Prefix: 48,
|
||||
NamedNetblocks: []anomaly.NamedNetblock{{Name: "office", Netblock: office()}},
|
||||
Alerts: queue,
|
||||
}),
|
||||
Now: midnight,
|
||||
ProcessLog: discard,
|
||||
@@ -1298,10 +1721,18 @@ func newParams(dir string) state.Params {
|
||||
}
|
||||
}
|
||||
|
||||
// office is the named netblock of the anomaly counters of newParams.
|
||||
func office() netip.Prefix {
|
||||
return netip.MustParsePrefix("203.0.113.0/24")
|
||||
}
|
||||
|
||||
// fill puts a permanent ban an admin made, a ban for a broken limit and
|
||||
// one for a clear sign of attack, clients with counts and histories,
|
||||
// GeoJS answers, and alerts, as filledAlertsJSON holds them, into the
|
||||
// parts of params.
|
||||
// GeoJS answers, the blocklists' last tries and the copy of one, two
|
||||
// verdicts of a DNSBL zone, and the AbuseIPDB checks spent today with two
|
||||
// scores, as filledReputationJSON holds them, and alerts
|
||||
// and anomaly counters, as filledAlertsJSON holds them, into the parts of
|
||||
// params.
|
||||
func fill(params state.Params) {
|
||||
now := midnight()
|
||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
||||
@@ -1314,9 +1745,10 @@ func fill(params state.Params) {
|
||||
bans.Notes{RuleID: "env-file", Target: "path"})
|
||||
|
||||
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.Count(netip.MustParsePrefix(c), now, whole)
|
||||
}
|
||||
|
||||
params.Limiter.CountBytes(client, now, 8, whole)
|
||||
params.Limiter.AddToHistory(client, now, ratelimit.Request{
|
||||
Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5,
|
||||
})
|
||||
@@ -1333,6 +1765,32 @@ func fill(params state.Params) {
|
||||
},
|
||||
})
|
||||
|
||||
// The copy of drop.txt was fetched an hour ago, and the fetches of it
|
||||
// and of tor.txt tried since failed.
|
||||
err := params.Lists.Load([]reputation.List{{
|
||||
URL: blocklistURL, Tried: now, Fetched: now.Add(-time.Hour),
|
||||
Lines: []string{
|
||||
"; Spamhaus DROP List 2026/10/05 - (c) 2026 The Spamhaus Project SLL",
|
||||
"203.0.113.0/24 ; SBL1",
|
||||
"2001:db8::/32 ; SBL2",
|
||||
},
|
||||
}, {URL: torURL, Tried: now}})
|
||||
if err != nil {
|
||||
panic(err) // the copy reads
|
||||
}
|
||||
|
||||
params.DNSBL.Load([]reputation.Verdict{
|
||||
{
|
||||
Zone: dnsblZone, Client: netip.MustParseAddr("2001:db8::1"),
|
||||
Fetched: now.Add(-2 * time.Hour),
|
||||
},
|
||||
{Zone: dnsblZone, Client: client.Addr(), Listed: true, Fetched: now.Add(-time.Hour)},
|
||||
})
|
||||
params.AbuseIPDB.Load(reputation.Checks{Day: now, Spent: 3, Scores: []reputation.Score{
|
||||
{Client: netip.MustParsePrefix("2001:db8::/64"), Fetched: now.Add(-2 * time.Hour)},
|
||||
{Client: client, Score: 100, Fetched: now.Add(-time.Hour)},
|
||||
}})
|
||||
|
||||
// An alert waiting, a repeat of it the cooldown holds back, another
|
||||
// alert waiting, and one past the two an hour, for the hour's summary.
|
||||
ban := alerts.Alert{
|
||||
@@ -1351,10 +1809,16 @@ func fill(params state.Params) {
|
||||
params.Alerts.Raise(alerts.Alert{
|
||||
Event: alerts.EventSourceFailure, Reason: "asking GeoJS failed",
|
||||
})
|
||||
|
||||
params.Anomalies.Count(now, anomaly.Request{
|
||||
Client: client.Addr(), ClientGroup: client, ASN: asn, Bytes: 8,
|
||||
})
|
||||
}
|
||||
|
||||
// permanentBan is the ban permanentBansJSON holds.
|
||||
func permanentBan() bans.Ban {
|
||||
score := int64(100)
|
||||
|
||||
return bans.Ban{
|
||||
Netblock: netip.MustParsePrefix("2001:db8::/64"),
|
||||
Start: midnight(),
|
||||
@@ -1367,6 +1831,9 @@ func permanentBan() bans.Ban {
|
||||
Limit: 1000,
|
||||
Window: "minute",
|
||||
Count: 1000.5,
|
||||
Reputation: []bans.ReputationHit{
|
||||
{Source: blocklistURL}, {Source: reputation.AbuseIPDBSource, Score: &score},
|
||||
},
|
||||
Request: bans.Request{
|
||||
Time: midnight(),
|
||||
Method: "GET",
|
||||
@@ -1542,10 +2009,10 @@ func edit(t *testing.T, dir, name, content string) {
|
||||
|
||||
// wantEqual checks that the entries read back from file are those
|
||||
// written.
|
||||
func wantEqual[E comparable](t *testing.T, file string, got, want []E) {
|
||||
func wantEqual[E any](t *testing.T, file string, got, want []E) {
|
||||
t.Helper()
|
||||
|
||||
if !slices.Equal(got, want) {
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("%s read back\n%+v\nwant\n%+v", file, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
+7
-5
@@ -7,9 +7,10 @@
|
||||
# request bans for good, that `sv stop` stops smallwebwaf in order, that
|
||||
# `docker stop` stops the container without having to kill it, and that
|
||||
# a new container on the same volume still refuses the banned client. The
|
||||
# containers, the volume and both images are removed however the script
|
||||
# ends. Building the app needs network access, for nixpkgs' binary cache.
|
||||
# script/check does not run this.
|
||||
# containers run with SWWAF_LOOKUP_SOURCE=off, so that no address is sent
|
||||
# to GeoJS. The containers, the volume and both images are removed however
|
||||
# the script ends. Building the app needs network access, for nixpkgs'
|
||||
# binary cache. script/check does not run this.
|
||||
set -eu
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
||||
@@ -63,12 +64,13 @@ logged() {
|
||||
}
|
||||
|
||||
# start_container: run the app's container, with the state files on the
|
||||
# volume and a rate limit of one request a minute, and wait until it is
|
||||
# healthy.
|
||||
# volume, a rate limit of one request a minute and no client looked up,
|
||||
# and wait until it is healthy.
|
||||
start_container() {
|
||||
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \
|
||||
--volume "$VOLUME:/var/lib/smallwebwaf" \
|
||||
--env SWWAF_RATE_LIMIT_PER_MINUTE=1 \
|
||||
--env SWWAF_LOOKUP_SOURCE=off \
|
||||
"$APP_IMAGE" >/dev/null
|
||||
wait_for "the health check did not pass" healthy
|
||||
address="$(docker port "$CONTAINER" 8080/tcp)"
|
||||
|
||||
Executable
+19
@@ -0,0 +1,19 @@
|
||||
#!/bin/sh
|
||||
# script/tidy: write go.mod and go.sum as `go mod tidy` writes them, which
|
||||
# the test phase of the Dockerfile checks. This builds the Dockerfile's
|
||||
# tidy-files stage, which holds the two files alone, and --output writes
|
||||
# them into the working tree. The build makes no image, so it has no tag.
|
||||
# --no-cache because `go mod tidy` asks the module proxy, whose answers a
|
||||
# cached layer would repeat.
|
||||
set -eu
|
||||
|
||||
ROOT="$(cd "$(dirname "$0")/.." && pwd -P)"
|
||||
|
||||
main() {
|
||||
cd "$ROOT"
|
||||
docker build --no-cache \
|
||||
--target tidy-files \
|
||||
--output "type=local,dest=$ROOT" .
|
||||
}
|
||||
|
||||
main "$@"
|
||||
Reference in New Issue
Block a user