Compare commits

..
1 Commits
Author SHA1 Message Date
clawbot a8e86c18db Alerts to a JSON webhook, with a cooldown and an hourly summary (closes #26)
check / check (push) Canceled after 0s
SWWAF_ALERT_WEBHOOK_URL gets one JSON POST per alert, in SPEC.md's
schema, with SWWAF_ALERT_WEBHOOK_HEADERS: ban and permanent_ban, with
the ban's notes; source_failure for GeoJS; file_error for a rule or
state file edit that does not parse and a failed state write.
SWWAF_ALERT_EVENTS chooses, SWWAF_ALERT_COOLDOWN holds back repeats,
and past SWWAF_ALERT_MAX_PER_HOUR the hour ends in one summary. A
bounded queue, retried with backoff, holds up no request; alerts.json
keeps it, the cooldowns and the hour. The ledger now reports whether it
made a ban, or made one permanent.

Judgement call: the summary's event is summary, which SPEC.md omits.
Judgement call: admin bans and observe mode raise no alert.

Model: opus-5-5
2026-10-07 00:12:21 +00:00
76 changed files with 1686 additions and 16470 deletions
-4
View File
@@ -61,10 +61,6 @@ linters:
desc: >- desc: >-
Test-support code belongs in test files and in packages whose Test-support code belongs in test files and in packages whose
directory name ends in test, not in the shipped binary. 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 # Only decisions already recorded in the Go package defaults are
# listed here. Every entry matches the module path exactly. # listed here. Every entry matches the module path exactly.
gomodguard_v2: gomodguard_v2:
+1 -25
View File
@@ -29,12 +29,6 @@ RUN go mod download
COPY . . 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 # 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. # after this step, and writing it into the image takes seconds.
RUN --mount=type=tmpfs,target=/root/.cache/go-build \ RUN --mount=type=tmpfs,target=/root/.cache/go-build \
@@ -42,25 +36,7 @@ RUN --mount=type=tmpfs,target=/root/.cache/go-build \
{ echo "--- Rerunning with -v for details ---"; \ { echo "--- Rerunning with -v for details ---"; \
go test -timeout 90s -race -v ./...; exit 1; } go test -timeout 90s -race -v ./...; exit 1; }
# Tidy stage: `go mod tidy` in the test phase's Go, so that the files it # Build stage. Nothing is wanted from the two phases above; the copies
# 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 # are what make BuildKit build them first, so the image, which needs this
# stage, cannot be produced unless lint and test passed. # stage, cannot be produced unless lint and test passed.
# #
+3 -7
View File
@@ -1,10 +1,9 @@
.PHONY: bootstrap setup test lint fmt fmt-check tidy check docker hooks build \ .PHONY: bootstrap setup test lint fmt fmt-check check docker hooks build run \
run example-app example-app
# Makefile targets are thin shims; the implementations live in script/ # Makefile targets are thin shims; the implementations live in script/
# per the scripts-to-rule-them-all pattern (see the Entrypoints section # per the scripts-to-rule-them-all pattern (see the Entrypoints section
# of README.md). tidy writes go.mod and go.sum as `go mod tidy` does, # of README.md). build and run are for working on the code by hand;
# 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. # example-app checks the image with an app built on it.
bootstrap: bootstrap:
@@ -25,9 +24,6 @@ fmt:
fmt-check: fmt-check:
@script/fmt-check @script/fmt-check
tidy:
@script/tidy
check: check:
@script/check @script/check
+323 -1086
View File
File diff suppressed because it is too large Load Diff
+1 -4
View File
@@ -5,8 +5,6 @@ go 1.26.0
require ( require (
github.com/fsnotify/fsnotify v1.10.1 github.com/fsnotify/fsnotify v1.10.1
github.com/hashicorp/golang-lru/v2 v2.0.7 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 github.com/prometheus/client_golang v1.24.1
) )
@@ -18,7 +16,6 @@ require (
github.com/prometheus/client_model v0.6.2 // indirect github.com/prometheus/client_model v0.6.2 // indirect
github.com/prometheus/common v0.70.1 // indirect github.com/prometheus/common v0.70.1 // indirect
github.com/prometheus/procfs v0.21.1 // indirect github.com/prometheus/procfs v0.21.1 // indirect
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba // indirect golang.org/x/sys v0.47.0 // indirect
golang.org/x/sys v0.48.0 // indirect
google.golang.org/protobuf v1.36.11 // indirect google.golang.org/protobuf v1.36.11 // indirect
) )
+10 -12
View File
@@ -2,6 +2,8 @@ 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/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 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= 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 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo= github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
@@ -12,12 +14,10 @@ github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJn
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= 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 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw= 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 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
github.com/oschwald/maxminddb-golang/v2 v2.7.0 h1:ZcAr3GYc2LYC8aec2mCMX9+QOF0EolH3jDFKRV/Z1+U= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/oschwald/maxminddb-golang/v2 v2.7.0/go.mod h1:DuKJLbbug6TXC0yJXgs1MWifvXHmudRWzMobMIUu04g= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU= 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_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE=
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk= github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
@@ -26,17 +26,15 @@ github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi
github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc= 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 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI=
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY= github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= 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 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ= go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
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 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= 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=
+202 -605
View File
File diff suppressed because it is too large Load Diff
+53 -784
View File
@@ -26,18 +26,13 @@ import (
// clock of the test's own, which starts at 2000-01-01T00:00:00Z, the start // clock of the test's own, which starts at 2000-01-01T00:00:00Z, the start
// of an hour: a wait lasts exactly as long as it should, however slowly // of an hour: a wait lasts exactly as long as it should, however slowly
// the test process runs, and synctest.Wait returns once the queue has // the test process runs, and synctest.Wait returns once the queue has
// done all it can before time passes. The stand-ins for the webhook, // done all it can before time passes. The stand-in for the webhook
// Slack and ntfy answer without the network, since a request waiting on // answers without the network, since a request waiting on the network
// the network would keep that clock from moving on. // would keep that clock from moving on.
const ( const (
// webhookURL is where the alerts are posted, slackURL the Slack // webhookURL is where the alerts are posted.
// incoming webhook, and ntfyURL the ntfy topic. ntfyToken is the ntfy
// token of the tests that set one.
webhookURL = "https://alerts.example/smallwebwaf?team=ops" webhookURL = "https://alerts.example/smallwebwaf?team=ops"
slackURL = "https://hooks.slack.example/services/T0123/B4567/abcdef"
ntfyURL = "https://ntfy.example/smallwebwaf-alerts"
ntfyToken = "tk_0123456789abcdefghijklmnopq"
// instance is the instance name every alert gives. // instance is the instance name every alert gives.
instance = "fsn1app1/gitea" instance = "fsn1app1/gitea"
// started is when each test starts, as an alert gives it, and // started is when each test starts, as an alert gives it, and
@@ -123,24 +118,7 @@ func TestOnlyTheChosenEventsAreSent(t *testing.T) {
}) })
} }
func TestWouldSendOnlyForTheChosenEvents(t *testing.T) { func TestNothingIsQueuedWithoutAWebhook(t *testing.T) {
t.Parallel()
params := newParams()
params.Events = []string{alerts.EventSourceFailure, alerts.EventFileError}
q := alerts.New(params)
if q.WouldSend(alerts.EventBan, netblock(1)) ||
q.WouldSend(alerts.EventPermanentBan, netblock(1)) {
t.Error("a ban alert would be sent, though SWWAF_ALERT_EVENTS leaves it out")
}
if !q.WouldSend(alerts.EventFileError, netip.Prefix{}) {
t.Error("a file_error alert would not be sent")
}
}
func TestRaiseDoesNothingWithoutADestination(t *testing.T) {
t.Parallel() t.Parallel()
params := newParams() params := newParams()
@@ -149,47 +127,8 @@ func TestRaiseDoesNothingWithoutADestination(t *testing.T) {
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)}) q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
// No alert waits, no cooldown has started, and the hour counts none. if waiting := q.Snapshot().Waiting; len(waiting) != 0 {
want := alerts.State{ t.Errorf("%d alerts wait, want none", len(waiting))
Cooldowns: []alerts.Cooldown{},
Hour: alerts.Hour{HeldBack: map[string]int{}},
Waiting: map[string][]alerts.Alert{},
}
if got := q.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("state %+v, want %+v", got, want)
}
}
func TestWouldSendNothingWithoutADestination(t *testing.T) {
t.Parallel()
params := newParams()
params.WebhookURL = nil
q := alerts.New(params)
if q.WouldSend(alerts.EventBan, netblock(1)) {
t.Error("an alert would be sent with no webhook set")
}
}
func TestWouldSendWithOnlySlackOrOnlyNtfySet(t *testing.T) {
t.Parallel()
onlySlack := newParams()
onlySlack.WebhookURL = nil
onlySlack.SlackURL = parseURL(slackURL)
onlyNtfy := newParams()
onlyNtfy.WebhookURL = nil
onlyNtfy.NtfyURL = parseURL(ntfyURL)
for setting, params := range map[string]alerts.Params{
"SWWAF_ALERT_SLACK_WEBHOOK_URL": onlySlack,
"SWWAF_ALERT_NTFY_URL": onlyNtfy,
} {
if !alerts.New(params).WouldSend(alerts.EventBan, netblock(1)) {
t.Errorf("a ban alert would not be sent with only %s set", setting)
}
} }
} }
@@ -245,65 +184,6 @@ func TestRepeatWithinTheCooldownIsHeldBackAndCountedInTheNext(t *testing.T) {
}) })
} }
func TestFileErrorAndSourceFailureRepeatOnlyForTheSameFileOrSource(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
webhook, q := start(t, params)
fileError := func(file string) alerts.Alert {
return alerts.Alert{
Event: alerts.EventFileError,
Detail: map[string]any{"file": file, "error": "line 2: an error"},
}
}
sourceFailure := func(source string) alerts.Alert {
return alerts.Alert{
Event: alerts.EventSourceFailure, Detail: map[string]any{"source": source},
}
}
// Another file, or another source, is no repeat.
q.Raise(fileError("/rules.d/50-a.rules"))
q.Raise(fileError("/rules.d/50-b.rules"))
q.Raise(fileError("/rules.d/50-a.rules"))
q.Raise(sourceFailure("geojs"))
q.Raise(sourceFailure("abuseipdb"))
q.Raise(sourceFailure("geojs"))
synctest.Wait()
// Each alert is named by its file, or its source.
got := make([]string, 0, len(webhook.received()))
for _, request := range webhook.received() {
detail, _ := request.alert["detail"].(map[string]any)
file, _ := detail["file"].(string)
source, _ := detail["source"].(string)
got = append(got, file+source)
}
want := []string{
"/rules.d/50-a.rules", "/rules.d/50-b.rules", "geojs", "abuseipdb",
}
if !slices.Equal(got, want) {
t.Errorf("the webhook was sent alerts for %v, want %v", got, want)
}
wantCounts(t, q, 4, 0, 2, 0)
// alerts.json keeps each file's cooldown: a new queue holds back
// the next for the first file, and sends the one for a third.
after := alerts.New(params)
after.Load(roundTrip(t, q.Snapshot()))
after.Raise(fileError("/rules.d/50-a.rules"))
after.Raise(fileError("/rules.d/50-c.rules"))
waiting := after.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || waiting[0].Detail["file"] != "/rules.d/50-c.rules" {
t.Errorf("after loading, alerts wait %+v, want the one for 50-c.rules", waiting)
}
})
}
func TestNoCooldownSendsEveryRepeat(t *testing.T) { func TestNoCooldownSendsEveryRepeat(t *testing.T) {
t.Parallel() t.Parallel()
@@ -381,117 +261,6 @@ func TestAlertsPastTheHourlyLimitAreRolledIntoOneSummary(t *testing.T) {
}) })
} }
func TestRepeatsBeforeAnAlertPastTheHourlyLimitAreGivenByTheSummary(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
params.MaxPerHour = 1
webhook, q := start(t, params)
raise := func() {
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
}
// The hour's one alert, and two repeats the cooldown holds back.
raise()
raise()
raise()
// Once the cooldown has run out, the next is past the hourly limit.
time.Sleep(cooldown)
raise()
// 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()
synctest.Wait()
wantEvents(t, webhook, alerts.EventBan, alerts.EventSummary, alerts.EventBan)
got := webhook.received()
if len(got) == 3 {
detail, _ := got[1].alert["detail"].(map[string]any)
summaryRepeats := got[1].alert["suppressed_repeats"]
lastRepeats := got[2].alert["suppressed_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)
}
}
wantCounts(t, q, 3, 0, 3, 0)
})
}
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) { func TestFailedRequestIsSentAgainWithBackoff(t *testing.T) {
t.Parallel() t.Parallel()
@@ -543,89 +312,12 @@ func TestFailedRequestIsSentAgainWithBackoff(t *testing.T) {
wantCounts(t, q, 1, int64(len(want)), 0, 0) wantCounts(t, q, 1, int64(len(want)), 0, 0)
waiting := q.Snapshot().Waiting[alerts.DestinationWebhook] if waiting := q.Snapshot().Waiting; len(waiting) != 0 {
if len(waiting) != 0 {
t.Errorf("%d alerts still wait, want none", len(waiting)) t.Errorf("%d alerts still wait, want none", len(waiting))
} }
}) })
} }
func TestRefusedAlertIsGivenUpAndTheNextSent(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
log := &lockedBuffer{}
params.ProcessLog = slog.New(slog.NewJSONHandler(log, nil))
webhook, q := start(t, params)
// 429 and 408 are failures, and the alert is sent again; 400 refuses
// it, and it is given up.
webhook.set(http.StatusTooManyRequests)
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
synctest.Wait()
webhook.set(http.StatusRequestTimeout)
time.Sleep(time.Second)
synctest.Wait()
webhook.set(refusing)
time.Sleep(2 * time.Second)
synctest.Wait()
// The next alert is sent at once.
webhook.set(answering)
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(2)})
time.Sleep(time.Minute)
synctest.Wait()
// Each request, by when it was sent, and the netblock of its alert.
got := make([]string, 0, len(webhook.received()))
for _, request := range webhook.received() {
block, _ := request.alert["netblock"].(string)
got = append(got, request.at.Sub(midnight()).String()+" "+block)
}
want := []string{
"0s " + netblock(1).String(), "1s " + netblock(1).String(),
"3s " + netblock(1).String(), "3s " + netblock(2).String(),
}
if !slices.Equal(got, want) {
t.Errorf("requests %v, want %v", got, want)
}
wantCounts(t, q, 1, 3, 0, 1)
if !strings.Contains(log.String(),
`"msg":"gave up an alert SWWAF_ALERT_WEBHOOK_URL refused"`) {
t.Errorf("process log %q names no alert given up", log.String())
}
})
}
func TestFailedRequestIsLoggedWithoutTheURL(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := newParams()
log := &lockedBuffer{}
params.ProcessLog = slog.New(slog.NewJSONHandler(log, nil))
webhook, q := start(t, params)
webhook.set(hanging)
// The request is abandoned after 10 seconds, with an error from the
// HTTP client, which names the URL.
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
time.Sleep(11 * time.Second)
synctest.Wait()
logged := log.String()
if !strings.Contains(logged,
`"msg":"sending an alert to SWWAF_ALERT_WEBHOOK_URL failed"`) ||
strings.Contains(logged, "alerts.example") || strings.Contains(logged, "team=ops") {
t.Errorf("process log %q names no failure, or names the URL", logged)
}
})
}
func TestFullQueueDropsTheOldestAndRaiseNeverWaits(t *testing.T) { func TestFullQueueDropsTheOldestAndRaiseNeverWaits(t *testing.T) {
t.Parallel() t.Parallel()
@@ -652,7 +344,7 @@ func TestFullQueueDropsTheOldestAndRaiseNeverWaits(t *testing.T) {
wantCounts(t, q, 0, 0, 0, 1) wantCounts(t, q, 0, 0, 0, 1)
waiting := q.Snapshot().Waiting[alerts.DestinationWebhook] waiting := q.Snapshot().Waiting
if len(waiting) != alerts.QueueSize || waiting[0].Netblock != netblock(1) { if len(waiting) != alerts.QueueSize || waiting[0].Netblock != netblock(1) {
t.Fatalf("%d alerts wait, the first for %s, want %d, the first for %s", t.Fatalf("%d alerts wait, the first for %s, want %d, the first for %s",
len(waiting), waiting[0].Netblock, alerts.QueueSize, netblock(1)) len(waiting), waiting[0].Netblock, alerts.QueueSize, netblock(1))
@@ -702,8 +394,7 @@ func TestStateLoadedIntoANewQueueCarriesOn(t *testing.T) {
after.Load(roundTrip(t, before.Snapshot())) after.Load(roundTrip(t, before.Snapshot()))
// The new queue sends the alert waiting, holds back the repeat as // The new queue sends the alert waiting, holds back the repeat as
// the cooldown still runs, and sends the summary of the hour, which // the cooldown still runs, and sends the summary of the hour.
// gives both repeats, as the cooldown has run out.
after.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)}) after.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
synctest.Wait() synctest.Wait()
wantEvents(t, webhook, alerts.EventBan) wantEvents(t, webhook, alerts.EventBan)
@@ -712,283 +403,30 @@ func TestStateLoadedIntoANewQueueCarriesOn(t *testing.T) {
synctest.Wait() synctest.Wait()
wantEvents(t, webhook, alerts.EventBan, alerts.EventSummary) wantEvents(t, webhook, alerts.EventBan, alerts.EventSummary)
summary := webhook.received()[1].alert detail, _ := webhook.received()[1].alert["detail"].(map[string]any)
detail, _ := summary["detail"].(map[string]any) if detail["count"] != float64(1) {
t.Errorf("the summary counts %v alerts, want 1", detail["count"])
}
if detail["count"] != float64(1) || summary["suppressed_repeats"] != float64(2) { // The cooldown has run out, and the next one gives both repeats.
t.Errorf("the summary counts %v alerts and %v repeats, want 1 and 2", after.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
detail["count"], summary["suppressed_repeats"]) 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)
} }
}) })
} }
func TestSlackAndNtfyAreSentAMessageForEachEvent(t *testing.T) { // How the stand-in for the webhook answers.
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := withSlackAndNtfy(newParams())
params.WebhookURL = nil
standIns, q := startAll(t, params)
for _, each := range anAlertForEachEvent() {
q.Raise(each.alert)
}
synctest.Wait()
want := anAlertForEachEvent()
slack := standIns[alerts.DestinationSlack].received()
ntfy := standIns[alerts.DestinationNtfy].received()
if len(slack) != len(want) || len(ntfy) != len(want) {
t.Fatalf("Slack was sent %d messages and ntfy %d, want %d each",
len(slack), len(ntfy), len(want))
}
for i, each := range want {
wantSlackMessage(t, slack[i], "*"+each.title+"*\n"+each.text)
wantNtfyMessage(t, ntfy[i], each.title, each.priorityAndTag, each.text)
}
})
}
func TestSlackAndNtfyAreSentTheSummaryAndTheRepeatsHeldBack(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := withSlackAndNtfy(newParams())
params.WebhookURL = nil
params.MaxPerHour = 1
standIns, q := startAll(t, params)
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,
// 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"})
time.Sleep(time.Hour)
synctest.Wait()
q.Raise(ban)
synctest.Wait()
slack := standIns[alerts.DestinationSlack].received()
ntfy := standIns[alerts.DestinationNtfy].received()
if len(slack) != 3 || len(ntfy) != 3 {
t.Fatalf("Slack was sent %d messages and ntfy %d, want 3 each",
len(slack), len(ntfy))
}
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; 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")
wantNtfyMessage(t, ntfy[2], instance+": ban", "default no_entry",
"a ban\nnetblock: 203.0.113.1/32")
})
}
func TestSlackMessageEscapesAmpersandsAndAngleBrackets(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := withSlackAndNtfy(newParams())
params.Instance = "app<1>"
standIns, q := startAll(t, params)
q.Raise(alerts.Alert{
Event: alerts.EventFileError, Reason: "<!channel> & <https://x.example|y>",
})
synctest.Wait()
got := standIns[alerts.DestinationSlack].received()
if len(got) != 1 {
t.Fatalf("Slack was sent %d messages, want 1", len(got))
}
wantSlackMessage(t, got[0], "*app&lt;1&gt;: file_error*\n"+
"&lt;!channel&gt; &amp; &lt;https://x.example|y&gt;")
wantBodies(t, standIns[alerts.DestinationNtfy], "<!channel> & <https://x.example|y>")
})
}
func TestNtfyTokenIsSentToNtfyAlone(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name, token string
want []string
}{
{"with a token", ntfyToken, []string{"Bearer " + ntfyToken}},
{"without one", "", nil},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := withSlackAndNtfy(newParams())
params.NtfyToken = tc.token
standIns, q := startAll(t, params)
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(1)})
synctest.Wait()
for destination, want := range map[string][]string{
alerts.DestinationWebhook: nil,
alerts.DestinationSlack: nil,
alerts.DestinationNtfy: tc.want,
} {
got := standIns[destination].received()
if len(got) != 1 {
t.Fatalf("%s was sent %d requests, want 1", destination, len(got))
}
authorization := got[0].header.Values("Authorization")
if !slices.Equal(authorization, want) {
t.Errorf("%s was sent Authorization %q, want %q", destination,
authorization, want)
}
}
})
})
}
}
func TestDestinationThatDoesNotAnswerHoldsUpNeitherOther(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := withSlackAndNtfy(newParams())
log := &lockedBuffer{}
params.ProcessLog = slog.New(slog.NewJSONHandler(log, nil))
standIns, q := startAll(t, params)
standIns[alerts.DestinationSlack].set(hanging)
for n := range 3 {
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(n)})
}
synctest.Wait()
// With no time passed, the webhook and ntfy have taken every alert,
// while Slack has not answered the first.
for destination, want := range map[string]alerts.Counts{
alerts.DestinationWebhook: {Sent: 3},
alerts.DestinationSlack: {},
alerts.DestinationNtfy: {Sent: 3},
} {
wantDestinationCounts(t, q, destination, want)
}
if got := standIns[alerts.DestinationSlack].received(); len(got) != 1 {
t.Errorf("Slack was sent %d requests, want 1", len(got))
}
// Slack's request is abandoned after 10 seconds, and logged without
// its URL; Slack, which answers again, is sent every alert a second
// later.
standIns[alerts.DestinationSlack].set(answering)
time.Sleep(11 * time.Second)
synctest.Wait()
wantDestinationCounts(t, q, alerts.DestinationSlack,
alerts.Counts{Sent: 3, Failed: 1})
logged := log.String()
if !strings.Contains(logged,
`"msg":"sending an alert to SWWAF_ALERT_SLACK_WEBHOOK_URL failed"`) ||
strings.Contains(logged, "hooks.slack.example") ||
strings.Contains(logged, "T0123") {
t.Errorf("process log %q names no failure, or names the URL", logged)
}
})
}
func TestDropsAreCountedForEachDestination(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := withSlackAndNtfy(newParams())
params.MaxPerHour = 0
standIns, q := startAll(t, params)
standIns[alerts.DestinationSlack].set(hanging)
standIns[alerts.DestinationNtfy].set(refusing)
// Each alert is raised once each destination has done all it can
// with those before it.
for n := range alerts.QueueSize + 1 {
q.Raise(alerts.Alert{Event: alerts.EventBan, Netblock: netblock(n)})
synctest.Wait()
}
// Slack, which does not answer the first alert, drops it from its
// full queue for the last; ntfy refuses each, which is given up; and
// the webhook takes every one.
const all = alerts.QueueSize + 1
for destination, want := range map[string]alerts.Counts{
alerts.DestinationWebhook: {Sent: all},
alerts.DestinationSlack: {Dropped: 1},
alerts.DestinationNtfy: {Failed: all, Dropped: all},
} {
wantDestinationCounts(t, q, destination, want)
}
})
}
func TestEachDestinationIsSentOnlyTheAlertsWaitingForIt(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
params := withSlackAndNtfy(newParams())
params.WebhookURL = nil
standIns, q := startAll(t, params)
waiting := func(reason string) alerts.Alert {
return alerts.Alert{
Instance: instance, Time: midnight(), Event: alerts.EventFileError,
Reason: reason,
}
}
// As read from alerts.json. The alert waiting for the webhook, which
// is not set, is dropped.
q.Load(roundTrip(t, alerts.State{Waiting: map[string][]alerts.Alert{
alerts.DestinationWebhook: {waiting("first")},
alerts.DestinationSlack: {waiting("second"), waiting("third")},
alerts.DestinationNtfy: {waiting("fourth")},
}}))
synctest.Wait()
wantBodies(t, standIns[alerts.DestinationSlack],
`{"text":"*fsn1app1/gitea: file_error*\nsecond"}`,
`{"text":"*fsn1app1/gitea: file_error*\nthird"}`)
wantBodies(t, standIns[alerts.DestinationNtfy], "fourth")
// alerts.json then lists Slack and ntfy alone, with no alert waiting.
want := map[string][]alerts.Alert{
alerts.DestinationSlack: {}, alerts.DestinationNtfy: {},
}
if got := q.Snapshot().Waiting; !reflect.DeepEqual(got, want) {
t.Errorf("alerts waiting %v, want %v", got, want)
}
})
}
// How a stand-in for a destination answers: with a status, or, hanging,
// not at all, until the request is abandoned.
const ( const (
answering = http.StatusNoContent answering = iota // with 204
failing = http.StatusServiceUnavailable failing // with 503
refusing = http.StatusBadRequest hanging // not at all, until the request is abandoned
hanging = 0
) )
// standIn is a stand-in for a destination. It notes each request it is // standIn is a stand-in for the webhook. It notes each request it is
// sent. // sent.
type standIn struct { type standIn struct {
mu sync.Mutex mu sync.Mutex
@@ -996,16 +434,14 @@ type standIn struct {
requests []post requests []post
} }
// post is a request a destination was sent: when, its method, URL and // post is a request the webhook was sent: when, its method, URL and
// headers, its body, and that body read as a JSON object, which for the // headers, the alert it carried, and whether the webhook answered it with
// webhook is the alert it carried, and whether the destination answered // a 2xx status.
// it with a 2xx status.
type post struct { type post struct {
at time.Time at time.Time
method string method string
url string url string
header http.Header header http.Header
body string
alert map[string]any alert map[string]any
answered bool answered bool
} }
@@ -1039,18 +475,21 @@ func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) {
answers := s.answers answers := s.answers
s.requests = append(s.requests, post{ s.requests = append(s.requests, post{
at: time.Now(), method: r.Method, url: r.URL.String(), header: r.Header.Clone(), at: time.Now(), method: r.Method, url: r.URL.String(), header: r.Header.Clone(),
body: string(body), alert: alert, answered: answers == answering, alert: alert, answered: answers == answering,
}) })
s.mu.Unlock() s.mu.Unlock()
if answers == hanging { switch answers {
case failing:
w.WriteHeader(http.StatusServiceUnavailable)
case hanging:
<-r.Context().Done() <-r.Context().Done()
} else { default:
w.WriteHeader(answers) w.WriteHeader(http.StatusNoContent)
} }
} }
// set sets how the stand-in answers: with the status answers, or hanging. // set sets how the stand-in answers.
func (s *standIn) set(answers int) { func (s *standIn) set(answers int) {
s.mu.Lock() s.mu.Lock()
defer s.mu.Unlock() defer s.mu.Unlock()
@@ -1093,8 +532,13 @@ func (b *lockedBuffer) String() string {
// every event, the default cooldown and hourly limit, and the bubble's // every event, the default cooldown and hourly limit, and the bubble's
// clock in UTC. // clock in UTC.
func newParams() alerts.Params { func newParams() alerts.Params {
webhook, err := url.Parse(webhookURL)
if err != nil {
panic(err)
}
return alerts.Params{ return alerts.Params{
WebhookURL: parseURL(webhookURL), WebhookURL: webhook,
Events: alerts.Events(), Events: alerts.Events(),
Cooldown: cooldown, Cooldown: cooldown,
MaxPerHour: 60, MaxPerHour: 60,
@@ -1104,48 +548,14 @@ func newParams() alerts.Params {
} }
} }
// withSlackAndNtfy returns params with Slack at slackURL and ntfy at
// ntfyURL set as well.
func withSlackAndNtfy(params alerts.Params) alerts.Params {
params.SlackURL = parseURL(slackURL)
params.NtfyURL = parseURL(ntfyURL)
return params
}
// parseURL returns rawURL, parsed.
func parseURL(rawURL string) *url.URL {
parsed, err := url.Parse(rawURL)
if err != nil {
panic(err)
}
return parsed
}
// start returns a stand-in for the webhook that answers, and a Queue that // start returns a stand-in for the webhook that answers, and a Queue that
// sends to it, run until the test ends. // sends to it, run until the test ends.
func start(t *testing.T, params alerts.Params) (*standIn, *alerts.Queue) { func start(t *testing.T, params alerts.Params) (*standIn, *alerts.Queue) {
t.Helper() t.Helper()
standIns, q := startAll(t, params) webhook := &standIn{}
return standIns[alerts.DestinationWebhook], q
}
// startAll returns, by destination, a stand-in that answers for each
// destination params sets, and a Queue that sends to them, run until the
// test ends.
func startAll(t *testing.T, params alerts.Params) (map[string]*standIn, *alerts.Queue) {
t.Helper()
q := alerts.New(params) q := alerts.New(params)
standIns := map[string]*standIn{} q.SetTransport(webhook)
for _, destination := range q.DestinationsSet() {
standIns[destination] = &standIn{answers: answering}
q.SetTransport(destination, standIns[destination])
}
ctx, stop := context.WithCancel(t.Context()) ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{}) stopped := make(chan struct{})
@@ -1160,7 +570,7 @@ func startAll(t *testing.T, params alerts.Params) (map[string]*standIn, *alerts.
<-stopped <-stopped
}) })
return standIns, q return webhook, q
} }
// midnight is when each test starts. // midnight is when each test starts.
@@ -1219,158 +629,17 @@ func wantEvents(t *testing.T, webhook *standIn, want ...string) {
} }
} }
// wantCounts checks the alerts q counts as sent to the webhook, the // wantCounts checks the alerts q counts as sent, the requests it counts as
// requests to it that it counts as failed, the alerts it counts as held // failed, and the alerts it counts as held back and as dropped.
// back, and those it counts as dropped for the webhook.
func wantCounts( func wantCounts(
t *testing.T, q *alerts.Queue, sent, failed, suppressed, dropped int64, t *testing.T, q *alerts.Queue, sent, failed, suppressed, dropped int64,
) { ) {
t.Helper() t.Helper()
wantDestinationCounts(t, q, alerts.DestinationWebhook, if q.Sent() != sent || q.Failed() != failed || q.Suppressed() != suppressed ||
alerts.Counts{Sent: sent, Failed: failed, Dropped: dropped}) q.Dropped() != dropped {
t.Errorf("counts sent %d, failed %d, suppressed %d and dropped %d, "+
if q.Suppressed() != suppressed { "want %d, %d, %d and %d", q.Sent(), q.Failed(), q.Suppressed(), q.Dropped(),
t.Errorf("%d alerts held back, want %d", q.Suppressed(), suppressed) sent, failed, suppressed, dropped)
}
}
// wantDestinationCounts checks q's counts for destination.
func wantDestinationCounts(
t *testing.T, q *alerts.Queue, destination string, want alerts.Counts,
) {
t.Helper()
if got := q.Counts(destination); got != want {
t.Errorf("%s counts %+v, want %+v", destination, got, want)
}
}
// eventMessage is an alert, and the message Slack and ntfy are sent for
// it: its title, the priority and the tag ntfy is sent, with a space
// between them, and its text.
type eventMessage struct {
alert alerts.Alert
title, priorityAndTag, text string
}
// anAlertForEachEvent returns an alert for each event SWWAF_ALERT_EVENTS
// names, the first in observe mode, each with its message.
func anAlertForEachEvent() []eventMessage {
return []eventMessage{
{
alerts.Alert{
Event: alerts.EventBan, Client: netip.MustParseAddr("203.0.113.9"),
Netblock: netip.MustParsePrefix("203.0.113.0/24"), Country: "DE",
Reason: "requests per hour over the limit of 10000",
Detail: map[string]any{"mode": "observe"},
},
"fsn1app1/gitea: ban", "default no_entry",
"requests per hour over the limit of 10000\nclient: 203.0.113.9\n" +
"netblock: 203.0.113.0/24\ncountry: DE\nmode: observe",
},
{
alerts.Alert{
Event: alerts.EventPermanentBan, Client: netip.MustParseAddr("198.51.100.7"),
Netblock: netip.MustParsePrefix("198.51.100.7/32"), Country: "FR",
Reason: "matched the rule env-file",
},
"fsn1app1/gitea: permanent_ban", "high no_entry",
"matched the rule env-file\nclient: 198.51.100.7\n" +
"netblock: 198.51.100.7/32\ncountry: FR",
},
{
alerts.Alert{
Event: alerts.EventWAFBlock, Client: netip.MustParseAddr("192.0.2.1"),
Netblock: netip.MustParsePrefix("192.0.2.1/32"),
Reason: "refused by the Core Rule Set",
},
"fsn1app1/gitea: waf_block", "default shield",
"refused by the Core Rule Set\nclient: 192.0.2.1\nnetblock: 192.0.2.1/32",
},
{
alerts.Alert{
Event: alerts.EventAnomaly, Netblock: netip.MustParsePrefix("192.0.2.0/24"),
Reason: "requests per minute over the threshold of 5000",
},
"fsn1app1/gitea: anomaly", "high chart_with_upwards_trend",
"requests per minute over the threshold of 5000\nnetblock: 192.0.2.0/24",
},
{
alerts.Alert{
Event: alerts.EventReputationHit, Client: netip.MustParseAddr("192.0.2.2"),
Netblock: netip.MustParsePrefix("192.0.2.2/32"),
Reason: "listed by a DNS blocklist",
},
"fsn1app1/gitea: reputation_hit", "low label",
"listed by a DNS blocklist\nclient: 192.0.2.2\nnetblock: 192.0.2.2/32",
},
{
alerts.Alert{
Event: alerts.EventSourceFailure, Reason: "asking GeoJS failed",
Detail: map[string]any{"source": "geojs"},
},
"fsn1app1/gitea: source_failure", "high warning",
"asking GeoJS failed\nsource: geojs",
},
{
alerts.Alert{
Event: alerts.EventFileError,
Reason: "a rule file has an error, and the rules stay as they were",
Detail: map[string]any{
"file": "/etc/smallwebwaf/rules.d/50-app.rules", "error": "line 2: no action",
},
},
"fsn1app1/gitea: file_error", "high warning",
"a rule file has an error, and the rules stay as they were\n" +
"file: /etc/smallwebwaf/rules.d/50-app.rules\nerror: line 2: no action",
},
}
}
// wantSlackMessage checks that request, one Slack was sent, posted the
// message text, as JSON.
func wantSlackMessage(t *testing.T, request post, text string) {
t.Helper()
if request.method != http.MethodPost || request.url != slackURL ||
request.header.Get("Content-Type") != "application/json" ||
!reflect.DeepEqual(request.alert, map[string]any{"text": text}) {
t.Errorf("Slack was sent %s %s, Content-Type %q, %s, want POST %s, "+
"application/json, the text %q", request.method, request.url,
request.header.Get("Content-Type"), request.body, slackURL, text)
}
}
// wantNtfyMessage checks that request, one ntfy was sent, posted the
// message text with the title, and with the priority and the tag
// priorityAndTag gives, with a space between them.
func wantNtfyMessage(t *testing.T, request post, title, priorityAndTag, text string) {
t.Helper()
got := []string{
request.method, request.url, request.header.Get("Title"),
request.header.Get("Priority") + " " + request.header.Get("Tags"), request.body,
}
want := []string{http.MethodPost, ntfyURL, title, priorityAndTag, text}
if !slices.Equal(got, want) {
t.Errorf("ntfy was sent the method, URL, title, priority and tag, and text "+
"%q, want %q", got, want)
}
}
// wantBodies checks the bodies of the requests the stand-in was sent, in
// order.
func wantBodies(t *testing.T, s *standIn, want ...string) {
t.Helper()
got := make([]string, 0, len(want))
for _, request := range s.received() {
got = append(got, request.body)
}
if !slices.Equal(got, want) {
t.Errorf("bodies %q, want %q", got, want)
} }
} }
+5 -9
View File
@@ -2,15 +2,11 @@ package alerts
import "net/http" import "net/http"
// QueueSize is the most alerts that wait to be sent to a destination. // QueueSize is the most alerts that wait to be sent.
const QueueSize = queueSize const QueueSize = queueSize
// SetTransport has q's requests to the destination name go through // SetTransport has q's requests to the webhook go through transport
// transport instead of the network. // instead of the network.
func (q *Queue) SetTransport(name string, transport http.RoundTripper) { func (q *Queue) SetTransport(transport http.RoundTripper) {
for _, d := range q.destinations { q.httpClient.Transport = transport
if d.name == name {
d.httpClient.Transport = transport
}
}
} }
-408
View File
@@ -1,408 +0,0 @@
// 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,
}
}
-238
View File
@@ -1,238 +0,0 @@
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
})
}
+6 -10
View File
@@ -2,7 +2,6 @@ package bans_test
import ( import (
"net/netip" "net/netip"
"reflect"
"testing" "testing"
"time" "time"
@@ -67,15 +66,12 @@ func TestReasonOfTheBansSmallwebwafMakes(t *testing.T) {
ledger := bans.New(defaultRules()) ledger := bans.New(defaultRules())
limit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(), limit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
bans.Notes{Kind: "requests", Limit: 1000, Window: "minute"}) bans.Notes{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(), attack, _ := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
bans.Notes{RuleID: "git-dir", Target: "path"}) bans.Notes{RuleID: "git-dir", Target: "path"})
for _, tc := range []struct{ got, want string }{ for _, tc := range []struct{ got, want string }{
{limit.Reason, "requests per minute over the limit of 1000"}, {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"}, {attack.Reason, "matched the rule git-dir"},
} { } {
if tc.got != tc.want { if tc.got != tc.want {
@@ -118,7 +114,7 @@ func TestLiftedBanForALimitRefusesNothingAndMakesNoBanLonger(t *testing.T) {
} }
held := ledger.Bans(netblock) held := ledger.Bans(netblock)
if len(held) != 2 || !reflect.DeepEqual(held[0], lifted) { if len(held) != 2 || held[0] != lifted {
t.Errorf("the ledger holds %+v, want the lifted ban and the new one", held) t.Errorf("the ledger holds %+v, want the lifted ban and the new one", held)
} }
} }
@@ -138,7 +134,7 @@ func TestLiftedBanForAnAttackRefusesNothingAndMakesNoBanLonger(t *testing.T) {
now := midnight().Add(2 * time.Hour) now := midnight().Add(2 * time.Hour)
_, banned, _ := ledger.Find(netblock.Addr(), now) _, banned := ledger.Find(netblock.Addr(), now)
if banned { if banned {
t.Error("the lifted ban refuses") t.Error("the lifted ban refuses")
} }
@@ -213,7 +209,7 @@ func TestAdminsBanIsMadeWhileAnotherLasts(t *testing.T) {
got := ledger.BanForAdmin(netip.MustParsePrefix("203.0.113.9/24"), now, time.Time{}, got := ledger.BanForAdmin(netip.MustParsePrefix("203.0.113.9/24"), now, time.Time{},
"probes for logins") "probes for logins")
if !reflect.DeepEqual(got, want) { if got != want {
t.Errorf("the admin's ban is\n%+v\nwant\n%+v", got, want) t.Errorf("the admin's ban is\n%+v\nwant\n%+v", got, want)
} }
@@ -224,8 +220,8 @@ func TestAdminsBanIsMadeWhileAnotherLasts(t *testing.T) {
} }
// It refuses once the ban for the limit has ended. // It refuses once the ban for the limit has ended.
ban, banned, _ := ledger.Find(netblock.Addr(), midnight().Add(2*time.Hour)) ban, banned := ledger.Find(netblock.Addr(), midnight().Add(2*time.Hour))
if !banned || !reflect.DeepEqual(ban, want) { if !banned || ban != want {
t.Errorf("after the limit's ban the netblock is under %+v (%t), want %+v", t.Errorf("after the limit's ban the netblock is under %+v (%t), want %+v",
ban, banned, want) ban, banned, want)
} }
+23 -125
View File
@@ -1,8 +1,8 @@
// Package bans is the ban ledger: the bans smallwebwaf makes on the // Package bans is the ban ledger: the bans smallwebwaf makes on the
// netblocks of clients that break a rate limit or a byte limit or show a // netblocks of clients that break a rate limit or show a clear sign of
// clear sign of attack, and those an admin makes, with their notes, as // attack, and those an admin makes, with their notes, as the "Bans"
// the "Bans" section of SPEC.md describes. The bans are kept in memory, // section of SPEC.md describes. The bans are kept in memory, and written
// and written to bans.json and read from it by the state package. // to bans.json and read from it by the state package.
package bans package bans
import ( import (
@@ -91,39 +91,22 @@ func (b Ban) ActiveAt(now time.Time) bool {
// //
//nolint:tagliatelle // the state files use snake_case, as the request log does //nolint:tagliatelle // the state files use snake_case, as the request log does
type Notes struct { type Notes struct {
// ASN, ASName and Country are the client's AS number, AS name and // Country is the client's country, when it was looked up.
// country, when they were looked up: when the request that caused the
// ban was made, or when GeoJS answered about the client afterwards.
ASN string `json:"asn"`
ASName string `json:"as_name"`
Country string `json:"country"` Country string `json:"country"`
// Kind, Limit, Window and Count are, for a ban for a broken limit, // Limit, Window and Count are, for a ban for a broken limit, the limit
// what the limit was on, "requests" for a rate limit or "bytes" for a // that was broken, its window, "minute", "hour" or "day", and the
// byte limit, the limit that was broken, its window, "minute", "hour" // count reached: the client's requests in the window, the one that
// or "day", and the count reached: the client's requests, or bytes, in // broke the limit included. These are the requests that counted
// the window, those of the request that broke the limit included. // toward the ban, and the window is the time over which they came.
// 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"` Limit int64 `json:"limit,omitempty"`
Window string `json:"window,omitempty"` Window string `json:"window,omitempty"`
Count float64 `json:"count,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 // 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. // of the rule file rule that matched, and its target.
RuleID string `json:"rule_id,omitempty"` RuleID string `json:"rule_id,omitempty"`
Target string `json:"target,omitempty"` Target string `json:"target,omitempty"`
// Reputation is the reputation sources that listed the client when // Request is the request that broke the limit, or that was the clear
// the request that caused the ban was made, in the order the request // sign of attack.
// 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"` Request Request `json:"request"`
// Requests is how many requests the netblock has sent since it was // 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 // first seen, and Refused how many of them the ban has refused so
@@ -135,15 +118,6 @@ type Notes struct {
EarlierBans EarlierBans `json:"earlier_bans"` 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. // EarlierBans counts a netblock's bans before a ban, by cause.
type EarlierBans struct { type EarlierBans struct {
Limit int `json:"limit"` Limit int `json:"limit"`
@@ -244,19 +218,17 @@ func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool, bool) {
} }
// Find is Check without counting the request among those the ban // Find is Check without counting the request among those the ban
// refused, and without making the ban permanent: in observe mode a ban // refused: in observe mode a ban refuses nothing.
// refuses nothing. The last result reports whether Check would have made func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool) {
// the ban permanent.
func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool, bool) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
ban := l.active(client, now) ban := l.active(client, now)
if ban == nil { if ban == nil {
return Ban{}, false, false return Ban{}, false
} }
return *ban, true, ban.Cause == CauseAttack && !ban.Permanent() return *ban, true
} }
// activeBan returns the ban in bans, a netblock's bans oldest first, that // activeBan returns the ban in bans, a netblock's bans oldest first, that
@@ -282,20 +254,14 @@ func activeBan(bans []Ban, now time.Time) *Ban {
// active, as when two of its requests break a limit at once, that ban is // 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 // returned with false, and no other is made. The ledger fills in the
// notes' Refused and EarlierBans itself, and gives the ban the reason // notes' Refused and EarlierBans itself, and gives the ban the reason
// "<Kind> per <Window> over the limit of <Limit>", from the notes, such // "requests per <Window> over the limit of <Limit>", from the notes.
// as "requests per minute over the limit of 1000".
func (l *Ledger) BanForLimit( func (l *Ledger) BanForLimit(
netblock netip.Prefix, now time.Time, notes Notes, netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) { ) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, true) reason := fmt.Sprintf("requests per %s over the limit of %d",
} notes.Window, notes.Limit)
// WouldBanForLimit returns what BanForLimit would, without making the ban: return l.ban(netblock, now, CauseLimit, reason, notes)
// what observe mode would have done.
func (l *Ledger) WouldBanForLimit(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, false)
} }
// BanForAttack bans netblock at now for a clear sign of attack, with // BanForAttack bans netblock at now for a clear sign of attack, with
@@ -306,48 +272,7 @@ func (l *Ledger) WouldBanForLimit(
func (l *Ledger) BanForAttack( func (l *Ledger) BanForAttack(
netblock netip.Prefix, now time.Time, notes Notes, netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) { ) (Ban, bool) {
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, true) return l.ban(netblock, now, CauseAttack, "matched the rule "+notes.RuleID, notes)
}
// WouldBanForAttack returns what BanForAttack would, without making the
// ban: what observe mode would have done.
func (l *Ledger) WouldBanForAttack(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, false)
}
// WouldBePermanent reports whether a ban on netblock for cause, CauseLimit
// or CauseAttack, made at now would be permanent, as BanForLimit or
// BanForAttack would make it. It works out nothing else of the ban.
func (l *Ledger) WouldBePermanent(
netblock netip.Prefix, now time.Time, cause string,
) bool {
l.mu.Lock()
defer l.mu.Unlock()
var held []Ban
if bans, found := l.netblocks.Peek(netblock); found {
held = *bans
}
if cause == CauseAttack {
return l.attackExpiry(held, now).IsZero()
}
return l.limitExpiry(held, now).IsZero()
}
// limitReason is the reason of a ban for a broken limit, with notes.
func limitReason(notes Notes) string {
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
// notes.
func attackReason(notes Notes) string {
return "matched the rule " + notes.RuleID
} }
// BanForAdmin bans netblock at now for an admin, with reason, until // BanForAdmin bans netblock at now for an admin, with reason, until
@@ -439,28 +364,6 @@ func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
return slices.Clone(*bans) return slices.Clone(*bans)
} }
// 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 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()
bans, found := l.netblocks.Peek(netblock)
if !found {
return
}
for i := range *bans {
notes := &(*bans)[i].Notes
if notes.ASN == "" && notes.ASName == "" && notes.Country == "" {
notes.ASN, notes.ASName, notes.Country = asn, asName, country
}
}
}
// Made returns how many bans for cause have been made since the start: // Made returns how many bans for cause have been made since the start:
// for CauseLimit and CauseAttack, by the ledger; for CauseAdmin, by an // for CauseLimit and CauseAttack, by the ledger; for CauseAdmin, by an
// admin, with BanForAdmin or in an edit of bans.json, as LoadEdit counts // admin, with BanForAdmin or in an edit of bans.json, as LoadEdit counts
@@ -588,10 +491,9 @@ func (l *Ledger) holds(netblock netip.Prefix, start time.Time) bool {
// ban bans netblock at now for cause, with reason and notes, as // ban bans netblock at now for cause, with reason and notes, as
// BanForLimit and BanForAttack describe, and returns the ban, and whether // BanForLimit and BanForAttack describe, and returns the ban, and whether
// it made it. Unless keep is true, the ban is not made, only returned: it // it made it.
// is the ban that would have been made.
func (l *Ledger) ban( func (l *Ledger) ban(
netblock netip.Prefix, now time.Time, cause, reason string, notes Notes, keep bool, netblock netip.Prefix, now time.Time, cause, reason string, notes Notes,
) (Ban, bool) { ) (Ban, bool) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
@@ -619,10 +521,6 @@ func (l *Ledger) ban(
ban.Expires = l.limitExpiry(held, now) ban.Expires = l.limitExpiry(held, now)
} }
if !keep {
return ban, true
}
l.add(ban) l.add(ban)
l.made[cause]++ l.made[cause]++
l.markChanged() l.markChanged()
+11 -133
View File
@@ -2,7 +2,6 @@ package bans_test
import ( import (
"net/netip" "net/netip"
"reflect"
"strings" "strings"
"testing" "testing"
"time" "time"
@@ -131,13 +130,13 @@ func TestBrokenLimitDuringABanMakesNoOther(t *testing.T) {
again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{}) again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
if made || !reflect.DeepEqual(again, first) || len(ledger.Bans(netblock)) != 1 { if made || again != first || len(ledger.Bans(netblock)) != 1 {
t.Errorf("a limit broken during a ban gave %+v, made %t, and %d bans, "+ 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) "want %+v, not made, and 1", again, made, len(ledger.Bans(netblock)), first)
} }
again, made = ledger.BanForAttack(netblock, midnight().Add(time.Minute), bans.Notes{}) again, made = ledger.BanForAttack(netblock, midnight().Add(time.Minute), bans.Notes{})
if made || !reflect.DeepEqual(again, first) { if made || again != first {
t.Errorf("an attack during a ban gave %+v, made %t, want %+v, not made", t.Errorf("an attack during a ban gave %+v, made %t, want %+v, not made",
again, made, first) again, made, first)
} }
@@ -182,17 +181,17 @@ func TestFindCountsNothing(t *testing.T) {
netblock := netip.MustParsePrefix("203.0.113.9/32") netblock := netip.MustParsePrefix("203.0.113.9/32")
ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5}) ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5})
got, banned, _ := ledger.Find(netblock.Addr(), ban.Expires.Add(-time.Nanosecond)) got, banned := ledger.Find(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
if !banned || !reflect.DeepEqual(got, ban) { if !banned || got != ban {
t.Errorf("find during the ban gives %+v and %t, want %+v", got, banned, ban) t.Errorf("find during the ban gives %+v and %t, want %+v", got, banned, ban)
} }
_, banned, _ = ledger.Find(netblock.Addr(), ban.Expires) _, banned = ledger.Find(netblock.Addr(), ban.Expires)
if banned { if banned {
t.Error("the ban did not end") t.Error("the ban did not end")
} }
if notes := ledger.Bans(netblock)[0].Notes; !reflect.DeepEqual(notes, ban.Notes) { if notes := ledger.Bans(netblock)[0].Notes; notes != ban.Notes {
t.Errorf("the notes are %+v, want them unchanged, %+v", notes, ban.Notes) t.Errorf("the notes are %+v, want them unchanged, %+v", notes, ban.Notes)
} }
} }
@@ -248,7 +247,7 @@ func TestFullLedgerDropsTheEarlierBanOfTheNetblockBannedAgain(t *testing.T) {
second, _ := ledger.BanForLimit(netblock, first.Expires, bans.Notes{}) second, _ := ledger.BanForLimit(netblock, first.Expires, bans.Notes{})
held := ledger.Bans(netblock) held := ledger.Bans(netblock)
if len(held) != 1 || !reflect.DeepEqual(held[0], second) || if len(held) != 1 || held[0] != second ||
held[0].Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) { held[0].Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("the ledger holds %+v, want only the second ban, "+ t.Errorf("the ledger holds %+v, want only the second ban, "+
"with 1 earlier ban for a limit", held) "with 1 earlier ban for a limit", held)
@@ -273,16 +272,12 @@ func TestRequestDuringAnAttackBanMakesItPermanent(t *testing.T) {
wantChanged(t, ledger, true) wantChanged(t, ledger, true)
// In observe mode the ban refuses nothing, and stays as it is, while // In observe mode the ban refuses nothing, and stays as it is.
// Find tells that the request would have made it permanent. got, _ := ledger.Find(netblock.Addr(), midnight().Add(time.Hour))
got, _, wouldMakePermanent := ledger.Find(netblock.Addr(), midnight().Add(time.Hour)) if got.Permanent() {
if got.Permanent() || ledger.Bans(netblock)[0].Permanent() || !wouldMakePermanent { t.Fatal("a request found under the ban made it permanent")
t.Fatalf("a request found under the ban left it %+v, would have made it "+
"permanent %t, want it as it was, and true", got, wouldMakePermanent)
} }
wantChanged(t, ledger, false)
// A request it refuses makes it permanent, says so, and makes // A request it refuses makes it permanent, says so, and makes
// bans.json due. // bans.json due.
got, _, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(time.Hour)) got, _, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(time.Hour))
@@ -333,84 +328,6 @@ func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
} }
} }
func TestWouldBanGivesTheBanWithoutMakingIt(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
wantChanged(t, ledger, true)
// While the first ban lasts, none would be made.
during, would := ledger.WouldBanForAttack(netblock, midnight(), bans.Notes{})
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{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)
if !wouldAttack || !attack.Expires.Equal(first.Expires.Add(7*day)) ||
attack.Reason != "matched the rule git-dir" || !wouldLimit ||
!limit.Expires.Equal(first.Expires.Add(3*time.Hour)) ||
limit.Reason != "requests per minute over the limit of 1" {
t.Errorf("would ban with %+v and %+v, want seven days for the attack and "+
"three hours for the limit", attack, limit)
}
if len(ledger.Bans(netblock)) != 1 || ledger.Made(bans.CauseLimit) != 1 ||
ledger.Made(bans.CauseAttack) != 0 {
t.Errorf("the ledger holds %+v, want the first ban alone", ledger.Bans(netblock))
}
wantChanged(t, ledger, false)
// The ban made is the one that would have been.
made, _ := ledger.BanForLimit(netblock, first.Expires, limitNotes)
if !reflect.DeepEqual(made, limit) {
t.Errorf("the ban made is %+v, want %+v", made, limit)
}
}
func TestWouldBePermanentAnswersAsTheBanWouldBeMade(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
now := midnight()
// Five bans for a limit in a row, of 1, 3, 9, 27 and 81 hours, are not
// permanent. The sixth, of 243 hours, would be, while a first ban for
// an attack would not.
for i := range 5 {
if ledger.WouldBePermanent(netblock, now, bans.CauseLimit) {
t.Fatalf("ban %d for a limit would be permanent", i+1)
}
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
now = ban.Expires
}
if !ledger.WouldBePermanent(netblock, now, bans.CauseLimit) {
t.Error("the sixth ban for a limit would not be permanent")
}
if ledger.WouldBePermanent(netblock, now, bans.CauseAttack) {
t.Error("a first ban for an attack would be permanent")
}
// Once a first ban for an attack has ended, the next would be permanent.
attack, _ := ledger.BanForAttack(netblock, now, bans.Notes{})
if !ledger.WouldBePermanent(netblock, attack.Expires, bans.CauseAttack) {
t.Error("a second ban for an attack would not be permanent")
}
}
func TestAttackBanDoesNotLengthenTheNextBanForALimit(t *testing.T) { func TestAttackBanDoesNotLengthenTheNextBanForALimit(t *testing.T) {
t.Parallel() t.Parallel()
@@ -455,45 +372,6 @@ func TestRequestTextsAreCutTo256Bytes(t *testing.T) {
} }
} }
func TestLookupFillsTheNotesOfTheNetblocksBansWithoutOne(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
other := netip.MustParsePrefix("198.51.100.7/32")
// A ban made with the client's lookup, one made before it came, after
// the first ended, and one on another netblock.
ledger.BanForLimit(netblock, midnight(), bans.Notes{
ASN: "AS64497", ASName: "Other Net", Country: "FR",
})
ledger.BanForLimit(netblock, midnight().Add(time.Hour), bans.Notes{})
ledger.BanForLimit(other, midnight(), bans.Notes{})
ledger.AddLookup(netblock, "AS64496", "Example Net", "DE")
held := ledger.Bans(netblock)
if len(held) != 2 {
t.Fatalf("%s has %d bans, want 2", netblock, len(held))
}
for i, want := range []bans.Notes{
{ASN: "AS64497", ASName: "Other Net", Country: "FR"},
{ASN: "AS64496", ASName: "Example Net", Country: "DE"},
} {
got := held[i].Notes
if got.ASN != want.ASN || got.ASName != want.ASName || got.Country != want.Country {
t.Errorf("ban %d's notes give %q, %q and %q, want %q, %q and %q", i+1,
got.ASN, got.ASName, got.Country, want.ASN, want.ASName, want.Country)
}
}
if notes := ledger.Bans(other)[0].Notes; notes.ASN != "" || notes.Country != "" {
t.Errorf("the ban on %s has %q and %q, want neither",
other, notes.ASN, notes.Country)
}
}
// defaultRules are the rules at the settings' defaults. // defaultRules are the rules at the settings' defaults.
func defaultRules() bans.Rules { func defaultRules() bans.Rules {
return bans.Rules{ return bans.Rules{
+3 -4
View File
@@ -2,7 +2,6 @@ package bans_test
import ( import (
"net/netip" "net/netip"
"reflect"
"slices" "slices"
"strings" "strings"
"testing" "testing"
@@ -151,7 +150,7 @@ func TestPermanentBanStartedBeforeAnEndedOneRefuses(t *testing.T) {
now := midnight().Add(2 * time.Hour) now := midnight().Add(2 * time.Hour)
client := netip.MustParseAddr("203.0.113.9") client := netip.MustParseAddr("203.0.113.9")
ban, banned, _ := ledger.Find(client, now) ban, banned := ledger.Find(client, now)
if !banned || !ban.Permanent() { if !banned || !ban.Permanent() {
t.Errorf("find gives %+v and %t, want the permanent ban", ban, banned) t.Errorf("find gives %+v and %t, want the permanent ban", ban, banned)
} }
@@ -225,7 +224,7 @@ func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) {
ledger.Load([]bans.Ban{later, earlier}) ledger.Load([]bans.Ban{later, earlier})
held := ledger.Snapshot() held := ledger.Snapshot()
if len(held) != 1 || !reflect.DeepEqual(held[0], later) { if len(held) != 1 || held[0] != later {
t.Errorf("the ledger holds %+v, want only the ban that began later", held) t.Errorf("the ledger holds %+v, want only the ban that began later", held)
} }
} }
@@ -268,7 +267,7 @@ func TestLoadReplacesTheBansHeld(t *testing.T) {
bans.Notes{}) bans.Notes{})
want := []bans.Ban{first, second, kept} want := []bans.Ban{first, second, kept}
if got := ledger.Snapshot(); !reflect.DeepEqual(got, want) { if got := ledger.Snapshot(); !slices.Equal(got, want) {
t.Errorf("the ledger holds %+v, want %+v", got, want) t.Errorf("the ledger holds %+v, want %+v", got, want)
} }
} }
+65 -838
View File
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
-233
View File
@@ -1,233 +0,0 @@
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
}
-379
View File
@@ -1,379 +0,0 @@
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)
}
}
+66 -129
View File
@@ -1,8 +1,7 @@
// Package lookup looks up each client's AS number and country, through // Package lookup looks up each client's country through the GeoJS web
// the GeoJS web service or in the lookup database, the IPinfo Lite file // service, and keeps the answers in memory, for at most 100,000 clients
// SWWAF_LOOKUP_DB_PATH names. GeoJS's answers are kept in memory, for at // and for 7 days each. The answers are written to lookups.json and read
// most 100,000 clients and for 7 days each, and are written to // from it by the state package.
// lookups.json and read from it by the state package.
package lookup package lookup
import ( import (
@@ -15,7 +14,6 @@ import (
"net/http" "net/http"
"net/netip" "net/netip"
"slices" "slices"
"strconv"
"strings" "strings"
"sync" "sync"
"time" "time"
@@ -25,10 +23,9 @@ import (
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
) )
// URL is GeoJS's endpoint for an address's place and network. Asked about // URL is GeoJS's country endpoint. Asked about several addresses at once,
// several addresses at once, comma separated in its ip parameter, it // comma separated in its ip parameter, it answers with a list.
// answers with a list. const URL = "https://get.geojs.io/v1/ip/country.json"
const URL = "https://get.geojs.io/v1/ip/geo.json"
const ( const (
// keepFor is how long an answer is used instead of asking GeoJS again. // keepFor is how long an answer is used instead of asking GeoJS again.
@@ -43,8 +40,9 @@ const (
maxWaiting = 10000 maxWaiting = 10000
// maxPerRequest is how many addresses one request to GeoJS asks about. // maxPerRequest is how many addresses one request to GeoJS asks about.
maxPerRequest = 200 maxPerRequest = 200
// unknownASN is the AS number GeoJS gives when it knows none. // timeout is how long a new client waits for its answer, and how long
unknownASN = 64512 // a request to GeoJS may take before it is abandoned.
timeout = time.Second
// After a failure GeoJS is not asked again for a second, and for // After a failure GeoJS is not asked again for a second, and for
// retryDelayFactor times as long after each further failure in a row, // retryDelayFactor times as long after each further failure in a row,
// up to five minutes. // up to five minutes.
@@ -64,16 +62,6 @@ var (
type Params struct { type Params struct {
// URL is where GeoJS is asked, normally URL. // URL is where GeoJS is asked, normally URL.
URL string URL string
// Timeout is how long a request waits for its client's first answer,
// and how long a request to GeoJS may take before it is abandoned
// (SWWAF_LOOKUP_TIMEOUT).
Timeout time.Duration
// Wait is true when a setting needs each request's answer before the
// request goes on. Otherwise no request waits for one.
Wait bool
// Answered, unless nil, is given each answer GeoJS gives, once it is
// kept.
Answered func(Answer)
// Now tells the time, normally time.Now. // Now tells the time, normally time.Now.
Now func() time.Time Now func() time.Time
// ProcessLog receives GeoJS's failures. // ProcessLog receives GeoJS's failures.
@@ -85,14 +73,11 @@ type Params struct {
Alerts *alerts.Queue Alerts *alerts.Queue
} }
// GeoJS looks up clients' AS numbers and countries through GeoJS. At most // GeoJS looks up clients' countries through GeoJS. At most one request
// one request to GeoJS is under way at a time, and it asks about every // to GeoJS is under way at a time, and it asks about every client waiting,
// client waiting, up to maxPerRequest. It is safe for concurrent use. // up to maxPerRequest. It is safe for concurrent use.
type GeoJS struct { type GeoJS struct {
url string url string
timeout time.Duration
wait bool
answered func(Answer)
now func() time.Time now func() time.Time
processLog *slog.Logger processLog *slog.Logger
metrics *metrics.Metrics metrics *metrics.Metrics
@@ -114,18 +99,11 @@ type GeoJS struct {
retryAt time.Time retryAt time.Time
} }
// Answer is what GeoJS or the lookup database said about a client: its AS // Answer is what GeoJS said about a client, as lookups.json holds it: its
// number, such as AS64496, and the AS's name, both "" when the source knows // country, "" when GeoJS cannot place it, when GeoJS said so, and when
// no AS number for it; its country, "" when the source cannot place it; // the answer was last used.
// 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 { type Answer struct {
Client netip.Prefix `json:"client"` Client netip.Prefix `json:"client"`
ASN string `json:"asn"`
ASName string `json:"as_name"`
Country string `json:"country"` Country string `json:"country"`
Answered time.Time `json:"answered"` Answered time.Time `json:"answered"`
Used time.Time `json:"used"` Used time.Time `json:"used"`
@@ -150,9 +128,6 @@ func New(params Params) *GeoJS {
return &GeoJS{ return &GeoJS{
url: params.URL, url: params.URL,
timeout: params.Timeout,
wait: params.Wait,
answered: params.Answered,
now: params.Now, now: params.Now,
processLog: params.ProcessLog, processLog: params.ProcessLog,
metrics: params.Metrics, metrics: params.Metrics,
@@ -167,24 +142,23 @@ func New(params Params) *GeoJS {
} }
} }
// LookUp returns the answer GeoJS gave about client, with its country as // Country returns the country GeoJS places client in, as a two-letter
// a two-letter code in capitals, or the zero Answer when there is none // code in capitals, or "" when the country cannot be found: GeoJS cannot
// yet. An answer is kept for 7 days. Without one, the client is asked // place the client, or has not answered in time. An answer is kept for 7
// about in the background, and, while Wait is set, the request waits up // days. Without one, a client waits up to timeout for it, unless it has
// to Timeout for the answer, unless the client has gone without one // gone without one before; until GeoJS answers, the client is asked about
// before. ctx is the context of the client's request, and ends the wait // again in the background. ctx is the context of the client's request,
// when it ends. // and ends the wait when it ends.
// //
// GeoJS is asked about the client's first address, which is the client's // 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 // own address for IPv4, and an address in the same place for an IPv6 /64.
// netblock. func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
func (g *GeoJS) LookUp(ctx context.Context, client netip.Prefix) Answer { country, asked := g.answerOrWait(ctx, client)
answer, asked := g.answerOrWait(ctx, client)
if asked == nil { if asked == nil {
return answer return country
} }
timer := time.NewTimer(g.timeout) timer := time.NewTimer(timeout)
defer timer.Stop() defer timer.Stop()
select { select {
@@ -196,7 +170,7 @@ func (g *GeoJS) LookUp(ctx context.Context, client netip.Prefix) Answer {
g.mu.Lock() g.mu.Lock()
defer g.mu.Unlock() defer g.mu.Unlock()
answer, found := g.kept(client) country, found := g.kept(client)
if !found { if !found {
g.metrics.GeoJSUnanswered.Inc() g.metrics.GeoJSUnanswered.Inc()
} }
@@ -206,15 +180,7 @@ func (g *GeoJS) LookUp(ctx context.Context, client netip.Prefix) Answer {
w.late = true w.late = true
} }
return answer return country
}
// Kept returns client's answer, if one is kept, without asking GeoJS.
func (g *GeoJS) Kept(client netip.Prefix) (Answer, bool) {
g.mu.Lock()
defer g.mu.Unlock()
return g.kept(client)
} }
// Snapshot returns every answer kept, sorted by client, as lookups.json // Snapshot returns every answer kept, sorted by client, as lookups.json
@@ -266,13 +232,13 @@ func (g *GeoJS) Load(answers []Answer) {
// nil when there is nothing to wait for. // nil when there is nothing to wait for.
func (g *GeoJS) answerOrWait( func (g *GeoJS) answerOrWait(
ctx context.Context, client netip.Prefix, ctx context.Context, client netip.Prefix,
) (Answer, <-chan struct{}) { ) (string, <-chan struct{}) {
g.mu.Lock() g.mu.Lock()
defer g.mu.Unlock() defer g.mu.Unlock()
answer, found := g.kept(client) country, found := g.kept(client)
if found { if found {
return answer, nil return country, nil
} }
w, waiting := g.waiting[client] w, waiting := g.waiting[client]
@@ -283,14 +249,10 @@ func (g *GeoJS) answerOrWait(
g.ask(ctx) g.ask(ctx)
if !g.wait {
return Answer{}, nil // the answer is not needed before the request goes on
}
if w == nil { if w == nil {
g.metrics.GeoJSUnanswered.Inc() g.metrics.GeoJSUnanswered.Inc()
return Answer{}, nil // too many clients wait already return "", nil // too many clients wait already
} }
if !g.asking { if !g.asking {
@@ -301,25 +263,25 @@ func (g *GeoJS) answerOrWait(
if w.late { if w.late {
g.metrics.GeoJSUnanswered.Inc() g.metrics.GeoJSUnanswered.Inc()
return Answer{}, nil return "", nil
} }
return Answer{}, w.asked return "", w.asked
} }
// kept returns client's answer, if GeoJS gave it less than keepFor ago, // kept returns client's answer, if GeoJS gave it less than keepFor ago,
// and notes that it was used. // and notes that it was used.
func (g *GeoJS) kept(client netip.Prefix) (Answer, bool) { func (g *GeoJS) kept(client netip.Prefix) (string, bool) {
now := g.now() now := g.now()
kept, found := g.answers.Get(client) kept, found := g.answers.Get(client)
if !found || now.Sub(kept.Answered) >= keepFor { if !found || now.Sub(kept.Answered) >= keepFor {
return Answer{}, false return "", false
} }
kept.Used = now kept.Used = now
return *kept, true return kept.Country, true
} }
// ask starts asking GeoJS about the waiting clients, unless a request to // ask starts asking GeoJS about the waiting clients, unless a request to
@@ -337,8 +299,7 @@ func (g *GeoJS) ask(ctx context.Context) {
} }
// askAboutWaiting asks GeoJS about the waiting clients, one request at a // askAboutWaiting asks GeoJS about the waiting clients, one request at a
// time, until none is left or GeoJS fails. Each answer kept is given to // time, until none is left or GeoJS fails.
// Answered, outside the lock, since Answered takes locks of its own.
func (g *GeoJS) askAboutWaiting(ctx context.Context) { func (g *GeoJS) askAboutWaiting(ctx context.Context) {
for { for {
clients := g.nextClients() clients := g.nextClients()
@@ -346,16 +307,8 @@ func (g *GeoJS) askAboutWaiting(ctx context.Context) {
return return
} }
given, err := g.request(ctx, clients) countries, err := g.request(ctx, clients)
kept, answered := g.keep(clients, given, err) if !g.keep(clients, countries, err) {
if g.answered != nil {
for _, answer := range kept {
g.answered(answer)
}
}
if !answered {
return return
} }
} }
@@ -387,35 +340,32 @@ func (g *GeoJS) nextClients() []netip.Prefix {
return clients return clients
} }
// keep notes how a request to GeoJS about clients ended, given being the // keep notes how a request to GeoJS about clients ended, and reports
// answer for each address GeoJS's answer names. It returns the answers it // whether GeoJS answered about all of them. Each client whose address
// kept, and reports whether GeoJS answered about all of the clients. Each // GeoJS's answer names gets its answer, with no country when GeoJS gave
// client whose address GeoJS's answer names gets its answer. An answer // none. An answer that leaves an address out is a failure. After a
// that leaves an address out is a failure. After a failure GeoJS is left // failure GeoJS is left alone for a while, and every client still waiting
// alone for a while, and every client still waiting stops waiting and is // stops waiting and is asked about once GeoJS is asked again.
// asked about once GeoJS is asked again.
func (g *GeoJS) keep( func (g *GeoJS) keep(
clients []netip.Prefix, given map[netip.Addr]Answer, err error, clients []netip.Prefix, countries map[netip.Addr]string, err error,
) ([]Answer, bool) { ) bool {
g.mu.Lock() g.mu.Lock()
defer g.mu.Unlock() defer g.mu.Unlock()
now := g.now() now := g.now()
kept := make([]Answer, 0, len(clients))
leftOut := 0 leftOut := 0
for _, client := range clients { for _, client := range clients {
answer, named := given[client.Addr()] country, named := countries[client.Addr()]
if !named { if !named {
leftOut++ leftOut++
continue continue
} }
answer.Client, answer.Answered, answer.Used = client, now, now g.answers.Add(client, &Answer{
g.answers.Add(client, &answer) Client: client, Country: country, Answered: now, Used: now,
kept = append(kept, answer) })
close(g.waiting[client].asked) close(g.waiting[client].asked)
delete(g.waiting, client) delete(g.waiting, client)
} }
@@ -450,28 +400,26 @@ func (g *GeoJS) keep(
}, },
}) })
return kept, false return false
} }
g.retryDelay = 0 g.retryDelay = 0
return kept, true return true
} }
// request asks GeoJS about clients in one request, and returns the answer // request asks GeoJS about clients in one request, and returns the
// for each address GeoJS's answer names: its AS number and the AS's name, // country it gave, in capitals, for each address its answer names.
// both "" for the AS number 64512, which GeoJS gives when it knows none,
// and its country, in capitals.
func (g *GeoJS) request( func (g *GeoJS) request(
ctx context.Context, clients []netip.Prefix, ctx context.Context, clients []netip.Prefix,
) (map[netip.Addr]Answer, error) { ) (map[netip.Addr]string, error) {
addrs := make([]string, 0, len(clients)) addrs := make([]string, 0, len(clients))
for _, client := range clients { for _, client := range clients {
addrs = append(addrs, client.Addr().String()) addrs = append(addrs, client.Addr().String())
} }
ctx, cancel := context.WithTimeout(ctx, g.timeout) ctx, cancel := context.WithTimeout(ctx, timeout)
defer cancel() defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, g.url, http.NoBody) req, err := http.NewRequestWithContext(ctx, http.MethodGet, g.url, http.NoBody)
@@ -498,12 +446,9 @@ func (g *GeoJS) request(
return nil, fmt.Errorf("%w %s", errStatus, res.Status) return nil, fmt.Errorf("%w %s", errStatus, res.Status)
} }
//nolint:tagliatelle // GeoJS's own names
var answers []struct { var answers []struct {
IP string `json:"ip"` IP string `json:"ip"`
ASN int64 `json:"asn"` Country string `json:"country"`
ASName string `json:"organization_name"`
CountryCode string `json:"country_code"`
} }
err = json.NewDecoder(io.LimitReader(res.Body, maxResponseBytes)).Decode(&answers) err = json.NewDecoder(io.LimitReader(res.Body, maxResponseBytes)).Decode(&answers)
@@ -511,22 +456,14 @@ func (g *GeoJS) request(
return nil, fmt.Errorf("read GeoJS's answer: %w", err) return nil, fmt.Errorf("read GeoJS's answer: %w", err)
} }
given := make(map[netip.Addr]Answer, len(answers)) countries := make(map[netip.Addr]string, len(answers))
for _, item := range answers { for _, item := range answers {
addr, err := netip.ParseAddr(item.IP) addr, err := netip.ParseAddr(item.IP)
if err != nil { if err == nil {
continue countries[addr] = strings.ToUpper(item.Country)
} }
answer := Answer{Country: strings.ToUpper(item.CountryCode)}
if item.ASN != 0 && item.ASN != unknownASN {
answer.ASN = "AS" + strconv.FormatInt(item.ASN, 10)
answer.ASName = item.ASName
}
given[addr] = answer
} }
return given, nil return countries, nil
} }
+16 -187
View File
@@ -23,13 +23,9 @@ import (
const ( const (
// germany is where the stand-in for GeoJS places every address but // germany is where the stand-in for GeoJS places every address but
// unplaced, and asNumber, kept as asn, and asName the AS it gives them. // unplaced.
germany = "DE" germany = "DE"
asNumber = 64496 // unplaced is the address it cannot place.
asn = "AS64496"
asName = "Example Net"
// unplaced is the address it cannot place, for which it gives the AS
// number 64512 and the AS name Unknown, as GeoJS does.
unplaced = "192.0.2.1" unplaced = "192.0.2.1"
// leftOut is the address it leaves out of its answer when // leftOut is the address it leaves out of its answer when
// answeringWithoutLeftOut. // answeringWithoutLeftOut.
@@ -87,7 +83,7 @@ func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
var earlier sync.WaitGroup var earlier sync.WaitGroup
earlier.Go(func() { g.LookUp(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) }) earlier.Go(func() { g.Country(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) })
defer earlier.Wait() defer earlier.Wait()
waitForRequests(t, geojs, 1) waitForRequests(t, geojs, 1)
@@ -119,58 +115,6 @@ 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) { func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
t.Parallel() t.Parallel()
@@ -243,96 +187,6 @@ func TestCountryIsKeptInCapitals(t *testing.T) {
}) })
} }
func TestAnswerHoldsTheASNumberTheASNameAndTheCountry(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
_, clock, g := start()
placed := netip.MustParsePrefix("203.0.113.9/32")
notPlaced := netip.MustParsePrefix(unplaced + "/32")
now := clock.Now()
// For the client it cannot place, GeoJS gives the AS number 64512
// and the AS name Unknown, which count as unknown.
for client, want := range map[netip.Prefix]lookup.Answer{
placed: {
Client: placed, ASN: asn, ASName: asName, Country: germany,
Answered: now, Used: now,
},
notPlaced: {Client: notPlaced, Answered: now, Used: now},
} {
got := g.LookUp(t.Context(), client)
if got != want {
t.Errorf("answer for %s\n%+v\nwant\n%+v", client, got, want)
}
}
})
}
func TestWithoutWaitTheRequestGoesOnAtOnceAndTheAnswerIsGivenWhenItComes(
t *testing.T,
) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
var (
mu sync.Mutex
given []lookup.Answer
)
geojs := &standIn{answers: answeringSlowly}
clock := newClock()
m := metrics.New(1, "app")
g := lookup.New(lookup.Params{
URL: lookup.URL,
Timeout: timeout,
Answered: func(answer lookup.Answer) {
mu.Lock()
defer mu.Unlock()
given = append(given, answer)
},
Now: clock.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: m,
Alerts: alerts.New(alerts.Params{}),
})
g.SetTransport(geojs)
client := netip.MustParsePrefix("203.0.113.9/32")
// The request goes on at once, without an answer, and GeoJS is asked
// about the client, which it answers most of a second later.
began := time.Now()
got := g.LookUp(t.Context(), client)
if took := time.Since(began); took != 0 || got != (lookup.Answer{}) {
t.Errorf("waited %s for %+v, want no wait and no answer", took, got)
}
waitForRequests(t, geojs, 1)
wantAsked(t, geojs, 0, "203.0.113.9")
time.Sleep(timeout)
synctest.Wait()
now := clock.Now()
want := lookup.Answer{
Client: client, ASN: asn, ASName: asName, Country: germany,
Answered: now, Used: now,
}
mu.Lock()
if !slices.Equal(given, []lookup.Answer{want}) {
t.Errorf("answers given %+v, want only %+v", given, want)
}
mu.Unlock()
wantCountry(t, g, client, germany)
wantUnanswered(t, m, 0)
})
}
func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) { func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
t.Parallel() t.Parallel()
@@ -343,11 +197,9 @@ func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
geojs := &standIn{answers: hanging} geojs := &standIn{answers: hanging}
g := lookup.New(lookup.Params{ g := lookup.New(lookup.Params{
URL: lookup.URL, URL: lookup.URL,
Timeout: timeout,
Wait: true,
Now: time.Now, Now: time.Now,
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)), ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
Metrics: metrics.New(1, "app"), Metrics: metrics.New(1),
Alerts: alerts.New(alerts.Params{}), Alerts: alerts.New(alerts.Params{}),
}) })
g.SetTransport(geojs) g.SetTransport(geojs)
@@ -391,7 +243,7 @@ func TestFailureRaisesASourceFailureAlertOncePerCooldown(t *testing.T) {
wantCountry(t, g, clients(), "") wantCountry(t, g, clients(), "")
wantRequests(t, geojs, 2) wantRequests(t, geojs, 2)
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook] waiting := queue.Snapshot().Waiting
if len(waiting) != 1 || !reflect.DeepEqual(waiting[0], want) { if len(waiting) != 1 || !reflect.DeepEqual(waiting[0], want) {
t.Errorf("alerts waiting %+v, want only %+v", waiting, want) t.Errorf("alerts waiting %+v, want only %+v", waiting, want)
} }
@@ -546,11 +398,9 @@ func TestClientsWithoutAnAnswerAreCounted(t *testing.T) {
t.Parallel() t.Parallel()
synctest.Test(t, func(t *testing.T) { synctest.Test(t, func(t *testing.T) {
m := metrics.New(1, "app") m := metrics.New(1)
g := lookup.New(lookup.Params{ g := lookup.New(lookup.Params{
URL: lookup.URL, URL: lookup.URL,
Timeout: timeout,
Wait: true,
Now: time.Now, Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler), ProcessLog: slog.New(slog.DiscardHandler),
Metrics: m, Metrics: m,
@@ -644,24 +494,21 @@ func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) {
} }
} }
list := make([]map[string]any, 0, len(addrs)) list := make([]map[string]string, 0, len(addrs))
for _, addr := range addrs { for _, addr := range addrs {
item := map[string]any{ country := germany
"ip": addr, "asn": asNumber, "organization_name": asName,
"country_code": germany,
}
switch { switch {
case addr == unplaced: case addr == unplaced:
item = map[string]any{"ip": addr, "asn": 64512, "organization_name": "Unknown"} country = ""
case addr == leftOut && answers == answeringWithoutLeftOut: case addr == leftOut && answers == answeringWithoutLeftOut:
continue continue
case answers == answeringInLowerCase: case answers == answeringInLowerCase:
item["country_code"] = strings.ToLower(germany) country = strings.ToLower(germany)
} }
list = append(list, item) list = append(list, map[string]string{"ip": addr, "country": country})
} }
var answer any = list var answer any = list
@@ -719,8 +566,7 @@ func (c *testClock) advance(d time.Duration) {
} }
// start returns a stand-in for GeoJS that answers, a clock, and a GeoJS // start returns a stand-in for GeoJS that answers, a clock, and a GeoJS
// asking the stand-in by that clock, for which a request waits for its // asking the stand-in by that clock.
// client's first answer.
func start() (*standIn, *testClock, *lookup.GeoJS) { func start() (*standIn, *testClock, *lookup.GeoJS) {
geojs, clock, g, _ := startWithAlerts() geojs, clock, g, _ := startWithAlerts()
@@ -732,7 +578,7 @@ func start() (*standIn, *testClock, *lookup.GeoJS) {
// cooldown, by the same clock. // cooldown, by the same clock.
func startWithAlerts() (*standIn, *testClock, *lookup.GeoJS, *alerts.Queue) { func startWithAlerts() (*standIn, *testClock, *lookup.GeoJS, *alerts.Queue) {
geojs := &standIn{} geojs := &standIn{}
clock := newClock() clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
queue := alerts.New(alerts.Params{ queue := alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"}, WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(), Events: alerts.Events(),
@@ -741,11 +587,9 @@ func startWithAlerts() (*standIn, *testClock, *lookup.GeoJS, *alerts.Queue) {
}) })
g := lookup.New(lookup.Params{ g := lookup.New(lookup.Params{
URL: lookup.URL, URL: lookup.URL,
Timeout: timeout,
Wait: true,
Now: clock.Now, Now: clock.Now,
ProcessLog: slog.New(slog.DiscardHandler), ProcessLog: slog.New(slog.DiscardHandler),
Metrics: metrics.New(1, "app"), Metrics: metrics.New(1),
Alerts: queue, Alerts: queue,
}) })
g.SetTransport(geojs) g.SetTransport(geojs)
@@ -753,11 +597,6 @@ func startWithAlerts() (*standIn, *testClock, *lookup.GeoJS, *alerts.Queue) {
return geojs, clock, g, queue return geojs, clock, g, queue
} }
// newClock returns a clock set to the start of a day.
func newClock() *testClock {
return &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
}
// newClients returns what returns a new IPv4 client each time it is // newClients returns what returns a new IPv4 client each time it is
// called. // called.
func newClients() func() netip.Prefix { func newClients() func() netip.Prefix {
@@ -774,7 +613,7 @@ func newClients() func() netip.Prefix {
func wantCountry(t *testing.T, g *lookup.GeoJS, client netip.Prefix, want string) { func wantCountry(t *testing.T, g *lookup.GeoJS, client netip.Prefix, want string) {
t.Helper() t.Helper()
got := g.LookUp(t.Context(), client).Country got := g.Country(t.Context(), client)
if got != want { if got != want {
t.Errorf("%s is in %q, want %q", client, got, want) t.Errorf("%s is in %q, want %q", client, got, want)
} }
@@ -819,16 +658,6 @@ 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, // waitForRequests waits until g has done all it can before time passes,
// checks that GeoJS has had count requests, and returns the addresses each // checks that GeoJS has had count requests, and returns the addresses each
// asked about. // asked about.
-83
View File
@@ -1,83 +0,0 @@
// 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)
}
}
+2 -5
View File
@@ -26,11 +26,8 @@ func TestSnapshotHoldsEachAnswerAndWhenItWasLastUsed(t *testing.T) {
wantCountry(t, g, placed, germany) wantCountry(t, g, placed, germany)
want := []lookup.Answer{ want := []lookup.Answer{
{Client: notPlaced, Answered: asked, Used: asked}, {Client: notPlaced, Country: "", Answered: asked, Used: asked},
{ {Client: placed, Country: germany, Answered: asked, Used: asked.Add(time.Hour)},
Client: placed, ASN: asn, ASName: asName, Country: germany,
Answered: asked, Used: asked.Add(time.Hour),
},
} }
if got := g.Snapshot(); !slices.Equal(got, want) { if got := g.Snapshot(); !slices.Equal(got, want) {
t.Errorf("snapshot\n%+v\nwant\n%+v", got, want) t.Errorf("snapshot\n%+v\nwant\n%+v", got, want)
-155
View File
@@ -1,155 +0,0 @@
package metrics
import (
"sync"
"github.com/prometheus/client_golang/prometheus"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// other is the label value under which the countries or AS numbers
// outside the busiest are counted.
const other = "other"
// busiest are the metrics by one thing the lookup finds of the client,
// its country or its AS number, for requests whose client's is known.
// The topN busiest countries or AS numbers, by their requests since the
// start, have series of their own, and the others are counted under
// other, so that there are never more than topN + 1 series. One that
// drops out of the busiest loses its series, and its next requests are
// counted under other; one that becomes one of them gets a series that
// counts from then on. Each series therefore only ever goes up.
type busiest struct {
topN int
requests *prometheus.CounterVec
requestBytes *prometheus.CounterVec
responseBytes *prometheus.CounterVec
// refused are, by country, the requests the country lists refused; nil
// by AS number.
refused *prometheus.CounterVec
mu sync.Mutex
// seen is each country's or AS number's requests since the start, by
// which they are ranked.
seen map[string]int64
// top are those with series of their own.
top map[string]bool
}
// newCountries returns the metrics by country, with series of their own
// for the topN busiest countries.
func newCountries(topN int) *busiest {
countries := newBusiest(topN, "country", "the client's country")
countries.refused = counterVec("smallwebwaf_country_list_refusals_total",
"Requests the country lists refused, by the client's country.",
[]string{"country"})
return countries
}
// newASNs returns the metrics by AS number, with series of their own for
// the topN busiest AS numbers.
func newASNs(topN int) *busiest {
return newBusiest(topN, "asn", "the client's AS number")
}
// newBusiest returns the metrics by label, which is described as
// description, with series of their own for the topN busiest values.
func newBusiest(topN int, label, description string) *busiest {
by := []string{label}
return &busiest{
topN: topN,
requests: counterVec("smallwebwaf_"+label+"_requests_total",
"Requests, by "+description+".", by),
requestBytes: counterVec("smallwebwaf_"+label+"_request_bytes_total",
"Request body bytes, by "+description+".", by),
responseBytes: counterVec("smallwebwaf_"+label+"_response_bytes_total",
"Response body bytes, by "+description+".", by),
seen: map[string]int64{},
top: map[string]bool{},
}
}
// Describe and Collect make the metrics a prometheus.Collector, so that
// they are registered together.
func (b *busiest) Describe(ch chan<- *prometheus.Desc) {
for _, vec := range b.vecs() {
vec.Describe(ch)
}
}
// Collect is the other half of prometheus.Collector, with Describe.
func (b *busiest) Collect(ch chan<- prometheus.Metric) {
for _, vec := range b.vecs() {
vec.Collect(ch)
}
}
// vecs returns the metrics: by AS number, those of requests and bytes; by
// country, the refusals by the country lists as well.
func (b *busiest) vecs() []*prometheus.CounterVec {
vecs := []*prometheus.CounterVec{b.requests, b.requestBytes, b.responseBytes}
if b.refused != nil {
vecs = append(vecs, b.refused)
}
return vecs
}
// add counts a request from its log line, whose client's country or AS
// number, value, is known.
func (b *busiest) add(value string, line *requestlog.Line) {
b.mu.Lock()
defer b.mu.Unlock()
b.seen[value]++
label := b.label(value)
b.requests.WithLabelValues(label).Inc()
b.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes))
b.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes))
if b.refused != nil && line.Action == requestlog.ActionCountryDenied {
b.refused.WithLabelValues(label).Inc()
}
}
// label returns the label a request from value is counted under: value
// while it is one of the busiest, other while it is not. A value busier
// than the least busy of them takes its place, and that one's series are
// dropped.
func (b *busiest) label(value string) string {
if b.top[value] {
return value
}
if len(b.top) < b.topN {
b.top[value] = true
return value
}
least := ""
for top := range b.top {
if least == "" || b.seen[top] < b.seen[least] {
least = top
}
}
if b.seen[value] <= b.seen[least] {
return other
}
delete(b.top, least)
for _, vec := range b.vecs() {
vec.DeleteLabelValues(least)
}
b.top[value] = true
return value
}
+116
View File
@@ -0,0 +1,116 @@
package metrics
import (
"sync"
"github.com/prometheus/client_golang/prometheus"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// other is the label under which the countries outside the busiest are
// counted.
const other = "other"
// countries are the metrics by the client's country, for requests whose
// client's country is known. The topN busiest countries, by their requests
// since the start, have series of their own, and the others are counted
// under other, so that there are never more than topN + 1 series. A
// country that drops out of the busiest loses its series, and its next
// requests are counted under other; one that becomes one of them gets a
// series that counts from then on. Each series therefore only ever goes
// up.
type countries struct {
topN int
requests *prometheus.CounterVec
requestBytes *prometheus.CounterVec
responseBytes *prometheus.CounterVec
// refused are the requests the country lists refused.
refused *prometheus.CounterVec
mu sync.Mutex
// seen is each country's requests since the start, by which the
// countries are ranked. GeoJS gives two-letter codes, so it holds at
// most a few hundred.
seen map[string]int64
// top are the countries with series of their own.
top map[string]bool
}
// newCountries returns the metrics by country, with series of their own
// for the topN busiest countries.
func newCountries(topN int) *countries {
byCountry := []string{"country"}
return &countries{
topN: topN,
requests: counterVec("smallwebwaf_country_requests_total",
"Requests, by the client's country.", byCountry),
requestBytes: counterVec("smallwebwaf_country_request_bytes_total",
"Request body bytes, by the client's country.", byCountry),
responseBytes: counterVec("smallwebwaf_country_response_bytes_total",
"Response body bytes, by the client's country.", byCountry),
refused: counterVec("smallwebwaf_country_list_refusals_total",
"Requests the country lists refused, by the client's country.",
byCountry),
seen: map[string]int64{},
top: map[string]bool{},
}
}
// add counts a request from its log line, whose country is known.
func (c *countries) add(line *requestlog.Line) {
c.mu.Lock()
defer c.mu.Unlock()
c.seen[line.Country]++
label := c.label(line.Country)
c.requests.WithLabelValues(label).Inc()
c.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes))
c.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes))
if line.Action == requestlog.ActionCountryDenied {
c.refused.WithLabelValues(label).Inc()
}
}
// label returns the label a request from country is counted under: the
// country while it is one of the busiest, other while it is not. A
// country busier than the least busy of them takes its place, and that
// country's series are dropped.
func (c *countries) label(country string) string {
if c.top[country] {
return country
}
if len(c.top) < c.topN {
c.top[country] = true
return country
}
least := ""
for top := range c.top {
if least == "" || c.seen[top] < c.seen[least] {
least = top
}
}
if c.seen[country] <= c.seen[least] {
return other
}
delete(c.top, least)
for _, vec := range []*prometheus.CounterVec{
c.requests, c.requestBytes, c.responseBytes, c.refused,
} {
vec.DeleteLabelValues(least)
}
c.top[country] = true
return country
}
+53 -200
View File
@@ -6,7 +6,6 @@ package metrics
import ( import (
"net/http" "net/http"
"strconv" "strconv"
"strings"
"time" "time"
"github.com/prometheus/client_golang/prometheus" "github.com/prometheus/client_golang/prometheus"
@@ -14,18 +13,15 @@ import (
"github.com/prometheus/client_golang/prometheus/promhttp" "github.com/prometheus/client_golang/prometheus/promhttp"
"sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/remotelog" "sneak.berlin/go/smallwebwaf/internal/remotelog"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
) )
// Metrics are smallwebwaf's metrics. They are safe for concurrent use. // Metrics are smallwebwaf's metrics. They are safe for concurrent use.
type Metrics struct { type Metrics struct {
// registry gives every metric registered with it the label instance. registry *prometheus.Registry
registry prometheus.Registerer
handler http.Handler handler http.Handler
inFlight prometheus.Gauge inFlight prometheus.Gauge
@@ -37,17 +33,14 @@ type Metrics struct {
rateLimitHits *prometheus.CounterVec rateLimitHits *prometheus.CounterVec
sizeAndTimeLimitHits *prometheus.CounterVec sizeAndTimeLimitHits *prometheus.CounterVec
offences *prometheus.CounterVec offences *prometheus.CounterVec
// ruleMatches are made by AddRules, and reputationHits by // ruleMatches are made by AddRules.
// AddReputation. ruleMatches *prometheus.CounterVec
ruleMatches *prometheus.CounterVec countries *countries
reputationHits *prometheus.CounterVec
countries *busiest
asns *busiest
// GeoJSRequests are the requests to GeoJS, and GeoJSFailures those // GeoJSRequests are the requests to GeoJS, and GeoJSFailures those
// that failed. GeoJSUnanswered are the requests that needed their // that failed. GeoJSUnanswered are the requests whose client counted
// client's answer, for a setting that acts on it, and went on without // as coming from an unknown country because GeoJS had not answered
// it because GeoJS had not given it in time. // about it in time.
GeoJSRequests prometheus.Counter GeoJSRequests prometheus.Counter
GeoJSFailures prometheus.Counter GeoJSFailures prometheus.Counter
GeoJSUnanswered prometheus.Counter GeoJSUnanswered prometheus.Counter
@@ -61,18 +54,14 @@ type Metrics struct {
} }
// New returns the metrics, with the Go runtime's and the process's own. // New returns the metrics, with the Go runtime's and the process's own.
// topN is how many countries and how many AS numbers get series of their // topN is how many countries get series of their own
// own (SWWAF_METRICS_TOP_N). Every metric carries instanceName // (SWWAF_METRICS_TOP_N).
// (SWWAF_INSTANCE_NAME) as its label instance. func New(topN int) *Metrics {
func New(topN int, instanceName string) *Metrics {
byStatus := []string{"status_class", "action"} byStatus := []string{"status_class", "action"}
byFile := []string{"file"} byFile := []string{"file"}
registry := prometheus.NewRegistry()
m := &Metrics{ m := &Metrics{
registry: prometheus.WrapRegistererWith( registry: prometheus.NewRegistry(),
prometheus.Labels{"instance": instanceName}, registry),
handler: promhttp.HandlerFor(registry, promhttp.HandlerOpts{}),
inFlight: prometheus.NewGauge(prometheus.GaugeOpts{ inFlight: prometheus.NewGauge(prometheus.GaugeOpts{
Name: "smallwebwaf_requests_in_flight", Name: "smallwebwaf_requests_in_flight",
Help: "Requests under way.", Help: "Requests under way.",
@@ -94,16 +83,14 @@ func New(topN int, instanceName string) *Metrics {
Help: "How long requests passed to the app took, from then to their end.", Help: "How long requests passed to the app took, from then to their end.",
}), }),
rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total", rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total",
"Requests that broke a rate limit or a byte limit, by its window and "+ "Requests that broke a rate limit, by its window.",
"its kind, requests or bytes.", []string{"window"}),
[]string{"window", "kind"}),
sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total", sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total",
"Requests that passed a size or time limit, by its setting.", "Requests that passed a size or time limit, by its setting.",
[]string{"limit"}), []string{"limit"}),
offences: counterVec("smallwebwaf_offences_total", offences: counterVec("smallwebwaf_offences_total",
"Offences, by kind.", []string{"kind"}), "Offences, by kind.", []string{"kind"}),
countries: newCountries(topN), countries: newCountries(topN),
asns: newASNs(topN),
GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{ GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_requests_total", Name: "smallwebwaf_geojs_requests_total",
Help: "Requests to GeoJS.", Help: "Requests to GeoJS.",
@@ -114,8 +101,8 @@ func New(topN int, instanceName string) *Metrics {
}), }),
GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{ GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_unanswered_total", Name: "smallwebwaf_geojs_unanswered_total",
Help: "Requests that needed their client's answer from GeoJS and " + Help: "Requests whose client counted as coming from an unknown " +
"went on without it, because GeoJS had not given it in time.", "country because GeoJS had not answered about it in time.",
}), }),
stateFileWrites: counterVec("smallwebwaf_state_file_writes_total", stateFileWrites: counterVec("smallwebwaf_state_file_writes_total",
"Writes of each state file.", byFile), "Writes of each state file.", byFile),
@@ -132,12 +119,16 @@ func New(topN int, instanceName string) *Metrics {
byFile), byFile),
} }
m.handler = promhttp.HandlerFor(m.registry, promhttp.HandlerOpts{})
m.registry.MustRegister( m.registry.MustRegister(
collectors.NewGoCollector(), collectors.NewGoCollector(),
collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}), collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}),
m.inFlight, m.requests, m.requestBytes, m.responseBytes, m.inFlight, m.requests, m.requestBytes, m.responseBytes,
m.requestDuration, m.upstreamDuration, m.requestDuration, m.upstreamDuration,
m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences, m.countries, m.asns, m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences,
m.countries.requests, m.countries.requestBytes, m.countries.responseBytes,
m.countries.refused,
m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered, m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered,
m.stateFileWrites, m.stateFileWriteFailures, m.stateFileWrites, m.stateFileWriteFailures,
m.stateFileLastWrite, m.stateFileSize, m.stateFileLastWrite, m.stateFileSize,
@@ -235,146 +226,46 @@ func (m *Metrics) AddRemoteLog(remote *remotelog.Sender) {
) )
} }
// AddLookupFile adds the metrics of the lookup database, read as the // AddAlerts adds the metrics of the alerts sent to
// metrics are asked for: when the file in use was read, which lastRead // SWWAF_ALERT_WEBHOOK_URL, read from queue as the metrics are asked for,
// returns, and the replacements of it that could not be read, which // with the destination webhook: the alerts sent, the requests to the
// readFailures returns. The lookup package's File, which has both, cannot // webhook that failed, and the alerts held back and dropped.
// be named here: that package counts GeoJS's requests in these metrics. func (m *Metrics) AddAlerts(queue *alerts.Queue) {
func (m *Metrics) AddLookupFile(lastRead func() time.Time, readFailures func() int) { webhook := prometheus.Labels{"destination": "webhook"}
m.registry.MustRegister( m.registry.MustRegister(
prometheus.NewGaugeFunc(prometheus.GaugeOpts{ prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_lookup_database_last_read_timestamp_seconds", Name: "smallwebwaf_alerts_sent_total",
Help: "When the lookup database in use was read, in seconds since 1970.", Help: "Alerts the destination took.",
ConstLabels: webhook,
}, func() float64 { }, func() float64 {
return float64(lastRead().Unix()) return float64(queue.Sent())
}), }),
prometheus.NewCounterFunc(prometheus.CounterOpts{ prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_lookup_database_read_failures_total", Name: "smallwebwaf_alerts_failed_total",
Help: "Replacements of the lookup database that could not be read.", Help: "Requests to the destination that failed.",
ConstLabels: webhook,
}, func() float64 { }, func() float64 {
return float64(readFailures()) return float64(queue.Failed())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_suppressed_total",
Help: "Alerts held back: repeats within SWWAF_ALERT_COOLDOWN, and " +
"alerts past SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary.",
ConstLabels: webhook,
}, func() float64 {
return float64(queue.Suppressed())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_dropped_total",
Help: "Alerts dropped, the oldest first, from a full queue.",
ConstLabels: webhook,
}, func() float64 {
return float64(queue.Dropped())
}), }),
) )
} }
// 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
// held back, which are the same for every destination, and those
// dropped. With no destination set, it adds none.
func (m *Metrics) AddAlerts(queue *alerts.Queue) {
for _, name := range queue.DestinationsSet() {
destination := prometheus.Labels{"destination": name}
m.registry.MustRegister(
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_sent_total",
Help: "Alerts the destination took.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Counts(name).Sent)
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_failed_total",
Help: "Requests to the destination that failed.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Counts(name).Failed)
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_suppressed_total",
Help: "Alerts held back: repeats within SWWAF_ALERT_COOLDOWN, and " +
"alerts past SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Suppressed())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_dropped_total",
Help: "Alerts dropped, the oldest first, from a full queue, and " +
"alerts given up as the destination refused them.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Counts(name).Dropped)
}),
)
}
}
// ServeHTTP answers with the metrics in the Prometheus text format. // ServeHTTP answers with the metrics in the Prometheus text format.
func (m *Metrics) ServeHTTP(w http.ResponseWriter, r *http.Request) { func (m *Metrics) ServeHTTP(w http.ResponseWriter, r *http.Request) {
m.handler.ServeHTTP(w, r) m.handler.ServeHTTP(w, r)
@@ -405,15 +296,7 @@ func (m *Metrics) RequestEnded(
} }
if line.LimitHit != "" { if line.LimitHit != "" {
// The log line names a byte limit's window with _bytes after it. m.rateLimitHits.WithLabelValues(line.LimitHit).Inc()
window, isBytes := strings.CutSuffix(line.LimitHit, "_bytes")
kind := ratelimit.KindRequests
if isBytes {
kind = ratelimit.KindBytes
}
m.rateLimitHits.WithLabelValues(window, kind).Inc()
} }
if limit != "" { if limit != "" {
@@ -425,11 +308,7 @@ func (m *Metrics) RequestEnded(
} }
if line.Country != "" { if line.Country != "" {
m.countries.add(line.Country, line) m.countries.add(line)
}
if line.ASN != "" {
m.asns.add(line.ASN, line)
} }
} }
@@ -470,32 +349,6 @@ func (m *Metrics) StateFileEditSetAside(name string) {
m.stateFileEditsSetAside.WithLabelValues(name).Inc() 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 // statusClass returns the class of status, such as 2xx, or none when no
// status was sent. // status was sent.
func statusClass(status int) string { func statusClass(status int) string {
+1 -1
View File
@@ -282,7 +282,7 @@ func (rq *request) showClient() {
answer := clientAnswer{Bans: state.BanEntries(rq.h.ledger.Covering(addr))} answer := clientAnswer{Bans: state.BanEntries(rq.h.ledger.Covering(addr))}
client, seen := rq.h.limiter.Client(rq.h.clientGroup(addr)) client, seen := rq.h.limiter.Client(clientGroup(addr))
if seen { if seen {
answer.Client = &client answer.Client = &client
} }
+3 -3
View File
@@ -4,7 +4,7 @@ import (
"encoding/json" "encoding/json"
"net/http" "net/http"
"net/netip" "net/netip"
"reflect" "slices"
"strconv" "strconv"
"strings" "strings"
"testing" "testing"
@@ -48,7 +48,7 @@ func TestAdminEndpointsAreOffWhileTheTokenIsUnset(t *testing.T) {
} }
} }
if after := server.Ledger.Snapshot(); !reflect.DeepEqual(after, before) { if after := server.Ledger.Snapshot(); !slices.Equal(after, before) {
t.Errorf("the bans are now\n%+v\nwant them unchanged\n%+v", 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(); !reflect.DeepEqual(after, before) { if after := server.Ledger.Snapshot(); !slices.Equal(after, before) {
t.Errorf("%s %s without the token changed the bans to\n%+v\nfrom\n%+v", t.Errorf("%s %s without the token changed the bans to\n%+v\nfrom\n%+v",
e.method, e.path, after, before) e.method, e.path, after, before)
} }
+27 -127
View File
@@ -16,7 +16,6 @@ import (
const ( const (
alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL" alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL"
alertMaxPerHour = "SWWAF_ALERT_MAX_PER_HOUR"
// alertInstance is the instance every alert of these tests gives. // alertInstance is the instance every alert of these tests gives.
alertInstance = "fsn1app1/gitea" alertInstance = "fsn1app1/gitea"
) )
@@ -40,10 +39,19 @@ func TestBanForABrokenLimitRaisesABanAlertWithItsNotes(t *testing.T) {
clk.advance(time.Minute) clk.advance(time.Minute)
s.get(client, http.StatusForbidden, requestlog.ActionBanned) s.get(client, http.StatusForbidden, requestlog.ActionBanned)
wantAlerts(t, queue, banAlert(alerts.EventBan, start, client, bans.Ban{ wantAlerts(t, queue, alerts.Alert{
Netblock: netblock, Cause: bans.CauseLimit, Instance: alertInstance,
Reason: "requests per minute over the limit of 1", Notes: ban.Notes, Time: start,
}, requestlog.FormatTime(start.Add(time.Hour)))) Event: alerts.EventBan,
Client: netip.MustParseAddr(client),
Netblock: netblock,
Reason: "requests per minute over the limit of 1",
Detail: map[string]any{
"cause": bans.CauseLimit,
"ban_expires": requestlog.FormatTime(start.Add(time.Hour)),
"notes": ban.Notes,
},
})
if ban.Notes.Limit != 1 || ban.Notes.Request.Path != "/" { if ban.Notes.Limit != 1 || ban.Notes.Request.Path != "/" {
t.Errorf("the alert's notes are %+v, want those of the broken limit", ban.Notes) t.Errorf("the alert's notes are %+v, want those of the broken limit", ban.Notes)
@@ -93,108 +101,20 @@ func TestAttackBanRaisesABanAlertThenAPermanentBanAlert(t *testing.T) {
) )
} }
func TestObserveModeRaisesTheBanAlertsItWouldHave(t *testing.T) { func TestObserveModeRaisesNoBanAlert(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
mode: observe,
rateLimitPerMinute: "2",
rulesDir: writeRules(t, testRules),
})
start := clk.Now()
// A ban for a clear sign of attack, which a request under it would make
// permanent.
group := netip.MustParsePrefix(ipv6Group)
attackBan, _ := server.Ledger.BanForAttack(group, start, bans.Notes{RuleID: "probe"})
// The third request breaks the limit, and so does the fourth, within the
// cooldown, which raises nothing. The probe is a clear sign of attack.
for range 4 {
s.get(client, http.StatusOK, requestlog.ActionForward)
}
s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
line := s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
// No ban is made, and none made permanent.
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)
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 3 || queue.Suppressed() != 0 {
t.Fatalf("%d alerts wait and %d are held back, want 3 and 0: %+v",
len(waiting), queue.Suppressed(), waiting)
}
limitNotes, _ := waiting[0].Detail["notes"].(bans.Notes)
attackNotes, _ := waiting[1].Detail["notes"].(bans.Notes)
if limitNotes.Limit != 2 || limitNotes.Request.Path != "/" ||
attackNotes.Request.Path != "/.env" {
t.Errorf("the notes are %+v and %+v, want those of the broken limit and "+
"of the probe", limitNotes, attackNotes)
}
// Each alert is the one enforce mode would have raised, with mode
// observe in its detail.
want := []alerts.Alert{
banAlert(alerts.EventBan, start, client, bans.Ban{
Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit,
Reason: "requests per minute over the limit of 2", Notes: limitNotes,
}, requestlog.FormatTime(start.Add(time.Hour))),
attackAlert(alerts.EventBan, start, otherClient, bans.Ban{
Netblock: netip.MustParsePrefix(otherClient + "/32"), Notes: attackNotes,
}, requestlog.FormatTime(start.Add(7*24*time.Hour))),
attackAlert(alerts.EventPermanentBan, start, ipv6Client, attackBan, permanent),
}
for _, alert := range want {
alert.Detail["mode"] = observe
}
wantAlerts(t, queue, want...)
}
func TestObserveModeWorksOutABanOnlyWhenItsAlertWouldBeSent(t *testing.T) {
t.Parallel() t.Parallel()
s, _, _, queue := startWithAlerts(t, map[string]string{ s, _, _, queue := startWithAlerts(t, map[string]string{
mode: observe, mode: "observe",
rateLimitPerMinute: "2", rateLimitPerMinute: "1",
rulesDir: writeRules(t, testRules), rulesDir: writeRules(t, testRules),
alertMaxPerHour: "2",
}) })
// The client's third request breaks the limit, and raises the first s.get(client, http.StatusOK, requestlog.ActionForward)
// alert of the hour. Its fourth is within the cooldown. s.get(client, http.StatusOK, requestlog.ActionForward)
for range 4 { s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusOK, requestlog.ActionForward)
}
// The other client's first probe raises the second. Its second probe is wantAlerts(t, queue)
// within the cooldown.
for range 2 {
s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
}
// The IPv6 client's third request breaks the limit past the two alerts
// an hour.
for range 3 {
s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
}
// Had the ban been worked out for any of the requests within the
// cooldown or past the two an hour, its alert would have been raised,
// held back and counted.
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 2 || queue.Suppressed() != 0 {
t.Errorf("%d alerts wait and %d are held back, want 2 and 0: %+v",
len(waiting), queue.Suppressed(), waiting)
}
} }
// startWithAlerts is startWithClock with alerts to a webhook, which is // startWithAlerts is startWithClock with alerts to a webhook, which is
@@ -204,16 +124,7 @@ func startWithAlerts(
) (*sender, *clock, *proxy.Server, *alerts.Queue) { ) (*sender, *clock, *proxy.Server, *alerts.Queue) {
t.Helper() t.Helper()
return startAppWithAlerts(t, func(http.ResponseWriter, *http.Request) {}, env) app := startApp(t, func(http.ResponseWriter, *http.Request) {})
}
// 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)} clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
settings := map[string]string{ settings := map[string]string{
trustedProxies: trustLocalhost, trustedProxies: trustLocalhost,
@@ -227,10 +138,10 @@ func startAppWithAlerts(
return &sender{t: t, addr: addr, out: out}, clk, server, queue return &sender{t: t, addr: addr, out: out}, clk, server, queue
} }
// banAlert returns the alert for event, raised by a request from client at // attackAlert returns the alert for event, raised by a request from client
// the time raised, for ban, with its netblock, cause, reason and notes, // at the time raised, for ban, a ban for the probe rule of testRules,
// which ends at expires, as the log line gives it. // which ends at expires, as the log line gives it.
func banAlert( func attackAlert(
event string, raised time.Time, client string, ban bans.Ban, expires string, event string, raised time.Time, client string, ban bans.Ban, expires string,
) alerts.Alert { ) alerts.Alert {
return alerts.Alert{ return alerts.Alert{
@@ -239,29 +150,18 @@ func banAlert(
Event: event, Event: event,
Client: netip.MustParseAddr(client), Client: netip.MustParseAddr(client),
Netblock: ban.Netblock, Netblock: ban.Netblock,
Reason: ban.Reason, Reason: "matched the rule probe",
Detail: map[string]any{ Detail: map[string]any{
"cause": ban.Cause, "ban_expires": expires, "notes": ban.Notes, "cause": bans.CauseAttack, "ban_expires": expires, "notes": ban.Notes,
}, },
} }
} }
// attackAlert is banAlert for a ban for the probe rule of testRules, with
// the netblock and the notes of ban.
func attackAlert(
event string, raised time.Time, client string, ban bans.Ban, expires string,
) alerts.Alert {
return banAlert(event, raised, client, bans.Ban{
Netblock: ban.Netblock, Cause: bans.CauseAttack, Reason: "matched the rule probe",
Notes: ban.Notes,
}, expires)
}
// wantAlerts checks the alerts waiting in queue, in order. // wantAlerts checks the alerts waiting in queue, in order.
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) { func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
t.Helper() t.Helper()
got := queue.Snapshot().Waiting[alerts.DestinationWebhook] got := queue.Snapshot().Waiting
if len(got) != len(want) { if len(got) != len(want) {
t.Fatalf("%d alerts wait, want %d: %+v", len(got), len(want), got) t.Fatalf("%d alerts wait, want %d: %+v", len(got), len(want), got)
} }
-30
View File
@@ -1,30 +0,0 @@
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()
}
-377
View File
@@ -1,377 +0,0 @@
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)
}
}
+52 -178
View File
@@ -6,7 +6,6 @@ import (
"sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
) )
@@ -19,182 +18,83 @@ func (rq *request) banResponse(action string) *refusal {
// banned reports whether a ban on a netblock the client is in covers the // banned reports whether a ban on a netblock the client is in covers the
// request at now, and notes for the log line when that ban ends. A // request at now, and notes for the log line when that ban ends. A
// request that makes the ban permanent, or in observe mode would have, // request that makes the ban permanent raises the alert for it.
// raises the alert for it.
func (rq *request) banned(now time.Time) bool { func (rq *request) banned(now time.Time) bool {
check := rq.h.ledger.Check var (
if rq.h.config.Observe { ban bans.Ban
check = rq.h.ledger.Find // the ban refuses nothing, and stays as it is banned bool
} madePermanent bool
)
ban, banned, madePermanent := check(rq.client, now) if rq.h.config.Observe {
if banned { ban, banned = rq.h.ledger.Find(rq.client, now) // the ban refuses nothing
rq.line.BanExpires = banExpires(ban) } else {
ban, banned, madePermanent = rq.h.ledger.Check(rq.client, now)
} }
if madePermanent { if madePermanent {
ban.Expires = time.Time{} // the ban made permanent, which Find leaves as it is
rq.alertBan(ban) rq.alertBan(ban)
} }
if banned {
rq.line.BanExpires = banExpires(ban)
}
return banned return banned
} }
// limitBroken counts the request for the rate limits at now, notes the // 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 // client's counts for the log line, and reports whether the request takes
// the client over a rate limit, as its limit percentage lowers it, which // the client over a limit. In enforce mode such a request bans the
// breaks it. // client's netblock, and sets the client's counters back to zero; in
// observe mode it does neither.
func (rq *request) limitBroken(now time.Time) bool { func (rq *request) limitBroken(now time.Time) bool {
counts, hit, over := rq.h.limiter.Count(rq.h.clientGroup(rq.client), now, group := clientGroup(rq.client)
rq.limitPercent.percent)
counts, hit, over := rq.h.limiter.Count(group, now)
rq.line.Counts = counts rq.line.Counts = counts
if over { if !over {
rq.banForLimit(now, hit, rq.h.config.BanResponse) return false
} }
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 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 rq.line.Offence = requestlog.OffenceLimit
netblock := rq.h.netblock(rq.client)
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseLimit) {
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.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 { if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForLimit(netblock, now, notes) return true
if wouldBan {
rq.alertBan(ban)
}
return
} }
ban, made := rq.h.ledger.BanForLimit(netblock, now, notes) netblock := rq.h.netblock(rq.client)
rq.h.limiter.Reset(rq.h.clientGroup(rq.client)) ban, made := rq.h.ledger.BanForLimit(netblock, now, bans.Notes{
Country: rq.line.Country,
Limit: hit.Limit,
Window: hit.Window,
Count: hit.Requests,
Request: rq.noted(now),
Requests: rq.netblockRequests(netblock),
})
rq.h.limiter.Reset(group)
rq.line.BanExpires = banExpires(ban) rq.line.BanExpires = banExpires(ban)
if made { if made {
rq.alertBan(ban) rq.alertBan(ban)
} }
return true
} }
// banForAttack bans the client's netblock at now for a clear sign of // banForAttack bans the client's netblock at now for a clear sign of
// attack, the match of rule, a ban rule. In observe mode it makes no ban, // attack, the match of rule, a ban rule.
// and raises the alert for the ban it would have made, if that alert
// would be sent.
func (rq *request) banForAttack(now time.Time, rule rules.Rule) { func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
netblock := rq.h.netblock(rq.client) netblock := rq.h.netblock(rq.client)
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseAttack) { ban, made := rq.h.ledger.BanForAttack(netblock, now, bans.Notes{
return Country: rq.line.Country,
} RuleID: rule.ID,
Target: rule.Target,
notes := bans.Notes{ Request: rq.noted(now),
ASN: rq.line.ASN, Requests: rq.netblockRequests(netblock),
ASName: rq.line.ASName, })
Country: rq.line.Country,
RuleID: rule.ID,
Target: rule.Target,
Reputation: rq.reputation,
Request: rq.noted(now, rq.h.config.BanResponse),
Requests: rq.netblockRequests(netblock),
}
if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForAttack(netblock, now, notes)
if wouldBan {
rq.alertBan(ban)
}
return
}
ban, made := rq.h.ledger.BanForAttack(netblock, now, notes)
rq.line.BanExpires = banExpires(ban) rq.line.BanExpires = banExpires(ban)
if made { if made {
@@ -202,62 +102,36 @@ func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
} }
} }
// wouldAlertBan reports whether the alert for a ban on netblock for cause
// made at now would be sent. In observe mode the ban the request would
// have made is worked out only then, at most once per
// SWWAF_ALERT_COOLDOWN and never with no webhook set: its notes count the
// netblock's requests, which can mean going through every client.
func (rq *request) wouldAlertBan(
netblock netip.Prefix, now time.Time, cause string,
) bool {
event := alerts.EventBan
if rq.h.ledger.WouldBePermanent(netblock, now, cause) {
event = alerts.EventPermanentBan
}
return rq.h.alerts.WouldSend(event, netblock)
}
// alertBan raises the alert for ban, which the request made, or made // alertBan raises the alert for ban, which the request made, or made
// permanent: permanent_ban for a permanent ban, ban for another. Its // permanent: permanent_ban for a permanent ban, ban for another. Its
// detail gives the ban's cause, when it ends, and its notes, and in // detail gives the ban's cause, when it ends, and its notes.
// observe mode, where ban is the ban that would have been made, or made
// permanent, mode, observe.
func (rq *request) alertBan(ban bans.Ban) { func (rq *request) alertBan(ban bans.Ban) {
event := alerts.EventBan event := alerts.EventBan
if ban.Permanent() { if ban.Permanent() {
event = alerts.EventPermanentBan event = alerts.EventPermanentBan
} }
detail := map[string]any{
"cause": ban.Cause, "ban_expires": banExpires(ban), "notes": ban.Notes,
}
if rq.h.config.Observe {
detail["mode"] = "observe"
}
rq.h.alerts.Raise(alerts.Alert{ rq.h.alerts.Raise(alerts.Alert{
Event: event, Event: event,
Client: rq.client, Client: rq.client,
Netblock: ban.Netblock, Netblock: ban.Netblock,
ASN: ban.Notes.ASN,
ASName: ban.Notes.ASName,
Country: ban.Notes.Country, Country: ban.Notes.Country,
Reason: ban.Reason, Reason: ban.Reason,
Detail: detail, Detail: map[string]any{
"cause": ban.Cause, "ban_expires": banExpires(ban), "notes": ban.Notes,
},
}) })
} }
// noted is the request, at now, with status, what the client was sent, or // noted is the request, refused at now with SWWAF_BAN_RESPONSE, as the
// in observe mode would have been, as the notes of the ban it makes keep // notes of the ban it makes keep it.
// it. func (rq *request) noted(now time.Time) bans.Request {
func (rq *request) noted(now time.Time, status int) bans.Request {
return bans.Request{ return bans.Request{
Time: now, Time: now,
Method: rq.in.Method, Method: rq.in.Method,
Host: rq.in.Host, Host: rq.in.Host,
Path: rq.in.URL.RequestURI(), Path: rq.in.URL.RequestURI(),
Status: status, Status: rq.h.config.BanResponse,
UserAgent: rq.in.UserAgent(), UserAgent: rq.in.UserAgent(),
} }
} }
@@ -278,7 +152,7 @@ func (h *handler) netblock(client netip.Addr) netip.Prefix {
return netip.PrefixFrom(addr, h.config.BanScopeV4Prefix).Masked() return netip.PrefixFrom(addr, h.config.BanScopeV4Prefix).Masked()
} }
return h.clientGroup(addr) return clientGroup(addr)
} }
// banExpires is when ban ends, as the log line gives it: a time, or // banExpires is when ban ends, as the log line gives it: a time, or
+4 -14
View File
@@ -7,7 +7,6 @@ import (
"maps" "maps"
"net/http" "net/http"
"net/netip" "net/netip"
"reflect"
"slices" "slices"
"sync" "sync"
"testing" "testing"
@@ -166,14 +165,9 @@ func TestBanCoversTheClientsNetblock(t *testing.T) {
[]string{otherClient, exempt}, []string{"203.0.112.9", allowed}, []string{otherClient, exempt}, []string{"203.0.112.9", allowed},
}, },
{ {
"an IPv6 /64, by default", nil, "2001:db8:5::1", "an IPv6 /64", nil, "2001:db8:5::1",
[]string{"2001:db8:5::ffff:1"}, []string{"2001:db8:5:1::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.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -287,10 +281,7 @@ func TestBanNotes(t *testing.T) {
Cause: bans.CauseLimit, Cause: bans.CauseLimit,
Reason: "requests per minute over the limit of 1", Reason: "requests per minute over the limit of 1",
Notes: bans.Notes{ Notes: bans.Notes{
ASN: asnDE,
ASName: asNameDE,
Country: "DE", Country: "DE",
Kind: "requests",
Limit: 1, Limit: 1,
Window: minute, Window: minute,
Count: 2, Count: 2,
@@ -313,7 +304,7 @@ func TestBanNotes(t *testing.T) {
ledger := server.Ledger ledger := server.Ledger
got := ledger.Bans(netblock) got := ledger.Bans(netblock)
if len(got) != 1 || !reflect.DeepEqual(got[0], want) { if len(got) != 1 || got[0] != want {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want) t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
} }
@@ -371,9 +362,8 @@ func (c *clock) advance(d time.Duration) {
// startWithClock starts smallwebwaf in front of an app that answers 200, // startWithClock starts smallwebwaf in front of an app that answers 200,
// with the settings in env on top of trusting localhost's // with the settings in env on top of trusting localhost's
// X-Forwarded-For, clients' AS numbers and countries looked up at // X-Forwarded-For, clients' countries looked up at geojsURL, and a clock
// geojsURL, and a clock set to midnight, the start of a bucket in every // set to midnight, the start of a bucket in every window.
// window.
func startWithClock( func startWithClock(
t *testing.T, geojsURL string, env map[string]string, t *testing.T, geojsURL string, env map[string]string,
) (*sender, *clock, *proxy.Server) { ) (*sender, *clock, *proxy.Server) {
-119
View File
@@ -1,119 +0,0 @@
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
}
-504
View File
@@ -1,504 +0,0 @@
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)
}
}
-42
View File
@@ -103,48 +103,6 @@ func (b *responseBody) Close() error {
return b.body.Close() 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 // limitBody returns body, cut off with an *http.MaxBytesError after
// maxBytes, or unchanged if maxBytes is zero, which is off. // maxBytes, or unchanged if maxBytes is zero, which is off.
func limitBody(body io.ReadCloser, maxBytes int64) io.ReadCloser { func limitBody(body io.ReadCloser, maxBytes int64) io.ReadCloser {
-584
View File
@@ -1,584 +0,0 @@
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
}
+7 -5
View File
@@ -76,14 +76,16 @@ func scheme(r *http.Request, peerTrusted bool) string {
return proto 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 // clientGroup is the client a request is counted toward: its IPv4
// address, or its IPv6 group, the netblock its IPv6 address is in of the // address, or the /64 its IPv6 address is in, since one abuser usually
// length SWWAF_IPV6_GROUP_PREFIX sets, a /64 by default, since one abuser // holds a whole /64. An IPv4 address in IPv6 form counts as IPv4.
// usually holds a whole /64. An IPv4 address in IPv6 form counts as IPv4. func clientGroup(addr netip.Addr) netip.Prefix {
func (h *handler) clientGroup(addr netip.Addr) netip.Prefix {
addr = addr.Unmap() addr = addr.Unmap()
if addr.Is6() { if addr.Is6() {
return netip.PrefixFrom(addr, h.config.IPv6GroupPrefix).Masked() return netip.PrefixFrom(addr, ipv6GroupPrefix).Masked()
} }
return netip.PrefixFrom(addr, addr.BitLen()) return netip.PrefixFrom(addr, addr.BitLen())
+26 -6
View File
@@ -1,17 +1,31 @@
package proxy package proxy
import ( import (
"context"
"net/netip"
"slices" "slices"
) )
// countryDenied reports whether the country lists refuse the request, by // countryDenied reports whether the country lists refuse the request.
// the client's country as it was looked up. A client without a country, // The client's country is looked up only while a list is set, and never
// or whose country cannot be found, is refused only by // for a client on a private, loopback or link-local address, which has
// SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES. // no country. A client without a country, or whose country cannot be
func (rq *request) countryDenied() bool { // found, is refused only by SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES. ctx is
// the request's own context.
func (rq *request) countryDenied(ctx context.Context) bool {
denied := rq.h.config.DeniedCountries denied := rq.h.config.DeniedCountries
allowed := rq.h.config.ExclusivelyAllowedCountries allowed := rq.h.config.ExclusivelyAllowedCountries
country := rq.line.Country
if len(denied) == 0 && len(allowed) == 0 {
return false
}
var country string
if hasCountry(rq.client) {
country = rq.h.geojs.Country(ctx, clientGroup(rq.client))
}
rq.line.Country = country
if slices.Contains(denied, country) { if slices.Contains(denied, country) {
return true return true
@@ -19,3 +33,9 @@ func (rq *request) countryDenied() bool {
return len(allowed) > 0 && !slices.Contains(allowed, country) return len(allowed) > 0 && !slices.Contains(allowed, country)
} }
// hasCountry reports whether addr can be placed in a country: private,
// loopback and link-local addresses cannot.
func hasCountry(addr netip.Addr) bool {
return !addr.IsPrivate() && !addr.IsLoopback() && !addr.IsLinkLocalUnicast()
}
+34 -99
View File
@@ -56,21 +56,16 @@ func TestCountryLists(t *testing.T) {
maps.Copy(env, tc.env) maps.Copy(env, tc.env)
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env) addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env)
// The AS number GeoJS gives unplaced, 64512, counts as unknown. for i, sent := range []struct{ client, country string }{
for i, sent := range []struct{ client, asn, asName, country string }{ {fromDE, "DE"}, {fromKP, "KP"}, {unplaced, ""},
{fromDE, asnDE, asNameDE, "DE"}, {fromKP, asnKP, asNameKP, "KP"},
{unplaced, "", "", ""},
} { } {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, sent.client) req.Header.Set(forwardedFor, sent.client)
got := do(t, req) got := do(t, req)
line := out.requestLines(t, i+1)[i] line := out.requestLines(t, i+1)[i]
if line.ASN != sent.asn || line.ASName != sent.asName || if line.Country != sent.country {
line.Country != sent.country { t.Errorf("log line has country %q, want %q", line.Country, sent.country)
t.Errorf("log line has %q, %q and %q, want %q, %q and %q",
line.ASN, line.ASName, line.Country,
sent.asn, sent.asName, sent.country)
} }
if slices.Contains(tc.refused, sent.client) { if slices.Contains(tc.refused, sent.client) {
@@ -135,7 +130,7 @@ func TestRequestRefusedByCountryIsNotCounted(t *testing.T) {
return return
} }
answer := []geojsAnswer{{IP: r.URL.Query().Get("ip"), CountryCode: "DE"}} answer := []map[string]string{{"ip": r.URL.Query().Get("ip"), "country": "DE"}}
err := json.NewEncoder(w).Encode(answer) err := json.NewEncoder(w).Encode(answer)
if err != nil { if err != nil {
@@ -179,15 +174,20 @@ func TestRequestRefusedByCountryIsNotCounted(t *testing.T) {
wantStatus(t, got, http.StatusOK) wantStatus(t, got, http.StatusOK)
} }
func TestPrivateAddressIsNeverLookedUp(t *testing.T) { func TestCountryNotLookedUpWithoutAListOrForAPrivateAddress(t *testing.T) {
t.Parallel() t.Parallel()
for _, tc := range []struct { for _, tc := range []struct {
name string name string
env map[string]string env map[string]string
clients []string // "" sends no X-Forwarded-For: the client is 127.0.0.1
}{ }{
{"no setting needs the lookup", nil}, {"no country list is set", nil, []string{fromKP, fromDE}},
{"a country list is set", map[string]string{deniedCountries: "kp"}}, {
"private, loopback and link-local addresses",
map[string]string{deniedCountries: "kp"},
[]string{"10.0.0.5", "192.168.1.9", "fd00::5", "", "169.254.0.9", "fe80::9"},
},
} { } {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -198,10 +198,7 @@ func TestPrivateAddressIsNeverLookedUp(t *testing.T) {
maps.Copy(env, tc.env) maps.Copy(env, tc.env)
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env) addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env)
// "" sends no X-Forwarded-For: the client is 127.0.0.1. for i, sent := range tc.clients {
for i, sent := range []string{
"10.0.0.5", "192.168.1.9", "fd00::5", "", "169.254.0.9", "fe80::9",
} {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
if sent != "" { if sent != "" {
req.Header.Set(forwardedFor, sent) req.Header.Set(forwardedFor, sent)
@@ -212,26 +209,15 @@ func TestPrivateAddressIsNeverLookedUp(t *testing.T) {
line := out.requestLines(t, i+1)[i] line := out.requestLines(t, i+1)[i]
wantLine(t, line, http.StatusOK, requestlog.ActionForward) wantLine(t, line, http.StatusOK, requestlog.ActionForward)
for _, field := range []string{"asn", "as_name", "country"} { country, present := line.fields["country"]
value, present := line.fields[field] if !present || country != "" {
if !present || value != "" { t.Errorf("log line for %q has country %v, want an empty one",
t.Errorf("log line for %q has %s %v, want an empty one", line.ClientIP, country)
line.ClientIP, field, value)
}
} }
} }
// GeoJS is asked about up to 200 waiting clients at once, so once it if len(asked()) != 0 {
// has been asked about fromDE, which comes last, it has been asked t.Errorf("GeoJS was asked about %v, want nothing", asked())
// about every client before it that waited for an answer.
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, fromDE)
wantStatus(t, do(t, req), http.StatusOK)
waitUntil(func() bool { return slices.Contains(asked(), fromDE) })
if got := asked(); !slices.Equal(got, []string{fromDE}) {
t.Errorf("GeoJS was asked about %v, want %s alone", got, fromDE)
} }
}) })
} }
@@ -278,31 +264,18 @@ func TestExclusiveListRefusesAPrivateAddressUnlessAllowed(t *testing.T) {
} }
} }
// startGeoJS starts a stand-in for GeoJS, which places fromDE and fromKP, // startGeoJS starts a stand-in for GeoJS, which places fromDE and fromKP
// each in an AS of its own, and no other address. It returns its URL, and // and no other address. It returns its URL, and what returns the
// what returns the addresses it has been asked about. // addresses it has been asked about.
func startGeoJS(t *testing.T) (string, func() []string) { func startGeoJS(t *testing.T) (string, func() []string) {
t.Helper() t.Helper()
geojsURL, asked, release := startHeldGeoJS(t) places := map[string]string{fromDE: "DE", fromKP: "KP"}
release()
return geojsURL, asked var asked struct {
} mu sync.Mutex
addrs []string
// startHeldGeoJS is startGeoJS for a stand-in that answers nothing until }
// release is called. Each request to it waits until then.
func startHeldGeoJS(t *testing.T) (string, func() []string, func()) {
t.Helper()
var (
asked struct {
mu sync.Mutex
addrs []string
}
released = make(chan struct{})
once sync.Once
)
geojs := httptest.NewServer(http.HandlerFunc( geojs := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) { func(w http.ResponseWriter, r *http.Request) {
@@ -312,11 +285,11 @@ func startHeldGeoJS(t *testing.T) (string, func() []string, func()) {
asked.addrs = append(asked.addrs, addrs...) asked.addrs = append(asked.addrs, addrs...)
asked.mu.Unlock() asked.mu.Unlock()
<-released answers := make([]map[string]string, 0, len(addrs))
answers := make([]geojsAnswer, 0, len(addrs))
for _, addr := range addrs { for _, addr := range addrs {
answers = append(answers, answerAbout(addr)) answers = append(answers, map[string]string{
"ip": addr, "country": places[addr],
})
} }
err := json.NewEncoder(w).Encode(answers) err := json.NewEncoder(w).Encode(answers)
@@ -326,48 +299,10 @@ func startHeldGeoJS(t *testing.T) (string, func() []string, func()) {
})) }))
t.Cleanup(geojs.Close) t.Cleanup(geojs.Close)
release := func() { once.Do(func() { close(released) }) }
// Run before geojs.Close, which waits for every request to be answered.
t.Cleanup(release)
return geojs.URL, func() []string { return geojs.URL, func() []string {
asked.mu.Lock() asked.mu.Lock()
defer asked.mu.Unlock() defer asked.mu.Unlock()
return slices.Clone(asked.addrs) return slices.Clone(asked.addrs)
}, release
}
// The AS numbers and names the stand-in for GeoJS gives fromDE and
// fromKP, as they are logged.
const (
asnDE = "AS64496"
asNameDE = "Example Net"
asnKP = "AS64511"
asNameKP = "Other Net"
)
// geojsAnswer is an answer of GeoJS about one address, with the fields
// smallwebwaf reads.
//
//nolint:tagliatelle // GeoJS's own names
type geojsAnswer struct {
IP string `json:"ip"`
ASN int `json:"asn"`
ASName string `json:"organization_name"`
CountryCode string `json:"country_code,omitempty"`
}
// answerAbout is what the stand-in for GeoJS answers about addr: for an
// address it cannot place, the AS number 64512 and the AS name Unknown
// with no country, as GeoJS does.
func answerAbout(addr string) geojsAnswer {
switch addr {
case fromDE:
return geojsAnswer{IP: addr, ASN: 64496, ASName: asNameDE, CountryCode: "DE"}
case fromKP:
return geojsAnswer{IP: addr, ASN: 64511, ASName: asNameKP, CountryCode: "KP"}
} }
return geojsAnswer{IP: addr, ASN: 64512, ASName: "Unknown"}
} }
+2 -24
View File
@@ -24,9 +24,7 @@ func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
start := clk.Now() start := clk.Now()
// Two let through, one over the limit, which bans the client, and one // Two let through, one over the limit, which bans the client, and one
// refused under that ban, for which the client is not looked up. GeoJS // refused under that ban, for which the country is not looked up.
// answers about the client at its first request, and its later ones
// use that answer.
s.get(fromDE, http.StatusOK, requestlog.ActionForward) s.get(fromDE, http.StatusOK, requestlog.ActionForward)
clk.advance(time.Second) clk.advance(time.Second)
s.get(fromDE, http.StatusOK, requestlog.ActionForward) s.get(fromDE, http.StatusOK, requestlog.ActionForward)
@@ -37,10 +35,8 @@ func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
want := ratelimit.History{ want := ratelimit.History{
FirstSeen: start, FirstSeen: start,
LastSeen: start.Add(2 * time.Second), LastSeen: start.Add(2 * time.Second),
ASN: asnDE,
ASName: asNameDE,
Country: "DE", Country: "DE",
LookedUp: start, LookedUp: start.Add(time.Second),
Requests: 4, Requests: 4,
Forwarded: 2, Forwarded: 2,
Refused: 2, Refused: 2,
@@ -56,24 +52,6 @@ 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) { func TestHistoryCountsTheBodiesEachWay(t *testing.T) {
t.Parallel() t.Parallel()
-71
View File
@@ -1,71 +0,0 @@
package proxy
import (
"context"
"net/http"
"net/netip"
"sneak.berlin/go/smallwebwaf/internal/lookup"
)
// The headers in which the app is passed the client's AS number and
// 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, 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
}
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))
}
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)
h.ledger.AddLookup(h.netblock(answer.Client.Addr()),
answer.ASN, answer.ASName, answer.Country)
}
// setLookupHeaders sets the headers in which the app is passed the
// client's AS number and country, leaving out one that is unknown.
func setLookupHeaders(header http.Header, asn, country string) {
if asn != "" {
header.Set(asnHeader, asn)
}
if country != "" {
header.Set(countryHeader, country)
}
}
// canBePlaced reports whether a lookup can place addr: private, loopback
// and link-local addresses have no AS number or country.
func canBePlaced(addr netip.Addr) bool {
return !addr.IsPrivate() && !addr.IsLoopback() && !addr.IsLinkLocalUnicast()
}
-378
View File
@@ -1,378 +0,0 @@
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"
)
// asnAndCountry is what a lookup gives a client: its AS number, AS name
// 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()
// The stand-in for GeoJS answers only once released. A request that
// waited for it would wait an hour, and get no answer within
// waitLimit.
geojsURL, asked, release := startHeldGeoJS(t)
s, _, server := startWithClock(t, geojsURL, map[string]string{
lookupTimeout: "1h",
rateLimitPerMinute: "1",
})
// fromDE's second request breaks the limit and bans it, and fromKP
// comes too. None waits for GeoJS.
for _, line := range []logLine{
s.get(fromDE, http.StatusOK, requestlog.ActionForward),
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited),
s.get(fromKP, http.StatusOK, requestlog.ActionForward),
} {
got := asnAndCountry{line.ASN, line.ASName, line.Country}
if got != (asnAndCountry{}) {
t.Errorf("log line has %+v before GeoJS answered, want nothing", got)
}
}
// Once GeoJS answers, each answer reaches the client's history, and
// fromDE's reaches the notes of its ban.
release()
netblock := netip.MustParsePrefix(fromDE + "/32")
waitUntil(func() bool {
return historyOf(t, server, fromDE).ASN != "" &&
historyOf(t, server, fromKP).ASN != "" &&
server.Ledger.Bans(netblock)[0].Notes.ASN != ""
})
de := asnAndCountry{asnDE, asNameDE, "DE"}
for addr, want := range map[string]asnAndCountry{
fromDE: de, fromKP: {asnKP, asNameKP, "KP"},
} {
h := historyOf(t, server, addr)
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != want {
t.Errorf("%s's history has %+v, want %+v", addr, got, want)
}
}
notes := server.Ledger.Bans(netblock)[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)
}
// GeoJS was asked about each client once, fromKP after fromDE, whose
// request was under way when fromKP came.
if got := asked(); !slices.Equal(got, []string{fromDE, fromKP}) {
t.Errorf("GeoJS was asked about %v, want %s and %s", got, fromDE, fromKP)
}
}
func TestASNumberAndNameInTheLogLineTheHistoryTheBanNotesAndTheAlert(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
addr, out, server, queue := startProxyWithAlerts(t, app.URL, geojsURL, clk.Now,
map[string]string{
trustedProxies: trustLocalhost,
alertWebhookURL: "https://alerts.example/smallwebwaf",
rateLimitPerMinute: "1",
})
s := &sender{t: t, addr: addr, out: out}
// The answer is kept before the requests, so GeoJS is not asked, and
// gives no answer of its own.
netblock := netip.MustParsePrefix(fromDE + "/32")
server.GeoJS.Load([]lookup.Answer{{
Client: netblock, ASN: asnDE, ASName: asNameDE, Country: "DE",
Answered: clk.Now(), Used: clk.Now(),
}})
want := asnAndCountry{asnDE, asNameDE, "DE"}
for _, line := range []logLine{
s.get(fromDE, http.StatusOK, requestlog.ActionForward),
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited),
} {
if got := (asnAndCountry{line.ASN, line.ASName, line.Country}); got != want {
t.Errorf("log line has %+v, want %+v", got, want)
}
}
h := historyOf(t, server, fromDE)
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != want ||
!h.LookedUp.Equal(clk.Now()) {
t.Errorf("history has %+v, looked up at %s; want %+v, at %s",
got, h.LookedUp, want, clk.Now())
}
notes := server.Ledger.Bans(netblock)[0].Notes
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != want {
t.Errorf("the ban's notes have %+v, want %+v", got, want)
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 {
t.Fatalf("alerts waiting %+v, want the ban's alone", waiting)
}
alert := waiting[0]
if got := (asnAndCountry{alert.ASN, alert.ASName, alert.Country}); got != want {
t.Errorf("the ban's alert has %+v, want %+v", got, want)
}
}
func TestLookupSourceOffLooksNoClientUp(t *testing.T) {
t.Parallel()
geojsURL, asked := startGeoJS(t)
s, clk, server := startWithClock(t, geojsURL, map[string]string{lookupSource: "off"})
// Even an answer kept from before is not used.
server.GeoJS.Load([]lookup.Answer{{
Client: netip.MustParsePrefix(fromDE + "/32"), ASN: asnDE, ASName: asNameDE,
Country: "DE", Answered: clk.Now(), Used: clk.Now(),
}})
for _, from := range []string{fromDE, fromKP} {
line := s.get(from, http.StatusOK, requestlog.ActionForward)
got := asnAndCountry{line.ASN, line.ASName, line.Country}
if got != (asnAndCountry{}) {
t.Errorf("log line has %+v, want nothing", got)
}
}
if h := historyOf(t, server, fromDE); h.ASN != "" || !h.LookedUp.IsZero() {
t.Errorf("history has %q, looked up at %s, want no lookup", h.ASN, h.LookedUp)
}
if len(asked()) != 0 {
t.Errorf("GeoJS was asked about %v, want nothing", asked())
}
}
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()
var (
mu sync.Mutex
got [][2][]string // each request's X-Client-ASN and X-Client-Country
)
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
mu.Lock()
defer mu.Unlock()
got = append(got, [2][]string{
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,
addLookupHeaders: "true",
})
s := &sender{t: t, addr: addr, out: out}
// 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.
for _, from := range []string{fromDE, unplaced, "10.0.0.8"} {
s.requestWithHeader(from, "/", clientsOwnLookupHeaders,
http.StatusOK, requestlog.ActionForward)
}
mu.Lock()
defer mu.Unlock()
want := [][2][]string{{{asnDE}, {"DE"}}, {nil, nil}, {nil, nil}}
if !slices.EqualFunc(got, want, func(a, b [2][]string) bool {
return slices.Equal(a[0], b[0]) && slices.Equal(a[1], b[1])
}) {
t.Errorf("the app was passed %v, want %v", got, want)
}
}
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)
for !done() && time.Now().Before(deadline) {
time.Sleep(pollInterval)
}
}
+45 -106
View File
@@ -148,8 +148,8 @@ func TestMetricsCountTheTraffic(t *testing.T) {
wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound) wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound)
out.requestLines(t, 2) out.requestLines(t, 2)
forward := `{action="forward",instance="app",status_class="2xx"}` forward := `{action="forward",status_class="2xx"}`
notFound := `{action="admin",instance="app",status_class="4xx"}` notFound := `{action="admin",status_class="4xx"}`
// The request for the metrics is itself under way. // The request for the metrics is itself under way.
metrics := scrape(t, addr) metrics := scrape(t, addr)
@@ -159,13 +159,11 @@ func TestMetricsCountTheTraffic(t *testing.T) {
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+forward, 5) wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+forward, 5)
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+notFound, wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+notFound,
float64(len("Not Found\n"))) float64(len("Not Found\n")))
wantMetric(t, metrics, wantMetric(t, metrics, "smallwebwaf_request_duration_seconds_count", 2)
`smallwebwaf_request_duration_seconds_count{instance="app"}`, 2) wantMetric(t, metrics, "smallwebwaf_upstream_duration_seconds_count", 1)
wantMetric(t, metrics, wantMetric(t, metrics, "smallwebwaf_requests_in_flight", 1)
`smallwebwaf_upstream_duration_seconds_count{instance="app"}`, 1) metric(t, metrics, "go_goroutines")
wantMetric(t, metrics, `smallwebwaf_requests_in_flight{instance="app"}`, 1) metric(t, metrics, "process_start_time_seconds")
metric(t, metrics, `go_goroutines{instance="app"}`)
metric(t, metrics, `process_start_time_seconds{instance="app"}`)
// A request the app holds is under way until it ends. // A request the app holds is under way until it ends.
httpClient := newClient(t) httpClient := newClient(t)
@@ -182,7 +180,7 @@ func TestMetricsCountTheTraffic(t *testing.T) {
}() }()
<-arrived <-arrived
wantMetric(t, scrape(t, addr), `smallwebwaf_requests_in_flight{instance="app"}`, 2) wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 2)
releaseApp() releaseApp()
err := <-ended err := <-ended
@@ -191,7 +189,7 @@ func TestMetricsCountTheTraffic(t *testing.T) {
} }
out.requestLines(t, 5) out.requestLines(t, 5)
wantMetric(t, scrape(t, addr), `smallwebwaf_requests_in_flight{instance="app"}`, 1) wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 1)
} }
func TestMetricsCountLimitsAndBans(t *testing.T) { func TestMetricsCountLimitsAndBans(t *testing.T) {
@@ -221,16 +219,15 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
metrics := s.scrape(scraper) metrics := s.scrape(scraper)
wantMetric(t, metrics, wantMetric(t, metrics,
`smallwebwaf_requests_total{action="denied",instance="app",status_class="none"}`, 1) `smallwebwaf_requests_total{action="denied",status_class="none"}`, 1)
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+ wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 1)
`kind="requests",window="minute"}`, 1) wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 1)
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 1) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_permanent_bans", 0)
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 0)
clk.advance(time.Hour) clk.advance(time.Hour)
wantMetric(t, s.scrape(scraper), `smallwebwaf_active_bans{instance="app"}`, 0) wantMetric(t, s.scrape(scraper), "smallwebwaf_active_bans", 0)
// A limit broken again right after would ban for three hours, longer // A limit broken again right after would ban for three hours, longer
// than SWWAF_MAX_BAN_DURATION, so the ban is permanent. // than SWWAF_MAX_BAN_DURATION, so the ban is permanent.
@@ -238,14 +235,13 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
s.get(client, 0, requestlog.ActionRateLimited) s.get(client, 0, requestlog.ActionRateLimited)
metrics = s.scrape(scraper) metrics = s.scrape(scraper)
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+ wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 2)
`kind="requests",window="minute"}`, 2) wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 2)
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 2) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 2)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 2) wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 1)
// denied, client, and the scraper as of its earlier requests. // denied, client, and the scraper as of its earlier requests.
wantMetric(t, metrics, `smallwebwaf_tracked_clients{instance="app"}`, 3) wantMetric(t, metrics, "smallwebwaf_tracked_clients", 3)
} }
func TestMetricsCountTheBansAnAdminMakes(t *testing.T) { func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
@@ -259,7 +255,7 @@ func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
rateLimitExemptNets: scraper, rateLimitExemptNets: scraper,
}) })
const admins = `smallwebwaf_bans_made_total{cause="admin",instance="app"}` const admins = `smallwebwaf_bans_made_total{cause="admin"}`
wantMetric(t, s.scrape(scraper), admins, 0) wantMetric(t, s.scrape(scraper), admins, 0)
@@ -291,8 +287,7 @@ func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
metricsTopN: "2", metricsTopN: "2",
deniedCountries: "kp", deniedCountries: "kp",
} }
geojsURL, _ := startGeoJS(t) addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, env)
addr, out, server := startProxyWithClock(t, app.URL, geojsURL, time.Now, env)
// The answers are kept before the requests, so that none waits for // The answers are kept before the requests, so that none waits for
// GeoJS. // GeoJS.
@@ -324,82 +319,28 @@ func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
metrics := scrape(t, addr) metrics := scrape(t, addr)
lines++ lines++
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
`smallwebwaf_country_requests_total{country="KP",instance="app"}`, 3) wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`, 2)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 1)
`smallwebwaf_country_requests_total{country="DE",instance="app"}`, 2) wantMetric(t, metrics, `smallwebwaf_country_list_refusals_total{country="KP"}`, 3)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="KP"}`, 0)
`smallwebwaf_country_requests_total{country="other",instance="app"}`, 1) wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`, 6)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="KP"}`,
`smallwebwaf_country_list_refusals_total{country="KP",instance="app"}`, 3)
wantMetric(t, metrics,
`smallwebwaf_country_request_bytes_total{country="KP",instance="app"}`, 0)
wantMetric(t, metrics,
`smallwebwaf_country_request_bytes_total{country="DE",instance="app"}`, 6)
wantMetric(t, metrics,
`smallwebwaf_country_response_bytes_total{country="KP",instance="app"}`,
float64(3*len("Forbidden\n"))) float64(3*len("Forbidden\n")))
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="other"}`,
`smallwebwaf_country_response_bytes_total{country="other",instance="app"}`,
float64(len("hello"))) float64(len("hello")))
wantNoSeries(t, metrics, wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`)
`smallwebwaf_country_requests_total{country="FR",instance="app"}`)
// Once FR is busier than DE, it takes DE's place: its series counts // Once FR is busier than DE, it takes DE's place: its series counts
// from then on, and DE's is gone. // from then on, and DE's is gone.
send(fromFR, 3, http.StatusOK) send(fromFR, 3, http.StatusOK)
metrics = scrape(t, addr) metrics = scrape(t, addr)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
`smallwebwaf_country_requests_total{country="KP",instance="app"}`, 3) wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`, 2)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 2)
`smallwebwaf_country_requests_total{country="FR",instance="app"}`, 2) wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`)
wantMetric(t, metrics, wantNoSeries(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`)
`smallwebwaf_country_requests_total{country="other",instance="app"}`, 2)
wantNoSeries(t, metrics,
`smallwebwaf_country_requests_total{country="DE",instance="app"}`)
wantNoSeries(t, metrics,
`smallwebwaf_country_request_bytes_total{country="DE",instance="app"}`)
}
func TestMetricsByASNumberKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
s, clk, server := startWithClock(t, geojsURL, map[string]string{
metricsToken: token,
metricsTopN: "1",
})
// The answers are kept before the requests, so that GeoJS gives none
// of its own. Each client is in an AS of its own.
answer := func(addr, asn string) lookup.Answer {
return lookup.Answer{
Client: netip.MustParsePrefix(addr + "/32"), ASN: asn,
Answered: clk.Now(), Used: clk.Now(),
}
}
server.GeoJS.Load([]lookup.Answer{
answer(fromDE, "AS64501"), answer(fromKP, "AS64502"),
})
// With one AS number of its own, the other is counted as other. The
// metrics are asked for from a private address, which has no AS number.
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromKP, http.StatusOK, requestlog.ActionForward)
metrics := s.scrape("10.0.0.9")
wantMetric(t, metrics,
`smallwebwaf_asn_requests_total{asn="AS64501",instance="app"}`, 2)
wantMetric(t, metrics,
`smallwebwaf_asn_requests_total{asn="other",instance="app"}`, 1)
wantMetric(t, metrics,
`smallwebwaf_asn_request_bytes_total{asn="AS64501",instance="app"}`, 0)
wantMetric(t, metrics,
`smallwebwaf_asn_response_bytes_total{asn="other",instance="app"}`, 0)
wantNoSeries(t, metrics,
`smallwebwaf_asn_requests_total{asn="AS64502",instance="app"}`)
} }
func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) { func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) {
@@ -429,16 +370,16 @@ func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) {
deadline := time.Now().Add(waitLimit) deadline := time.Now().Add(waitLimit)
metrics := scrape(t, addr) metrics := scrape(t, addr)
for metric(t, metrics, `smallwebwaf_geojs_failures_total{instance="app"}`) == 0 && for metric(t, metrics, "smallwebwaf_geojs_failures_total") == 0 &&
time.Now().Before(deadline) { time.Now().Before(deadline) {
time.Sleep(pollInterval) time.Sleep(pollInterval)
metrics = scrape(t, addr) metrics = scrape(t, addr)
} }
wantMetric(t, metrics, `smallwebwaf_geojs_requests_total{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_geojs_requests_total", 1)
wantMetric(t, metrics, `smallwebwaf_geojs_failures_total{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_geojs_failures_total", 1)
wantMetric(t, metrics, `smallwebwaf_geojs_unanswered_total{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_geojs_unanswered_total", 1)
} }
// keptAnswer returns GeoJS's answer that the client at addr is in // keptAnswer returns GeoJS's answer that the client at addr is in
@@ -481,9 +422,8 @@ func (s *sender) scrape(from string) string {
// metric returns the value of series in metrics, which are in the // metric returns the value of series in metrics, which are in the
// Prometheus text format. series is a name and its labels in the order of // Prometheus text format. series is a name and its labels in the order of
// their names, such as // their names, such as smallwebwaf_offences_total{kind="limit"}. It fails
// smallwebwaf_offences_total{instance="app",kind="limit"}. It fails the // the test if there is no such series.
// test if there is no such series.
func metric(t *testing.T, metrics, series string) float64 { func metric(t *testing.T, metrics, series string) float64 {
t.Helper() t.Helper()
@@ -531,8 +471,7 @@ func wantNoSeries(t *testing.T, metrics, series string) {
func wantLimitHits(t *testing.T, addr, limit string, hits int) { func wantLimitHits(t *testing.T, addr, limit string, hits int) {
t.Helper() t.Helper()
series := `smallwebwaf_size_and_time_limit_hits_total{instance="app",limit="` + series := `smallwebwaf_size_and_time_limit_hits_total{limit="` + limit + `"}`
limit + `"}`
metrics := scrape(t, addr) metrics := scrape(t, addr)
if hits == 0 { if hits == 0 {
+1 -2
View File
@@ -5,7 +5,6 @@ import (
"io" "io"
"net/http" "net/http"
"net/netip" "net/netip"
"reflect"
"sync/atomic" "sync/atomic"
"testing" "testing"
"time" "time"
@@ -127,7 +126,7 @@ func TestObserveModeMakesNoBanAndKeepsTheBansItHas(t *testing.T) {
} }
got := server.Ledger.Snapshot() got := server.Ledger.Snapshot()
if len(got) != 1 || !reflect.DeepEqual(got[0], kept) { if len(got) != 1 || got[0] != kept {
t.Errorf("bans\n%+v\nwant only\n%+v", got, kept) t.Errorf("bans\n%+v\nwant only\n%+v", got, kept)
} }
} }
+5 -26
View File
@@ -6,6 +6,7 @@ import (
"errors" "errors"
"io" "io"
"net/http" "net/http"
"os"
"reflect" "reflect"
"slices" "slices"
"strings" "strings"
@@ -122,10 +123,10 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) { func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
t.Helper() t.Helper()
bytes := float64(sent + received) hostname, _ := os.Hostname()
want := withTimings(line, requestlog.Line{ want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: "app", Type: requestType, Time: line.Time, Instance: hostname,
ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host, ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
Path: rawPath, Query: rawQuery, Protocol: protocol, Path: rawPath, Query: rawQuery, Protocol: protocol,
Status: http.StatusTeapot, RequestBytes: int64(sent), Status: http.StatusTeapot, RequestBytes: int64(sent),
@@ -133,10 +134,7 @@ func wantRequestFields(t *testing.T, line logLine, host string, sent, received i
RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32", RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8", ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward, UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
Counts: ratelimit.Counts{ Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
Minute: 1, Hour: 1, Day: 1,
MinuteBytes: bytes, HourBytes: bytes, DayBytes: bytes,
},
}) })
if !reflect.DeepEqual(line.Line, want) { if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want) t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
@@ -320,7 +318,7 @@ func TestServerHasTheDefaultLimits(t *testing.T) {
server := proxy.New(proxy.Params{ server := proxy.New(proxy.Params{
Config: cfg, Config: cfg,
RequestLog: io.Discard, RequestLog: io.Discard,
ProcessLog: requestlog.NewProcessLogger(io.Discard, cfg.InstanceName, cfg.LogLevel), ProcessLog: requestlog.NewProcessLogger(io.Discard),
}) })
if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 || if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 ||
@@ -399,25 +397,6 @@ 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) { func TestLogsAnAnswerThatBrokeOff(t *testing.T) {
t.Parallel() t.Parallel()
+25 -120
View File
@@ -12,13 +12,11 @@ import (
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
) )
@@ -56,15 +54,9 @@ type Params struct {
RequestLog io.Writer RequestLog io.Writer
// ProcessLog receives the process's own messages. // ProcessLog receives the process's own messages.
ProcessLog *slog.Logger ProcessLog *slog.Logger
// GeoJSURL is where clients' AS numbers and countries are looked up // GeoJSURL is where clients' countries are looked up, normally
// while SWWAF_LOOKUP_SOURCE is geojs, normally lookup.URL. // lookup.URL. GeoJS is asked only while a country list is set.
GeoJSURL string 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 // 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, // limits, bans are made and run out, and GeoJS's answers are kept,
// normally time.Now in UTC, the time the state files give. // normally time.Now in UTC, the time the state files give.
@@ -73,30 +65,19 @@ type Params struct {
// against. // against.
Rules *rules.Files Rules *rules.Files
// Alerts receive the alert for each ban the proxy makes or makes // Alerts receive the alert for each ban the proxy makes or makes
// permanent, for each count over an anomaly threshold, for each request // permanent, and for GeoJS failing.
// 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 Alerts *alerts.Queue
} }
// Server is the server smallwebwaf runs, with the parts of the proxy // Server is the server smallwebwaf runs, with the parts of the proxy
// whose state the state files keep, the lookup database, nil unless // whose state the state files keep, and the metrics.
// 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 { type Server struct {
*http.Server *http.Server
Ledger *bans.Ledger Ledger *bans.Ledger
Limiter *ratelimit.Limiter Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS GeoJS *lookup.GeoJS
Anomalies *anomaly.Counters Metrics *metrics.Metrics
LookupFile *lookup.File
Lists *reputation.Lists
DNSBL *reputation.DNSBL
AbuseIPDB *reputation.AbuseIPDB
Metrics *metrics.Metrics
} }
// New returns the server smallwebwaf runs: each request it reads passes // New returns the server smallwebwaf runs: each request it reads passes
@@ -107,8 +88,7 @@ type Server struct {
// applies the timeouts and size limits from then on. // applies the timeouts and size limits from then on.
func New(params Params) *Server { func New(params Params) *Server {
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn) errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
m := metrics.New(params.Config.MetricsTopN, params.Config.InstanceName) m := metrics.New(params.Config.MetricsTopN)
lists, dnsbl, abuseIPDB := newReputation(params, m)
h := &handler{ h := &handler{
config: params.Config, config: params.Config,
requestLog: params.RequestLog, requestLog: params.RequestLog,
@@ -118,13 +98,10 @@ func New(params Params) *Server {
now: params.Now, now: params.Now,
metrics: m, metrics: m,
limiter: ratelimit.New(ratelimit.Limits{ limiter: ratelimit.New(ratelimit.Limits{
PerMinute: params.Config.RateLimitPerMinute, PerMinute: params.Config.RateLimitPerMinute,
PerHour: params.Config.RateLimitPerHour, PerHour: params.Config.RateLimitPerHour,
PerDay: params.Config.RateLimitPerDay, PerDay: params.Config.RateLimitPerDay,
BytesPerMinute: params.Config.BytesLimitPerMinute, }),
BytesPerHour: params.Config.BytesLimitPerHour,
BytesPerDay: params.Config.BytesLimitPerDay,
}, params.Config.MaxTrackedClients),
ledger: bans.New(bans.Rules{ ledger: bans.New(bans.Rules{
LimitBanDuration: params.Config.LimitBanDuration, LimitBanDuration: params.Config.LimitBanDuration,
LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow, LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow,
@@ -132,38 +109,16 @@ func New(params Params) *Server {
AttackBanDuration: params.Config.AttackBanDuration, AttackBanDuration: params.Config.AttackBanDuration,
MaxBans: params.Config.MaxBans, MaxBans: params.Config.MaxBans,
}), }),
anomalies: anomaly.New(anomaly.Params{ geojs: lookup.New(lookup.Params{
Client: params.Config.AnomalyClient, URL: params.GeoJSURL,
Net: params.Config.AnomalyNet, Now: params.Now,
ASN: params.Config.AnomalyASN, ProcessLog: params.ProcessLog,
Total: params.Config.AnomalyTotal, Metrics: m,
Watch: params.Config.AnomalyWatch, Alerts: params.Alerts,
NetV4Prefix: params.Config.AnomalyNetV4Prefix,
NetV6Prefix: params.Config.AnomalyNetV6Prefix,
NamedNetblocks: params.Config.WatchNets,
Alerts: params.Alerts,
}), }),
lookupFile: params.LookupFile, rules: params.Rules,
lists: lists, alerts: params.Alerts,
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, 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 || biasedThresholdsSet(params.Config),
Answered: h.addLookup,
Now: params.Now,
ProcessLog: params.ProcessLog,
Metrics: m,
Alerts: params.Alerts,
})
m.AddBansAndClients(h.ledger, h.limiter, params.Now) m.AddBansAndClients(h.ledger, h.limiter, params.Now)
m.AddRules(params.Rules) m.AddRules(params.Rules)
@@ -180,52 +135,13 @@ func New(params Params) *Server {
MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10), MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
ErrorLog: errorLog, ErrorLog: errorLog,
}, },
Ledger: h.ledger, Ledger: h.ledger,
Limiter: h.limiter, Limiter: h.limiter,
GeoJS: h.geojs, GeoJS: h.geojs,
Anomalies: h.anomalies, Metrics: m,
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 // handler is the proxy. It holds what every request shares; what belongs
// to one request is in a request. // to one request is in a request.
type handler struct { type handler struct {
@@ -239,11 +155,6 @@ type handler struct {
limiter *ratelimit.Limiter limiter *ratelimit.Limiter
ledger *bans.Ledger ledger *bans.Ledger
geojs *lookup.GeoJS geojs *lookup.GeoJS
anomalies *anomaly.Counters
lookupFile *lookup.File
lists *reputation.Lists
dnsbl *reputation.DNSBL
abuseIPDB *reputation.AbuseIPDB
rules *rules.Files rules *rules.Files
alerts *alerts.Queue alerts *alerts.Queue
} }
@@ -282,7 +193,6 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
// Once the request has ended, before its log line is written. // Once the request has ended, before its log line is written.
defer rq.addToHistory() defer rq.addToHistory()
defer rq.countAnomalies()
refused := rq.check(r.Context()) refused := rq.check(r.Context())
rq.checked = time.Now() rq.checked = time.Now()
@@ -301,10 +211,5 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return 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()) rq.forward(r.Context())
} }
+28 -74
View File
@@ -16,7 +16,6 @@ import (
"sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
@@ -61,18 +60,12 @@ const (
requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES" requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES"
responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES" responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES"
trustedProxies = "SWWAF_TRUSTED_PROXIES" trustedProxies = "SWWAF_TRUSTED_PROXIES"
ipv6GroupPrefix = "SWWAF_IPV6_GROUP_PREFIX"
maxTrackedClients = "SWWAF_MAX_TRACKED_CLIENTS"
allowNets = "SWWAF_ALLOW_NETS" allowNets = "SWWAF_ALLOW_NETS"
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS" rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
denyNets = "SWWAF_DENY_NETS" denyNets = "SWWAF_DENY_NETS"
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE" rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS" 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" deniedCountries = "SWWAF_DENIED_COUNTRIES"
allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES" allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES"
banResponse = "SWWAF_BAN_RESPONSE" banResponse = "SWWAF_BAN_RESPONSE"
@@ -210,8 +203,8 @@ func startProxy(t *testing.T, appURL string, env map[string]string) (string, *ou
return startProxyWithGeoJS(t, appURL, "", env) return startProxyWithGeoJS(t, appURL, "", env)
} }
// startProxyWithGeoJS is startProxy with clients' AS numbers and // startProxyWithGeoJS is startProxy with clients' countries looked up at
// countries looked up at geojsURL. // geojsURL.
func startProxyWithGeoJS( func startProxyWithGeoJS(
t *testing.T, appURL, geojsURL string, env map[string]string, t *testing.T, appURL, geojsURL string, env map[string]string,
) (string, *output) { ) (string, *output) {
@@ -224,9 +217,7 @@ func startProxyWithGeoJS(
// startProxyWithClock is startProxyWithGeoJS with requests counted and // startProxyWithClock is startProxyWithGeoJS with requests counted and
// bans made by the time now tells, and returns the server as well. Unless // bans made by the time now tells, and returns the server as well. Unless
// env sets SWWAF_RULES_DIR, it is an empty directory, of no rules, and // env sets SWWAF_RULES_DIR, it is an empty directory, of no rules.
// unless it sets SWWAF_INSTANCE_NAME, that is app, the label instance of
// every metric.
func startProxyWithClock( func startProxyWithClock(
t *testing.T, appURL, geojsURL string, now func() time.Time, t *testing.T, appURL, geojsURL string, now func() time.Time,
env map[string]string, env map[string]string,
@@ -239,52 +230,15 @@ func startProxyWithClock(
} }
// startProxyWithAlerts is startProxyWithClock, and returns the queue of // startProxyWithAlerts is startProxyWithClock, and returns the queue of
// the alerts the proxy raises as well, as newProxy makes them. // 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.
func startProxyWithAlerts( func startProxyWithAlerts(
t *testing.T, appURL, geojsURL string, now func() time.Time, t *testing.T, appURL, geojsURL string, now func() time.Time,
env map[string]string, env map[string]string,
) (string, *output, *proxy.Server, *alerts.Queue) { ) (string, *output, *proxy.Server, *alerts.Queue) {
t.Helper() t.Helper()
server, out, alertQueue := newProxy(t, appURL, geojsURL, now, env) settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir()}
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",
}
if geojsURL == "" {
settings[lookupSource] = "off"
}
maps.Copy(settings, env) maps.Copy(settings, env)
cfg, err := config.FromEnvironment(func(name string) (string, bool) { cfg, err := config.FromEnvironment(func(name string) (string, bool) {
@@ -297,7 +251,7 @@ func newProxy(
} }
out := &output{} out := &output{}
processLog := requestlog.NewProcessLogger(out, cfg.InstanceName, cfg.LogLevel) processLog := requestlog.NewProcessLogger(out)
ruleFiles, err := rules.Load(rules.Params{ ruleFiles, err := rules.Load(rules.Params{
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog, Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
@@ -316,30 +270,30 @@ func newProxy(
ProcessLog: processLog, 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{ server := proxy.New(proxy.Params{
Config: cfg, Config: cfg,
RequestLog: out, RequestLog: out,
ProcessLog: processLog, ProcessLog: processLog,
GeoJSURL: geojsURL, GeoJSURL: geojsURL,
AbuseIPDBURL: abuseIPDBURL, Now: now,
LookupFile: lookupFile, Rules: ruleFiles,
Now: now, Alerts: alertQueue,
Rules: ruleFiles,
Alerts: alertQueue,
}) })
return server, out, 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
} }
// newClient returns an HTTP client that sends requests as they are made, // newClient returns an HTTP client that sends requests as they are made,
+1 -45
View File
@@ -72,56 +72,12 @@ 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) { func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
t.Parallel() t.Parallel()
const denied = "192.0.2.50" // in SWWAF_DENY_NETS const denied = "192.0.2.50" // in SWWAF_DENY_NETS
geojsURL, _ := startGeoJS(t) s, _, server := startWithClock(t, "", map[string]string{
s, _, server := startWithClock(t, geojsURL, map[string]string{
rateLimitPerMinute: "1", rateLimitPerMinute: "1",
rateLimitExemptPaths: "/assets/,/favicon.ico", rateLimitExemptPaths: "/assets/,/favicon.ico",
denyNets: denied, denyNets: denied,
-104
View File
@@ -1,104 +0,0 @@
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
+29 -152
View File
@@ -3,7 +3,6 @@ package proxy
import ( import (
"context" "context"
"errors" "errors"
"io"
"net/http" "net/http"
"net/http/httptrace" "net/http/httptrace"
"net/http/httputil" "net/http/httputil"
@@ -16,10 +15,6 @@ import (
"sync/atomic" "sync/atomic"
"time" "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/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
) )
@@ -53,29 +48,7 @@ type request struct {
client netip.Addr client netip.Addr
peer netip.Addr peer netip.Addr
peerTrusted bool peerTrusted bool
// lookedUp is true once the client's AS number and country have been start time.Time
// 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 // checked is when the checks were done, and upstreamStart when the
// request was handed to the app. // request was handed to the app.
checked time.Time checked time.Time
@@ -86,9 +59,6 @@ type request struct {
refused atomic.Pointer[refusal] refused atomic.Pointer[refusal]
// complete is true once the app's whole answer has been passed on. // complete is true once the app's whole answer has been passed on.
complete bool 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 // mu guards what follows. The timeouts run on goroutines of their
// own, and the transport starts and stops them, and notes the times // own, and the transport starts and stops them, and notes the times
@@ -144,7 +114,7 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
RequestID: requestID(r, peerTrusted), RequestID: requestID(r, peerTrusted),
PeerIP: peer.String(), PeerIP: peer.String(),
ForwardedFor: strings.Join(forwardedFor, ", "), ForwardedFor: strings.Join(forwardedFor, ", "),
ClientGroup: h.clientGroup(client).String(), ClientGroup: clientGroup(client).String(),
ContentType: r.Header.Get("Content-Type"), ContentType: r.Header.Get("Content-Type"),
RequestHeaders: requestHeaders(r, h.config.LogRequestHeaders), RequestHeaders: requestHeaders(r, h.config.LogRequestHeaders),
HasAuthorization: len(r.Header.Values("Authorization")) > 0, HasAuthorization: len(r.Header.Values("Authorization")) > 0,
@@ -222,18 +192,14 @@ func (rq *request) check(ctx context.Context) *refusal {
// checkClient runs the checks on the request's client, and returns the // checkClient runs the checks on the request's client, and returns the
// action of the first that refuses the request, or "" when none does. A // action of the first that refuses the request, or "" when none does. A
// client in SWWAF_ALLOW_NETS skips them, and is not looked up. For any // client in SWWAF_ALLOW_NETS skips them. For any other client,
// other client, SWWAF_DENY_NETS comes first, then a ban on its netblock, // SWWAF_DENY_NETS comes first, then a ban on its netblock, so that a
// so that a client either refuses is not looked up, then the lookup of // client either refuses is not looked up, and then the country lists; a
// its AS number and country, then the country lists, then the blocklists, // request any of them refuses is not counted for the rate limits. Then
// then the DNSBL zones' verdicts, and then AbuseIPDB's score; a request // come the rate limits, unless the client is in
// any of them refuses is not counted for the rate limits. Then come the // SWWAF_RATE_LIMIT_EXEMPT_NETS or the request's path is exempt under
// rate limits, unless the client is in SWWAF_RATE_LIMIT_EXEMPT_NETS or the // SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that every other request is counted,
// request's path is exempt under SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that // 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 { func (rq *request) checkClient(ctx context.Context) string {
cfg := rq.h.config cfg := rq.h.config
if isInside(rq.client, cfg.AllowNets) { if isInside(rq.client, cfg.AllowNets) {
@@ -250,29 +216,13 @@ func (rq *request) checkClient(ctx context.Context) string {
return requestlog.ActionBanned return requestlog.ActionBanned
} }
rq.lookUp(ctx) if rq.countryDenied(ctx) {
if rq.countryDenied() {
return requestlog.ActionCountryDenied return requestlog.ActionCountryDenied
} }
if rq.blocklistDenied() { exempt := isInside(rq.client, cfg.RateLimitExemptNets) ||
return requestlog.ActionDenied pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
} if !exempt && rq.limitBroken(now) {
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 return requestlog.ActionRateLimited
} }
@@ -337,9 +287,7 @@ func (rq *request) forward(ctx context.Context) {
// rewrite makes the request the app receives: the client's request, // rewrite makes the request the app receives: the client's request,
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and // unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
// the request's id set, without any X-Client-ASN or X-Client-Country the // the request's id set.
// 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) { func (rq *request) rewrite(pr *httputil.ProxyRequest) {
upstream := rq.h.config.UpstreamURL upstream := rq.h.config.UpstreamURL
pr.Out.URL.Scheme = upstream.Scheme pr.Out.URL.Scheme = upstream.Scheme
@@ -349,12 +297,6 @@ func (rq *request) rewrite(pr *httputil.ProxyRequest) {
pr.Out.URL.RawQuery = pr.In.URL.RawQuery pr.Out.URL.RawQuery = pr.In.URL.RawQuery
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted) setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
pr.Out.Header.Set(requestIDHeader, rq.line.RequestID) 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)
}
} }
// modifyResponse looks at the app's answer before ReverseProxy passes it // modifyResponse looks at the app's answer before ReverseProxy passes it
@@ -365,18 +307,11 @@ func (rq *request) modifyResponse(res *http.Response) error {
if res.StatusCode == http.StatusSwitchingProtocols { if res.StatusCode == http.StatusSwitchingProtocols {
// An upgraded connection, such as a WebSocket, is not cut by the // An upgraded connection, such as a WebSocket, is not cut by the
// timeouts. ReverseProxy writes this answer straight to the // timeouts. ReverseProxy writes this answer straight to the
// connection it takes over, not through rq.out, and then copies // connection it takes over, not through rq.out.
// what passes each way through res.Body, the connection to the app.
rq.stopTimers() rq.stopTimers()
rq.out.status = res.StatusCode rq.out.status = res.StatusCode
rq.line.Websocket = true rq.line.Websocket = true
conn, ok := res.Body.(io.ReadWriteCloser)
if ok {
rq.upgraded = &upgradedConn{ReadWriteCloser: conn}
res.Body = rq.upgraded
}
return nil return nil
} }
@@ -479,7 +414,10 @@ func (rq *request) finish() {
line.ResponseContentType = header.Get("Content-Type") line.ResponseContentType = header.Get("Content-Type")
line.CacheControl = header.Get("Cache-Control") line.CacheControl = header.Get("Cache-Control")
line.Location = header.Get("Location") line.Location = header.Get("Location")
line.RequestBytes = rq.requestBytes()
if rq.body != nil {
line.RequestBytes = rq.body.bytes.Load()
}
// limit is the setting whose size or time limit the request passed. // limit is the setting whose size or time limit the request passed.
var limit string var limit string
@@ -536,85 +474,24 @@ func timing(start, end time.Time) *float64 {
} }
// addToHistory adds the request, which has ended, to its client's // addToHistory adds the request, which has ended, to its client's
// history, and then the lookup's answer about the client, as // history.
// 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() { func (rq *request) addToHistory() {
var requestBytes int64
if rq.body != nil {
requestBytes = rq.body.bytes.Load()
}
forwarded := !rq.upstreamStart.IsZero() forwarded := !rq.upstreamStart.IsZero()
rq.h.limiter.AddToHistory(rq.h.clientGroup(rq.client), rq.h.now(), ratelimit.Request{ rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{
Country: rq.line.Country,
Forwarded: forwarded, Forwarded: forwarded,
Refused: !forwarded && rq.refused.Load() != nil, Refused: !forwarded && rq.refused.Load() != nil,
Status: rq.out.status, Status: rq.out.status,
RequestBytes: rq.requestBytes(), RequestBytes: requestBytes,
ResponseBytes: rq.out.bytes, ResponseBytes: rq.out.bytes,
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit, BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
Attack: rq.attack,
RuleBlocked: rq.ruleBlocked,
}) })
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, _ := 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 // clientRequestDeadline is when the client must have sent its whole
+1 -4
View File
@@ -119,10 +119,7 @@ func wantFullLine(t *testing.T, line logLine) {
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound, ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
CacheControl: "no-store", Location: "/elsewhere", CacheControl: "no-store", Location: "/elsewhere",
Action: requestlog.ActionForward, Action: requestlog.ActionForward,
// Its 3 bytes in and 5 out, each way counted by default. Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
Counts: ratelimit.Counts{
Minute: 1, Hour: 1, Day: 1, MinuteBytes: 8, HourBytes: 8, DayBytes: 8,
},
}) })
if !reflect.DeepEqual(line.Line, want) { if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want) t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
+9 -10
View File
@@ -5,7 +5,6 @@ import (
"net/netip" "net/netip"
"os" "os"
"path/filepath" "path/filepath"
"reflect"
"slices" "slices"
"testing" "testing"
"time" "time"
@@ -74,7 +73,7 @@ func TestEachRuleAction(t *testing.T) {
} }
got := server.Ledger.Bans(netblock) got := server.Ledger.Bans(netblock)
if len(got) != 1 || !reflect.DeepEqual(got[0], want) { if len(got) != 1 || got[0] != want {
t.Fatalf("bans\n%+v\nwant\n%+v", got, want) t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
} }
@@ -198,15 +197,15 @@ func TestMetricsCountRuleMatchesAndBansForAnAttack(t *testing.T) {
metrics := s.scrape(scraper) metrics := s.scrape(scraper)
wantMetric(t, metrics, wantMetric(t, metrics,
`smallwebwaf_rule_matches_total{action="block",instance="app",rule_id="blocked"}`, 1) `smallwebwaf_rule_matches_total{action="block",rule_id="blocked"}`, 1)
wantMetric(t, metrics, wantMetric(t, metrics,
`smallwebwaf_rule_matches_total{action="ban",instance="app",rule_id="probe"}`, 1) `smallwebwaf_rule_matches_total{action="ban",rule_id="probe"}`, 1)
wantMetric(t, metrics, `smallwebwaf_rules_loaded{instance="app"}`, 2) wantMetric(t, metrics, "smallwebwaf_rules_loaded", 2)
wantMetric(t, metrics, `smallwebwaf_requests_total{action="rule_blocked",`+ wantMetric(t, metrics,
`instance="app",status_class="4xx"}`, 1) `smallwebwaf_requests_total{action="rule_blocked",status_class="4xx"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack",instance="app"}`, 1) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 0) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 0)
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
} }
// writeRules writes content as a rule file into a new directory, and // writeRules writes content as a rule file into a new directory, and
+5 -8
View File
@@ -10,10 +10,8 @@ import (
// checkRules checks the request against the rules of the rule files at // checkRules checks the request against the rules of the rule files at
// now, notes the ids of those it matches in the log line, and returns the // now, notes the ids of those it matches in the log line, and returns the
// action of the rule that refuses it, ActionRuleBlocked for a block rule // 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 // and ActionBanned for a ban rule, or "" when none does. In enforce mode
// the client's netblock for a clear sign of attack, or in observe mode // a ban rule bans the client's netblock for a clear sign of attack.
// 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 { func (rq *request) checkRules(now time.Time) string {
matched := rq.h.rules.Match(rq.in) matched := rq.h.rules.Match(rq.in)
@@ -29,12 +27,11 @@ func (rq *request) checkRules(now time.Time) string {
// Only the last rule matched can refuse the request. // Only the last rule matched can refuse the request.
switch last := matched[len(matched)-1]; last.Action { switch last := matched[len(matched)-1]; last.Action {
case rules.ActionBlock: case rules.ActionBlock:
rq.ruleBlocked = true
return requestlog.ActionRuleBlocked return requestlog.ActionRuleBlocked
case rules.ActionBan: case rules.ActionBan:
rq.attack = true if !rq.h.config.Observe {
rq.banForAttack(now, last) rq.banForAttack(now, last)
}
return requestlog.ActionBanned return requestlog.ActionBanned
default: default:
+7 -42
View File
@@ -11,15 +11,15 @@ import (
func TestHistoryKeepsEveryRequest(t *testing.T) { func TestHistoryKeepsEveryRequest(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{}, tableSize) limiter := ratelimit.New(ratelimit.Limits{})
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
for i, r := range []ratelimit.Request{ for i, r := range []ratelimit.Request{
{Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100}, {Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100},
{Forwarded: true, Status: 101}, {Forwarded: true, Status: 101},
{Forwarded: true, Status: 304, RequestBytes: 5}, {Forwarded: true, Status: 304, RequestBytes: 5},
{Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true}, {Country: "FR", Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true},
{Forwarded: true, Status: 502, ResponseBytes: 12}, {Forwarded: true, Status: 502, ResponseBytes: 12},
// Closed without an answer: refused, and no response. // Closed without an answer: refused, and no response.
{Refused: true, Status: 0}, {Refused: true, Status: 0},
@@ -33,6 +33,8 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
want := ratelimit.History{ want := ratelimit.History{
FirstSeen: start, FirstSeen: start,
LastSeen: start.Add(6 * time.Minute), LastSeen: start.Add(6 * time.Minute),
Country: "FR",
LookedUp: start.Add(3 * time.Minute),
Requests: 7, Requests: 7,
Forwarded: 4, Forwarded: 4,
Refused: 2, Refused: 2,
@@ -50,47 +52,10 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
} }
} }
func TestLookupReachesTheHistoryOfAClientInTheTable(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{}, tableSize)
client := netip.MustParsePrefix("203.0.113.9/32")
other := netip.MustParsePrefix("198.51.100.7/32")
start := midnight()
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
limiter.AddLookup(client, start, "AS64496", "Example Net", "DE")
// A later answer replaces it, and one for a client the table does not
// hold adds no client.
limiter.AddLookup(client, start.Add(time.Hour), "AS64497", "Other Net", "FR")
limiter.AddLookup(other, start, "AS64496", "Example Net", "DE")
want := ratelimit.History{
FirstSeen: start,
LastSeen: start,
ASN: "AS64497",
ASName: "Other Net",
Country: "FR",
LookedUp: start.Add(time.Hour),
Requests: 1,
Forwarded: 1,
}
got := historyOf(t, limiter, client)
if got != want {
t.Errorf("history\n%+v\nwant\n%+v", got, want)
}
if clients := limiter.Snapshot(); len(clients) != 1 {
t.Errorf("the table holds %+v, want %s alone", clients, client)
}
}
func TestResetKeepsTheHistory(t *testing.T) { func TestResetKeepsTheHistory(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}, tableSize) limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit})
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -109,7 +74,7 @@ func TestResetKeepsTheHistory(t *testing.T) {
func TestRequestsAddsUpTheClientsInsideTheNetblock(t *testing.T) { func TestRequestsAddsUpTheClientsInsideTheNetblock(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{}, tableSize) limiter := ratelimit.New(ratelimit.Limits{})
for client, requests := range map[string]int{ for client, requests := range map[string]int{
"198.51.100.9/32": 2, "198.51.100.9/32": 2,
+97 -230
View File
@@ -1,10 +1,9 @@
// Package ratelimit keeps the table of clients: each client's requests // Package ratelimit keeps the table of clients: each client's requests
// and bytes counted over a minute, an hour and a day, as the "Counting // counted over a minute, an hour and a day, as the "Counting method"
// method" section of SPEC.md describes, which tell when a request takes // section of SPEC.md describes, which tell when a request takes the client
// the client over a rate limit or a byte limit, and each client's history // over a rate limit, and each client's history since it was first seen.
// since it was first seen. At most SWWAF_MAX_TRACKED_CLIENTS clients are // At most 20,000 clients are kept, in memory, and written to clients.json
// kept, in memory, and written to clients.json and read from it by the // and read from it by the state package.
// state package.
package ratelimit package ratelimit
import ( import (
@@ -17,32 +16,26 @@ import (
"github.com/hashicorp/golang-lru/v2/simplelru" "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 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 // Limits are the most requests a client may make in a minute, an hour and
// a day, and the most bytes. Zero is no limit. // a day. Zero is no limit.
type Limits struct { type Limits struct {
PerMinute int64 PerMinute int64
PerHour int64 PerHour int64
PerDay int64 PerDay int64
BytesPerMinute int64
BytesPerHour int64
BytesPerDay int64
} }
// Limiter counts each client's requests and bytes against the limits, and // Limiter counts each client's requests against the limits, and keeps
// keeps its history. It is safe for concurrent use. // its history. It is safe for concurrent use.
type Limiter struct { type Limiter struct {
// windows are the minute, the hour and the day, in the order of // windows are the minute, the hour and the day, in the order of
// Client.buckets and Client.byteBuckets. // Client.buckets.
windows [3]window windows [3]window
mu sync.Mutex mu sync.Mutex
@@ -50,23 +43,17 @@ type Limiter struct {
} }
// Client is a client in the table, as clients.json holds it: its buckets // Client is a client in the table, as clients.json holds it: its buckets
// of requests and of bytes in each window, and its history. // in each window, and its history.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Client struct { type Client struct {
Client netip.Prefix `json:"client"` Client netip.Prefix `json:"client"`
Minute Buckets `json:"minute"` Minute Buckets `json:"minute"`
Hour Buckets `json:"hour"` Hour Buckets `json:"hour"`
Day Buckets `json:"day"` Day Buckets `json:"day"`
MinuteBytes Buckets `json:"minute_bytes"` History History `json:"history"`
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, or the // Buckets are a client's two buckets in one window: the requests in the
// bytes, in the bucket under way, which began at Start, and in the bucket // bucket under way, which began at Start, and in the bucket before it.
// before it.
type Buckets struct { type Buckets struct {
Start time.Time `json:"start"` Start time.Time `json:"start"`
Current int64 `json:"current"` Current int64 `json:"current"`
@@ -79,12 +66,8 @@ type Buckets struct {
type History struct { type History struct {
FirstSeen time.Time `json:"first_seen"` FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"` LastSeen time.Time `json:"last_seen"`
// ASN, ASName and Country are the client's AS number, AS name and // Country is the client's country as it was last looked up, and
// country as last looked up, each empty when the lookup could not // LookedUp when that was; both are empty while it never was.
// 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"` Country string `json:"country,omitempty"`
LookedUp time.Time `json:"looked_up,omitzero"` LookedUp time.Time `json:"looked_up,omitzero"`
// Requests are all the client's requests: Forwarded those passed to // Requests are all the client's requests: Forwarded those passed to
@@ -113,19 +96,15 @@ type Responses struct {
} }
// Offences are a client's offences, by kind. // Offences are a client's offences, by kind.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Offences struct { type Offences struct {
// Limit is its requests that broke a rate limit or a byte limit, // Limit is its requests that broke a rate limit.
// Attack those that matched a ban rule, a clear sign of attack, and Limit int64 `json:"limit"`
// 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. // Request is what a client's history keeps of one of its requests.
type Request struct { type Request struct {
// Country is the client's country, when the request looked it up.
Country string
// Forwarded is true for a request passed to the app, Refused for one // Forwarded is true for a request passed to the app, Refused for one
// refused before anything reached it, a 401 at smallwebwaf's own // refused before anything reached it, a 401 at smallwebwaf's own
// endpoints included. Both are false for any other request smallwebwaf // endpoints included. Both are false for any other request smallwebwaf
@@ -138,19 +117,12 @@ type Request struct {
// and of its response. // and of its response.
RequestBytes int64 RequestBytes int64
ResponseBytes int64 ResponseBytes int64
// BrokeLimit is true for a request that broke a rate limit or a byte // BrokeLimit is true for a request that broke a rate limit.
// limit, Attack for one that matched a ban rule, and RuleBlocked for BrokeLimit bool
// one a block rule refused.
BrokeLimit bool
Attack bool
RuleBlocked bool
} }
// New returns a Limiter for limits, with no client counted yet, whose // New returns a Limiter for limits, with no client counted yet.
// table holds at most maxClients clients (SWWAF_MAX_TRACKED_CLIENTS). Past func New(limits Limits) *Limiter {
// 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) clients, err := simplelru.NewLRU[netip.Prefix, *Client](maxClients, nil)
if err != nil { if err != nil {
panic(err) // NewLRU fails only for a size below one panic(err) // NewLRU fails only for a size below one
@@ -158,74 +130,62 @@ func New(limits Limits, maxClients int) *Limiter {
return &Limiter{ return &Limiter{
windows: [3]window{ windows: [3]window{
{ {name: "minute", length: time.Minute, limit: limits.PerMinute},
name: "minute", length: time.Minute, {name: "hour", length: time.Hour, limit: limits.PerHour},
limit: limits.PerMinute, byteLimit: limits.BytesPerMinute, {name: "day", length: day, limit: limits.PerDay},
},
{
name: "hour", length: time.Hour,
limit: limits.PerHour, byteLimit: limits.BytesPerHour,
},
{
name: "day", length: day,
limit: limits.PerDay, byteLimit: limits.BytesPerDay,
},
}, },
clients: clients, clients: clients,
} }
} }
// Hit is a request that takes a client over a rate limit, or whose bytes // Hit is a request that takes a client over a rate limit.
// take it over a byte limit.
type Hit struct { type Hit struct {
// Kind is KindRequests for a rate limit, KindBytes for a byte limit.
Kind string
// Window is "minute", "hour" or "day". // Window is "minute", "hour" or "day".
Window string Window string
// Limit is the window's limit, as the client's percentage of it. // Limit is the window's limit.
Limit int64 Limit int64
// Count is the client's requests, or bytes, counted in the window, // Requests is the client's requests counted in the window, this one
// this request's included. // included.
Count float64 Requests float64
} }
// Counts are a client's requests and bytes in the minute, the hour and // Counts are a client's requests in the minute, the hour and the day that
// the day that end at a request, that request's included. // end at a request, that request included.
//
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
type Counts struct { type Counts struct {
Minute float64 `json:"minute"` Minute float64 `json:"minute"`
Hour float64 `json:"hour"` Hour float64 `json:"hour"`
Day float64 `json:"day"` 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 // Count counts a request from client at now, in every window, whether or
// not it is refused, and returns the client's counts in each window. It // not it is refused, and returns the client's requests in each window. It
// reports whether the request takes the client over a rate limit, of // reports whether the request takes the client over a limit, and the
// which the client gets the percentage percent, rounded down, and the hit: // window whose limit it goes over, the shortest if it is over several.
// the window whose limit it goes over, the shortest if it is over func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) {
// several. A limit that is off stays off. l.mu.Lock()
func (l *Limiter) Count( defer l.mu.Unlock()
client netip.Prefix, now time.Time, percent int64,
) (Counts, Hit, bool) { var (
return l.count(client, now, 1, 0, percent) 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]}
}
}
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, // Reset sets client's counts in every window back to zero. Its history
// at now, in every window, and returns the client's counts in each window. // keeps its totals.
// 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 of requests and of bytes in every window
// back to zero. Its history keeps its totals.
func (l *Limiter) Reset(client netip.Prefix) { func (l *Limiter) Reset(client netip.Prefix) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
@@ -233,7 +193,6 @@ func (l *Limiter) Reset(client netip.Prefix) {
c, seen := l.clients.Peek(client) c, seen := l.clients.Peek(client)
if seen { if seen {
c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{} c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{}
c.MinuteBytes, c.HourBytes, c.DayBytes = Buckets{}, Buckets{}, Buckets{}
} }
} }
@@ -250,6 +209,11 @@ func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
h.LastSeen = now h.LastSeen = now
if r.Country != "" {
h.Country = r.Country
h.LookedUp = now
}
h.Requests++ h.Requests++
if r.Forwarded { if r.Forwarded {
h.Forwarded++ h.Forwarded++
@@ -266,33 +230,6 @@ func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
if r.BrokeLimit { if r.BrokeLimit {
h.Offences.Limit++ 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
// 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,
) {
l.mu.Lock()
defer l.mu.Unlock()
c, held := l.clients.Peek(client)
if !held {
return
}
h := &c.History
h.ASN, h.ASName, h.Country, h.LookedUp = asn, asName, country, lookedUp
} }
// Requests returns how many requests the clients inside netblock have // Requests returns how many requests the clients inside netblock have
@@ -375,11 +312,12 @@ func (l *Limiter) Load(clients []Client, now time.Time) {
l.clients.Purge() l.clients.Purge()
for _, c := range clients { for _, c := range clients {
for i, w := range l.windows { for i, b := range c.buckets() {
for _, b := range []*Buckets{c.buckets()[i], c.byteBuckets()[i]} { // The window that ends at now covers neither bucket once it
if b.Passed(now, w.length) { // begins after the bucket under way has ended.
*b = Buckets{} length := l.windows[i].length
} if !now.Add(-length).Before(b.Start.Add(length)) {
*b = Buckets{}
} }
} }
@@ -387,51 +325,6 @@ func (l *Limiter) Load(clients []Client, now time.Time) {
} }
} }
// 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 // get returns client's entry in the table, a new one if it has none, and
// makes it the most recently seen. // makes it the most recently seen.
func (l *Limiter) get(client netip.Prefix) *Client { func (l *Limiter) get(client netip.Prefix) *Client {
@@ -444,48 +337,30 @@ func (l *Limiter) get(client netip.Prefix) *Client {
return c return c
} }
// buckets returns c's buckets of requests in the minute, the hour and the // buckets returns c's buckets in the minute, the hour and the day.
// day.
func (c *Client) buckets() [3]*Buckets { func (c *Client) buckets() [3]*Buckets {
return [3]*Buckets{&c.Minute, &c.Hour, &c.Day} return [3]*Buckets{&c.Minute, &c.Hour, &c.Day}
} }
// byteBuckets returns c's buckets of bytes in the minute, the hour and the // window is a length of time over which requests are counted, and the
// day. // most requests a client may make in it.
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 { type window struct {
name string name string
length time.Duration length time.Duration
limit int64 limit int64
byteLimit int64
} }
// percentOf returns the percentage percent of limit, rounded down. It is // add counts a request at now in a window of length, and returns the
// written as limit's hundreds times percent, plus the rest's share, since // client's requests in the window that ends at now: those in the bucket
// limit*percent can overflow for a byte limit. // under way, and those in the bucket before it weighted by how much of
func percentOf(limit, percent int64) int64 { // that bucket the window still covers.
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 // 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 // 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 // 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 // was set back, and the buckets start afresh: otherwise the bucket before
// would keep its full weight until the clock caught up. // would keep its full weight until the clock caught up.
func (b *Buckets) Add(now time.Time, length time.Duration, n int64) float64 { func (b *Buckets) add(now time.Time, length time.Duration) float64 {
if now.Before(b.Start.Add(-time.Second)) { if now.Before(b.Start.Add(-time.Second)) {
*b = Buckets{} *b = Buckets{}
} }
@@ -502,7 +377,7 @@ func (b *Buckets) Add(now time.Time, length time.Duration, n int64) float64 {
b.Current = 0 b.Current = 0
} }
b.Current += n b.Current++
elapsed := max(now.Sub(b.Start), 0) elapsed := max(now.Sub(b.Start), 0)
covered := 1 - float64(elapsed)/float64(length) covered := 1 - float64(elapsed)/float64(length)
@@ -510,14 +385,6 @@ func (b *Buckets) Add(now time.Time, length time.Duration, n int64) float64 {
return float64(b.Previous)*covered + float64(b.Current) 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 // add counts a response with status in its class. A status of 0, for
// nothing sent, is not a response. // nothing sent, is not a response.
func (r *Responses) add(status int) { func (r *Responses) add(status int) {
+17 -208
View File
@@ -1,7 +1,6 @@
package ratelimit_test package ratelimit_test
import ( import (
"math"
"net/netip" "net/netip"
"testing" "testing"
"time" "time"
@@ -12,14 +11,6 @@ import (
// limit is the limit the tests set. // limit is the limit the tests set.
const limit = 3 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. // The windows, as Count names them.
const ( const (
minute = "minute" minute = "minute"
@@ -41,7 +32,7 @@ func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) {
t.Run(tc.window, func(t *testing.T) { t.Run(tc.window, func(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(tc.limits, tableSize) limiter := ratelimit.New(tc.limits)
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
quarter := tc.length / 4 quarter := tc.length / 4
@@ -66,204 +57,43 @@ func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) {
func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) { func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: limit}, tableSize) limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: limit})
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
for range limit { for range limit {
_, _, over := limiter.Count(client, start, whole) _, _, over := limiter.Count(client, start)
if over { if over {
t.Fatal("a request within the limit is over it") t.Fatal("a request within the limit is over it")
} }
} }
// Over both limits; the minute's is named, with the four requests. // Over both limits; the minute's is named, with the four requests.
_, hit, over := limiter.Count(client, start, whole) _, hit, over := limiter.Count(client, start)
want := ratelimit.Hit{ want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1}
Kind: ratelimit.KindRequests, Window: minute, Limit: limit, Count: limit + 1,
}
if !over || hit != want { if !over || hit != want {
t.Errorf("request over the limit gives %+v and %t, want %+v and true", t.Errorf("request over the limit gives %+v and %t, want %+v and true",
hit, over, want) 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) { func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{}, tableSize) limiter := ratelimit.New(ratelimit.Limits{})
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
for range 3 { for range 3 {
limiter.Count(client, start, whole) limiter.Count(client, start)
} }
// A quarter into the next hour, the minute has only this request. The // 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 // 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 // requests, which count 2.25, and this one: 3.25. The day covers all
// four. // four.
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4), whole) counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4))
want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4} want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4}
if counts != want { if counts != want {
@@ -274,7 +104,7 @@ func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
func TestResetSetsTheCountsBackToZero(t *testing.T) { func TestResetSetsTheCountsBackToZero(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerDay: limit}, tableSize) limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerDay: limit})
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -296,7 +126,7 @@ func TestResetSetsTheCountsBackToZero(t *testing.T) {
func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) { func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize) limiter := ratelimit.New(ratelimit.Limits{PerHour: limit})
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -315,8 +145,7 @@ func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) {
func TestRefusedRequestsCount(t *testing.T) { func TestRefusedRequestsCount(t *testing.T) {
t.Parallel() 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") refused := netip.MustParsePrefix("203.0.113.9/32")
within := netip.MustParsePrefix("203.0.113.10/32") within := netip.MustParsePrefix("203.0.113.10/32")
start := midnight() start := midnight()
@@ -349,7 +178,7 @@ func TestRefusedRequestsCount(t *testing.T) {
func TestRequestCountedLateGoesInTheBucketUnderWay(t *testing.T) { func TestRequestCountedLateGoesInTheBucketUnderWay(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}, tableSize) limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit})
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -365,7 +194,7 @@ func TestRequestCountedLateGoesInTheBucketUnderWay(t *testing.T) {
func TestClockSetBackStartsTheBucketsAfresh(t *testing.T) { func TestClockSetBackStartsTheBucketsAfresh(t *testing.T) {
t.Parallel() t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize) limiter := ratelimit.New(ratelimit.Limits{PerHour: limit})
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
@@ -388,12 +217,12 @@ func TestClockSetBackStartsTheBucketsAfresh(t *testing.T) {
wantCount(t, limiter, client, setBack, hour) wantCount(t, limiter, client, setBack, hour)
} }
func TestKeepsAtMostMaxClientsDroppingTheLeastRecentlySeen(t *testing.T) { func TestKeepsAtMost20000ClientsDroppingTheLeastRecentlySeen(t *testing.T) {
t.Parallel() t.Parallel()
const maxClients = 3 const maxClients = 20000
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 1}, maxClients) limiter := ratelimit.New(ratelimit.Limits{PerMinute: 1})
now := midnight() now := midnight()
clients := make([]netip.Prefix, maxClients+1) clients := make([]netip.Prefix, maxClients+1)
@@ -415,11 +244,6 @@ func TestKeepsAtMostMaxClientsDroppingTheLeastRecentlySeen(t *testing.T) {
// One client more drops the least recently seen, the second, which // One client more drops the least recently seen, the second, which
// starts afresh, while the first is kept. // starts afresh, while the first is kept.
wantCount(t, limiter, clients[maxClients], now, "") 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[1], now, "")
wantCount(t, limiter, clients[0], now, minute) wantCount(t, limiter, clients[0], now, minute)
} }
@@ -437,24 +261,9 @@ func wantCount(
) { ) {
t.Helper() t.Helper()
_, hit, _ := limiter.Count(client, now, whole) _, hit, _ := limiter.Count(client, now)
if hit.Window != want { if hit.Window != want {
t.Errorf("request from %s at %s is over %q, want %q", t.Errorf("request from %s at %s is over %q, want %q",
client, now.Format(time.RFC3339), hit.Window, want) 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 -21
View File
@@ -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"} 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{}, tableSize) limiter := ratelimit.New(ratelimit.Limits{})
for _, i := range []int{2, 3, 0, 1} { for _, i := range []int{2, 3, 0, 1} {
limiter.Count(netip.MustParsePrefix(want[i]), midnight(), whole) limiter.Count(netip.MustParsePrefix(want[i]), midnight())
} }
snapshot := limiter.Snapshot() snapshot := limiter.Snapshot()
@@ -43,7 +43,7 @@ func TestLoadedCountsCarryOn(t *testing.T) {
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
before := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize) before := ratelimit.New(ratelimit.Limits{PerHour: limit})
for range limit { for range limit {
wantCount(t, before, client, start, "") 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 // Loaded into a new limiter, as across a restart, the client has no
// fresh allowance. // fresh allowance.
later := start.Add(time.Minute) later := start.Add(time.Minute)
after := ratelimit.New(ratelimit.Limits{PerHour: limit}, tableSize) after := ratelimit.New(ratelimit.Limits{PerHour: limit})
after.Load(before.Snapshot(), later) after.Load(before.Snapshot(), later)
wantCount(t, after, client, later, hour) wantCount(t, after, client, later, hour)
} }
@@ -62,47 +62,40 @@ func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
start := midnight() start := midnight()
limiter := ratelimit.New(ratelimit.Limits{}, tableSize) limiter := ratelimit.New(ratelimit.Limits{})
limiter.Count(client, start, whole) limiter.Count(client, start)
limiter.CountBytes(client, start, 5, whole)
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true}) limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
loaded := func(now time.Time) ratelimit.Client { loaded := func(now time.Time) ratelimit.Client {
t.Helper() t.Helper()
after := ratelimit.New(ratelimit.Limits{}, tableSize) after := ratelimit.New(ratelimit.Limits{})
after.Load(limiter.Snapshot(), now) after.Load(limiter.Snapshot(), now)
return after.Snapshot()[0] return after.Snapshot()[0]
} }
// Two minutes on, the window that ends then covers neither of the // Two minutes on, the window that ends then covers neither of the
// minute's buckets, of requests and of bytes, which are emptied; the // minute's buckets, which are emptied; the hour's and the day's stay,
// hour's and the day's stay, and so does the history. // and so does the history.
got := loaded(start.Add(2 * time.Minute)) got := loaded(start.Add(2 * time.Minute))
if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 || if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 ||
got.Day.Current != 1 || got.History.Requests != 1 { got.Day.Current != 1 || got.History.Requests != 1 {
t.Errorf("loaded two minutes on as %+v", got) 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. // A moment before, the window still covers some of the earlier one.
got = loaded(start.Add(2*time.Minute - time.Nanosecond)) got = loaded(start.Add(2*time.Minute - time.Nanosecond))
if got.Minute.Current != 1 || got.MinuteBytes.Current != 5 { if got.Minute.Current != 1 {
t.Errorf("loaded just under two minutes on with minute buckets %+v and %+v", t.Errorf("loaded just under two minutes on with minute buckets %+v",
got.Minute, got.MinuteBytes) got.Minute)
} }
} }
func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) { func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) {
t.Parallel() t.Parallel()
const maxClients = 3 const maxClients = 20000
// clients.json lists the clients by address. Here each was last seen // 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 // a second before the one listed before it, so the last listed is the
@@ -116,7 +109,7 @@ func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) {
addr = addr.Next() addr = addr.Next()
} }
limiter := ratelimit.New(ratelimit.Limits{}, maxClients) limiter := ratelimit.New(ratelimit.Limits{})
limiter.Load(clients, midnight()) limiter.Load(clients, midnight())
got := limiter.Snapshot() got := limiter.Snapshot()
-328
View File
@@ -1,328 +0,0 @@
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())
}
-663
View File
@@ -1,663 +0,0 @@
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()
}
-332
View File
@@ -1,332 +0,0 @@
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
}
-737
View File
@@ -1,737 +0,0 @@
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)
}
}
-33
View File
@@ -1,33 +0,0 @@
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})
}
-509
View File
@@ -1,509 +0,0 @@
// 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
}
-610
View File
@@ -1,610 +0,0 @@
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)
}
}
+10 -34
View File
@@ -35,9 +35,7 @@ const (
// rule. // rule.
ActionRuleBlocked = "rule_blocked" ActionRuleBlocked = "rule_blocked"
// ActionDenied is a request refused because its client is in // ActionDenied is a request refused because its client is in
// SWWAF_DENY_NETS, in a blocklist while SWWAF_BLOCKLIST_ACTION is deny, // SWWAF_DENY_NETS.
// or listed by a DNSBL zone, or scored a hit by AbuseIPDB, while
// SWWAF_REPUTATION_ACTION is deny.
ActionDenied = "denied" ActionDenied = "denied"
// ActionCountryDenied is a request refused for its client's country. // ActionCountryDenied is a request refused for its client's country.
ActionCountryDenied = "country_denied" ActionCountryDenied = "country_denied"
@@ -47,7 +45,7 @@ const (
) )
// OffenceLimit is the offence a request line names for a request that // OffenceLimit is the offence a request line names for a request that
// broke a rate limit, or whose bytes broke a byte limit. // broke a rate limit.
const OffenceLimit = "limit" const OffenceLimit = "limit"
// timeLayout is RFC 3339 with milliseconds. // timeLayout is RFC 3339 with milliseconds.
@@ -82,14 +80,11 @@ type Line struct {
// Request detail. RequestID is the X-Request-ID a trusted proxy sent, // Request detail. RequestID is the X-Request-ID a trusted proxy sent,
// or a new one, and is sent on to the app. ForwardedFor is the // or a new one, and is sent on to the app. ForwardedFor is the
// X-Forwarded-For header as received. ClientGroup is the netblock the // X-Forwarded-For header as received. ClientGroup is the netblock the
// client is counted as. ASN, ASName and Country are the client's AS // client is counted as.
// number, AS name and country, as looked up.
RequestID string `json:"request_id"` RequestID string `json:"request_id"`
PeerIP string `json:"peer_ip"` PeerIP string `json:"peer_ip"`
ForwardedFor string `json:"forwarded_for,omitempty"` ForwardedFor string `json:"forwarded_for,omitempty"`
ClientGroup string `json:"client_group"` ClientGroup string `json:"client_group"`
ASN string `json:"asn"`
ASName string `json:"as_name"`
Country string `json:"country"` Country string `json:"country"`
ContentType string `json:"content_type,omitempty"` ContentType string `json:"content_type,omitempty"`
// ContentLength is the length of its body the request announced. // ContentLength is the length of its body the request announced.
@@ -119,30 +114,14 @@ type Line struct {
// ActionBanned, ActionCountryDenied, ActionRateLimited or // ActionBanned, ActionCountryDenied, ActionRateLimited or
// ActionRuleBlocked. // ActionRuleBlocked.
WouldAction string `json:"would_action,omitempty"` WouldAction string `json:"would_action,omitempty"`
// LimitPercent and LimitPercentSetting are, for a request the rate // Counts are the client's requests as the rate limits counted them
// limits counted whose client a biased threshold gives a percentage of // with this one, for a request they counted.
// 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"` Counts ratelimit.Counts `json:"counts,omitzero"`
// RuleIDs are the ids of the rule file rules the request matched. // RuleIDs are the ids of the rule file rules the request matched.
RuleIDs []string `json:"rule_ids,omitempty"` RuleIDs []string `json:"rule_ids,omitempty"`
// LimitHit is the window whose limit the request went over, named as // LimitHit is the window whose rate limit the request went over:
// Counts names its count: minute, hour or day for a rate limit, and // minute, hour or day.
// minute_bytes, hour_bytes or day_bytes for a byte limit.
LimitHit string `json:"limit_hit,omitempty"` 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 is the offence the request was held as, OffenceLimit.
Offence string `json:"offence,omitempty"` Offence string `json:"offence,omitempty"`
// BanExpires is when the ban the request made, or was refused under, // BanExpires is when the ban the request made, or was refused under,
@@ -192,12 +171,9 @@ func Milliseconds(d time.Duration) float64 {
// NewProcessLogger returns the logger for the process's own messages: // NewProcessLogger returns the logger for the process's own messages:
// JSON lines on w, marked "type":"process", with the time in the same form // 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. // as a request line's.
// It writes only the messages at level, SWWAF_LOG_LEVEL, or more severe; func NewProcessLogger(w io.Writer) *slog.Logger {
// 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{ handler := slog.NewJSONHandler(w, &slog.HandlerOptions{
Level: level,
ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr { ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr {
if attr.Key == slog.TimeKey && len(groups) == 0 { if attr.Key == slog.TimeKey && len(groups) == 0 {
return slog.String(slog.TimeKey, FormatTime(attr.Value.Time())) return slog.String(slog.TimeKey, FormatTime(attr.Value.Time()))
@@ -207,5 +183,5 @@ func NewProcessLogger(w io.Writer, instanceName string, level slog.Level) *slog.
}, },
}) })
return slog.New(handler).With("type", "process", "instance", instanceName) return slog.New(handler).With("type", "process")
} }
+4 -52
View File
@@ -3,8 +3,6 @@ package requestlog_test
import ( import (
"bytes" "bytes"
"encoding/json" "encoding/json"
"log/slog"
"slices"
"strings" "strings"
"testing" "testing"
"time" "time"
@@ -67,13 +65,12 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
} }
} }
func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) { func TestProcessLinesAreMarkedProcess(t *testing.T) {
t.Parallel() t.Parallel()
var out bytes.Buffer var out bytes.Buffer
requestlog.NewProcessLogger(&out, "fsn1app1/gitea", slog.LevelInfo).Info("starting", requestlog.NewProcessLogger(&out).Info("starting", "version", "v1")
"version", "v1")
var fields map[string]any var fields map[string]any
@@ -82,9 +79,8 @@ func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
t.Fatalf("decode %q: %v", out.String(), err) t.Fatalf("decode %q: %v", out.String(), err)
} }
if fields["type"] != "process" || fields["instance"] != "fsn1app1/gitea" || if fields["type"] != "process" || fields["msg"] != "starting" ||
fields["msg"] != "starting" || fields["level"] != "INFO" || fields["level"] != "INFO" || fields["version"] != "v1" {
fields["version"] != "v1" {
t.Errorf("process line %v", fields) t.Errorf("process line %v", fields)
} }
@@ -97,47 +93,3 @@ func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
t.Errorf("process line time %q, want now in UTC with milliseconds", timeText) 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)
}
})
}
}
+12 -15
View File
@@ -129,7 +129,7 @@ func Load(params Params) (*Files, error) {
return f, nil return f, nil
} }
rules, _, err := read(params.Dir) rules, err := read(params.Dir)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -227,9 +227,9 @@ func (f *Files) readAfterChanges(
// readAgain reads the rule files again, in place of the rules loaded, or // readAgain reads the rule files again, in place of the rules loaded, or
// logs the error that keeps the rules as they were, and raises a // logs the error that keeps the rules as they were, and raises a
// file_error alert for it, for the file it is in. // file_error alert for it.
func (f *Files) readAgain() { func (f *Files) readAgain() {
rules, path, err := read(f.params.Dir) rules, err := read(f.params.Dir)
if err != nil { if err != nil {
const kept = "a rule file has an error, and the rules stay as they were" const kept = "a rule file has an error, and the rules stay as they were"
@@ -238,7 +238,7 @@ func (f *Files) readAgain() {
f.params.Alerts.Raise(alerts.Alert{ f.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError, Event: alerts.EventFileError,
Reason: kept, Reason: kept,
Detail: map[string]any{"file": path, "error": err.Error()}, Detail: map[string]any{"error": err.Error()},
}) })
f.params.ProcessLog.Error(kept, "error", err.Error()) f.params.ProcessLog.Error(kept, "error", err.Error())
@@ -257,14 +257,13 @@ func (f *Files) logRead(count int) {
} }
// read returns the rules of every rule file in dir, in the order of the // read returns the rules of every rule file in dir, in the order of the
// files' names, and then of their lines, or an error, with the path of the // files' names, and then of their lines. A file whose name starts with a
// rule file it is in, or dir. A file whose name starts with a dot, such as // dot, such as an editor's lock file .#50-app.rules, is not a rule file,
// an editor's lock file .#50-app.rules, is not a rule file, as a shell's // as a shell's *.rules would not match it.
// *.rules would not match it. func read(dir string) ([]Rule, error) {
func read(dir string) ([]Rule, string, error) {
entries, err := os.ReadDir(dir) entries, err := os.ReadDir(dir)
if err != nil { if err != nil {
return nil, dir, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err) return nil, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err)
} }
var rules []Rule var rules []Rule
@@ -278,15 +277,13 @@ func read(dir string) ([]Rule, string, error) {
continue continue
} }
path := filepath.Join(dir, name) rules, err = readFile(filepath.Join(dir, name), rules, places)
rules, err = readFile(path, rules, places)
if err != nil { if err != nil {
return nil, path, err return nil, err
} }
} }
return rules, "", nil return rules, nil
} }
// readFile appends the rules of the rule file at path to rules. places // readFile appends the rules of the rule file at path to rules. places
+3 -4
View File
@@ -383,14 +383,13 @@ func TestBrokenEditKeepsTheRulesAsTheyWere(t *testing.T) {
t.Errorf("logged %v, want an error %q", line, want) t.Errorf("logged %v, want an error %q", line, want)
} }
// The error is raised as a file_error alert too, for the file. // The error is raised as a file_error alert too.
wantFileError := func() { wantFileError := func() {
t.Helper() t.Helper()
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook] waiting := queue.Snapshot().Waiting
if len(waiting) != 1 || waiting[0].Event != alerts.EventFileError || if len(waiting) != 1 || waiting[0].Event != alerts.EventFileError ||
waiting[0].Reason != hasError || waiting[0].Detail["error"] != want || waiting[0].Reason != hasError || waiting[0].Detail["error"] != want {
waiting[0].Detail["file"] != filepath.Join(dir, firstFile) {
t.Errorf("alerts waiting %+v, want one file_error alert for %q", waiting, want) t.Errorf("alerts waiting %+v, want one file_error alert for %q", waiting, want)
} }
} }
+27 -89
View File
@@ -1,7 +1,6 @@
// Package smallwebwaf runs the smallwebwaf process: it reads the settings, // Package smallwebwaf runs the smallwebwaf process: it reads the settings,
// the rule files, the lookup database and the state files, serves requests // the rule files and the state files, serves requests until it is told to
// until it is told to stop, and then stops in an orderly way, writing the // stop, and then stops in an orderly way, writing the state files.
// state files.
package smallwebwaf package smallwebwaf
import ( import (
@@ -21,7 +20,6 @@ import (
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/remotelog" "sneak.berlin/go/smallwebwaf/internal/remotelog"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
"sneak.berlin/go/smallwebwaf/internal/state" "sneak.berlin/go/smallwebwaf/internal/state"
@@ -66,14 +64,11 @@ func Main(version string) int {
}) })
} }
// Run reads the settings, the rule files, the lookup database and the // Run reads the settings, the rule files and the state files, then serves
// state files, then serves requests until ctx is done. It returns the // requests until ctx is done. It returns the process's exit status, 1
// process's exit status, 1 when smallwebwaf cannot start. // when smallwebwaf cannot start.
func Run(ctx context.Context, params Params) int { func Run(ctx context.Context, params Params) int {
// Until the settings are read, the one message is an invalid setting's processLog := requestlog.NewProcessLogger(params.Stdout)
// error, which every SWWAF_LOG_LEVEL lets through.
processLog := requestlog.NewProcessLogger(params.Stdout,
config.InstanceName(params.LookupEnv), slog.LevelError)
cfg, err := config.FromEnvironment(params.LookupEnv) cfg, err := config.FromEnvironment(params.LookupEnv)
if err != nil { if err != nil {
@@ -91,11 +86,8 @@ func Run(ctx context.Context, params Params) int {
if cfg.LogRemoteURL != nil { if cfg.LogRemoteURL != nil {
remote = newRemoteLogSender(cfg) remote = newRemoteLogSender(cfg)
stdout = io.MultiWriter(params.Stdout, remote) stdout = io.MultiWriter(params.Stdout, remote)
} processLog = requestlog.NewProcessLogger(stdout)
processLog = requestlog.NewProcessLogger(stdout, cfg.InstanceName, cfg.LogLevel)
if remote != nil {
stopSending := startSending(ctx, remote, processLog) stopSending := startSending(ctx, remote, processLog)
defer stopSending() defer stopSending()
} }
@@ -117,17 +109,23 @@ func Run(ctx context.Context, params Params) int {
return 1 return 1
} }
server, err := newServer(cfg, stdout, processLog, now, ruleFiles, alertQueue) server := proxy.New(proxy.Params{
if err != nil { Config: cfg,
processLog.Error("cannot use the lookup database", "error", err.Error()) RequestLog: stdout,
ProcessLog: processLog,
return 1 GeoJSURL: lookup.URL,
} Now: now,
Rules: ruleFiles,
Alerts: alertQueue,
})
if remote != nil { if remote != nil {
server.Metrics.AddRemoteLog(remote) server.Metrics.AddRemoteLog(remote)
} }
if cfg.AlertWebhookURL != nil {
server.Metrics.AddAlerts(alertQueue)
}
files, err := loadStateFiles(cfg, server, alertQueue, now, processLog) files, err := loadStateFiles(cfg, server, alertQueue, now, processLog)
if err != nil { if err != nil {
processLog.Error("cannot use the state files", "error", err.Error()) processLog.Error("cannot use the state files", "error", err.Error())
@@ -148,50 +146,7 @@ func Run(ctx context.Context, params Params) int {
"address", listener.Addr().String(), "address", listener.Addr().String(),
"settings", cfg) "settings", cfg)
return serve(ctx, server, listener, files, ruleFiles, alertQueue, processLog) return serve(ctx, server.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 // loadStateFiles reads the state files into the parts of server and into
@@ -207,28 +162,21 @@ func loadStateFiles(
Ledger: server.Ledger, Ledger: server.Ledger,
Limiter: server.Limiter, Limiter: server.Limiter,
GeoJS: server.GeoJS, GeoJS: server.GeoJS,
Lists: server.Lists,
DNSBL: server.DNSBL,
AbuseIPDB: server.AbuseIPDB,
Alerts: alertQueue, Alerts: alertQueue,
Anomalies: server.Anomalies,
Now: now, Now: now,
ProcessLog: processLog, ProcessLog: processLog,
Metrics: server.Metrics, Metrics: server.Metrics,
}) })
} }
// newAlertQueue returns the queue of the alerts to the webhook, Slack and // newAlertQueue returns the queue of the alerts to SWWAF_ALERT_WEBHOOK_URL,
// ntfy, with the settings for them. // with the settings for it.
func newAlertQueue( func newAlertQueue(
cfg *config.Config, now func() time.Time, processLog *slog.Logger, cfg *config.Config, now func() time.Time, processLog *slog.Logger,
) *alerts.Queue { ) *alerts.Queue {
return alerts.New(alerts.Params{ return alerts.New(alerts.Params{
WebhookURL: cfg.AlertWebhookURL, WebhookURL: cfg.AlertWebhookURL,
WebhookHeaders: cfg.AlertWebhookHeaders, WebhookHeaders: cfg.AlertWebhookHeaders,
SlackURL: cfg.AlertSlackWebhookURL,
NtfyURL: cfg.AlertNtfyURL,
NtfyToken: cfg.AlertNtfyToken,
Events: cfg.AlertEvents, Events: cfg.AlertEvents,
Cooldown: cfg.AlertCooldown, Cooldown: cfg.AlertCooldown,
MaxPerHour: cfg.AlertMaxPerHour, MaxPerHour: cfg.AlertMaxPerHour,
@@ -277,13 +225,11 @@ func startSending(
// serve serves requests on listener, writes the state files as they are // 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 // due, takes in an admin's edits of them, reads the rule files again as
// they change, and the lookup database when it is replaced, fetches the // they change, and sends the alerts, until ctx is done. Then it gives the
// lists the settings name by URL as they are due, and sends the alerts, // requests in progress shutdownTimeout to finish, and writes every state
// until ctx is done. Then it gives the requests in progress // file, alerts.json with the alerts still waiting.
// shutdownTimeout to finish, and writes every state file, alerts.json with
// the alerts still waiting.
func serve( func serve(
ctx context.Context, server *proxy.Server, listener net.Listener, ctx context.Context, server *http.Server, listener net.Listener,
files *state.Files, ruleFiles *rules.Files, alertQueue *alerts.Queue, files *state.Files, ruleFiles *rules.Files, alertQueue *alerts.Queue,
processLog *slog.Logger, processLog *slog.Logger,
) int { ) int {
@@ -299,12 +245,6 @@ func serve(
written := inBackground(func() { files.Run(writing) }) written := inBackground(func() { files.Run(writing) })
watched := inBackground(func() { files.Watch(writing) }) watched := inBackground(func() { files.Watch(writing) })
rulesWatched := inBackground(func() { ruleFiles.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) }) alertsSent := inBackground(func() { alertQueue.Run(writing) })
select { select {
@@ -347,8 +287,6 @@ func serve(
<-written <-written
<-watched <-watched
<-rulesWatched <-rulesWatched
<-lookupFileWatched
<-listsFetched
<-alertsSent <-alertsSent
err = files.WriteAll() err = files.WriteAll()
+25 -537
View File
@@ -18,7 +18,6 @@ import (
"testing" "testing"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
"sneak.berlin/go/smallwebwaf/internal/smallwebwaf" "sneak.berlin/go/smallwebwaf/internal/smallwebwaf"
) )
@@ -38,22 +37,12 @@ const (
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY" stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL" stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
rulesDir = "SWWAF_RULES_DIR" rulesDir = "SWWAF_RULES_DIR"
lookupSource = "SWWAF_LOOKUP_SOURCE" adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
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
// adminSecret is the SWWAF_ADMIN_TOKEN the tests set. // adminSecret is the SWWAF_ADMIN_TOKEN the tests set.
adminSecret = "fedcba9876543210fedcba9876543210" adminSecret = "fedcba9876543210fedcba9876543210"
// instance is the SWWAF_INSTANCE_NAME the tests set where they look at
// it.
instance = "fsn1app1/gitea"
// greeting is what the tests' app answers. // greeting is what the tests' app answers.
greeting = "hello from the app" 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. // output collects what smallwebwaf writes on stdout.
@@ -110,16 +99,12 @@ func (o *output) text() string {
} }
// run runs smallwebwaf with the settings in env until ctx is done, and // run runs smallwebwaf with the settings in env until ctx is done, and
// returns its exit status. SWWAF_LOOKUP_SOURCE is off unless env sets it, // returns its exit status.
// so that no test sends GeoJS its clients' addresses.
func run(ctx context.Context, env map[string]string, out *output) int { func run(ctx context.Context, env map[string]string, out *output) int {
return smallwebwaf.Run(ctx, smallwebwaf.Params{ return smallwebwaf.Run(ctx, smallwebwaf.Params{
Version: testVersion, Version: testVersion,
LookupEnv: func(name string) (string, bool) { LookupEnv: func(name string) (string, bool) {
value, ok := env[name] value, ok := env[name]
if !ok && name == "SWWAF_LOOKUP_SOURCE" {
return "off", true
}
return value, ok return value, ok
}, },
@@ -132,10 +117,7 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
out := &output{} out := &output{}
status := run(t.Context(), map[string]string{ status := run(t.Context(), map[string]string{"SWWAF_REQUEST_MAX_BYTES": "lots"}, out)
"SWWAF_REQUEST_MAX_BYTES": "lots",
instanceName: instance,
}, out)
if status != 1 { if status != 1 {
t.Errorf("exit status %d, want 1", status) t.Errorf("exit status %d, want 1", status)
} }
@@ -143,8 +125,7 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
line := out.line(t, "msg", "invalid setting") line := out.line(t, "msg", "invalid setting")
message, _ := line["error"].(string) message, _ := line["error"].(string)
if line["type"] != "process" || line["instance"] != instance || if line["type"] != "process" || line["level"] != "ERROR" ||
line["level"] != "ERROR" ||
!strings.HasPrefix(message, "SWWAF_REQUEST_MAX_BYTES: ") { !strings.HasPrefix(message, "SWWAF_REQUEST_MAX_BYTES: ") {
t.Errorf("start refused with %v", line) t.Errorf("start refused with %v", line)
} }
@@ -155,7 +136,7 @@ func TestShortTokenStopsTheStartUnshown(t *testing.T) {
const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use
for _, name := range []string{adminToken, metricsToken} { for _, name := range []string{adminToken, "SWWAF_METRICS_TOKEN"} {
t.Run(name, func(t *testing.T) { t.Run(name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -249,94 +230,6 @@ func TestServesUntilToldToStop(t *testing.T) {
out.line(t, "msg", "stopped") 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()
const token = "0123456789abcdef0123456789abcdef"
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
metricsToken: token,
instanceName: instance,
}
var metrics string
out := runUntilStopped(t, env, func(url string) {
metrics = metricsText(t, url+"_smallwebwaf/metrics", token)
})
// The process's lines from its start to its stop, and the request's.
for line := range strings.Lines(out.text()) {
var fields map[string]any
err := json.Unmarshal([]byte(line), &fields)
if err != nil || fields["instance"] != instance {
t.Errorf("line %q (%v), want instance %s", line, err, instance)
}
}
// Each series, Go's and the process's included; the other lines are
// the comments.
for line := range strings.Lines(metrics) {
if !strings.HasPrefix(line, "#") &&
!strings.Contains(line, `instance="fsn1app1/gitea"`) {
t.Errorf("series %q, without instance=\"fsn1app1/gitea\"", line)
}
}
}
func TestStateKeptAcrossRestarts(t *testing.T) { func TestStateKeptAcrossRestarts(t *testing.T) {
t.Parallel() t.Parallel()
@@ -549,212 +442,6 @@ func TestRulesDirThatDoesNotExistStopsTheStart(t *testing.T) {
": no such file or directory") ": 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) { func TestEveryLineIsAlsoSentToTheRemoteLogEndpoint(t *testing.T) {
t.Parallel() t.Parallel()
@@ -830,8 +517,7 @@ func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
rulesDir: t.TempDir(), rulesDir: t.TempDir(),
"SWWAF_LOG_REMOTE_URL": "syslog+tls://" + endpoint.Addr().String(), "SWWAF_LOG_REMOTE_URL": "syslog+tls://" + endpoint.Addr().String(),
"SWWAF_LOG_REMOTE_BUFFER": "1", "SWWAF_LOG_REMOTE_BUFFER": "1",
metricsToken: token, "SWWAF_METRICS_TOKEN": token,
instanceName: instance,
} }
out := runUntilStopped(t, env, func(url string) { out := runUntilStopped(t, env, func(url string) {
@@ -841,17 +527,16 @@ func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
// last. // last.
metrics := metricsText(t, url+"_smallwebwaf/metrics", token) metrics := metricsText(t, url+"_smallwebwaf/metrics", token)
for _, series := range []string{ for _, series := range []string{
`smallwebwaf_remote_log_lines_sent_total{instance="fsn1app1/gitea"} 0`, "smallwebwaf_remote_log_lines_sent_total 0",
`smallwebwaf_remote_log_buffer_depth{instance="fsn1app1/gitea"} 1`, "smallwebwaf_remote_log_buffer_depth 1",
} { } {
if !strings.Contains(metrics, "\n"+series+"\n") { if !strings.Contains(metrics, "\n"+series+"\n") {
t.Errorf("no %q in the metrics:\n%s", series, metrics) t.Errorf("no %q in the metrics:\n%s", series, metrics)
} }
} }
const dropped = "\nsmallwebwaf_remote_log_lines_dropped_total" + if strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total 0\n") ||
`{instance="fsn1app1/gitea"} ` !strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total ") {
if strings.Contains(metrics, dropped+"0\n") || !strings.Contains(metrics, dropped) {
t.Errorf("no line dropped in the metrics:\n%s", metrics) t.Errorf("no line dropped in the metrics:\n%s", metrics)
} }
@@ -861,13 +546,6 @@ func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
}) })
out.line(t, "type", "request") out.line(t, "type", "request")
// While the lines are sent, the process's lines give the instance name
// too.
line := out.line(t, "msg", "starting")
if line["instance"] != instance {
t.Errorf("start logged with instance %v, want %s", line["instance"], instance)
}
} }
func TestBanIsAlertedAndAnAlertNotSentIsKeptAcrossARestart(t *testing.T) { func TestBanIsAlertedAndAnAlertNotSentIsKeptAcrossARestart(t *testing.T) {
@@ -890,7 +568,6 @@ func TestBanIsAlertedAndAnAlertNotSentIsKeptAcrossARestart(t *testing.T) {
rulesDir: rules, rulesDir: rules,
"SWWAF_ALERT_WEBHOOK_URL": webhook.url, "SWWAF_ALERT_WEBHOOK_URL": webhook.url,
"SWWAF_ALERT_WEBHOOK_HEADERS": "Authorization:Bearer " + adminSecret, "SWWAF_ALERT_WEBHOOK_HEADERS": "Authorization:Bearer " + adminSecret,
instanceName: instance,
} }
runUntilStopped(t, env, func(url string) { runUntilStopped(t, env, func(url string) {
@@ -914,7 +591,7 @@ func TestBanIsAlertedAndAnAlertNotSentIsKeptAcrossARestart(t *testing.T) {
// alerts.json keeps it as smallwebwaf stops, and once started again, // alerts.json keeps it as smallwebwaf stops, and once started again,
// smallwebwaf sends it. // smallwebwaf sends it.
var file struct { var file struct {
Waiting map[string][]struct { Waiting []struct {
Event string `json:"event"` Event string `json:"event"`
} `json:"waiting"` } `json:"waiting"`
} }
@@ -926,10 +603,8 @@ func TestBanIsAlertedAndAnAlertNotSentIsKeptAcrossARestart(t *testing.T) {
err = json.Unmarshal(data, &file) err = json.Unmarshal(data, &file)
} }
waiting := file.Waiting["webhook"] if err != nil || len(file.Waiting) != 1 || file.Waiting[0].Event != "permanent_ban" {
if err != nil || len(waiting) != 1 || waiting[0].Event != "permanent_ban" { t.Fatalf("alerts.json holds %s (%v), want the permanent_ban alert waiting", data, err)
t.Fatalf("alerts.json holds %s (%v), want the permanent_ban alert waiting for "+
"the webhook", data, err)
} }
// It counts the alert sent in the metrics, read here from a client the // It counts the alert sent in the metrics, read here from a client the
@@ -939,18 +614,24 @@ func TestBanIsAlertedAndAnAlertNotSentIsKeptAcrossARestart(t *testing.T) {
webhook.failing.Store(false) webhook.failing.Store(false)
env["SWWAF_ALLOW_NETS"] = localhost env["SWWAF_ALLOW_NETS"] = localhost
env[metricsToken] = token env["SWWAF_METRICS_TOKEN"] = token
runUntilStopped(t, env, func(url string) { runUntilStopped(t, env, func(url string) {
webhook.waitFor(t, "permanent_ban", true) webhook.waitFor(t, "permanent_ban", true)
const ofWebhook = `{destination="webhook",instance="fsn1app1/gitea"}` // As long as that takes, so that a slow test process cannot fail
// the test.
const sent = "\nsmallwebwaf_alerts_sent_total{destination=\"webhook\"} 1\n"
metrics := metricsWith(t, url+"_smallwebwaf/metrics", token, metrics := metricsText(t, url+"_smallwebwaf/metrics", token)
"\nsmallwebwaf_alerts_sent_total"+ofWebhook+" 1\n") for !strings.Contains(metrics, sent) {
time.Sleep(pollInterval)
metrics = metricsText(t, url+"_smallwebwaf/metrics", token)
}
for _, series := range []string{"failed", "suppressed", "dropped"} { for _, series := range []string{"failed", "suppressed", "dropped"} {
zero := "\nsmallwebwaf_alerts_" + series + "_total" + ofWebhook + " 0\n" zero := "\nsmallwebwaf_alerts_" + series + "_total{destination=\"webhook\"} 0\n"
if !strings.Contains(metrics, zero) { if !strings.Contains(metrics, zero) {
t.Errorf("no %q in the metrics:\n%s", zero, metrics) t.Errorf("no %q in the metrics:\n%s", zero, metrics)
} }
@@ -958,75 +639,6 @@ func TestBanIsAlertedAndAnAlertNotSentIsKeptAcrossARestart(t *testing.T) {
}) })
} }
func TestBanIsAlertedToSlackAndNtfy(t *testing.T) {
t.Parallel()
const (
ntfyToken = "tk_0123456789abcdefghijklmnopq"
token = "abcdef0123456789abcdef0123456789"
client = "203.0.113.9"
)
slack, ntfy := startDestination(t), startDestination(t)
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
trustedProxies: localhost + "/32",
rateLimitPerDay: "1",
// The metrics are read from 127.0.0.1, which no limit counts.
"SWWAF_ALLOW_NETS": localhost + "/32",
metricsToken: token,
instanceName: instance,
"SWWAF_ALERT_SLACK_WEBHOOK_URL": slack.url,
"SWWAF_ALERT_NTFY_URL": ntfy.url,
"SWWAF_ALERT_NTFY_TOKEN": ntfyToken,
}
runUntilStopped(t, env, func(url string) {
// The client's second request breaks the day limit, and bans it;
// Slack and ntfy are each sent the alert.
wantStatus(t, url, client, http.StatusOK)
wantStatus(t, url, client, http.StatusForbidden)
var message struct {
Text string `json:"text"`
}
slackPost := slack.firstPost(t)
err := json.Unmarshal([]byte(slackPost.body), &message)
if err != nil || !strings.HasPrefix(message.Text, "*fsn1app1/gitea: ban*\n") ||
!strings.Contains(message.Text, "\nclient: "+client+"\n") {
t.Errorf("Slack was sent %s", slackPost.body)
}
ntfyPost := ntfy.firstPost(t)
if ntfyPost.header.Get("Title") != "fsn1app1/gitea: ban" ||
ntfyPost.header.Get("Authorization") != "Bearer "+ntfyToken ||
!strings.Contains(ntfyPost.body, "\nclient: "+client+"\n") {
t.Errorf("ntfy was sent %s, with the headers %v", ntfyPost.body,
ntfyPost.header)
}
// The metrics count it for each, and give no series for the
// webhook, which is not set. As long as that takes, so that a slow
// test process cannot fail the test.
sent := []string{
"\nsmallwebwaf_alerts_sent_total{destination=\"slack\"," +
"instance=\"fsn1app1/gitea\"} 1\n",
"\nsmallwebwaf_alerts_sent_total{destination=\"ntfy\"," +
"instance=\"fsn1app1/gitea\"} 1\n",
}
metrics := metricsWith(t, url+"_smallwebwaf/metrics", token, sent...)
if strings.Contains(metrics, `destination="webhook"`) {
t.Errorf("the metrics give the webhook:\n%s", metrics)
}
})
}
func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) { func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
@@ -1168,7 +780,7 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
"SWWAF_REQUEST_MAX_BYTES": "100M", "SWWAF_REQUEST_MAX_BYTES": "100M",
"SWWAF_RESPONSE_MAX_BYTES": "5G", "SWWAF_RESPONSE_MAX_BYTES": "5G",
"SWWAF_ALLOW_NETS": "", "SWWAF_ALLOW_NETS": "",
rateLimitExemptNets: "", "SWWAF_RATE_LIMIT_EXEMPT_NETS": "",
"SWWAF_DENY_NETS": "", "SWWAF_DENY_NETS": "",
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000", "SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
"SWWAF_RATE_LIMIT_PER_HOUR": "10000", "SWWAF_RATE_LIMIT_PER_HOUR": "10000",
@@ -1286,70 +898,6 @@ func metricsText(t *testing.T, url, token string) string {
return string(body) return string(body)
} }
// metricsWith asks for the metrics at url with token until they hold each
// of series, as long as that takes, so that a slow test process cannot
// fail the test, and returns them.
func metricsWith(t *testing.T, url, token string, series ...string) string {
t.Helper()
for {
metrics := metricsText(t, url, token)
missing := slices.ContainsFunc(series, func(one string) bool {
return !strings.Contains(metrics, one)
})
if !missing {
return metrics
}
time.Sleep(pollInterval)
}
}
// 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 // askAsAdmin sends a request with method to url, with body and
// adminSecret, and checks that it is answered 200. // adminSecret, and checks that it is answered 200.
func askAsAdmin(t *testing.T, method, url, body string) { func askAsAdmin(t *testing.T, method, url, body string) {
@@ -1461,66 +1009,6 @@ func saveUntilAnswered(t *testing.T, path, content, url, from string, status int
} }
} }
// destination is a stand-in for SWWAF_ALERT_SLACK_WEBHOOK_URL or
// SWWAF_ALERT_NTFY_URL. It notes each request it is sent, and answers
// 200.
type destination struct {
url string
mu sync.Mutex
posts []destinationPost
}
// destinationPost is a request a destination was sent: its headers and
// its body.
type destinationPost struct {
header http.Header
body string
}
// startDestination starts a destination that takes every alert.
func startDestination(t *testing.T) *destination {
t.Helper()
d := &destination{}
server := httptest.NewServer(http.HandlerFunc(
func(_ http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
d.mu.Lock()
d.posts = append(d.posts, destinationPost{
header: r.Header.Clone(), body: string(body),
})
d.mu.Unlock()
}))
t.Cleanup(server.Close)
d.url = server.URL + "/alerts"
return d
}
// firstPost waits until the destination has been sent a request, and
// returns the first. It waits as long as that takes, so that a slow test
// process cannot fail the test.
func (d *destination) firstPost(t *testing.T) destinationPost {
t.Helper()
for {
d.mu.Lock()
if len(d.posts) > 0 {
post := d.posts[0]
d.mu.Unlock()
return post
}
d.mu.Unlock()
time.Sleep(pollInterval)
}
}
// webhook is a stand-in for SWWAF_ALERT_WEBHOOK_URL. It notes each alert // webhook is a stand-in for SWWAF_ALERT_WEBHOOK_URL. It notes each alert
// it is sent, and answers 204, or 503 while failing. // it is sent, and answers 204, or 503 while failing.
type webhook struct { type webhook struct {
+72 -304
View File
@@ -1,16 +1,12 @@
// Package state keeps smallwebwaf's state in JSON files in // Package state keeps smallwebwaf's state in JSON files in
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes: // SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
// bans.json holds the bans, clients.json each client's counters and // bans.json holds the bans, clients.json each client's counters and
// history, lookups.json GeoJS's answers, reputation.json the last try and // history, lookups.json GeoJS's answers, and alerts.json the cooldowns,
// last good copy of each list fetched from a URL, the DNSBL zones' // the hour under way and the alerts waiting. Load reads them at start,
// verdicts, and AbuseIPDB's scores and checks spent, and alerts.json the // Watch takes in an admin's edit of one while smallwebwaf runs, and Run
// cooldowns, the hour under way, the alerts waiting for each destination // and WriteAll write them. The disk is read and written outside the
// and the anomaly counters. Load // parts' locks, which are held only to take a snapshot or to put in what
// reads them at start, Watch takes in an admin's edit of one while // a file holds, so that no request waits on the disk.
// 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
// snapshot or to put in what a file holds, so that no request waits on
// the disk.
package state package state
import ( import (
@@ -22,22 +18,18 @@ import (
"fmt" "fmt"
"io/fs" "io/fs"
"log/slog" "log/slog"
"maps"
"net/netip" "net/netip"
"os" "os"
"path/filepath" "path/filepath"
"slices"
"sync" "sync"
"time" "time"
"github.com/fsnotify/fsnotify" "github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/alerts" "sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/reputation"
) )
// version is the version of the files' format, the only one read. // version is the version of the files' format, the only one read.
@@ -49,23 +41,17 @@ const fileMode = 0o600
// The state files' names. // The state files' names.
const ( const (
bansJSON = "bans.json" bansJSON = "bans.json"
clientsJSON = "clients.json" clientsJSON = "clients.json"
lookupsJSON = "lookups.json" lookupsJSON = "lookups.json"
reputationJSON = "reputation.json" alertsJSON = "alerts.json"
alertsJSON = "alerts.json"
) )
var ( var (
errVersion = errors.New("unknown version") errVersion = errors.New("unknown version")
// errMissing is for an entry without a field it needs. // errMissing is for an entry without a field it needs.
errMissing = errors.New("has no") errMissing = errors.New("has no")
errCause = errors.New("is not limit, attack or admin") 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`)
) )
// Params are what Load needs. // Params are what Load needs.
@@ -77,17 +63,13 @@ type Params struct {
// is (SWWAF_STATE_COUNTER_INTERVAL). // is (SWWAF_STATE_COUNTER_INTERVAL).
WriteDelay time.Duration WriteDelay time.Duration
CounterInterval time.Duration CounterInterval time.Duration
// Ledger, Limiter, GeoJS, Lists, DNSBL, AbuseIPDB, Alerts and Anomalies // Ledger, Limiter, GeoJS and Alerts hold the state. Alerts also
// hold the state. Alerts also receive a file_error alert for an edit set // receive a file_error alert for an edit set aside, and for a write
// aside, and for a write that fails while smallwebwaf runs. // that fails while smallwebwaf runs.
Ledger *bans.Ledger Ledger *bans.Ledger
Limiter *ratelimit.Limiter Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS GeoJS *lookup.GeoJS
Lists *reputation.Lists Alerts *alerts.Queue
DNSBL *reputation.DNSBL
AbuseIPDB *reputation.AbuseIPDB
Alerts *alerts.Queue
Anomalies *anomaly.Counters
// Now tells the time by which the counters' buckets run out, normally // Now tells the time by which the counters' buckets run out, normally
// time.Now in UTC. // time.Now in UTC.
Now func() time.Time Now func() time.Time
@@ -145,24 +127,12 @@ type lookupsFile struct {
Lookups []lookup.Answer `json:"lookups"` 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. // 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 { type alertsFile struct {
Version int `json:"version"` Version int `json:"version"`
Cooldowns []alerts.Cooldown `json:"cooldowns"` Cooldowns []alerts.Cooldown `json:"cooldowns"`
Hour alerts.Hour `json:"hour"` Hour alerts.Hour `json:"hour"`
Waiting map[string][]alerts.Alert `json:"waiting"` Waiting []alerts.Alert `json:"waiting"`
AnomalyCounters []anomaly.Counter `json:"anomaly_counters"`
} }
// stateFile is the struct of a state file. Once the file is decoded, its // stateFile is the struct of a state file. Once the file is decoded, its
@@ -175,10 +145,10 @@ type stateFile interface {
} }
// Load checks that files can be written in Dir, and reads the state files // Load checks that files can be written in Dir, and reads the state files
// in it into the parts of Params that hold the state. A missing file is // in it into the ledger, the limiter and GeoJS. A missing file is empty
// empty state, as on a first start. A file that does not parse, has an // state, as on a first start. A file that does not parse, has an unknown
// unknown version, or has an entry without a field it needs, is an error // version, or has an entry without a field it needs, is an error that
// that names the file and, where the JSON decoder tells it, the line and // names the file and, where the JSON decoder tells it, the line and
// column, or else the entry. // column, or else the entry.
func Load(params Params) (*Files, error) { func Load(params Params) (*Files, error) {
err := checkWritable(params.Dir) err := checkWritable(params.Dir)
@@ -191,17 +161,16 @@ func Load(params Params) (*Files, error) {
bansRead, bansErr := f.read(bansJSON) bansRead, bansErr := f.read(bansJSON)
clientsRead, clientsErr := f.read(clientsJSON) clientsRead, clientsErr := f.read(clientsJSON)
lookupsRead, lookupsErr := f.read(lookupsJSON) lookupsRead, lookupsErr := f.read(lookupsJSON)
reputationRead, reputationErr := f.read(reputationJSON)
alertsRead, alertsErr := f.read(alertsJSON) alertsRead, alertsErr := f.read(alertsJSON)
err = errors.Join(bansErr, clientsErr, lookupsErr, reputationErr, alertsErr) err = errors.Join(bansErr, clientsErr, lookupsErr, alertsErr)
if err != nil { if err != nil {
return nil, err return nil, err
} }
params.ProcessLog.Info("read the state files", "directory", params.Dir, params.ProcessLog.Info("read the state files", "directory", params.Dir,
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead, "bans", bansRead, "clients", clientsRead, "lookups", lookupsRead,
"lists", reputationRead, "alerts_waiting", alertsRead) "alerts_waiting", alertsRead)
return f, nil return f, nil
} }
@@ -228,13 +197,9 @@ func (f *Files) Run(ctx context.Context) {
case <-bansDue: case <-bansDue:
bansDue = nil bansDue = nil
f.logFailure(bansJSON, f.writeFile(bansJSON)) f.logFailure(f.writeFile(bansJSON))
case <-interval.C: case <-interval.C:
for _, name := range []string{ f.logFailure(f.WriteAll())
bansJSON, clientsJSON, lookupsJSON, reputationJSON, alertsJSON,
} {
f.logFailure(name, f.writeFile(name))
}
} }
} }
} }
@@ -243,7 +208,7 @@ func (f *Files) Run(ctx context.Context) {
// fails does not keep the others from being written. // fails does not keep the others from being written.
func (f *Files) WriteAll() error { func (f *Files) WriteAll() error {
return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON), return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON),
f.writeFile(lookupsJSON), f.writeFile(reputationJSON), f.writeFile(alertsJSON)) f.writeFile(lookupsJSON), f.writeFile(alertsJSON))
} }
// Watch watches Dir until ctx is done, and takes in an admin's edit of a // Watch watches Dir until ctx is done, and takes in an admin's edit of a
@@ -278,7 +243,7 @@ func (f *Files) Watch(ctx context.Context) {
return return
case event := <-watcher.Events: case event := <-watcher.Events:
switch name := filepath.Base(event.Name); name { switch name := filepath.Base(event.Name); name {
case bansJSON, clientsJSON, lookupsJSON, reputationJSON, alertsJSON: case bansJSON, clientsJSON, lookupsJSON, alertsJSON:
f.fileChanged(name) f.fileChanged(name)
} }
case err = <-watcher.Errors: case err = <-watcher.Errors:
@@ -288,9 +253,9 @@ func (f *Files) Watch(ctx context.Context) {
} }
} }
// logFailure logs a write of the state file name that failed, and raises // logFailure logs a write that failed, and raises a file_error alert for
// a file_error alert for it. // it.
func (f *Files) logFailure(name string, err error) { func (f *Files) logFailure(err error) {
if err != nil { if err != nil {
const failed = "writing the state files failed" const failed = "writing the state files failed"
@@ -299,9 +264,7 @@ func (f *Files) logFailure(name string, err error) {
f.params.Alerts.Raise(alerts.Alert{ f.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError, Event: alerts.EventFileError,
Reason: failed, Reason: failed,
Detail: map[string]any{ Detail: map[string]any{"error": err.Error()},
"file": filepath.Join(f.params.Dir, name), "error": err.Error(),
},
}) })
f.params.ProcessLog.Error(failed, "error", err.Error()) f.params.ProcessLog.Error(failed, "error", err.Error())
} }
@@ -424,29 +387,18 @@ func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
f.params.GeoJS.Load(file.Lookups) f.params.GeoJS.Load(file.Lookups)
entries = len(file.Lookups) entries = len(file.Lookups)
case reputationJSON: case alertsJSON:
var file reputationFile var file alertsFile
err := parse(path, data, &file) err := parse(path, data, &file)
if err != nil { if err != nil {
return 0, err return 0, err
} }
err = f.params.Lists.Load(file.Lists) f.params.Alerts.Load(alerts.State{
if err != nil { Cooldowns: file.Cooldowns, Hour: file.Hour, Waiting: file.Waiting,
return 0, fmt.Errorf("%s: %w", path, err) })
} entries = len(file.Waiting)
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) f.sums[name] = sha256.Sum256(data)
@@ -454,41 +406,6 @@ func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
return entries, nil 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 {
Waiting json.RawMessage `json:"waiting"`
}
if json.Unmarshal(data, &written) == nil &&
bytes.HasPrefix(written.Waiting, []byte("[")) {
return 0, fmt.Errorf("%s: %w", path, errWaitingList)
}
var file alertsFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
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)
}
return entries, nil
}
// writeFile writes the state file name from what smallwebwaf holds. An // writeFile writes the state file name from what smallwebwaf holds. An
// edit made since smallwebwaf last read or wrote the file is taken in // edit made since smallwebwaf last read or wrote the file is taken in
// first, so that it is not overwritten, or set aside if it does not // first, so that it is not overwritten, or set aside if it does not
@@ -564,37 +481,32 @@ func (f *Files) setAside(name string, parseErr error) error {
func (f *Files) encode(name string) ([]byte, error) { func (f *Files) encode(name string) ([]byte, error) {
switch name { switch name {
case bansJSON: case bansJSON:
return encodeIndented(bansFile{ file := bansFile{Version: version, Bans: BanEntries(f.params.Ledger.Snapshot())}
Version: version, Bans: BanEntries(f.params.Ledger.Snapshot()),
}) data, err := json.MarshalIndent(file, "", " ")
if err != nil {
return nil, err
}
return append(data, '\n'), nil
case clientsJSON: case clientsJSON:
return encodeOnePerLine("clients", f.params.Limiter.Snapshot()) return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
case lookupsJSON: case lookupsJSON:
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot()) 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 default: // alerts.json
held := f.params.Alerts.Snapshot() held := f.params.Alerts.Snapshot()
file := alertsFile{
return encodeIndented(alertsFile{
Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour, Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour,
Waiting: held.Waiting, AnomalyCounters: f.params.Anomalies.Snapshot(), Waiting: held.Waiting,
}) }
}
}
// encodeIndented encodes file, a state file's struct, indented for an data, err := json.MarshalIndent(file, "", " ")
// admin to read and edit. if err != nil {
func encodeIndented(file any) ([]byte, error) { return nil, err
data, err := json.MarshalIndent(file, "", " ") }
if err != nil {
return nil, err
}
return append(data, '\n'), nil return append(data, '\n'), nil
}
} }
// BanEntries returns held as bans.json lists them, an empty list for // BanEntries returns held as bans.json lists them, an empty list for
@@ -678,8 +590,8 @@ func (f *bansFile) check(data []byte) error {
} }
// check refuses a client without its address, which would count nobody's // check refuses a client without its address, which would count nobody's
// requests, or with requests or bytes in a window but no start, which // requests, or with requests in a window but no start, which would drop
// would drop them and give the client a fresh allowance. // them and give the client a fresh allowance.
func (f *clientsFile) check([]byte) error { func (f *clientsFile) check([]byte) error {
for i, client := range f.Clients { for i, client := range f.Clients {
switch { switch {
@@ -691,12 +603,6 @@ func (f *clientsFile) check([]byte) error {
return missing(i, "hour.start") return missing(i, "hour.start")
case countsWithoutStart(client.Day): case countsWithoutStart(client.Day):
return missing(i, "day.start") 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")
} }
} }
@@ -734,92 +640,9 @@ func (f *lookupsFile) check(data []byte) error {
return nil 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, // 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 // which would hold back no repeat, and an alert waiting without its event
// another name than webhook, slack or ntfy, most likely misspelt, an // or its time.
// alert waiting without its event or its time, and an anomaly counter as
// checkAnomalyCounters does.
func (f *alertsFile) check([]byte) error { func (f *alertsFile) check([]byte) error {
for i, cooldown := range f.Cooldowns { for i, cooldown := range f.Cooldowns {
switch { switch {
@@ -830,75 +653,20 @@ func (f *alertsFile) check([]byte) error {
} }
} }
for _, destination := range slices.Sorted(maps.Keys(f.Waiting)) { for i, alert := range f.Waiting {
if !slices.Contains(alerts.Destinations(), destination) { switch {
return fmt.Errorf("waiting %q %w", destination, errDestination) case alert.Event == "":
} return fmt.Errorf("waiting %w", missing(i, "event"))
case alert.Time.IsZero():
for i, alert := range f.Waiting[destination] { return fmt.Errorf("waiting %w", missing(i, "time"))
switch {
case alert.Event == "":
return fmt.Errorf("waiting %s %w", destination, missing(i, "event"))
case alert.Time.IsZero():
return fmt.Errorf("waiting %s %w", destination, missing(i, "time"))
}
}
}
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 return nil
} }
// missingFromCounter returns the first field counter, an anomaly counter, // countsWithoutStart reports whether b holds requests but no start, which
// needs and has not, or "" when it has them all: what tells it from the // places them in time.
// 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 { func countsWithoutStart(b ratelimit.Buckets) bool {
return b.Start.IsZero() && (b.Current != 0 || b.Previous != 0) return b.Start.IsZero() && (b.Current != 0 || b.Previous != 0)
} }
File diff suppressed because it is too large Load Diff
+5 -7
View File
@@ -7,10 +7,9 @@
# request bans for good, that `sv stop` stops smallwebwaf in order, that # request bans for good, that `sv stop` stops smallwebwaf in order, that
# `docker stop` stops the container without having to kill it, and 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 # a new container on the same volume still refuses the banned client. The
# containers run with SWWAF_LOOKUP_SOURCE=off, so that no address is sent # containers, the volume and both images are removed however the script
# to GeoJS. The containers, the volume and both images are removed however # ends. Building the app needs network access, for nixpkgs' binary cache.
# the script ends. Building the app needs network access, for nixpkgs' # script/check does not run this.
# binary cache. script/check does not run this.
set -eu set -eu
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
@@ -64,13 +63,12 @@ logged() {
} }
# start_container: run the app's container, with the state files on the # start_container: run the app's container, with the state files on the
# volume, a rate limit of one request a minute and no client looked up, # volume and a rate limit of one request a minute, and wait until it is
# and wait until it is healthy. # healthy.
start_container() { start_container() {
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \ docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \
--volume "$VOLUME:/var/lib/smallwebwaf" \ --volume "$VOLUME:/var/lib/smallwebwaf" \
--env SWWAF_RATE_LIMIT_PER_MINUTE=1 \ --env SWWAF_RATE_LIMIT_PER_MINUTE=1 \
--env SWWAF_LOOKUP_SOURCE=off \
"$APP_IMAGE" >/dev/null "$APP_IMAGE" >/dev/null
wait_for "the health check did not pass" healthy wait_for "the health check did not pass" healthy
address="$(docker port "$CONTAINER" 8080/tcp)" address="$(docker port "$CONTAINER" 8080/tcp)"
-19
View File
@@ -1,19 +0,0 @@
#!/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 "$@"