9 Commits
Author SHA1 Message Date
clawbot 2421cdc273 Lower limits for listed AS numbers and countries (closes #21)
check / check (push) Waiting to run
SWWAF_ASN_LIMIT_PERCENT and SWWAF_COUNTRY_LIMIT_PERCENT give the clients
of the AS numbers and countries they list that percentage of every rate
and byte limit, rounded down; SWWAF_ASN_BYTES_PERCENT and
SWWAF_COUNTRY_BYTES_PERCENT take its place for the byte limits of those
they list; SWWAF_UNKNOWN_LIMIT_PERCENT (100) covers clients without a
country. The lowest applies. While one lowers a limit, a request waits
for its client's lookup, and SWWAF_LOOKUP_SOURCE=off stops the start. Log
lines give limit_percent and bytes_percent with their settings; ban
notes, and so alerts, give the broken limit's.

Judgement call: a client without a country is unknown, whatever its AS number.
Judgement call: bytes_percent and its setting are log fields SPEC does not name.
Rule suppressed: funlen on FromEnvironment, one line per setting.

Model: opus-5-5
2026-10-07 14:22:01 +02:00
clawbot f35e3ddfe8 Byte limits per client over a minute, an hour and a day (closes #20)
check / check (push) Waiting to run
SWWAF_BYTES_LIMIT_PER_MINUTE, _PER_HOUR and _PER_DAY (10G, 20G, 50G)
and SWWAF_BYTES_COUNT (both). A request's bytes are counted once its
answer has ended, for a request passed to the app that the rate limits
count; what a WebSocket carries each way, once it closes. Bytes over a
limit ban the client as a broken rate limit does, and cut nothing
short. clients.json keeps the byte buckets, the log line's counts carry
the byte totals, ban notes say what the limit is on, and the limit hits
metric is labelled by kind.

Judgement call: limit_hit names a byte window minute_bytes, hour_bytes
or day_bytes, as counts names the byte totals.
Judgement call: in observe mode, the bytes of a request enforce mode
would have refused are not counted.

Model: opus-5-5
2026-10-07 13:13:06 +02:00
clawbot 0dc26041dc make tidy, and make check failing on an untidy go.mod or go.sum (closes #99)
check / check (push) Waiting to run
script/tidy runs `go mod tidy` in a new tidy stage of the Dockerfile,
on the test phase's Go image, and writes go.mod and go.sum back into the
working tree with `docker build --output`. The test phase now runs
`go mod tidy -diff` before the tests and fails naming `make tidy`; it
comes first because a missing go.sum line otherwise fails the tests with
Go's own message.

Judgement call: script/tidy's build has no tag, unlike the other builds
in script/: it makes no image.

Model: opus-5-5
2026-10-07 10:30:46 +02:00
clawbot c80753c56e Look clients up in the IPinfo Lite file with SWWAF_LOOKUP_SOURCE=file (closes #22)
check / check (push) Waiting to run
SWWAF_LOOKUP_SOURCE=file looks every client up in the file
SWWAF_LOOKUP_DB_PATH names, without GeoJS. file without the path, the
path with another source, or a file that cannot be read stops the start.
The file is read whole into memory, so overwriting it in place cannot
disturb a lookup, and read again 2 seconds after its last change; a
replacement that cannot be read is logged, counted and sent as a
file_error alert, and the old one stays in use. Metrics give when it
was read and the failed reads. Tests write their databases through
internal/lookup/lookuptest.

Deviation: go.mod and go.sum written by hand; go runs only through make.
Judgement call: the 2-second wait, as the rule files have.

Model: opus-5-5
2026-10-07 10:04:10 +02:00
clawbot 26f4abef7f AS number and country looked up for every client (closes #95)
check / check (push) Waiting to run
GeoJS's geo.json is asked about every new visitor unless
SWWAF_LOOKUP_SOURCE is off. A request waits for its client's first
answer only while a country list or SWWAF_ADD_LOOKUP_HEADERS needs it;
otherwise the answer reaches the client's history and ban notes when it
comes. The AS number and name go beside the country in the request log,
history, ban notes, alerts and lookups.json, with metrics by AS number;
64512 counts as unknown. A client's own X-Client-ASN and
X-Client-Country never reach the app, whatever the setting says, and
make example-app sends no address to GeoJS.

Judgement call: AS numbers are written AS64496, as SPEC's settings write them.
Judgement call: SWWAF_LOOKUP_TIMEOUT is added, default 1s, and cannot be off.

Model: opus-5-5
2026-10-07 08:46:01 +02:00
clawbot f35cbd01cf Instance name on process log lines and every metric (closes #91)
check / check (push) Waiting to run
Process log lines carry instance, as request lines do; the instance name
is read before the other settings, so the line saying a setting is
invalid carries it too. Every metric, Go's and the process's included,
carries the label instance, set once on the registry. README.md says so,
and that Prometheus keeps it as exported_instance unless the scrape sets
honor_labels. An instance name that is not valid UTF-8 stops the start,
as the metrics library panics on such a label.

Tests that read metrics expect the label; one helper replaces the alert
tests' loops that wait for them.

Judgement call: the label is named instance, as in the log lines and
alerts, although Prometheus gives each target a label of that name.

Model: opus-5-5
2026-10-07 06:48:09 +02:00
clawbot 70a8ea1b92 Alerts to Slack and ntfy, each destination with its own queue (closes #90)
check / check (push) Waiting to run
Each alert is posted as a message to the Slack incoming webhook
SWWAF_ALERT_SLACK_WEBHOOK_URL names, and published to the ntfy topic
SWWAF_ALERT_NTFY_URL names, with SWWAF_ALERT_NTFY_TOKEN as a bearer
token and a priority and tag by event. The cooldown and the hourly
limit stay shared; past them, each destination has its own bounded
queue and backoff, and its own sent, failed and dropped counts.
alerts.json keeps the alerts waiting by destination; one whose waiting
is still a list stops the start, saying what to change. A control
character in the ntfy token, or in the instance name ntfy is sent,
stops the start.

Judgement call: messages also give the detail's file, source, error and mode.
Judgement call: alerts_suppressed_total is the same for every destination.

Model: opus-5-5
2026-10-07 05:54:34 +02:00
clawbot 432097ee3f Alerts to a JSON webhook, with a cooldown and an hourly summary (closes #26)
check / check (push) Waiting to run
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, in observe mode too, marked mode observe and worked
out only when the alert would be sent; source_failure for GeoJS;
file_error for a rule or state file with an error. SWWAF_ALERT_EVENTS
chooses; SWWAF_ALERT_COOLDOWN holds back repeats by netblock, file or
source; past SWWAF_ALERT_MAX_PER_HOUR the hour ends in one summary. A
bounded queue, retried with backoff, holds up no request; a 4xx other
than 408 and 429 gives the alert up. alerts.json keeps the queue, the
cooldowns and the hour. Nothing shows the URL's path or query.

Judgement call: the summary's event is summary, which SPEC.md omits.
Judgement call: an admin's ban raises no alert.

Model: opus-5-5
2026-10-07 04:21:24 +02:00
clawbot 5d6f6ffaf9 Admin endpoints for bans and clients on the single listener (closes #27)
check / check (push) Successful in 4m26s
SWWAF_ADMIN_TOKEN, or its _FILE form, opens GET and POST
/_smallwebwaf/bans, DELETE /_smallwebwaf/bans/<client> and GET
/_smallwebwaf/clients/<ip>. Unset, they answer 404; a missing or wrong
token gets 401, in observe mode too. They go through every check, as
the metrics do. POST takes a netblock, not IPv4-mapped and without a
zone, or a client's address, a duration or permanent, and a reason, and
makes an admin ban even while another lasts. DELETE lifts every active
ban covering the address, kept and marked lifted. Bans come back as
bans.json entries; a client as clients.json holds it, with its bans.

Judgement call: answers leave out bans.json's version field.
Judgement call: DELETE takes an address, not a netblock.
Rule suppressed: gosec G304 on a test reading bans.json.

Model: opus-5-5
2026-10-07 01:13:16 +02:00
62 changed files with 10127 additions and 1120 deletions
+4
View File
@@ -61,6 +61,10 @@ 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:
+25 -1
View File
@@ -29,6 +29,12 @@ 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 \
@@ -36,7 +42,25 @@ 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; }
# Build stage. Nothing is wanted from the two phases above; the copies # Tidy stage: `go mod tidy` in the test phase's Go, so that the files it
# writes pass the test phase's check. Nothing else depends on it, so only
# script/tidy, which names the stage after it, builds it.
#
# golang 1.27.1-trixie, 2026-09-19
FROM golang@sha256:3b77fc618ec235a1ab412de7737f120dd507c57e8d87de4cbb7994fb94275ed5 AS tidy
WORKDIR /src
COPY . .
RUN go mod tidy
# go.mod and go.sum alone, which script/tidy writes into the working tree.
FROM scratch AS tidy-files
COPY --from=tidy /src/go.mod /src/go.sum /
# Build stage. Nothing is wanted from the lint and test phases; the copies
# are what make BuildKit build them first, so the image, which needs this # 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.
# #
+7 -3
View File
@@ -1,9 +1,10 @@
.PHONY: bootstrap setup test lint fmt fmt-check check docker hooks build run \ .PHONY: bootstrap setup test lint fmt fmt-check tidy check docker hooks build \
example-app run 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). build and run are for working on the code by hand; # of README.md). tidy writes go.mod and go.sum as `go mod tidy` does,
# which test checks. build and run are for working on the code by hand;
# example-app checks the image with an app built on it. # example-app checks the image with an app built on it.
bootstrap: bootstrap:
@@ -24,6 +25,9 @@ fmt:
fmt-check: fmt-check:
@script/fmt-check @script/fmt-check
tidy:
@script/tidy
check: check:
@script/check @script/check
+619 -216
View File
File diff suppressed because it is too large Load Diff
+4 -1
View File
@@ -5,6 +5,8 @@ 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
) )
@@ -16,6 +18,7 @@ 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
golang.org/x/sys v0.47.0 // indirect go4.org/netipx v0.0.0-20231129151722-fdeea329fbba // indirect
golang.org/x/sys v0.48.0 // indirect
google.golang.org/protobuf v1.36.11 // indirect google.golang.org/protobuf v1.36.11 // indirect
) )
+12 -10
View File
@@ -2,8 +2,6 @@ github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= github.com/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=
@@ -14,10 +12,12 @@ github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJn
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= github.com/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/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/oschwald/maxminddb-golang/v2 v2.7.0 h1:ZcAr3GYc2LYC8aec2mCMX9+QOF0EolH3jDFKRV/Z1+U=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/oschwald/maxminddb-golang/v2 v2.7.0/go.mod h1:DuKJLbbug6TXC0yJXgs1MWifvXHmudRWzMobMIUu04g=
github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU= github.com/prometheus/client_golang v1.24.1 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,15 +26,17 @@ github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi
github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc= github.com/prometheus/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.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0 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=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg=
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBseWJUpBw5I82+2U4M=
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= google.golang.org/protobuf v1.36.11 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=
+884
View File
@@ -0,0 +1,884 @@
// Package alerts sends alerts on bans, on a source that fails and on a
// file with an error to each destination set: to the webhook
// SWWAF_ALERT_WEBHOOK_URL names, each as one JSON object, as the "Alert
// webhook schema" section of SPEC.md describes, to the Slack incoming
// webhook SWWAF_ALERT_SLACK_WEBHOOK_URL names, as a message, and to the
// ntfy topic SWWAF_ALERT_NTFY_URL names. A repeat within
// SWWAF_ALERT_COOLDOWN is held back, and so is an alert past
// SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary. The others wait in a
// bounded queue of each destination's own, so that a destination that is
// slow or unreachable holds up neither the others nor any request. The
// state is written to alerts.json and read from it by the state package.
// Nothing logged names a destination's URL, whose path or query can carry
// a secret.
package alerts
import (
"bytes"
"cmp"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"maps"
"net/http"
"net/netip"
"net/url"
"slices"
"strings"
"sync"
"sync/atomic"
"time"
)
// The events an alert is for, as SWWAF_ALERT_EVENTS names them.
const (
// EventBan is a ban smallwebwaf made.
EventBan = "ban"
// EventPermanentBan is a permanent ban smallwebwaf made, or a ban it
// made permanent.
EventPermanentBan = "permanent_ban"
// EventWAFBlock, EventAnomaly and EventReputationHit come with the
// Core Rule Set, the anomaly thresholds and the reputation sources;
// nothing raises them yet.
EventWAFBlock = "waf_block"
EventAnomaly = "anomaly"
EventReputationHit = "reputation_hit"
// EventSourceFailure is GeoJS failing or refusing smallwebwaf.
EventSourceFailure = "source_failure"
// EventFileError is a rule file or state file edited while smallwebwaf
// runs that does not parse, a replacement of the lookup database that
// cannot be read, or a state file that cannot be written.
EventFileError = "file_error"
// EventSummary is the summary of the alerts an hour held back past
// SWWAF_ALERT_MAX_PER_HOUR. SWWAF_ALERT_EVENTS does not name it.
EventSummary = "summary"
)
// Events returns every event SWWAF_ALERT_EVENTS can name, which is its
// default.
func Events() []string {
return []string{
EventBan, EventPermanentBan, EventWAFBlock, EventAnomaly,
EventReputationHit, EventSourceFailure, EventFileError,
}
}
// The destinations alerts are sent to, as the metrics and alerts.json
// name them.
const (
// DestinationWebhook is the webhook SWWAF_ALERT_WEBHOOK_URL names.
DestinationWebhook = "webhook"
// DestinationSlack is the Slack incoming webhook
// SWWAF_ALERT_SLACK_WEBHOOK_URL names.
DestinationSlack = "slack"
// DestinationNtfy is the ntfy topic SWWAF_ALERT_NTFY_URL names.
DestinationNtfy = "ntfy"
)
// Destinations returns every destination alerts can be sent to.
func Destinations() []string {
return []string{DestinationWebhook, DestinationSlack, DestinationNtfy}
}
const (
// queueSize is the most alerts that wait to be sent to a destination.
// Past it, the oldest is dropped.
queueSize = 1000
// sendTimeout bounds one request to a destination.
sendTimeout = 10 * time.Second
// After a request to a destination fails, the alert is sent again a
// second later, and retryDelayFactor times as long after each further
// failure in a row, up to a minute.
firstRetryDelay = time.Second
retryDelayFactor = 2
maxRetryDelay = time.Minute
// maxAnswerBytes is the most of a destination's answer that is read.
maxAnswerBytes = 64 << 10
)
var (
errStatus = errors.New("the destination answered")
// errRefused is a 4xx answer other than 408 and 429: the destination
// refuses the alert itself, and would refuse it again.
errRefused = errors.New("the destination refused the alert, answering")
)
// Params are what New needs. With none of WebhookURL, SlackURL and
// NtfyURL set, no alert is sent.
type Params struct {
// WebhookURL is where each alert is posted as JSON
// (SWWAF_ALERT_WEBHOOK_URL), nil while it is unset. WebhookHeaders
// are sent with each (SWWAF_ALERT_WEBHOOK_HEADERS).
WebhookURL *url.URL
WebhookHeaders http.Header
// SlackURL is the Slack incoming webhook each alert is posted to as a
// message (SWWAF_ALERT_SLACK_WEBHOOK_URL), nil while it is unset.
SlackURL *url.URL
// NtfyURL is the ntfy topic each alert is published to
// (SWWAF_ALERT_NTFY_URL), nil while it is unset. NtfyToken, unless
// empty, is sent with each as a bearer token (SWWAF_ALERT_NTFY_TOKEN).
NtfyURL *url.URL
NtfyToken string
// Events are the events alerts are sent for (SWWAF_ALERT_EVENTS).
Events []string
// Cooldown is how long a repeat of an alert is held back
// (SWWAF_ALERT_COOLDOWN), 0 for no time. MaxPerHour is the most alerts
// sent in an hour (SWWAF_ALERT_MAX_PER_HOUR), 0 for no limit.
Cooldown time.Duration
MaxPerHour int
// Instance is SWWAF_INSTANCE_NAME, which every alert gives.
Instance string
// Now tells the time of an alert, normally time.Now in UTC.
Now func() time.Time
// ProcessLog receives the requests to a destination that fail.
ProcessLog *slog.Logger
}
// Alert is one alert, as the webhook is sent it and alerts.json holds it,
// with the fields of the "Alert webhook schema" section of SPEC.md. ASN,
// ASName and Country are, for a ban, the client's as the ban's notes give
// them.
//
//nolint:tagliatelle // SPEC.md's alert webhook schema names its fields in snake_case
type Alert struct {
Instance string `json:"instance"`
Time time.Time `json:"time"`
Event string `json:"event"`
Client netip.Addr `json:"client"`
Netblock netip.Prefix `json:"netblock"`
ASN string `json:"asn"`
ASName string `json:"as_name"`
Country string `json:"country"`
// Reason is a short sentence, and Detail what is particular to the
// event: for a file_error, its "file", and for a source_failure, its
// "source", which the cooldown tells repeats by.
Reason string `json:"reason"`
Detail map[string]any `json:"detail"`
// SuppressedRepeats is how many repeats of the alert the cooldown
// held back since the last one let through.
SuppressedRepeats int `json:"suppressed_repeats"`
}
// Cooldown is, for an event on a netblock, or about a file or a source,
// when the last alert let through was raised, and how many repeats the
// cooldown has held back since, as alerts.json holds it.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Cooldown struct {
Event string `json:"event"`
Netblock netip.Prefix `json:"netblock"`
File string `json:"file,omitempty"`
Source string `json:"source,omitempty"`
Sent time.Time `json:"sent"`
SuppressedRepeats int `json:"suppressed_repeats"`
}
// Hour is the hour under way, by the clock, as alerts.json holds it: when
// it started, how many alerts were let through in it, and how many were
// held back in it past MaxPerHour, by event, for its summary.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Hour struct {
Start time.Time `json:"start"`
Sent int `json:"sent"`
HeldBack map[string]int `json:"held_back"`
}
// State is what alerts.json holds: the cooldowns, the hour under way, and
// for each destination set, the alerts waiting to be sent to it, oldest
// first.
type State struct {
Cooldowns []Cooldown `json:"cooldowns"`
Hour Hour `json:"hour"`
Waiting map[string][]Alert `json:"waiting"`
}
// Counts are, for a destination, how many alerts it took, how many
// requests to it failed, and how many alerts were dropped from its full
// queue or given up as it refused them.
type Counts struct {
Sent, Failed, Dropped int64
}
// Queue takes the alerts raised, holds back those it must, and sends the
// others to each destination set, from a queue of the destination's own.
// It is safe for concurrent use.
type Queue struct {
params Params
// destinations are the destinations set, in the order of
// Destinations.
destinations []*destination
mu sync.Mutex
// cooldowns are the alerts last let through, by event and netblock,
// file or source.
cooldowns map[cooldownKey]*Cooldown
hour Hour
suppressed atomic.Int64
}
// destination is a destination set, with the alerts waiting to be sent
// to it. Its mu is taken after the Queue's, never before.
type destination struct {
// name is how the metrics and alerts.json name the destination, and
// setting the setting that is its URL, which the log names in place
// of the URL.
name string
setting string
url *url.URL
// message returns the body an alert is posted with, and the headers
// sent with it.
message func(alert *Alert) ([]byte, http.Header, error)
// httpClient follows no redirect: a redirect is a failure.
httpClient *http.Client
processLog *slog.Logger
// queued receives a value when an alert joins the queue, unless one
// waits already, so that run looks at the queue again.
queued chan struct{}
mu sync.Mutex
// waiting are the alerts waiting to be sent, oldest first.
waiting []*Alert
sent, failed, dropped atomic.Int64
}
// cooldownKey is what makes an alert a repeat of another: the same event
// on the same netblock, and about the same file or source, as its detail
// names them. Each is empty for an alert without one.
type cooldownKey struct {
event string
netblock netip.Prefix
file string
source string
}
// cooldownKeyOf returns what makes another alert a repeat of alert.
func cooldownKeyOf(alert *Alert) cooldownKey {
file, _ := alert.Detail["file"].(string)
source, _ := alert.Detail["source"].(string)
return cooldownKey{alert.Event, alert.Netblock, file, source}
}
// New returns a Queue with no alert yet.
func New(params Params) *Queue {
q := &Queue{
params: params,
cooldowns: map[cooldownKey]*Cooldown{},
hour: Hour{HeldBack: map[string]int{}},
}
if params.WebhookURL != nil {
q.addDestination(DestinationWebhook, "SWWAF_ALERT_WEBHOOK_URL",
params.WebhookURL, q.webhookMessage)
}
if params.SlackURL != nil {
q.addDestination(DestinationSlack, "SWWAF_ALERT_SLACK_WEBHOOK_URL",
params.SlackURL, slackMessage)
}
if params.NtfyURL != nil {
q.addDestination(DestinationNtfy, "SWWAF_ALERT_NTFY_URL",
params.NtfyURL, q.ntfyMessage)
}
return q
}
// Raise sends alert, which names its event and what is particular to it,
// unless no destination is set or SWWAF_ALERT_EVENTS leaves its event
// out. It gives alert the instance and the time. An alert that repeats
// the last one let through less than Cooldown before is held back and
// counted, and the next one let through gives that count. Past MaxPerHour
// alerts let through in the hour under way, by the clock, an alert is
// held back for that hour's summary instead, which is sent once the hour
// has ended; it starts no cooldown, and the repeats held back before it
// are given by the next alert let through. Raise never waits: an alert
// let through joins the queue of each destination, from which Run sends
// it, and with queueSize alerts waiting for a destination, the oldest is
// dropped.
func (q *Queue) Raise(alert Alert) {
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, alert.Event) {
return
}
q.mu.Lock()
defer q.mu.Unlock()
now := q.params.Now()
alert.Instance = q.params.Instance
alert.Time = now
if q.repeat(&alert, now) {
q.suppressed.Add(1)
return
}
q.endHour(now)
if q.params.MaxPerHour > 0 && q.hour.Sent >= q.params.MaxPerHour {
q.hour.HeldBack[alert.Event]++
q.suppressed.Add(1)
return
}
q.startCooldown(&alert, now)
q.hour.Sent++
q.queue(&alert)
}
// WouldSend reports whether Raise would let an alert for event on
// netblock through now: a destination is set, SWWAF_ALERT_EVENTS chooses
// event, no alert for event on netblock was let through less than
// Cooldown before, and fewer than MaxPerHour alerts have been let through
// in the hour under way. Unlike Raise, it counts nothing.
func (q *Queue) WouldSend(event string, netblock netip.Prefix) bool {
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, event) {
return false
}
q.mu.Lock()
defer q.mu.Unlock()
now := q.params.Now()
last, found := q.cooldowns[cooldownKey{event: event, netblock: netblock}]
if q.params.Cooldown > 0 && found && now.Sub(last.Sent) < q.params.Cooldown {
return false
}
q.endHour(now)
return q.params.MaxPerHour == 0 || q.hour.Sent < q.params.MaxPerHour
}
// Run sends the alerts waiting to each destination, from its own queue,
// as destination.run does, until ctx is done. It also ends each hour as
// Raise does, so that the hour's summary is sent as it ends. With no
// destination set, it returns at once.
func (q *Queue) Run(ctx context.Context) {
if len(q.destinations) == 0 {
return
}
var sending sync.WaitGroup
for _, d := range q.destinations {
sending.Go(func() { d.run(ctx) })
}
for {
q.mu.Lock()
untilHourEnds := q.hour.Start.Add(time.Hour).Sub(q.params.Now())
q.mu.Unlock()
select {
case <-ctx.Done():
sending.Wait()
return
case <-time.After(untilHourEnds):
q.mu.Lock()
q.endHour(q.params.Now())
q.mu.Unlock()
}
}
}
// Counts returns the counts of the destination name, all 0 for one not
// set.
func (q *Queue) Counts(name string) Counts {
for _, d := range q.destinations {
if d.name == name {
return Counts{
Sent: d.sent.Load(), Failed: d.failed.Load(), Dropped: d.dropped.Load(),
}
}
}
return Counts{}
}
// DestinationsSet returns the destinations set, in the order of
// Destinations.
func (q *Queue) DestinationsSet() []string {
names := make([]string, 0, len(q.destinations))
for _, d := range q.destinations {
names = append(names, d.name)
}
return names
}
// Suppressed is how many alerts were held back: by the cooldown, and past
// MaxPerHour. No destination is sent such an alert.
func (q *Queue) Suppressed() int64 {
return q.suppressed.Load()
}
// Snapshot returns the queue's state, as alerts.json holds it, with the
// cooldowns sorted by netblock, then by event, file and source.
func (q *Queue) Snapshot() State {
q.mu.Lock()
defer q.mu.Unlock()
state := State{
Cooldowns: make([]Cooldown, 0, len(q.cooldowns)),
Hour: q.hour,
Waiting: map[string][]Alert{},
}
state.Hour.HeldBack = maps.Clone(q.hour.HeldBack)
for _, cooldown := range q.cooldowns {
state.Cooldowns = append(state.Cooldowns, *cooldown)
}
slices.SortFunc(state.Cooldowns, func(a, b Cooldown) int {
return cmp.Or(a.Netblock.Compare(b.Netblock), cmp.Compare(a.Event, b.Event),
cmp.Compare(a.File, b.File), cmp.Compare(a.Source, b.Source))
})
for _, d := range q.destinations {
state.Waiting[d.name] = d.snapshot()
}
return state
}
// Load puts state, read from alerts.json, in place of the queue's state.
// Each cooldown's netblock is masked to its length, so that
// 203.0.113.9/24 is 203.0.113.0/24. The alerts waiting for a destination
// that is not set are dropped, and so are the oldest past queueSize
// alerts waiting for one that is.
func (q *Queue) Load(state State) {
q.mu.Lock()
defer q.mu.Unlock()
q.cooldowns = map[cooldownKey]*Cooldown{}
for _, cooldown := range state.Cooldowns {
cooldown.Netblock = cooldown.Netblock.Masked()
key := cooldownKey{cooldown.Event, cooldown.Netblock, cooldown.File, cooldown.Source}
q.cooldowns[key] = &cooldown
}
q.hour = state.Hour
q.hour.HeldBack = maps.Clone(state.Hour.HeldBack)
if q.hour.HeldBack == nil {
q.hour.HeldBack = map[string]int{}
}
for _, d := range q.destinations {
d.load(state.Waiting[d.name])
}
}
// addDestination adds a destination: its name, the setting that gives
// its URL, that URL, target, and message, which makes the messages sent
// to it.
func (q *Queue) addDestination(
name, setting string, target *url.URL,
message func(alert *Alert) ([]byte, http.Header, error),
) {
q.destinations = append(q.destinations, &destination{
name: name,
setting: setting,
url: target,
message: message,
httpClient: &http.Client{
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
},
},
processLog: q.params.ProcessLog,
queued: make(chan struct{}, 1),
})
}
// repeat reports whether alert, raised at now, repeats the last one let
// through less than Cooldown before, and counts it if it does.
func (q *Queue) repeat(alert *Alert, now time.Time) bool {
if q.params.Cooldown == 0 {
return false
}
last, found := q.cooldowns[cooldownKeyOf(alert)]
if !found || now.Sub(last.Sent) >= q.params.Cooldown {
return false
}
last.SuppressedRepeats++
return true
}
// startCooldown gives alert, let through at now, the count of the repeats
// held back since the last one let through, and notes alert as the last
// one let through.
func (q *Queue) startCooldown(alert *Alert, now time.Time) {
if q.params.Cooldown == 0 {
return
}
key := cooldownKeyOf(alert)
last, found := q.cooldowns[key]
if found {
alert.SuppressedRepeats = last.SuppressedRepeats
}
q.cooldowns[key] = &Cooldown{
Event: alert.Event, Netblock: alert.Netblock, File: key.file, Source: key.source,
Sent: now,
}
}
// endHour ends the hour under way, if now is past it: it queues that
// hour's summary when alerts were held back in it past MaxPerHour, and
// forgets the cooldowns that have run out with no repeat held back, which
// no alert needs any more.
func (q *Queue) endHour(now time.Time) {
start := now.Truncate(time.Hour)
if !start.After(q.hour.Start) {
return
}
heldBack := 0
for _, count := range q.hour.HeldBack {
heldBack += count
}
if heldBack > 0 {
q.queue(&Alert{
Instance: q.params.Instance,
Time: now,
Event: EventSummary,
Reason: fmt.Sprintf("%d alerts held back in the hour from %s, past the %d "+
"an hour SWWAF_ALERT_MAX_PER_HOUR allows", heldBack,
q.hour.Start.Format(time.RFC3339), q.params.MaxPerHour),
Detail: map[string]any{
"hour": q.hour.Start, "count": heldBack, "events": q.hour.HeldBack,
},
})
}
q.hour = Hour{Start: start, HeldBack: map[string]int{}}
for key, cooldown := range q.cooldowns {
if now.Sub(cooldown.Sent) >= q.params.Cooldown && cooldown.SuppressedRepeats == 0 {
delete(q.cooldowns, key)
}
}
}
// queue adds alert to the alerts waiting for each destination.
func (q *Queue) queue(alert *Alert) {
for _, d := range q.destinations {
d.add(alert)
}
}
// run sends the alerts waiting, oldest first, until ctx is done. An alert
// stays in the queue until the destination answers it with a 2xx status,
// or refuses it with a 4xx status other than 408 and 429: a refused alert
// is logged, counted as dropped, and given up, so that the next is sent.
// Any other request that fails is logged, and the alert sent again
// firstRetryDelay later, retryDelayFactor times as long after each
// further failure in a row, up to maxRetryDelay.
func (d *destination) run(ctx context.Context) {
var (
retryDelay time.Duration
retryAt time.Time
)
for {
alert := d.oldest()
var due <-chan time.Time // nil while no alert waits
if alert != nil {
due = time.After(time.Until(retryAt))
}
select {
case <-ctx.Done():
return
case <-d.queued:
case <-due:
err := d.send(ctx, alert)
switch {
case err == nil:
d.remove(alert)
d.sent.Add(1)
retryDelay = 0
retryAt = time.Time{}
case errors.Is(err, errRefused):
d.remove(alert)
d.failed.Add(1)
d.dropped.Add(1)
retryDelay = 0
retryAt = time.Time{}
d.processLog.Warn("gave up an alert "+d.setting+" refused",
"event", alert.Event, "error", err.Error())
case ctx.Err() == nil: // not cut off as smallwebwaf stops
d.failed.Add(1)
retryDelay = min(max(retryDelayFactor*retryDelay, firstRetryDelay),
maxRetryDelay)
retryAt = time.Now().Add(retryDelay)
d.processLog.Warn("sending an alert to "+d.setting+" failed",
"error", err.Error(), "sending_again_in", retryDelay.String())
}
}
}
}
// add adds alert to the alerts waiting, first dropping the oldest while
// queueSize wait, and has run look at the queue again.
func (d *destination) add(alert *Alert) {
d.mu.Lock()
defer d.mu.Unlock()
if len(d.waiting) == queueSize {
d.waiting = slices.Delete(d.waiting, 0, 1)
d.dropped.Add(1)
}
d.waiting = append(d.waiting, alert)
select {
case d.queued <- struct{}{}:
default: // a value waits already
}
}
// load puts waiting, read from alerts.json, in place of the alerts
// waiting, as add adds them.
func (d *destination) load(waiting []Alert) {
d.mu.Lock()
d.waiting = nil
d.mu.Unlock()
for _, alert := range waiting {
d.add(&alert)
}
}
// snapshot returns the alerts waiting, oldest first.
func (d *destination) snapshot() []Alert {
d.mu.Lock()
defer d.mu.Unlock()
waiting := make([]Alert, 0, len(d.waiting))
for _, alert := range d.waiting {
waiting = append(waiting, *alert)
}
return waiting
}
// oldest returns the oldest alert waiting, nil when none waits.
func (d *destination) oldest() *Alert {
d.mu.Lock()
defer d.mu.Unlock()
if len(d.waiting) == 0 {
return nil
}
return d.waiting[0]
}
// remove takes alert, which run has sent or given up, out of the queue,
// unless it has been dropped from it, or load has replaced the queue,
// since run took it. Only the oldest alert is ever dropped, so alert is
// the oldest if it is there at all.
func (d *destination) remove(alert *Alert) {
d.mu.Lock()
defer d.mu.Unlock()
if len(d.waiting) > 0 && d.waiting[0] == alert {
d.waiting = slices.Delete(d.waiting, 0, 1)
}
}
// send posts alert to the destination, as message makes it, and returns
// an error unless the destination answers with a 2xx status: one that
// wraps errRefused for a 4xx status other than 408 and 429. No error
// names the destination's URL, whose path or query can carry a secret.
func (d *destination) send(ctx context.Context, alert *Alert) error {
body, header, err := d.message(alert)
if err != nil {
return err
}
ctx, cancel := context.WithTimeout(ctx, sendTimeout)
defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodPost, d.url.String(),
bytes.NewReader(body))
if err != nil {
return fmt.Errorf("make the request: %w", err)
}
maps.Copy(req.Header, header)
res, err := d.httpClient.Do(req)
if err != nil {
// The client's error names the URL: only what went wrong is kept.
if urlErr, ok := errors.AsType[*url.Error](err); ok {
return urlErr.Err
}
return err
}
defer func() {
_ = res.Body.Close()
}()
// Read, so that the connection can be used again.
_, _ = io.Copy(io.Discard, io.LimitReader(res.Body, maxAnswerBytes))
switch status := res.StatusCode; {
case status >= http.StatusOK && status < http.StatusMultipleChoices:
return nil
case status >= http.StatusBadRequest && status < http.StatusInternalServerError &&
status != http.StatusRequestTimeout && status != http.StatusTooManyRequests:
return fmt.Errorf("%w %s", errRefused, res.Status)
default:
return fmt.Errorf("%w %s", errStatus, res.Status)
}
}
// webhookMessage returns alert as JSON, for the webhook, and the headers
// sent with it: WebhookHeaders, and its Content-Type.
func (q *Queue) webhookMessage(alert *Alert) ([]byte, http.Header, error) {
body, err := json.Marshal(alert)
if err != nil {
return nil, nil, fmt.Errorf("encode the alert: %w", err)
}
header := http.Header{}
maps.Copy(header, q.params.WebhookHeaders)
header.Set("Content-Type", "application/json")
return body, header, nil
}
// slackMessage returns alert as a message for a Slack incoming webhook,
// in JSON: its title in bold, then its text, with &, < and > escaped, as
// Slack asks, so that nothing in them is read as a link or a mention.
func slackMessage(alert *Alert) ([]byte, http.Header, error) {
escape := strings.NewReplacer("&", "&amp;", "<", "&lt;", ">", "&gt;").Replace
body, err := json.Marshal(map[string]string{
"text": "*" + escape(title(alert)) + "*\n" + escape(text(alert)),
})
if err != nil {
return nil, nil, fmt.Errorf("encode the message: %w", err)
}
return body, http.Header{"Content-Type": {"application/json"}}, nil
}
// ntfyMessage returns alert's text, as the message published to ntfy,
// and the headers sent with it: its title, the priority and the tag of
// its event, and NtfyToken, unless it is empty, as a bearer token.
func (q *Queue) ntfyMessage(alert *Alert) ([]byte, http.Header, error) {
header := http.Header{
"Title": {title(alert)},
"Priority": {ntfyPriority(alert.Event)},
"Tags": {ntfyTag(alert.Event)},
}
if q.params.NtfyToken != "" {
header.Set("Authorization", "Bearer "+q.params.NtfyToken)
}
return []byte(text(alert)), header, nil
}
// ntfyPriority returns the priority an alert for event is published to
// ntfy with: high for an event the admin needs to look at.
func ntfyPriority(event string) string {
switch event {
case EventPermanentBan, EventAnomaly, EventSourceFailure, EventFileError:
return "high"
case EventReputationHit:
return "low"
default: // ban, waf_block and summary
return "default"
}
}
// ntfyTag returns the tag an alert for event is published to ntfy with,
// which ntfy shows as an emoji.
func ntfyTag(event string) string {
switch event {
case EventBan, EventPermanentBan:
return "no_entry"
case EventWAFBlock:
return "shield"
case EventAnomaly:
return "chart_with_upwards_trend"
case EventReputationHit:
return "label"
case EventSourceFailure, EventFileError:
return "warning"
default: // summary
return "bar_chart"
}
}
// title returns the title of alert in Slack and ntfy: the instance and
// the event.
func title(alert *Alert) string {
return alert.Instance + ": " + alert.Event
}
// text returns the text of alert in Slack and ntfy: its reason, then a
// line for each of its client, netblock and country, the file, source,
// error and mode its detail gives, and its suppressed repeats, that it
// has.
func text(alert *Alert) string {
lines := []string{alert.Reason}
if alert.Client.IsValid() {
lines = append(lines, "client: "+alert.Client.String())
}
if alert.Netblock.IsValid() {
lines = append(lines, "netblock: "+alert.Netblock.String())
}
if alert.Country != "" {
lines = append(lines, "country: "+alert.Country)
}
for _, name := range []string{"file", "source", "error", "mode"} {
value, _ := alert.Detail[name].(string)
if value != "" {
lines = append(lines, name+": "+value)
}
}
if alert.SuppressedRepeats > 0 {
lines = append(lines, fmt.Sprintf("suppressed repeats: %d", alert.SuppressedRepeats))
}
return strings.Join(lines, "\n")
}
File diff suppressed because it is too large Load Diff
+16
View File
@@ -0,0 +1,16 @@
package alerts
import "net/http"
// QueueSize is the most alerts that wait to be sent to a destination.
const QueueSize = queueSize
// SetTransport has q's requests to the destination name go through
// transport instead of the network.
func (q *Queue) SetTransport(name string, transport http.RoundTripper) {
for _, d := range q.destinations {
if d.name == name {
d.httpClient.Transport = transport
}
}
}
+14 -11
View File
@@ -65,13 +65,16 @@ 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{Limit: 1000, Window: "minute"}) bans.Notes{Kind: "requests", Limit: 1000, Window: "minute"})
attack := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(), byteLimit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.3/32"),
midnight(), bans.Notes{Kind: "bytes", Limit: 10 << 30, Window: "hour"})
attack, _ := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
bans.Notes{RuleID: "git-dir", Target: "path"}) 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 {
@@ -101,12 +104,12 @@ func TestLiftedBanForALimitRefusesNothingAndMakesNoBanLonger(t *testing.T) {
// kept, and counted among the earlier bans. // kept, and counted among the earlier bans.
now := midnight().Add(30 * time.Minute) now := midnight().Add(30 * time.Minute)
_, banned := ledger.Check(netblock.Addr(), now) _, banned, _ := ledger.Check(netblock.Addr(), now)
if banned { if banned {
t.Error("the lifted ban refuses") t.Error("the lifted ban refuses")
} }
ban := ledger.BanForLimit(netblock, now, bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if ban.Expires.Sub(ban.Start) != time.Hour || if ban.Expires.Sub(ban.Start) != time.Hour ||
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) { ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("the next ban lasts %s with earlier bans %+v, want 1h and 1 for a limit", t.Errorf("the next ban lasts %s with earlier bans %+v, want 1h and 1 for a limit",
@@ -134,7 +137,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")
} }
@@ -145,7 +148,7 @@ func TestLiftedBanForAnAttackRefusesNothingAndMakesNoBanLonger(t *testing.T) {
} }
// The next clear sign of attack bans for seven days, as a first does. // The next clear sign of attack bans for seven days, as a first does.
ban := ledger.BanForAttack(netblock, now, bans.Notes{}) ban, _ := ledger.BanForAttack(netblock, now, bans.Notes{})
if ban.Expires.Sub(ban.Start) != 7*day { if ban.Expires.Sub(ban.Start) != 7*day {
t.Errorf("the next ban for an attack ends at %s, want seven days on", ban.Expires) t.Errorf("the next ban for an attack ends at %s, want seven days on", ban.Expires)
} }
@@ -155,7 +158,7 @@ func TestLoadEditCountsTheBansAnAdminMade(t *testing.T) {
t.Parallel() t.Parallel()
ledger := bans.New(defaultRules()) ledger := bans.New(defaultRules())
made := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(), made, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
bans.Notes{}) bans.Notes{})
atStart := bans.Ban{ atStart := bans.Ban{
Netblock: netip.MustParsePrefix("203.0.113.2/32"), Netblock: netip.MustParsePrefix("203.0.113.2/32"),
@@ -220,7 +223,7 @@ 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 || 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)
@@ -261,11 +264,11 @@ func TestLiftLiftsEveryActiveBanCoveringTheClient(t *testing.T) {
wantChanged(t, ledger, true) wantChanged(t, ledger, true)
if _, banned := ledger.Check(client, now); banned { if _, banned, _ := ledger.Check(client, now); banned {
t.Error("the client is still banned") t.Error("the client is still banned")
} }
if _, banned := ledger.Check(other.Addr(), now); !banned { if _, banned, _ := ledger.Check(other.Addr(), now); !banned {
t.Error("the other client's ban was lifted") t.Error("the other client's ban was lifted")
} }
+139 -42
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 show a clear sign of // netblocks of clients that break a rate limit or a byte limit or show a
// attack, and those an admin makes, with their notes, as the "Bans" // clear sign of attack, and those an admin makes, with their notes, as
// section of SPEC.md describes. The bans are kept in memory, and written // the "Bans" section of SPEC.md describes. The bans are kept in memory,
// to bans.json and read from it by the state package. // and written to bans.json and read from it by the state package.
package bans package bans
import ( import (
@@ -91,22 +91,35 @@ 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 {
// Country is the client's country, when it was looked up. // ASN, ASName and Country are the client's AS number, AS name and
// 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"`
// Limit, Window and Count are, for a ban for a broken limit, the limit // Kind, Limit, Window and Count are, for a ban for a broken limit,
// that was broken, its window, "minute", "hour" or "day", and the // what the limit was on, "requests" for a rate limit or "bytes" for a
// count reached: the client's requests in the window, the one that // byte limit, the limit that was broken, its window, "minute", "hour"
// broke the limit included. These are the requests that counted // or "day", and the count reached: the client's requests, or bytes, in
// toward the ban, and the window is the time over which they came. // the window, those of the request that broke the limit included.
// These are what counted toward the ban, and the window is the time
// over which they came.
Kind string `json:"kind,omitempty"`
Limit int64 `json:"limit,omitempty"` 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"`
// Request is the request that broke the limit, or that was the clear // Request is the request that broke the limit, or whose bytes broke
// sign of attack. // 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
@@ -194,39 +207,43 @@ func (l *Ledger) Changed() <-chan struct{} {
// a ban on a netblock client is in is active, and returns that ban, with // a ban on a netblock client is in is active, and returns that ban, with
// the request counted among those it refused. A ban for a clear sign of // the request counted among those it refused. A ban for a clear sign of
// attack is made permanent by the request: the netblock is malicious. // attack is made permanent by the request: the netblock is malicious.
func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool) { // The last result reports whether the request made the ban permanent.
func (l *Ledger) Check(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 return Ban{}, false, false
} }
ban.Notes.Requests++ ban.Notes.Requests++
ban.Notes.Refused++ ban.Notes.Refused++
if ban.Cause == CauseAttack && !ban.Permanent() { madePermanent := ban.Cause == CauseAttack && !ban.Permanent()
if madePermanent {
ban.Expires = time.Time{} ban.Expires = time.Time{}
l.markChanged() l.markChanged()
} }
return *ban, true return *ban, true, madePermanent
} }
// Find is Check without counting the request among those the ban // Find is Check without counting the request among those the ban
// refused: in observe mode a ban refuses nothing. // refused, and without making the ban permanent: in observe mode a ban
func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool) { // refuses nothing. The last result reports whether Check would have made
// 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 return Ban{}, false, false
} }
return *ban, true return *ban, true, ban.Cause == CauseAttack && !ban.Permanent()
} }
// 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
@@ -244,28 +261,80 @@ func activeBan(bans []Ban, now time.Time) *Ban {
} }
// BanForLimit bans netblock at now for a broken limit, with notes, and // BanForLimit bans netblock at now for a broken limit, with notes, and
// returns the ban. A first ban lasts LimitBanDuration. A ban made within // returns the ban, and true. A first ban lasts LimitBanDuration. A ban
// LimitBanRepeatWindow after the netblock's ban that ended last, other // made within LimitBanRepeatWindow after the netblock's ban that ended
// than one for a clear sign of attack or a lifted one, lasts repeatFactor // last, other than one for a clear sign of attack or a lifted one, lasts
// times as long as that one. A ban that would be longer than // repeatFactor times as long as that one. A ban that would be longer
// MaxBanDuration is permanent instead. If a ban on netblock is still // than MaxBanDuration is permanent instead. If a ban on netblock is still
// active, as when two of its requests break a limit at once, that ban is // active, as when two of its requests break a limit at once, that ban is
// returned and no other is made. The ledger fills in the notes' Refused // returned with false, and no other is made. The ledger fills in the
// and EarlierBans itself, and gives the ban the reason "requests per // notes' Refused and EarlierBans itself, and gives the ban the reason
// <Window> over the limit of <Limit>", from the notes. // "<Kind> per <Window> over the limit of <Limit>", from the notes, such
func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes) Ban { // as "requests per minute over the limit of 1000".
reason := fmt.Sprintf("requests per %s over the limit of %d", func (l *Ledger) BanForLimit(
notes.Window, notes.Limit) netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, true)
}
return l.ban(netblock, now, CauseLimit, reason, notes) // WouldBanForLimit returns what BanForLimit would, without making the ban:
// 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
// notes, and returns the ban, as BanForLimit does. A first ban lasts // notes, and returns the ban, and whether it made it, as BanForLimit
// AttackBanDuration; once the netblock has had one that was not lifted, // does. A first ban lasts AttackBanDuration; once the netblock has had
// the next is permanent. Its reason is "matched the rule <RuleID>". // one that was not lifted, the next is permanent. Its reason is "matched
func (l *Ledger) BanForAttack(netblock netip.Prefix, now time.Time, notes Notes) Ban { // the rule <RuleID>".
return l.ban(netblock, now, CauseAttack, "matched the rule "+notes.RuleID, notes) func (l *Ledger) BanForAttack(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, true)
}
// 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
@@ -357,6 +426,28 @@ 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
@@ -483,10 +574,12 @@ 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. // 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
// 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, netblock netip.Prefix, now time.Time, cause, reason string, notes Notes, keep bool,
) Ban { ) (Ban, bool) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
@@ -497,7 +590,7 @@ func (l *Ledger) ban(
if found { if found {
active := activeBan(*bans, now) active := activeBan(*bans, now)
if active != nil { if active != nil {
return *active return *active, false
} }
held = *bans held = *bans
@@ -513,11 +606,15 @@ 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()
return ban return ban, true
} }
// earlierBans returns how many bans a netblock with the bans held, oldest // earlierBans returns how many bans a netblock with the bans held, oldest
+179 -44
View File
@@ -21,7 +21,7 @@ func TestRepeatsTripleUntilPermanent(t *testing.T) {
// Each ban is followed by another as soon as it ends: 1, 3, 9, 27 and // Each ban is followed by another as soon as it ends: 1, 3, 9, 27 and
// 81 hours. // 81 hours.
for i, hours := range []int{1, 3, 9, 27, 81} { for i, hours := range []int{1, 3, 9, 27, 81} {
ban := ledger.BanForLimit(netblock, now, bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
length := time.Duration(hours) * time.Hour length := time.Duration(hours) * time.Hour
if !ban.Expires.Equal(now.Add(length)) || if !ban.Expires.Equal(now.Add(length)) ||
@@ -35,12 +35,12 @@ func TestRepeatsTripleUntilPermanent(t *testing.T) {
// The sixth would last 243 hours, more than seven days: it is // The sixth would last 243 hours, more than seven days: it is
// permanent, and never ends. // permanent, and never ends.
ban := ledger.BanForLimit(netblock, now, bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() { if !ban.Permanent() {
t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires) t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires)
} }
_, banned := ledger.Check(netblock.Addr(), now.Add(100*365*day)) _, banned, _ := ledger.Check(netblock.Addr(), now.Add(100*365*day))
if !banned { if !banned {
t.Error("a permanent ban ended") t.Error("a permanent ban ended")
} }
@@ -64,8 +64,8 @@ func TestRepeatWindowRunsOut(t *testing.T) {
ledger := bans.New(defaultRules()) ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32") netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{}) second, _ := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{})
if second.Expires.Sub(second.Start) != tc.want || if second.Expires.Sub(second.Start) != tc.want ||
second.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) { second.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
@@ -83,7 +83,7 @@ func TestFirstBanLongerThanTheMaximumIsPermanent(t *testing.T) {
rules.LimitBanDuration = rules.MaxBanDuration + time.Hour rules.LimitBanDuration = rules.MaxBanDuration + time.Hour
ledger := bans.New(rules) ledger := bans.New(rules)
ban := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(), ban, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
bans.Notes{}) bans.Notes{})
if !ban.Permanent() { if !ban.Permanent() {
t.Errorf("first ban ends at %s, want a permanent one", ban.Expires) t.Errorf("first ban ends at %s, want a permanent one", ban.Expires)
@@ -103,7 +103,7 @@ func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) {
now := midnight() now := midnight()
for i := range 14 { for i := range 14 {
ban := ledger.BanForLimit(netblock, now, bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Expires.After(ban.Start) { if !ban.Expires.After(ban.Start) {
t.Fatalf("ban %d starts at %s and ends at %s", i+1, ban.Start, ban.Expires) t.Fatalf("ban %d starts at %s and ends at %s", i+1, ban.Start, ban.Expires)
} }
@@ -111,7 +111,7 @@ func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) {
now = ban.Expires now = ban.Expires
} }
ban := ledger.BanForLimit(netblock, now, bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() { if !ban.Permanent() {
t.Errorf("15th ban ends at %s, want a permanent one", ban.Expires) t.Errorf("15th ban ends at %s, want a permanent one", ban.Expires)
} }
@@ -123,12 +123,22 @@ func TestBrokenLimitDuringABanMakesNoOther(t *testing.T) {
ledger := bans.New(defaultRules()) ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32") netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) first, made := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
again := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{}) if !made {
t.Error("the first ban was not made")
}
if again != first || len(ledger.Bans(netblock)) != 1 { again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
t.Errorf("a limit broken during a ban gave %+v and %d bans, want %+v and 1",
again, len(ledger.Bans(netblock)), first) if made || again != first || len(ledger.Bans(netblock)) != 1 {
t.Errorf("a limit broken during a ban gave %+v, made %t, and %d bans, "+
"want %+v, not made, and 1", again, made, len(ledger.Bans(netblock)), first)
}
again, made = ledger.BanForAttack(netblock, midnight().Add(time.Minute), bans.Notes{})
if made || again != first {
t.Errorf("an attack during a ban gave %+v, made %t, want %+v, not made",
again, made, first)
} }
} }
@@ -137,21 +147,21 @@ func TestCheckRefusesWhileTheBanLastsAndCountsTheRefusals(t *testing.T) {
ledger := bans.New(defaultRules()) ledger := bans.New(defaultRules())
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})
for range 3 { for range 3 {
got, banned := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond)) got, banned, _ := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond))
if !banned || got.Start != ban.Start { if !banned || got.Start != ban.Start {
t.Fatalf("check during the ban gives %+v and %t", got, banned) t.Fatalf("check during the ban gives %+v and %t", got, banned)
} }
} }
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight()) _, banned, _ := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight())
if banned { if banned {
t.Error("another netblock is banned") t.Error("another netblock is banned")
} }
_, banned = ledger.Check(netblock.Addr(), ban.Expires) _, banned, _ = ledger.Check(netblock.Addr(), ban.Expires)
if banned { if banned {
t.Error("the ban did not end") t.Error("the ban did not end")
} }
@@ -169,14 +179,14 @@ func TestFindCountsNothing(t *testing.T) {
ledger := bans.New(defaultRules()) ledger := bans.New(defaultRules())
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 || 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")
} }
@@ -198,7 +208,7 @@ func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) {
d := netip.MustParsePrefix("2001:db8::/64") d := netip.MustParsePrefix("2001:db8::/64")
now := midnight() now := midnight()
first := ledger.BanForLimit(a, now, bans.Notes{}) first, _ := ledger.BanForLimit(a, now, bans.Notes{})
ledger.BanForLimit(b, now, bans.Notes{}) ledger.BanForLimit(b, now, bans.Notes{})
ledger.BanForLimit(c, now, bans.Notes{}) ledger.BanForLimit(c, now, bans.Notes{})
@@ -233,8 +243,8 @@ func TestFullLedgerDropsTheEarlierBanOfTheNetblockBannedAgain(t *testing.T) {
ledger := bans.New(rules) ledger := bans.New(rules)
netblock := netip.MustParsePrefix("203.0.113.9/32") netblock := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
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 || held[0] != second || if len(held) != 1 || held[0] != second ||
@@ -251,7 +261,7 @@ func TestRequestDuringAnAttackBanMakesItPermanent(t *testing.T) {
netblock := netip.MustParsePrefix("203.0.113.9/32") netblock := netip.MustParsePrefix("203.0.113.9/32")
notes := bans.Notes{RuleID: "env-file", Target: "path"} notes := bans.Notes{RuleID: "env-file", Target: "path"}
ban := ledger.BanForAttack(netblock, midnight(), notes) ban, _ := ledger.BanForAttack(netblock, midnight(), notes)
if !ban.Expires.Equal(midnight().Add(7*day)) || ban.Cause != bans.CauseAttack || if !ban.Expires.Equal(midnight().Add(7*day)) || ban.Cause != bans.CauseAttack ||
ban.Notes.RuleID != "env-file" || ledger.Made(bans.CauseAttack) != 1 || ban.Notes.RuleID != "env-file" || ledger.Made(bans.CauseAttack) != 1 ||
ledger.Made(bans.CauseLimit) != 0 { ledger.Made(bans.CauseLimit) != 0 {
@@ -262,23 +272,31 @@ 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. // In observe mode the ban refuses nothing, and stays as it is, while
got, _ := ledger.Find(netblock.Addr(), midnight().Add(time.Hour)) // Find tells that the request would have made it permanent.
if got.Permanent() { got, _, wouldMakePermanent := ledger.Find(netblock.Addr(), midnight().Add(time.Hour))
t.Fatal("a request found under the ban made it permanent") if got.Permanent() || ledger.Bans(netblock)[0].Permanent() || !wouldMakePermanent {
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)
} }
// A request it refuses makes it permanent, and bans.json due. wantChanged(t, ledger, false)
got, _ = ledger.Check(netblock.Addr(), midnight().Add(time.Hour))
if !got.Permanent() || !ledger.Bans(netblock)[0].Permanent() { // A request it refuses makes it permanent, says so, and makes
t.Fatalf("after a request during the ban, it is %+v, want it permanent", got) // bans.json due.
got, _, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(time.Hour))
if !madePermanent || !got.Permanent() || !ledger.Bans(netblock)[0].Permanent() {
t.Fatalf("after a request during the ban, it is %+v, made permanent %t, "+
"want it made permanent", got, madePermanent)
} }
wantChanged(t, ledger, true) wantChanged(t, ledger, true)
_, banned := ledger.Check(netblock.Addr(), midnight().Add(100*365*day)) // The next request finds it permanent already.
if !banned { _, banned, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(100*365*day))
t.Error("the permanent ban ended") if !banned || madePermanent {
t.Errorf("a later request is banned %t, and made the ban permanent %t, "+
"want banned by the permanent ban", banned, madePermanent)
} }
} }
@@ -289,8 +307,8 @@ func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
netblock := netip.MustParsePrefix("203.0.113.9/32") netblock := netip.MustParsePrefix("203.0.113.9/32")
// A ban for a broken limit before does not count. // A ban for a broken limit before does not count.
first := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
second := ledger.BanForAttack(netblock, first.Expires, bans.Notes{}) second, _ := ledger.BanForAttack(netblock, first.Expires, bans.Notes{})
if second.Expires.Sub(second.Start) != 7*day { if second.Expires.Sub(second.Start) != 7*day {
t.Fatalf("the first ban for an attack lasts %s, want 7 days", t.Fatalf("the first ban for an attack lasts %s, want 7 days",
@@ -299,14 +317,14 @@ func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
// Once that has run out without a request, the netblock is served, and // Once that has run out without a request, the netblock is served, and
// its next clear sign of attack bans it for good. // its next clear sign of attack bans it for good.
_, banned := ledger.Check(netblock.Addr(), second.Expires) _, banned, _ := ledger.Check(netblock.Addr(), second.Expires)
if banned { if banned {
t.Fatal("the ban did not end") t.Fatal("the ban did not end")
} }
// Its notes show the earlier ban for an attack that makes it permanent, // Its notes show the earlier ban for an attack that makes it permanent,
// beside the one for a limit. // beside the one for a limit.
third := ledger.BanForAttack(netblock, second.Expires.Add(30*day), bans.Notes{}) third, _ := ledger.BanForAttack(netblock, second.Expires.Add(30*day), bans.Notes{})
if !third.Permanent() || if !third.Permanent() ||
third.Notes.EarlierBans != (bans.EarlierBans{Limit: 1, Attack: 1}) { third.Notes.EarlierBans != (bans.EarlierBans{Limit: 1, Attack: 1}) {
t.Errorf("the next ban for an attack is %+v, want a permanent one, "+ t.Errorf("the next ban for an attack is %+v, want a permanent one, "+
@@ -314,6 +332,84 @@ 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 || 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 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()
@@ -322,16 +418,16 @@ func TestAttackBanDoesNotLengthenTheNextBanForALimit(t *testing.T) {
// Three times the seven days would be permanent; a limit broken as the // Three times the seven days would be permanent; a limit broken as the
// ban for an attack ends bans for an hour, as a first broken limit does. // ban for an attack ends bans for an hour, as a first broken limit does.
attack := ledger.BanForAttack(netblock, midnight(), bans.Notes{}) attack, _ := ledger.BanForAttack(netblock, midnight(), bans.Notes{})
limit := ledger.BanForLimit(netblock, attack.Expires, bans.Notes{}) limit, _ := ledger.BanForLimit(netblock, attack.Expires, bans.Notes{})
if limit.Expires.Sub(limit.Start) != time.Hour || limit.Cause != bans.CauseLimit { if limit.Expires.Sub(limit.Start) != time.Hour || limit.Cause != bans.CauseLimit {
t.Errorf("the ban for a limit is %+v, want one of an hour", limit) t.Errorf("the ban for a limit is %+v, want one of an hour", limit)
} }
// And a request during the ban for a limit leaves it as it is. // And a request during the ban for a limit leaves it as it is.
got, _ := ledger.Check(netblock.Addr(), limit.Start) got, _, madePermanent := ledger.Check(netblock.Addr(), limit.Start)
if got.Permanent() { if got.Permanent() || madePermanent {
t.Error("a request during a ban for a limit made it permanent") t.Error("a request during a ban for a limit made it permanent")
} }
} }
@@ -346,7 +442,7 @@ func TestRequestTextsAreCutTo256Bytes(t *testing.T) {
Time: midnight(), Method: long, Host: long, Path: long, Status: 403, UserAgent: long, Time: midnight(), Method: long, Host: long, Path: long, Status: 403, UserAgent: long,
} }
ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request}) ban, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request})
cut := long[:256] cut := long[:256]
want := bans.Request{ want := bans.Request{
@@ -358,6 +454,45 @@ 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{
+12 -12
View File
@@ -42,7 +42,7 @@ func TestSnapshotListsEveryBanByNetblock(t *testing.T) {
high := netip.MustParsePrefix("203.0.113.10/32") high := netip.MustParsePrefix("203.0.113.10/32")
low := netip.MustParsePrefix("203.0.113.9/32") low := netip.MustParsePrefix("203.0.113.9/32")
first := ledger.BanForLimit(v6, midnight(), bans.Notes{}) first, _ := ledger.BanForLimit(v6, midnight(), bans.Notes{})
ledger.BanForLimit(high, midnight(), bans.Notes{}) ledger.BanForLimit(high, midnight(), bans.Notes{})
ledger.BanForLimit(low, midnight(), bans.Notes{}) ledger.BanForLimit(low, midnight(), bans.Notes{})
ledger.BanForLimit(v6, first.Expires, bans.Notes{}) ledger.BanForLimit(v6, first.Expires, bans.Notes{})
@@ -68,7 +68,7 @@ func TestLoadedBansCarryOn(t *testing.T) {
before := bans.New(defaultRules()) before := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32") netblock := netip.MustParsePrefix("203.0.113.9/32")
ban := before.BanForLimit(netblock, midnight(), bans.Notes{Limit: 1}) ban, _ := before.BanForLimit(netblock, midnight(), bans.Notes{Limit: 1})
// Loaded into a new ledger, as across a restart, the ban still refuses // Loaded into a new ledger, as across a restart, the ban still refuses
// while it lasts, and once it has ended a broken limit bans for three // while it lasts, and once it has ended a broken limit bans for three
@@ -76,12 +76,12 @@ func TestLoadedBansCarryOn(t *testing.T) {
after := bans.New(defaultRules()) after := bans.New(defaultRules())
after.Load(before.Snapshot()) after.Load(before.Snapshot())
_, banned := after.Check(netblock.Addr(), ban.Expires.Add(-time.Second)) _, banned, _ := after.Check(netblock.Addr(), ban.Expires.Add(-time.Second))
if !banned { if !banned {
t.Error("the loaded ban does not refuse") t.Error("the loaded ban does not refuse")
} }
again := after.BanForLimit(netblock, ban.Expires, bans.Notes{}) again, _ := after.BanForLimit(netblock, ban.Expires, bans.Notes{})
if again.Expires.Sub(again.Start) != 3*time.Hour || if again.Expires.Sub(again.Start) != 3*time.Hour ||
again.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) { again.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
t.Errorf("the next ban lasts %s with earlier bans %+v, want 3h and 1 for a limit", t.Errorf("the next ban lasts %s with earlier bans %+v, want 3h and 1 for a limit",
@@ -111,7 +111,7 @@ func TestLoadedBanRefusesEveryClientInItsNetblock(t *testing.T) {
"198.51.100.7": true, "198.51.100.7": true,
"198.51.100.8": false, "198.51.100.8": false,
} { } {
_, banned := ledger.Check(netip.MustParseAddr(client), midnight()) _, banned, _ := ledger.Check(netip.MustParseAddr(client), midnight())
if banned != want { if banned != want {
t.Errorf("%s is refused: %t, want %t", client, banned, want) t.Errorf("%s is refused: %t, want %t", client, banned, want)
} }
@@ -150,19 +150,19 @@ 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)
} }
ban, banned = ledger.Check(client, now) ban, banned, _ = ledger.Check(client, now)
if !banned || !ban.Permanent() { if !banned || !ban.Permanent() {
t.Errorf("the client is refused: %t, under %+v, want under the permanent ban", t.Errorf("the client is refused: %t, under %+v, want under the permanent ban",
banned, ban) banned, ban)
} }
// A limit broken now makes no shorter ban over the permanent one. // A limit broken now makes no shorter ban over the permanent one.
ban = ledger.BanForLimit(netblock, now, bans.Notes{}) ban, _ = ledger.BanForLimit(netblock, now, bans.Notes{})
if !ban.Permanent() || len(ledger.Bans(netblock)) != 2 { if !ban.Permanent() || len(ledger.Bans(netblock)) != 2 {
t.Errorf("a broken limit returned %+v and left the netblock %d bans, "+ t.Errorf("a broken limit returned %+v and left the netblock %d bans, "+
"want the permanent ban and 2", ban, len(ledger.Bans(netblock))) "want the permanent ban and 2", ban, len(ledger.Bans(netblock)))
@@ -194,7 +194,7 @@ func TestNextBanWorkedOutFromTheBanThatEndedLast(t *testing.T) {
// Once both have ended, a limit broken within the repeat window bans // Once both have ended, a limit broken within the repeat window bans
// for three times the 9 hours, and the notes count the two bans // for three times the 9 hours, and the notes count the two bans
// before the 9-hour one and it, for a limit, and the admin's. // before the 9-hour one and it, for a limit, and the admin's.
ban := ledger.BanForLimit(netblock, nineHours.Expires.Add(time.Hour), bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, nineHours.Expires.Add(time.Hour), bans.Notes{})
if ban.Expires.Sub(ban.Start) != 27*time.Hour || if ban.Expires.Sub(ban.Start) != 27*time.Hour ||
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 3, Admin: 1}) { ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 3, Admin: 1}) {
t.Errorf("the next ban lasts %s with earlier bans %+v, "+ t.Errorf("the next ban lasts %s with earlier bans %+v, "+
@@ -255,15 +255,15 @@ func TestLoadReplacesTheBansHeld(t *testing.T) {
// bans.json is taken in, that ban is lifted. // bans.json is taken in, that ban is lifted.
ledger.Load([]bans.Ban{kept}) ledger.Load([]bans.Ban{kept})
_, banned := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight()) _, banned, _ := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight())
if banned { if banned {
t.Error("a ban left out of the second load still refuses") t.Error("a ban left out of the second load still refuses")
} }
// The ledger holds one ban, so it makes two more without dropping any. // The ledger holds one ban, so it makes two more without dropping any.
first := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(), first, _ := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(),
bans.Notes{}) bans.Notes{})
second := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(), second, _ := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(),
bans.Notes{}) bans.Notes{})
want := []bans.Ban{first, second, kept} want := []bans.Ban{first, second, kept}
+494 -28
View File
@@ -20,8 +20,10 @@ import (
"strconv" "strconv"
"strings" "strings"
"time" "time"
"unicode"
"unicode/utf8" "unicode/utf8"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/remotelog" "sneak.berlin/go/smallwebwaf/internal/remotelog"
) )
@@ -32,9 +34,10 @@ type Config struct {
ListenAddr string ListenAddr string
// UpstreamURL is the app (SWWAF_UPSTREAM_URL). // UpstreamURL is the app (SWWAF_UPSTREAM_URL).
UpstreamURL *url.URL UpstreamURL *url.URL
// InstanceName is the name each request log line gives as instance // InstanceName is the name every log line and alert gives as instance,
// (SWWAF_INSTANCE_NAME), by default the host's name, which docker sets // and every metric carries as its label instance (SWWAF_INSTANCE_NAME),
// to the first 12 characters of the container's id. // by default the host's name, which docker sets to the first 12
// characters of the container's id.
InstanceName string InstanceName string
// Observe is true in observe mode, when SWWAF_MODE is observe rather // Observe is true in observe mode, when SWWAF_MODE is observe rather
// than enforce: a request that SWWAF_DENY_NETS, a ban, the country // than enforce: a request that SWWAF_DENY_NETS, a ban, the country
@@ -88,6 +91,28 @@ type Config struct {
// limits neither count nor refuse (SWWAF_RATE_LIMIT_EXEMPT_PATHS). // limits neither count nor refuse (SWWAF_RATE_LIMIT_EXEMPT_PATHS).
// Each starts with /. // Each starts with /.
RateLimitExemptPaths []string RateLimitExemptPaths []string
// BytesLimitPerMinute, BytesLimitPerHour and BytesLimitPerDay are the
// most bytes a client's requests may carry in a minute, an hour and a
// day (SWWAF_BYTES_LIMIT_PER_MINUTE, SWWAF_BYTES_LIMIT_PER_HOUR and
// SWWAF_BYTES_LIMIT_PER_DAY). BytesCount is which body bytes count
// toward them (SWWAF_BYTES_COUNT): response, request or both.
BytesLimitPerMinute int64
BytesLimitPerHour int64
BytesLimitPerDay int64
BytesCount string
// LookupSource is where each client's AS number and country are
// looked up (SWWAF_LOOKUP_SOURCE): geojs, file, or off for nowhere.
// LookupDBPath is the lookup database, the IPinfo Lite file looked up
// in while LookupSource is file (SWWAF_LOOKUP_DB_PATH), and "" for any
// other source. A request waits up to LookupTimeout for its client's
// first answer from GeoJS while a setting needs it
// (SWWAF_LOOKUP_TIMEOUT), which cannot be off. AddLookupHeaders is true
// when the app is passed the client's AS number and country in headers
// (SWWAF_ADD_LOOKUP_HEADERS).
LookupSource string
LookupDBPath string
LookupTimeout time.Duration
AddLookupHeaders bool
// DeniedCountries are the countries whose clients are refused // DeniedCountries are the countries whose clients are refused
// (SWWAF_DENIED_COUNTRIES). ExclusivelyAllowedCountries, when not // (SWWAF_DENIED_COUNTRIES). ExclusivelyAllowedCountries, when not
// empty, are the only countries whose clients are let through // empty, are the only countries whose clients are let through
@@ -95,6 +120,22 @@ type Config struct {
// capitals, as GeoJS gives them. // capitals, as GeoJS gives them.
DeniedCountries []string DeniedCountries []string
ExclusivelyAllowedCountries []string ExclusivelyAllowedCountries []string
// The biased thresholds. ASNLimitPercent and CountryLimitPercent give
// the clients of the AS numbers and the countries they list that
// percentage of every rate limit and byte limit
// (SWWAF_ASN_LIMIT_PERCENT and SWWAF_COUNTRY_LIMIT_PERCENT).
// ASNBytesPercent and CountryBytesPercent give those they list a
// percentage of the byte limits in place of that one
// (SWWAF_ASN_BYTES_PERCENT and SWWAF_COUNTRY_BYTES_PERCENT). Each holds
// percentages from 0 to 100, by AS number, written as AS64496, or by
// country, a two-letter code in capitals, as the lookup gives them.
// UnknownLimitPercent is the percentage of every limit a client without
// a country gets (SWWAF_UNKNOWN_LIMIT_PERCENT).
ASNLimitPercent map[string]int64
CountryLimitPercent map[string]int64
ASNBytesPercent map[string]int64
CountryBytesPercent map[string]int64
UnknownLimitPercent int64
// BanResponse is the status a refused client is answered with, 403 // BanResponse is the status a refused client is answered with, 403
// or 429, or 0 to close the connection without an answer // or 429, or 0 to close the connection without an answer
// (SWWAF_BAN_RESPONSE). It answers a banned client, a request that // (SWWAF_BAN_RESPONSE). It answers a banned client, a request that
@@ -135,8 +176,8 @@ type Config struct {
AdminToken string AdminToken string
// MetricsToken is the bearer token a scraper sends for the metrics // MetricsToken is the bearer token a scraper sends for the metrics
// (SWWAF_METRICS_TOKEN), "" while it is unset and the metrics are off. // (SWWAF_METRICS_TOKEN), "" while it is unset and the metrics are off.
// MetricsTopN is how many countries get series of their own in the // MetricsTopN is how many AS numbers and how many countries get series
// metrics (SWWAF_METRICS_TOP_N). // of their own in the metrics (SWWAF_METRICS_TOP_N).
MetricsToken string MetricsToken string
MetricsTopN int MetricsTopN int
// RulesDir is the directory of the rule files (SWWAF_RULES_DIR), read // RulesDir is the directory of the rule files (SWWAF_RULES_DIR), read
@@ -158,6 +199,27 @@ type Config struct {
LogRemoteBuffer int LogRemoteBuffer int
LogRemoteFacility int LogRemoteFacility int
LogRemoteAppName string LogRemoteAppName string
// AlertWebhookURL is where each alert is posted as JSON
// (SWWAF_ALERT_WEBHOOK_URL), nil while it is unset. AlertWebhookHeaders
// are sent with each (SWWAF_ALERT_WEBHOOK_HEADERS).
// AlertSlackWebhookURL is the Slack incoming webhook each alert is
// posted to as a message (SWWAF_ALERT_SLACK_WEBHOOK_URL), and
// AlertNtfyURL the ntfy topic each is published to
// (SWWAF_ALERT_NTFY_URL), each nil while it is unset; AlertNtfyToken,
// unless empty, is sent to ntfy with each (SWWAF_ALERT_NTFY_TOKEN).
// With none of the three URLs set, no alert is sent. AlertEvents are
// the events alerts are sent for (SWWAF_ALERT_EVENTS). A repeat of an
// alert within AlertCooldown is held back (SWWAF_ALERT_COOLDOWN), and
// so is an alert past AlertMaxPerHour in an hour, for the hour's
// summary (SWWAF_ALERT_MAX_PER_HOUR); 0 is off for both.
AlertWebhookURL *url.URL
AlertWebhookHeaders http.Header
AlertSlackWebhookURL *url.URL
AlertNtfyURL *url.URL
AlertNtfyToken string
AlertEvents []string
AlertCooldown time.Duration
AlertMaxPerHour int
// settings are the values read, as given or by default, and the // settings are the values read, as given or by default, and the
// files they were read from, for the log line at start. // files they were read from, for the log line at start.
@@ -168,6 +230,10 @@ type Config struct {
// off. // off.
const off = "off" const off = "off"
// fileSource is the SWWAF_LOOKUP_SOURCE that looks clients up in the
// lookup database, the file SWWAF_LOOKUP_DB_PATH names.
const fileSource = "file"
const ( const (
day = 24 * time.Hour day = 24 * time.Hour
kibibyte = 1 << 10 kibibyte = 1 << 10
@@ -176,7 +242,8 @@ const (
ipv4Bits = 32 ipv4Bits = 32
// minTokenLength is the fewest characters a token may have. // minTokenLength is the fewest characters a token may have.
minTokenLength = 32 minTokenLength = 32
// masked is what the log shows for a token that is set. // masked is what the log shows for a token that is set, and in place of
// a secret in another setting.
masked = "********" masked = "********"
// defaultListenAddr and defaultUpstreamURL are the defaults of // defaultListenAddr and defaultUpstreamURL are the defaults of
// SWWAF_LISTEN_ADDR and SWWAF_UPSTREAM_URL. // SWWAF_LISTEN_ADDR and SWWAF_UPSTREAM_URL.
@@ -208,6 +275,10 @@ var (
"is taken out of every request by Go's HTTP server, so it can never " + "is taken out of every request by Go's HTTP server, so it can never " +
"be logged") "be logged")
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too") errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
errNotLookupSource = errors.New("is not geojs, file or off")
errNeedsLookups = errors.New("it needs each client looked up")
errNeedsDBPath = errors.New("it names the file to look clients up in")
errDBPathUnused = errors.New("only file reads it")
errNotOver4K = errors.New("is not a size of more than 4K, such as 32K") errNotOver4K = errors.New("is not a size of more than 4K, such as 32K")
errNotDurationAboveZero = errors.New( errNotDurationAboveZero = errors.New(
"is not a duration above zero, such as 1h or 7d") "is not a duration above zero, such as 1h or 7d")
@@ -220,6 +291,7 @@ var (
"is not an absolute path, such as /var/lib/smallwebwaf") "is not an absolute path, such as /var/lib/smallwebwaf")
errShortToken = errors.New("is shorter than 32 characters") errShortToken = errors.New("is shorter than 32 characters")
errNotMode = errors.New("is not enforce or observe") errNotMode = errors.New("is not enforce or observe")
errNotBytesCount = errors.New("is not response, request or both")
errNotPathPrefix = errors.New( errNotPathPrefix = errors.New(
"is not a path prefix starting with /, such as /assets/") "is not a path prefix starting with /, such as /assets/")
errNotBoolean = errors.New("is not true or false") errNotBoolean = errors.New("is not true or false")
@@ -230,7 +302,25 @@ var (
errNotFacility = errors.New("is not a syslog facility such as local0 or daemon") errNotFacility = errors.New("is not a syslog facility such as local0 or daemon")
errNotAppName = errors.New( errNotAppName = errors.New(
"is not 1 to 48 printable ASCII characters without a space, such as gitea") "is not 1 to 48 printable ASCII characters without a space, such as gitea")
errSetTwice = errors.New("set only one of them") errSetTwice = errors.New("set only one of them")
errNotWebhookURL = errors.New(
"is not an http or https URL without a user or a fragment, " +
"such as https://alerts.example/smallwebwaf")
errNotWebhookHeader = errors.New(
"is not a header name followed by : and the header's value, " +
"such as Authorization:Bearer <token>")
errControlCharacter = errors.New(
"holds a control character, such as the carriage return of a Windows line end")
errNotAlertEvent = errors.New(
"is not ban, permanent_ban, waf_block, anomaly, reputation_hit, " +
"source_failure or file_error")
errNotNumberOrOff = errors.New("is not a whole number above zero, such as 60, or off")
errNotUTF8 = errors.New("is not valid UTF-8")
errNotASN = errors.New("is not an AS number such as AS64496")
errNotPercentItem = errors.New(
"is not a code, : and a percentage, such as AS64496:50 or cn:25")
errNotPercent = errors.New("is not a percentage, a whole number from 0 to 100")
errListedTwice = errors.New("is listed twice")
) )
// FromEnvironment reads the settings with lookupEnv, normally // FromEnvironment reads the settings with lookupEnv, normally
@@ -238,13 +328,14 @@ var (
// named by the setting's name with _FILE added names the file, which is // named by the setting's name with _FILE added names the file, which is
// read now (see lookup). A setting that is not set takes its default. A // read now (see lookup). A setting that is not set takes its default. A
// setting that is set but invalid is an error that names it. // setting that is set but invalid is an error that names it.
//
//nolint:funlen // one line for each setting, a list that grows with them
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) { func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
env := &environment{lookupEnv: lookupEnv} env := &environment{lookupEnv: lookupEnv}
hostname, _ := os.Hostname() // "" when the host has no name to give
cfg := &Config{ cfg := &Config{
ListenAddr: env.address("SWWAF_LISTEN_ADDR", defaultListenAddr), ListenAddr: env.address("SWWAF_LISTEN_ADDR", defaultListenAddr),
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", defaultUpstreamURL), UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", defaultUpstreamURL),
InstanceName: env.value("SWWAF_INSTANCE_NAME", hostname), InstanceName: env.instanceName(),
Observe: env.observe("SWWAF_MODE", "enforce"), Observe: env.observe("SWWAF_MODE", "enforce"),
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges), TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"), ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
@@ -263,9 +354,22 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"), RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"),
RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"), RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"),
RateLimitExemptPaths: env.pathPrefixes("SWWAF_RATE_LIMIT_EXEMPT_PATHS", ""), RateLimitExemptPaths: env.pathPrefixes("SWWAF_RATE_LIMIT_EXEMPT_PATHS", ""),
BytesLimitPerMinute: env.size("SWWAF_BYTES_LIMIT_PER_MINUTE", "10G"),
BytesLimitPerHour: env.size("SWWAF_BYTES_LIMIT_PER_HOUR", "20G"),
BytesLimitPerDay: env.size("SWWAF_BYTES_LIMIT_PER_DAY", "50G"),
BytesCount: env.bytesCount("SWWAF_BYTES_COUNT", "both"),
LookupSource: env.lookupSource("SWWAF_LOOKUP_SOURCE", "geojs"),
LookupDBPath: env.value("SWWAF_LOOKUP_DB_PATH", ""),
LookupTimeout: env.durationNotOff("SWWAF_LOOKUP_TIMEOUT", "1s"),
AddLookupHeaders: env.boolean("SWWAF_ADD_LOOKUP_HEADERS", "false"),
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""), DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
ExclusivelyAllowedCountries: env.countries( ExclusivelyAllowedCountries: env.countries(
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""), "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""),
ASNLimitPercent: env.percents("SWWAF_ASN_LIMIT_PERCENT", parseASN),
CountryLimitPercent: env.percents("SWWAF_COUNTRY_LIMIT_PERCENT", parseCountry),
ASNBytesPercent: env.percents("SWWAF_ASN_BYTES_PERCENT", parseASN),
CountryBytesPercent: env.percents("SWWAF_COUNTRY_BYTES_PERCENT", parseCountry),
UnknownLimitPercent: env.percent("SWWAF_UNKNOWN_LIMIT_PERCENT", "100"),
BanResponse: env.banResponse("SWWAF_BAN_RESPONSE", "403"), BanResponse: env.banResponse("SWWAF_BAN_RESPONSE", "403"),
LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"), LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"),
LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"), LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"),
@@ -278,26 +382,32 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"), StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"),
LogRequestHeaders: env.headerNames("SWWAF_LOG_REQUEST_HEADERS", LogRequestHeaders: env.headerNames("SWWAF_LOG_REQUEST_HEADERS",
"accept,accept-language,accept-encoding,content-type,origin,range"), "accept,accept-language,accept-encoding,content-type,origin,range"),
AdminToken: env.token("SWWAF_ADMIN_TOKEN"), AdminToken: env.token("SWWAF_ADMIN_TOKEN"),
MetricsToken: env.token("SWWAF_METRICS_TOKEN"), MetricsToken: env.token("SWWAF_METRICS_TOKEN"),
MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"), MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"),
RulesDir: env.value("SWWAF_RULES_DIR", "/etc/smallwebwaf/rules.d"), RulesDir: env.value("SWWAF_RULES_DIR", "/etc/smallwebwaf/rules.d"),
RulesEnabled: env.boolean("SWWAF_RULES_ENABLED", "true"), RulesEnabled: env.boolean("SWWAF_RULES_ENABLED", "true"),
LogRemoteURL: env.logRemoteURL("SWWAF_LOG_REMOTE_URL"), LogRemoteURL: env.logRemoteURL("SWWAF_LOG_REMOTE_URL"),
LogRemoteTLSCAs: env.certificates("SWWAF_LOG_REMOTE_TLS_CA_FILE"), LogRemoteTLSCAs: env.certificates("SWWAF_LOG_REMOTE_TLS_CA_FILE"),
LogRemoteBuffer: env.numberNotOff("SWWAF_LOG_REMOTE_BUFFER", "10000"), LogRemoteBuffer: env.numberNotOff("SWWAF_LOG_REMOTE_BUFFER", "10000"),
LogRemoteFacility: env.facility("SWWAF_LOG_REMOTE_FACILITY", "local0"), LogRemoteFacility: env.facility("SWWAF_LOG_REMOTE_FACILITY", "local0"),
AlertWebhookURL: env.webhookURL("SWWAF_ALERT_WEBHOOK_URL"),
AlertWebhookHeaders: env.webhookHeaders("SWWAF_ALERT_WEBHOOK_HEADERS"),
AlertSlackWebhookURL: env.webhookURL("SWWAF_ALERT_SLACK_WEBHOOK_URL"),
AlertNtfyURL: env.webhookURL("SWWAF_ALERT_NTFY_URL"),
AlertNtfyToken: env.secret("SWWAF_ALERT_NTFY_TOKEN"),
AlertEvents: env.alertEvents("SWWAF_ALERT_EVENTS",
strings.Join(alerts.Events(), ",")),
AlertCooldown: env.duration("SWWAF_ALERT_COOLDOWN", "15m"),
AlertMaxPerHour: env.numberOrOff("SWWAF_ALERT_MAX_PER_HOUR", "60"),
} }
cfg.LogRemoteAppName = env.appName("SWWAF_LOG_REMOTE_APP_NAME", cfg.LogRemoteAppName = env.appName("SWWAF_LOG_REMOTE_APP_NAME",
cfg.InstanceName, cfg.LogRemoteURL != nil) cfg.InstanceName, cfg.LogRemoteURL != nil)
for _, country := range cfg.ExclusivelyAllowedCountries { env.checkInstanceNameForNtfy(cfg.InstanceName, cfg.AlertNtfyURL != nil)
if slices.Contains(cfg.DeniedCountries, country) { env.checkLookupDBPath(cfg)
env.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", env.checkCountriesAndLookups(cfg)
fmt.Errorf("%q %w", country, errOnBothLists))
}
}
if env.err != nil { if env.err != nil {
return nil, env.err return nil, env.err
@@ -326,6 +436,16 @@ func ListenAddrAndUpstreamURL(
return listenAddr, upstreamURL, nil return listenAddr, upstreamURL, nil
} }
// InstanceName reads only SWWAF_INSTANCE_NAME, which may be given as a
// file, as FromEnvironment does, so that the line saying a setting is
// invalid carries it too. A file that cannot be read gives the default
// here, and FromEnvironment then stops the start over it.
func InstanceName(lookupEnv func(string) (string, bool)) string {
env := &environment{lookupEnv: lookupEnv}
return env.instanceName()
}
// privateRanges are the private address ranges, the default trusted // privateRanges are the private address ranges, the default trusted
// proxies. // proxies.
const privateRanges = "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16" const privateRanges = "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16"
@@ -473,6 +593,17 @@ func (e *environment) count(name, defaultValue string) int64 {
return count return count
} }
// bytesCount reads the setting that is which body bytes count toward the
// byte limits: response, request or both.
func (e *environment) bytesCount(name, defaultValue string) string {
value := e.value(name, defaultValue)
if value != "response" && value != "request" && value != "both" {
e.check(name, fmt.Errorf("%q %w", value, errNotBytesCount))
}
return value
}
// pathPrefixes reads a setting that is a list of path prefixes. // pathPrefixes reads a setting that is a list of path prefixes.
func (e *environment) pathPrefixes(name, defaultValue string) []string { func (e *environment) pathPrefixes(name, defaultValue string) []string {
prefixes, err := parsePathPrefixes(e.value(name, defaultValue)) prefixes, err := parsePathPrefixes(e.value(name, defaultValue))
@@ -489,6 +620,87 @@ func (e *environment) countries(name, defaultValue string) []string {
return countries return countries
} }
// percents reads a setting that is a list of AS numbers or countries,
// which parseCode reads, each with a percentage. It is empty by default.
func (e *environment) percents(
name string, parseCode func(string) (string, error),
) map[string]int64 {
percents, err := parsePercents(e.value(name, ""), parseCode)
e.check(name, err)
return percents
}
// percent reads a setting that is a percentage, from 0 to 100.
func (e *environment) percent(name, defaultValue string) int64 {
percent, err := parsePercent(e.value(name, defaultValue))
e.check(name, err)
return percent
}
// lookupSource reads the setting that is where clients are looked up:
// geojs, file, or off.
func (e *environment) lookupSource(name, defaultValue string) string {
source := e.value(name, defaultValue)
if source != "geojs" && source != fileSource && source != off {
e.check(name, fmt.Errorf("%q %w", source, errNotLookupSource))
}
return source
}
// checkLookupDBPath refuses SWWAF_LOOKUP_SOURCE=file without
// SWWAF_LOOKUP_DB_PATH, and SWWAF_LOOKUP_DB_PATH with any other source:
// one source at a time.
func (e *environment) checkLookupDBPath(cfg *Config) {
switch {
case cfg.LookupSource == fileSource && cfg.LookupDBPath == "":
e.check("SWWAF_LOOKUP_SOURCE", fmt.Errorf(
"is file while SWWAF_LOOKUP_DB_PATH is unset; %w", errNeedsDBPath))
case cfg.LookupSource != fileSource && cfg.LookupDBPath != "":
e.check("SWWAF_LOOKUP_DB_PATH", fmt.Errorf(
"is set while SWWAF_LOOKUP_SOURCE is %s; %w", cfg.LookupSource, errDBPathUnused))
}
}
// checkCountriesAndLookups refuses a country on both country lists, and,
// while SWWAF_LOOKUP_SOURCE is off, each setting that needs clients looked
// up: the country lists, SWWAF_ADD_LOOKUP_HEADERS, and the biased
// thresholds, of which SWWAF_UNKNOWN_LIMIT_PERCENT needs them only below
// 100, where it lowers a limit.
func (e *environment) checkCountriesAndLookups(cfg *Config) {
for _, country := range cfg.ExclusivelyAllowedCountries {
if slices.Contains(cfg.DeniedCountries, country) {
e.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES",
fmt.Errorf("%q %w", country, errOnBothLists))
}
}
if cfg.LookupSource != off {
return
}
for _, setting := range []struct {
name string
set bool
}{
{"SWWAF_DENIED_COUNTRIES", len(cfg.DeniedCountries) > 0},
{"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", len(cfg.ExclusivelyAllowedCountries) > 0},
{"SWWAF_ADD_LOOKUP_HEADERS", cfg.AddLookupHeaders},
{"SWWAF_ASN_LIMIT_PERCENT", len(cfg.ASNLimitPercent) > 0},
{"SWWAF_COUNTRY_LIMIT_PERCENT", len(cfg.CountryLimitPercent) > 0},
{"SWWAF_ASN_BYTES_PERCENT", len(cfg.ASNBytesPercent) > 0},
{"SWWAF_COUNTRY_BYTES_PERCENT", len(cfg.CountryBytesPercent) > 0},
{"SWWAF_UNKNOWN_LIMIT_PERCENT", cfg.UnknownLimitPercent < 100},
} {
if setting.set {
e.check(setting.name, fmt.Errorf("is set while SWWAF_LOOKUP_SOURCE is off; %w",
errNeedsLookups))
}
}
}
// headerNames reads a setting that is a list of header names, and // headerNames reads a setting that is a list of header names, and
// returns them in lower case. // returns them in lower case.
func (e *environment) headerNames(name, defaultValue string) []string { func (e *environment) headerNames(name, defaultValue string) []string {
@@ -613,6 +825,19 @@ func (e *environment) facility(name, defaultValue string) int {
return number return number
} }
// instanceName reads SWWAF_INSTANCE_NAME, by default the host's name. It
// must be valid UTF-8: the metrics library panics on a label that is not.
func (e *environment) instanceName() string {
hostname, _ := os.Hostname() // "" when the host has no name to give
value := e.value("SWWAF_INSTANCE_NAME", hostname)
if !utf8.ValidString(value) {
e.check("SWWAF_INSTANCE_NAME", fmt.Errorf("%q %w", value, errNotUTF8))
}
return value
}
// appName reads the setting that is the APP-NAME of the records the log // appName reads the setting that is the APP-NAME of the records the log
// lines are sent in, by default the instance name. Its value is checked // lines are sent in, by default the instance name. Its value is checked
// when it is set, and, while lines are sent, when it is the instance name. // when it is set, and, while lines are sent, when it is the instance name.
@@ -636,6 +861,83 @@ func (e *environment) appName(name, instanceName string, sending bool) string {
return value return value
} }
// checkInstanceNameForNtfy refuses an instance name that holds a control
// character while ntfySet, SWWAF_ALERT_NTFY_URL being set: ntfy is sent
// the instance name in a header, which cannot hold one.
func (e *environment) checkInstanceNameForNtfy(instanceName string, ntfySet bool) {
if ntfySet && strings.ContainsFunc(instanceName, unicode.IsControl) {
e.check("SWWAF_INSTANCE_NAME", fmt.Errorf(
"%q %w, and is sent to ntfy in a header while SWWAF_ALERT_NTFY_URL is set",
instanceName, errControlCharacter))
}
}
// webhookURL reads a setting that is a URL each alert is posted to:
// SWWAF_ALERT_WEBHOOK_URL, SWWAF_ALERT_SLACK_WEBHOOK_URL or
// SWWAF_ALERT_NTFY_URL. Unset or empty, it is nil, and no alert is posted
// there. The log shows ******** in place of its path and query, and an
// error shows none of it, since a webhook or an ntfy topic can carry its
// secret there.
func (e *environment) webhookURL(name string) *url.URL {
value, _ := e.lookup(name)
webhook, logged, err := parseWebhookURL(value)
e.settings = append(e.settings, slog.String(name, logged))
e.check(name, err)
return webhook
}
// webhookHeaders reads the setting that is the headers sent with each
// alert. The log shows each header's value as ********, since a header
// such as Authorization carries a secret.
func (e *environment) webhookHeaders(name string) http.Header {
value, _ := e.lookup(name)
headers, logged, err := parseWebhookHeaders(value)
e.settings = append(e.settings, slog.String(name, logged))
e.check(name, err)
return headers
}
// secret reads a setting that is a secret another service gave, such as
// an ntfy token, "" while it is unset. It is sent in a header, which
// cannot hold a control character, so one in it is an error. The log
// shows ******** in place of a value that is not empty, and an error
// shows none of it.
func (e *environment) secret(name string) string {
value, _ := e.lookup(name)
logged := ""
if value != "" {
logged = masked
}
e.settings = append(e.settings, slog.String(name, logged))
if strings.ContainsFunc(value, unicode.IsControl) {
e.check(name, errControlCharacter)
}
return value
}
// alertEvents reads the setting that is the events alerts are sent for.
func (e *environment) alertEvents(name, defaultValue string) []string {
events, err := parseAlertEvents(e.value(name, defaultValue))
e.check(name, err)
return events
}
// numberOrOff reads a setting that is a whole number above zero, or off,
// which is 0.
func (e *environment) numberOrOff(name, defaultValue string) int {
number, err := parseNumberOrOff(e.value(name, defaultValue))
e.check(name, err)
return number
}
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a // parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
// whole number of days such as 7d, or off. // whole number of days such as 7d, or off.
func parseDuration(value string) (time.Duration, error) { func parseDuration(value string) (time.Duration, error) {
@@ -900,13 +1202,12 @@ func parseCountries(value string) ([]string, error) {
return nil, err return nil, err
} }
known := strings.Fields(countryCodes)
countries := make([]string, 0, len(items)) countries := make([]string, 0, len(items))
for _, item := range items { for _, item := range items {
country := strings.ToUpper(item) country, err := parseCountry(item)
if !slices.Contains(known, country) { if err != nil {
return nil, fmt.Errorf("%q %w", item, errNotCountry) return nil, err
} }
countries = append(countries, country) countries = append(countries, country)
@@ -915,6 +1216,81 @@ func parseCountries(value string) ([]string, error) {
return countries, nil return countries, nil
} }
// parseCountry reads a country code in either case, and returns it in
// capitals.
func parseCountry(value string) (string, error) {
country := strings.ToUpper(value)
if !slices.Contains(strings.Fields(countryCodes), country) {
return "", fmt.Errorf("%q %w", value, errNotCountry)
}
return country, nil
}
// parseASN reads an AS number such as AS64496, in either case, and
// returns it as the lookup gives it: AS and the number, in capitals and
// without leading zeros.
func parseASN(value string) (string, error) {
digits, hasAS := strings.CutPrefix(strings.ToUpper(value), "AS")
number, err := strconv.ParseUint(digits, 10, 32)
if !hasAS || err != nil {
return "", fmt.Errorf("%q %w", value, errNotASN)
}
return "AS" + strconv.FormatUint(number, 10), nil
}
// parsePercents reads a comma-separated list of items, each an AS number
// or a country, which parseCode reads, then : and a percentage, such as
// AS64496:50 or cn:25, and returns each one's percentage. An empty value
// is an empty list. An AS number or country listed twice is an error.
func parsePercents(
value string, parseCode func(string) (string, error),
) (map[string]int64, error) {
items, err := parseList(value)
if err != nil {
return nil, err
}
percents := make(map[string]int64, len(items))
for _, item := range items {
codeText, percentText, found := strings.Cut(item, ":")
if !found {
return nil, fmt.Errorf("%q %w", item, errNotPercentItem)
}
code, err := parseCode(codeText)
if err != nil {
return nil, err
}
percent, err := parsePercent(percentText)
if err != nil {
return nil, err
}
if _, listed := percents[code]; listed {
return nil, fmt.Errorf("%q %w", codeText, errListedTwice)
}
percents[code] = percent
}
return percents, nil
}
// parsePercent reads a percentage, a whole number from 0 to 100.
func parsePercent(value string) (int64, error) {
percent, err := strconv.ParseInt(value, 10, 64)
if err != nil || percent < 0 || percent > 100 {
return 0, fmt.Errorf("%q %w", value, errNotPercent)
}
return percent, nil
}
// headerNameChars are the characters RFC 9110 allows in a header name: // headerNameChars are the characters RFC 9110 allows in a header name:
// letters, digits and these marks. // letters, digits and these marks.
const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" + const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" +
@@ -1051,6 +1427,96 @@ func parseFacility(value string) (int, error) {
return number, nil return number, nil
} }
// parseWebhookURL reads where each alert is posted: http or https, a
// host, and an optional port from 1 to 65535, path and query, without a
// user or a fragment. It returns the URL, and how the log shows it: its
// scheme and host, and ******** in place of its path and query, if it has
// either. An error shows no part of the value. An empty value is no URL.
func parseWebhookURL(value string) (*url.URL, string, error) {
if value == "" {
return nil, "", nil
}
webhook, err := url.Parse(value)
if err != nil {
return nil, "", errNotWebhookURL
}
port, err := strconv.ParseUint(webhook.Port(), 10, 16)
valid := (webhook.Scheme == "http" || webhook.Scheme == "https") &&
webhook.Hostname() != "" && (webhook.Port() == "" || (err == nil && port != 0)) &&
webhook.User == nil && webhook.Opaque == "" && webhook.Fragment == ""
if !valid {
return nil, "", errNotWebhookURL
}
logged := webhook.Scheme + "://" + webhook.Host
if webhook.Path != "" || webhook.RawQuery != "" {
logged += "/" + masked
}
return webhook, logged, nil
}
// parseWebhookHeaders reads a comma-separated list of headers, each its
// name, :, and its value, and returns them, and how the log shows them,
// with each value as ********. An error names the item by its place in
// the list, so that it shows no value. An empty value is an empty list.
func parseWebhookHeaders(value string) (http.Header, string, error) {
headers := http.Header{}
if strings.TrimSpace(value) == "" {
return headers, "", nil
}
logged := []string{}
for i, item := range strings.Split(value, ",") {
name, headerValue, found := strings.Cut(item, ":")
name = strings.TrimSpace(name)
if !found || !IsHeaderName(name) || strings.ContainsAny(headerValue, "\r\n\x00") {
return nil, "", fmt.Errorf("item %d %w", i+1, errNotWebhookHeader)
}
headers.Add(name, strings.TrimSpace(headerValue))
logged = append(logged, name+":"+masked)
}
return headers, strings.Join(logged, ","), nil
}
// parseAlertEvents reads a comma-separated list of the events alerts can
// be sent for.
func parseAlertEvents(value string) ([]string, error) {
events, err := parseList(value)
if err != nil {
return nil, err
}
for _, event := range events {
if !slices.Contains(alerts.Events(), event) {
return nil, fmt.Errorf("%q %w", event, errNotAlertEvent)
}
}
return events, nil
}
// parseNumberOrOff reads a whole number above zero, or off, which is 0.
func parseNumberOrOff(value string) (int, error) {
if value == off {
return 0, nil
}
n, err := strconv.Atoi(value)
if err != nil || n <= 0 {
return 0, fmt.Errorf("%q %w", value, errNotNumberOrOff)
}
return n, nil
}
// appNameMaxLength is the most characters RFC 5424 allows in an // appNameMaxLength is the most characters RFC 5424 allows in an
// APP-NAME. // APP-NAME.
const appNameMaxLength = 48 const appNameMaxLength = 48
+597 -16
View File
@@ -6,9 +6,11 @@ import (
"encoding/json" "encoding/json"
"log/slog" "log/slog"
"maps" "maps"
"net/http"
"net/netip" "net/netip"
"os" "os"
"path/filepath" "path/filepath"
"reflect"
"slices" "slices"
"strings" "strings"
"testing" "testing"
@@ -38,8 +40,21 @@ const (
rateLimitPerHour = "SWWAF_RATE_LIMIT_PER_HOUR" rateLimitPerHour = "SWWAF_RATE_LIMIT_PER_HOUR"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS" rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
bytesLimitPerMinute = "SWWAF_BYTES_LIMIT_PER_MINUTE"
bytesLimitPerHour = "SWWAF_BYTES_LIMIT_PER_HOUR"
bytesLimitPerDay = "SWWAF_BYTES_LIMIT_PER_DAY"
bytesCount = "SWWAF_BYTES_COUNT"
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"
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"
banResponse = "SWWAF_BAN_RESPONSE" banResponse = "SWWAF_BAN_RESPONSE"
limitBanDuration = "SWWAF_LIMIT_BAN_DURATION" limitBanDuration = "SWWAF_LIMIT_BAN_DURATION"
limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW" limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW"
@@ -62,6 +77,22 @@ const (
logRemoteBuffer = "SWWAF_LOG_REMOTE_BUFFER" logRemoteBuffer = "SWWAF_LOG_REMOTE_BUFFER"
logRemoteFacility = "SWWAF_LOG_REMOTE_FACILITY" logRemoteFacility = "SWWAF_LOG_REMOTE_FACILITY"
logRemoteAppName = "SWWAF_LOG_REMOTE_APP_NAME" logRemoteAppName = "SWWAF_LOG_REMOTE_APP_NAME"
alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL"
alertWebhookHeaders = "SWWAF_ALERT_WEBHOOK_HEADERS"
alertSlackWebhookURL = "SWWAF_ALERT_SLACK_WEBHOOK_URL"
alertNtfyURL = "SWWAF_ALERT_NTFY_URL"
alertNtfyToken = "SWWAF_ALERT_NTFY_TOKEN" //nolint:gosec // the setting's name
alertEvents = "SWWAF_ALERT_EVENTS"
alertCooldown = "SWWAF_ALERT_COOLDOWN"
alertMaxPerHour = "SWWAF_ALERT_MAX_PER_HOUR"
)
// defaultAlertEvents is the default of SWWAF_ALERT_EVENTS, and
// defaultAlertCooldown that of SWWAF_ALERT_COOLDOWN.
const (
defaultAlertEvents = "ban,permanent_ban,waf_block,anomaly,reputation_hit," +
"source_failure,file_error"
defaultAlertCooldown = "15m"
) )
// defaultLogRequestHeaders is the default of SWWAF_LOG_REQUEST_HEADERS. // defaultLogRequestHeaders is the default of SWWAF_LOG_REQUEST_HEADERS.
@@ -99,6 +130,16 @@ const (
// off switches a timeout, a size limit or a rate limit off. // off switches a timeout, a size limit or a rate limit off.
const off = "off" const off = "off"
// enabled is true, as a setting's value.
const enabled = "true"
// defaultLookupSource is the default of SWWAF_LOOKUP_SOURCE, and
// fileSource the source that is the lookup database.
const (
defaultLookupSource = "geojs"
fileSource = "file"
)
// environment is a set of environment variables, for FromEnvironment. // environment is a set of environment variables, for FromEnvironment.
type environment map[string]string type environment map[string]string
@@ -155,6 +196,13 @@ func TestDefaults(t *testing.T) {
RulesDir: "/etc/smallwebwaf/rules.d", RulesDir: "/etc/smallwebwaf/rules.d",
RulesEnabled: true, RulesEnabled: true,
}) })
wantLookupSettings(t, cfg, config.Config{
LookupSource: defaultLookupSource, LookupTimeout: time.Second,
})
wantByteLimitSettings(t, cfg, config.Config{
BytesLimitPerMinute: 10 << 30, BytesLimitPerHour: 20 << 30,
BytesLimitPerDay: 50 << 30, BytesCount: "both",
})
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" { if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
t.Errorf("%s is %s", upstreamURL, cfg.UpstreamURL) t.Errorf("%s is %s", upstreamURL, cfg.UpstreamURL)
@@ -185,6 +233,22 @@ func TestDefaults(t *testing.T) {
} }
} }
func TestNoSingleRequestBreaksAByteLimitAtTheDefaults(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{})
// The largest request body and the largest response, both counted.
largest := cfg.RequestMaxBytes + cfg.ResponseMaxBytes
for _, limit := range []int64{
cfg.BytesLimitPerMinute, cfg.BytesLimitPerHour, cfg.BytesLimitPerDay,
} {
if largest > limit {
t.Errorf("a request of %d bytes breaks the byte limit of %d", largest, limit)
}
}
}
func TestValuesAsSet(t *testing.T) { func TestValuesAsSet(t *testing.T) {
t.Parallel() t.Parallel()
@@ -267,6 +331,22 @@ func TestValuesAsSet(t *testing.T) {
wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE") wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE")
} }
func TestByteLimitSettingsAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
bytesLimitPerMinute: "512M",
bytesLimitPerHour: off,
bytesLimitPerDay: "100000",
bytesCount: "response",
})
wantByteLimitSettings(t, cfg, config.Config{
BytesLimitPerMinute: 512 << 20, BytesLimitPerHour: 0,
BytesLimitPerDay: 100000, BytesCount: "response",
})
}
func TestRateLimitExemptPathsAsSet(t *testing.T) { func TestRateLimitExemptPathsAsSet(t *testing.T) {
t.Parallel() t.Parallel()
@@ -477,6 +557,267 @@ func TestAppNameSetStopsTheStartWhileSending(t *testing.T) {
} }
} }
func TestAlertSettingsDefaults(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{})
if cfg.AlertWebhookURL != nil || len(cfg.AlertWebhookHeaders) != 0 ||
cfg.AlertSlackWebhookURL != nil || cfg.AlertNtfyURL != nil ||
cfg.AlertNtfyToken != "" ||
strings.Join(cfg.AlertEvents, ",") != defaultAlertEvents ||
cfg.AlertCooldown != 15*time.Minute || cfg.AlertMaxPerHour != 60 {
t.Errorf("alert settings %v, %v, %v, %v, %q, %v, %s and %d, want no URLs, "+
"no headers, no token, %s, 15m and 60", cfg.AlertWebhookURL,
cfg.AlertWebhookHeaders, cfg.AlertSlackWebhookURL, cfg.AlertNtfyURL,
cfg.AlertNtfyToken, cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour,
defaultAlertEvents)
}
}
func TestAlertSettingsAsSet(t *testing.T) {
t.Parallel()
const (
webhook = "https://alerts.example:8443/hooks/waf?team=ops"
slack = "https://hooks.slack.example/services/T0123/B4567/abcdef"
ntfy = "https://ntfy.example/smallwebwaf-alerts"
token = "tk_0123456789abcdefghijklmnopq"
)
cfg := fromEnvironment(t, environment{
alertWebhookURL: webhook,
alertWebhookHeaders: "Authorization: Bearer abc:def , x-team:ops",
alertSlackWebhookURL: slack,
alertNtfyURL: ntfy,
alertNtfyToken: token,
alertEvents: "ban, file_error",
alertCooldown: "1h",
alertMaxPerHour: "10",
})
headers := http.Header{"Authorization": {"Bearer abc:def"}, "X-Team": {"ops"}}
if cfg.AlertWebhookURL.String() != webhook ||
!reflect.DeepEqual(cfg.AlertWebhookHeaders, headers) ||
cfg.AlertSlackWebhookURL.String() != slack || cfg.AlertNtfyURL.String() != ntfy ||
cfg.AlertNtfyToken != token ||
!slices.Equal(cfg.AlertEvents, []string{"ban", "file_error"}) ||
cfg.AlertCooldown != time.Hour || cfg.AlertMaxPerHour != 10 {
t.Errorf("alert settings %v, %v, %v, %v, %q, %v, %s and %d", cfg.AlertWebhookURL,
cfg.AlertWebhookHeaders, cfg.AlertSlackWebhookURL, cfg.AlertNtfyURL,
cfg.AlertNtfyToken, cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour)
}
}
func TestAlertSettingsSetEmptyOrOff(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
alertWebhookURL: "", alertSlackWebhookURL: "", alertNtfyURL: "",
alertNtfyToken: "", alertEvents: "", alertCooldown: off, alertMaxPerHour: off,
})
if cfg.AlertWebhookURL != nil || cfg.AlertSlackWebhookURL != nil ||
cfg.AlertNtfyURL != nil || cfg.AlertNtfyToken != "" ||
len(cfg.AlertEvents) != 0 || cfg.AlertCooldown != 0 || cfg.AlertMaxPerHour != 0 {
t.Errorf("set empty or off, alert settings %v, %v, %v, %q, %v, %s and %d",
cfg.AlertWebhookURL, cfg.AlertSlackWebhookURL, cfg.AlertNtfyURL,
cfg.AlertNtfyToken, cfg.AlertEvents, cfg.AlertCooldown, cfg.AlertMaxPerHour)
}
}
func TestInvalidAlertSettingStopsTheStart(t *testing.T) {
t.Parallel()
wantStartStopped(t, []struct{ name, value string }{
{alertWebhookURL, "alerts.example/smallwebwaf"},
{alertWebhookURL, "ftp://alerts.example/"},
{alertWebhookURL, "https:///smallwebwaf"},
{alertWebhookURL, "https://user:password@alerts.example/"},
{alertWebhookURL, "https://alerts.example/#top"},
{alertWebhookURL, "https://alerts.example:0/"},
{alertWebhookURL, "https://alerts.example:65536/"},
{alertSlackWebhookURL, "hooks.slack.example/services/T0123"},
{alertSlackWebhookURL, "https://user:password@hooks.slack.example/"},
{alertNtfyURL, "ntfy://ntfy.example/smallwebwaf-alerts"},
{alertNtfyURL, "https://ntfy.example/smallwebwaf-alerts#top"},
{alertWebhookHeaders, "Authorization"},
{alertWebhookHeaders, "X Team:ops"},
{alertWebhookHeaders, ":ops"},
{alertWebhookHeaders, "X-Team:ops,"},
{alertWebhookHeaders, "X-Team:o\r\nps"},
{alertEvents, "bans"},
{alertEvents, "summary"},
{alertEvents, "ban,,file_error"},
{alertCooldown, "0"},
{alertCooldown, "soon"},
{alertMaxPerHour, "0"},
{alertMaxPerHour, "-1"},
{alertMaxPerHour, "1.5"},
})
}
func TestWebhookHeadersAreLoggedMaskedAndNeverShown(t *testing.T) {
t.Parallel()
const secret = "Bearer 0123456789abcdef"
cfg := fromEnvironment(t, environment{
alertWebhookHeaders: "Authorization:" + secret + ",X-Team:ops",
})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
logged := out.String()
if strings.Contains(logged, secret) || strings.Contains(logged, "ops") ||
!strings.Contains(logged,
`"`+alertWebhookHeaders+`":"Authorization:********,X-Team:********"`) {
t.Errorf("the headers are not logged masked: %s", logged)
}
// An item that is not a header is named by its place, not shown.
_, err := config.FromEnvironment(environment{
alertWebhookHeaders: "X-Team:ops," + secret,
}.lookupEnv)
want := alertWebhookHeaders + ": item 2 is not a header name followed by : " +
"and the header's value, such as Authorization:Bearer <token>"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestWebhookURLIsLoggedWithoutItsPathOrQueryAndNeverShown(t *testing.T) {
t.Parallel()
const secret = "T0123/B4567/abcdef"
for value, want := range map[string]string{
"https://hooks.example/services/" + secret: "https://hooks.example/********",
"https://hooks.example:8443?token=" + secret: "https://hooks.example:8443/********",
"http://[2001:db8::1]:8080": "http://[2001:db8::1]:8080",
} {
cfg := fromEnvironment(t, environment{alertWebhookURL: value})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
logged := out.String()
if strings.Contains(logged, secret) ||
!strings.Contains(logged, `"`+alertWebhookURL+`":"`+want+`"`) {
t.Errorf("%s is not logged as %s: %s", value, want, logged)
}
}
// A value that is not such a URL is not shown either.
for _, value := range []string{
"ftp://hooks.example/services/" + secret,
"https://hooks.example/services/%zz" + secret,
} {
_, err := config.FromEnvironment(environment{alertWebhookURL: value}.lookupEnv)
want := alertWebhookURL + ": is not an http or https URL without a user or " +
"a fragment, such as https://alerts.example/smallwebwaf"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
}
func TestSlackAndNtfySettingsAreLoggedWithoutTheirSecrets(t *testing.T) {
t.Parallel()
const token = "tk_0123456789abcdefghijklmnopq"
cfg := fromEnvironment(t, environment{
alertSlackWebhookURL: "https://hooks.slack.example/services/T0123/B4567/abcdef",
alertNtfyURL: "https://ntfy.example/smallwebwaf-alerts",
alertNtfyToken: token,
})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
logged := out.String()
for _, want := range []string{
`"` + alertSlackWebhookURL + `":"https://hooks.slack.example/********"`,
`"` + alertNtfyURL + `":"https://ntfy.example/********"`,
`"` + alertNtfyToken + `":"********"`,
} {
if !strings.Contains(logged, want) {
t.Errorf("no %s in the settings logged: %s", want, logged)
}
}
for _, secret := range []string{"T0123", "smallwebwaf-alerts", token} {
if strings.Contains(logged, secret) {
t.Errorf("%s in the settings logged: %s", secret, logged)
}
}
}
func TestNtfyTokenWithAControlCharacterStopsTheStartWithoutShowingIt(t *testing.T) {
t.Parallel()
// A file saved with Windows line ends keeps the carriage return.
for name, env := range map[string]environment{
"set": {alertNtfyToken: token + "\r"},
"in a file": {alertNtfyToken + "_FILE": writeFile(t, token+"\r\n")},
} {
_, err := config.FromEnvironment(env.lookupEnv)
want := alertNtfyToken + ": holds a control character, such as the " +
"carriage return of a Windows line end"
if err == nil || err.Error() != want {
t.Errorf("%s: error %v, want %s", name, err, want)
}
}
}
func TestInstanceNameWithAControlCharacterStopsTheStartOnlyWithNtfySet(t *testing.T) {
t.Parallel()
const name = "fsn1app1\r"
_, err := config.FromEnvironment(environment{
instanceName: name, alertNtfyURL: "https://ntfy.example/smallwebwaf-alerts",
}.lookupEnv)
want := instanceName + `: "fsn1app1\r" holds a control character, such as the ` +
`carriage return of a Windows line end, and is sent to ntfy in a header ` +
`while ` + alertNtfyURL + ` is set`
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
cfg := fromEnvironment(t, environment{instanceName: name})
if cfg.InstanceName != name {
t.Errorf("not sending to ntfy, %s is %q", instanceName, cfg.InstanceName)
}
}
func TestInstanceNameNotUTF8StopsTheStart(t *testing.T) {
t.Parallel()
// café saved in Latin-1.
const latin1 = "caf\xe9"
for name, env := range map[string]environment{
"set": {instanceName: latin1},
"in a file": {instanceName + "_FILE": writeFile(t, latin1+"\n")},
} {
_, err := config.FromEnvironment(env.lookupEnv)
want := instanceName + `: "caf\xe9" is not valid UTF-8`
if err == nil || err.Error() != want {
t.Errorf("%s: error %v, want %s", name, err, want)
}
}
}
func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) { func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
@@ -494,6 +835,178 @@ func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) {
} }
} }
func TestLookupSettingsAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{
lookupTimeout: "500ms", addLookupHeaders: enabled,
})
wantLookupSettings(t, cfg, config.Config{
LookupSource: defaultLookupSource, LookupTimeout: 500 * time.Millisecond,
AddLookupHeaders: true,
})
cfg = fromEnvironment(t, environment{lookupSource: off})
wantLookupSettings(t, cfg, config.Config{
LookupSource: off, LookupTimeout: time.Second,
})
}
func TestLookupDBPathGoesWithTheFileSourceAlone(t *testing.T) {
t.Parallel()
const path = "/var/lib/ipinfo/ipinfo_lite.mmdb"
cfg := fromEnvironment(t, environment{lookupSource: fileSource, lookupDBPath: path})
if cfg.LookupSource != fileSource || cfg.LookupDBPath != path {
t.Errorf("lookups from %q in %q, want file in %q",
cfg.LookupSource, cfg.LookupDBPath, path)
}
for _, tc := range []struct {
env environment
want string
}{
{
environment{lookupSource: fileSource},
lookupSource + ": is file while " + lookupDBPath +
" is unset; it names the file to look clients up in",
},
{
environment{lookupSource: fileSource, lookupDBPath: ""},
lookupSource + ": is file while " + lookupDBPath +
" is unset; it names the file to look clients up in",
},
{
environment{lookupDBPath: path},
lookupDBPath + ": is set while " + lookupSource + " is geojs; only file reads it",
},
{
environment{lookupSource: off, lookupDBPath: path},
lookupDBPath + ": is set while " + lookupSource + " is off; only file reads it",
},
} {
_, err := config.FromEnvironment(tc.env.lookupEnv)
if err == nil || err.Error() != tc.want {
t.Errorf("settings %v: error %v, want %s", tc.env, err, tc.want)
}
}
}
func TestSettingNeedingLookupsStopsTheStartWhileTheyAreOff(t *testing.T) {
t.Parallel()
for name, value := range map[string]string{
deniedCountries: "kp",
allowedCountries: "de",
addLookupHeaders: enabled,
asnLimitPercent: "AS64496:50",
countryLimitPercent: "cn:25",
asnBytesPercent: "AS64496:50",
countryBytesPercent: "cn:25",
unknownLimitPercent: "99",
} {
t.Run(name, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{
lookupSource: off, name: value,
}.lookupEnv)
if err == nil || !strings.HasPrefix(err.Error(), name+": ") ||
!strings.Contains(err.Error(), lookupSource+" is off") {
t.Errorf("error %v, want one naming %s and %s", err, name, lookupSource)
}
})
}
// Set empty, the lists need nothing looked up, and nor does
// SWWAF_UNKNOWN_LIMIT_PERCENT at 100, which lowers no limit.
fromEnvironment(t, environment{
lookupSource: off, deniedCountries: "", allowedCountries: "",
asnLimitPercent: "", countryLimitPercent: "", asnBytesPercent: "",
countryBytesPercent: "", unknownLimitPercent: "100",
})
}
func TestBiasedThresholdsAsSet(t *testing.T) {
t.Parallel()
cfg := fromEnvironment(t, environment{})
if len(cfg.ASNLimitPercent) != 0 || len(cfg.CountryLimitPercent) != 0 ||
len(cfg.ASNBytesPercent) != 0 || len(cfg.CountryBytesPercent) != 0 ||
cfg.UnknownLimitPercent != 100 {
t.Errorf("biased thresholds %v, %v, %v, %v and %d by default, "+
"want four empty lists and 100", cfg.ASNLimitPercent, cfg.CountryLimitPercent,
cfg.ASNBytesPercent, cfg.CountryBytesPercent, cfg.UnknownLimitPercent)
}
// AS numbers and countries in either case, an AS number with leading
// zeros, 0 and 100.
cfg = fromEnvironment(t, environment{
asnLimitPercent: "AS14061:50, as16276:0,AS045102:100",
countryLimitPercent: "cn:25,RU:50",
asnBytesPercent: "as16276:75",
countryBytesPercent: "ru:10",
unknownLimitPercent: "0",
})
for name, tc := range map[string]struct{ got, want map[string]int64 }{
asnLimitPercent: {
cfg.ASNLimitPercent,
map[string]int64{"AS14061": 50, "AS16276": 0, "AS45102": 100},
},
countryLimitPercent: {cfg.CountryLimitPercent, map[string]int64{"CN": 25, "RU": 50}},
asnBytesPercent: {cfg.ASNBytesPercent, map[string]int64{"AS16276": 75}},
countryBytesPercent: {cfg.CountryBytesPercent, map[string]int64{"RU": 10}},
} {
if !maps.Equal(tc.got, tc.want) {
t.Errorf("%s gave %v, want %v", name, tc.got, tc.want)
}
}
if cfg.UnknownLimitPercent != 0 {
t.Errorf("%s gave %d, want 0", unknownLimitPercent, cfg.UnknownLimitPercent)
}
}
func TestInvalidBiasedThresholdStopsTheStartSayingWhatIsWrong(t *testing.T) {
t.Parallel()
const (
notASN = " is not an AS number such as AS64496"
notItem = " is not a code, : and a percentage, such as AS64496:50 or cn:25"
notPercent = " is not a percentage, a whole number from 0 to 100"
)
for _, tc := range []struct{ name, value, want string }{
{asnLimitPercent, "14061:50", `"14061"` + notASN},
{asnLimitPercent, "AS4294967296:50", `"AS4294967296"` + notASN},
{asnLimitPercent, "AS14061", `"AS14061"` + notItem},
{asnLimitPercent, "AS14061:101", `"101"` + notPercent},
{asnLimitPercent, "AS14061:50,as14061:25", `"as14061" is listed twice`},
{
countryLimitPercent, "nk:25",
`"nk" is not a two-letter country code such as de or kp`,
},
{countryLimitPercent, "cn:25,CN:50", `"CN" is listed twice`},
{asnBytesPercent, "AS14061:-1", `"-1"` + notPercent},
{countryBytesPercent, "cn:50%", `"50%"` + notPercent},
{unknownLimitPercent, "101", `"101"` + notPercent},
{unknownLimitPercent, off, `"off"` + notPercent},
} {
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
t.Parallel()
_, err := config.FromEnvironment(environment{tc.name: tc.value}.lookupEnv)
want := tc.name + ": " + tc.want
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
})
}
}
func TestSizesAndOff(t *testing.T) { func TestSizesAndOff(t *testing.T) {
t.Parallel() t.Parallel()
@@ -615,6 +1128,12 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{rateLimitPerHour, "1.5"}, {rateLimitPerHour, "1.5"},
{rateLimitPerDay, "-1"}, {rateLimitPerDay, "lots"}, {rateLimitPerDay, "-1"}, {rateLimitPerDay, "lots"},
{rateLimitExemptPaths, "/assets/,,/static/"}, {rateLimitExemptPaths, "/assets/,,/static/"},
{bytesLimitPerMinute, "10GB"}, {bytesLimitPerHour, "0"},
{bytesLimitPerDay, "-1G"},
{bytesCount, "all"}, {bytesCount, "Both"}, {bytesCount, ""},
{lookupSource, "ipinfo"}, {lookupSource, "GeoJS"}, {lookupSource, ""},
{lookupTimeout, off}, {lookupTimeout, "0s"}, {lookupTimeout, "1"},
{addLookupHeaders, "yes"},
{deniedCountries, "nk"}, {deniedCountries, "nk"},
{deniedCountries, "kp,,ir"}, {deniedCountries, "kp,,ir"},
{deniedCountries, "prk"}, {deniedCountries, "prk"},
@@ -627,6 +1146,14 @@ func TestInvalidValueStopsTheStart(t *testing.T) {
{allowedCountries, "uk"}, {allowedCountries, "uk"},
{allowedCountries, "zz"}, {allowedCountries, "zz"},
{allowedCountries, "de,germany"}, {allowedCountries, "de,germany"},
{asnLimitPercent, "AS14061:50,,AS16276:50"}, {asnLimitPercent, "ASX:50"},
{asnLimitPercent, "AS14061:"}, {asnLimitPercent, "AS14061 :50"},
{asnLimitPercent, "AS14061:1.5"}, {asnLimitPercent, "AS-1:50"},
{countryLimitPercent, "cn"}, {countryLimitPercent, "cn:"},
{countryLimitPercent, "cn:25:50"}, {countryLimitPercent, "china:25"},
{asnBytesPercent, "AS14061:101"}, {countryBytesPercent, "su:50"},
{unknownLimitPercent, ""}, {unknownLimitPercent, "-1"},
{unknownLimitPercent, "50%"},
{metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"}, {metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"},
{logRequestHeaders, "accept,,origin"}, {logRequestHeaders, "accept;origin"}, {logRequestHeaders, "accept,,origin"}, {logRequestHeaders, "accept;origin"},
{logRequestHeaders, "accept language"}, {logRequestHeaders, "x-foo:"}, {logRequestHeaders, "accept language"}, {logRequestHeaders, "x-foo:"},
@@ -866,20 +1393,6 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
t.Parallel() t.Parallel()
cfg := fromEnvironment(t, environment{clientRequestTimeout: "45s"}) cfg := fromEnvironment(t, environment{clientRequestTimeout: "45s"})
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
var line struct {
Settings map[string]string `json:"settings"`
}
err := json.Unmarshal(out.Bytes(), &line)
if err != nil {
t.Fatalf("decode %s: %v", out.Bytes(), err)
}
hostname, _ := os.Hostname() hostname, _ := os.Hostname()
want := map[string]string{ want := map[string]string{
@@ -902,8 +1415,21 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
rateLimitPerHour: "10000", rateLimitPerHour: "10000",
rateLimitPerDay: "50000", rateLimitPerDay: "50000",
rateLimitExemptPaths: "", rateLimitExemptPaths: "",
bytesLimitPerMinute: "10G",
bytesLimitPerHour: "20G",
bytesLimitPerDay: "50G",
bytesCount: "both",
lookupSource: defaultLookupSource,
lookupDBPath: "",
lookupTimeout: "1s",
addLookupHeaders: "false",
deniedCountries: "", deniedCountries: "",
allowedCountries: "", allowedCountries: "",
asnLimitPercent: "",
countryLimitPercent: "",
asnBytesPercent: "",
countryBytesPercent: "",
unknownLimitPercent: "100",
banResponse: "403", banResponse: "403",
limitBanDuration: "1h", limitBanDuration: "1h",
limitBanRepeatWindow: "24h", limitBanRepeatWindow: "24h",
@@ -926,12 +1452,40 @@ func TestLogsEachSettingWithItsValue(t *testing.T) {
logRemoteBuffer: "10000", logRemoteBuffer: "10000",
logRemoteFacility: "local0", logRemoteFacility: "local0",
logRemoteAppName: hostname, logRemoteAppName: hostname,
alertWebhookURL: "",
alertWebhookHeaders: "",
alertSlackWebhookURL: "",
alertNtfyURL: "",
alertNtfyToken: "",
alertEvents: defaultAlertEvents,
alertCooldown: defaultAlertCooldown,
alertMaxPerHour: "60",
} }
if !maps.Equal(line.Settings, want) { if got := loggedSettings(t, cfg); !maps.Equal(got, want) {
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want) t.Errorf("logged settings\n%v\nwant\n%v", got, want)
} }
} }
// loggedSettings returns the settings as cfg logs them, each by its name.
func loggedSettings(t *testing.T, cfg *config.Config) map[string]string {
t.Helper()
var out bytes.Buffer
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
var line struct {
Settings map[string]string `json:"settings"`
}
err := json.Unmarshal(out.Bytes(), &line)
if err != nil {
t.Fatalf("decode %s: %v", out.Bytes(), err)
}
return line.Settings
}
// wantSettings checks the settings that are plain values. // wantSettings checks the settings that are plain values.
func wantSettings(t *testing.T, got *config.Config, want config.Config) { func wantSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper() t.Helper()
@@ -955,6 +1509,33 @@ func wantSettings(t *testing.T, got *config.Config, want config.Config) {
wantBanSettings(t, got, want) wantBanSettings(t, got, want)
} }
// wantByteLimitSettings checks the settings for the byte limits.
func wantByteLimitSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper()
if got.BytesLimitPerMinute != want.BytesLimitPerMinute ||
got.BytesLimitPerHour != want.BytesLimitPerHour ||
got.BytesLimitPerDay != want.BytesLimitPerDay ||
got.BytesCount != want.BytesCount {
t.Errorf("byte limits %d, %d and %d counting %s, want %d, %d and %d counting %s",
got.BytesLimitPerMinute, got.BytesLimitPerHour, got.BytesLimitPerDay,
got.BytesCount, want.BytesLimitPerMinute, want.BytesLimitPerHour,
want.BytesLimitPerDay, want.BytesCount)
}
}
// wantLookupSettings checks the settings for lookups.
func wantLookupSettings(t *testing.T, got *config.Config, want config.Config) {
t.Helper()
if got.LookupSource != want.LookupSource || got.LookupTimeout != want.LookupTimeout ||
got.AddLookupHeaders != want.AddLookupHeaders {
t.Errorf("lookups from %q, waited for %s, headers %t; want %q, %s, %t",
got.LookupSource, got.LookupTimeout, got.AddLookupHeaders,
want.LookupSource, want.LookupTimeout, want.AddLookupHeaders)
}
}
// wantBanSettings checks the settings for bans, the state files, the // wantBanSettings checks the settings for bans, the state files, the
// metrics and the rule files. // metrics and the rule files.
func wantBanSettings(t *testing.T, got *config.Config, want config.Config) { func wantBanSettings(t *testing.T, got *config.Config, want config.Config) {
+233
View File
@@ -0,0 +1,233 @@
package lookup
import (
"context"
"fmt"
"log/slog"
"net/netip"
"os"
"path/filepath"
"sync"
"time"
"github.com/fsnotify/fsnotify"
"github.com/oschwald/maxminddb-golang/v2"
"sneak.berlin/go/smallwebwaf/internal/alerts"
)
// quietTime is how long the lookup database must go without a change
// before it is read again, so that a file still being copied in is read
// only once whole.
const quietTime = 2 * time.Second
// FileParams are what OpenFile needs.
type FileParams struct {
// Path is the lookup database, the IPinfo Lite file in its .mmdb form
// (SWWAF_LOOKUP_DB_PATH).
Path string
// Now tells the time, normally time.Now.
Now func() time.Time
// ProcessLog receives each reading of the file, and why a replacement
// of it cannot be read.
ProcessLog *slog.Logger
// Alerts receive a file_error alert for each replacement that cannot
// be read.
Alerts *alerts.Queue
}
// File looks up clients' AS numbers and countries in the lookup database,
// held in memory, and reads it again when it is replaced. It is safe for
// concurrent use.
type File struct {
params FileParams
mu sync.Mutex
// reader is the database in use, and lastRead when it was read.
// readFailures are the replacements that could not be read.
reader *maxminddb.Reader
lastRead time.Time
readFailures int
}
// record is what the lookup database holds about a network, of the fields
// smallwebwaf reads.
type record struct {
ASN string `maxminddb:"asn"`
ASName string `maxminddb:"as_name"`
CountryCode string `maxminddb:"country_code"`
}
// OpenFile reads the lookup database. A file that cannot be read, or that
// is not a .mmdb file, is an error.
func OpenFile(params FileParams) (*File, error) {
reader, err := read(params.Path)
if err != nil {
return nil, err
}
f := &File{params: params}
f.use(reader)
return f, nil
}
// LookUp returns what the lookup database says about client: its AS
// number, such as AS64496, the AS's name, and its country, such as DE,
// each "" when the database does not give it, as for an address missing
// from it. The database is asked about the client's first address, as
// GeoJS is.
func (f *File) LookUp(client netip.Prefix) Answer {
f.mu.Lock()
reader := f.reader
f.mu.Unlock()
var found record
// A record that cannot be decoded places the client nowhere, as a
// missing one does.
err := reader.Lookup(client.Addr()).Decode(&found)
if err != nil {
found = record{}
}
return Answer{
Client: client,
ASN: found.ASN,
ASName: found.ASName,
Country: found.CountryCode,
Answered: f.params.Now(),
}
}
// LastRead returns when the lookup database in use was read.
func (f *File) LastRead() time.Time {
f.mu.Lock()
defer f.mu.Unlock()
return f.lastRead
}
// ReadFailures returns how many replacements of the lookup database could
// not be read.
func (f *File) ReadFailures() int {
f.mu.Lock()
defer f.mu.Unlock()
return f.readFailures
}
// Watch watches the directory of the lookup database until ctx is done,
// and reads the file again once it has gone without a change for
// quietTime, after it is replaced, written or removed, and after Watch
// starts watching. If the directory cannot be watched, that is logged, and
// the database read at start stays in use.
func (f *File) Watch(ctx context.Context) {
watcher, err := fsnotify.NewWatcher()
if err == nil {
defer func() {
_ = watcher.Close()
}()
err = watcher.Add(filepath.Dir(f.params.Path))
}
if err != nil {
f.params.ProcessLog.Error("cannot watch the lookup database for replacements",
"error", err.Error())
return
}
f.params.ProcessLog.Info("watching the lookup database for replacements",
"file", f.params.Path)
f.readAfterChanges(ctx, watcher.Events, watcher.Errors)
}
// readAfterChanges reads the lookup database again once quietTime has
// passed without a change to it from events, until ctx is done, and logs
// the errors from errs. A change to another file in its directory does not
// count. The wait starts at once, as if for a change, so that a file
// replaced after OpenFile read it, and before its directory was watched,
// is read too.
func (f *File) readAfterChanges(
ctx context.Context, events <-chan fsnotify.Event, errs <-chan error,
) {
path := filepath.Clean(f.params.Path)
quiet := time.NewTimer(quietTime)
defer quiet.Stop()
for {
select {
case <-ctx.Done():
return
case event := <-events:
if filepath.Clean(event.Name) == path {
quiet.Reset(quietTime)
}
case <-quiet.C:
f.readAgain()
case err := <-errs:
f.params.ProcessLog.Warn("watching the lookup database failed",
"error", err.Error())
}
}
}
// readAgain reads the lookup database again, in place of the one in use,
// or, if it cannot be read, counts that, raises a file_error alert for it
// and logs it, and the one in use stays in use.
func (f *File) readAgain() {
reader, err := read(f.params.Path)
if err != nil {
const kept = "the lookup database cannot be read, " +
"and the one read before stays in use"
f.mu.Lock()
f.readFailures++
f.mu.Unlock()
// Raised before it is logged, so that the alert is there once the
// log line is.
f.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError,
Reason: kept,
Detail: map[string]any{"file": f.params.Path, "error": err.Error()},
})
f.params.ProcessLog.Error(kept, "error", err.Error())
return
}
f.use(reader)
}
// use puts reader in use, in place of the database read before, and logs
// that the file was read.
func (f *File) use(reader *maxminddb.Reader) {
f.mu.Lock()
f.reader = reader
f.lastRead = f.params.Now()
f.mu.Unlock()
f.params.ProcessLog.Info("read the lookup database", "file", f.params.Path)
}
// read reads the lookup database at path. The whole file is read into
// memory, rather than mapped into it as the reader can, so that a file
// overwritten in place cannot change, or end, under a lookup.
func read(path string) (*maxminddb.Reader, error) {
data, err := os.ReadFile(path) //nolint:gosec // the file the admin names
if err != nil {
return nil, fmt.Errorf("SWWAF_LOOKUP_DB_PATH cannot be read: %w", err)
}
reader, err := maxminddb.OpenBytes(data)
if err != nil {
return nil, fmt.Errorf("SWWAF_LOOKUP_DB_PATH %s is not a .mmdb file: %w", path, err)
}
return reader, nil
}
+379
View File
@@ -0,0 +1,379 @@
package lookup
import (
"context"
"log/slog"
"net/netip"
"net/url"
"os"
"path/filepath"
"reflect"
"testing"
"testing/synctest"
"time"
"github.com/fsnotify/fsnotify"
"github.com/maxmind/mmdbwriter/mmdbtype"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
)
// testNetblock is the netblock the tests' lookup databases place, and
// testClient a client in it.
const (
testNetblock = "203.0.113.0/24"
testClient = "203.0.113.9/32"
)
func TestFilePlacesClientsAndCountsAnAddressMissingFromItAsUnknown(t *testing.T) {
t.Parallel()
germany := lookuptest.Network{ASN: "AS64496", ASName: "Example Net", Country: "DE"}
northKorea := lookuptest.Network{ASN: "AS64511", ASName: "Other Net", Country: "KP"}
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
lookuptest.Write(t, path, map[string]lookuptest.Network{
testNetblock: germany,
"2001:db8::/32": northKorea,
})
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
f, err := OpenFile(FileParams{
Path: path,
Now: func() time.Time { return now },
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: newQueue(),
})
if err != nil {
t.Fatalf("open %s: %v", path, err)
}
for client, want := range map[string]lookuptest.Network{
testClient: germany,
// An IPv6 client is its /64.
"2001:db8:1:2::/64": northKorea,
"198.51.100.7/32": {},
} {
prefix := netip.MustParsePrefix(client)
got := f.LookUp(prefix)
if got != (Answer{
Client: prefix, ASN: want.ASN, ASName: want.ASName, Country: want.Country,
Answered: now,
}) {
t.Errorf("%s has the answer %+v, want %+v, answered %s", client, got, want, now)
}
}
}
func TestRecordThatCannotBeReadPlacesTheClientNowhere(t *testing.T) {
t.Parallel()
// The AS number is a number, where a string belongs. The writer writes
// a record's fields in the order of their names, so as_name is read
// before the AS number fails.
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
lookuptest.WriteRecords(t, path, map[string]mmdbtype.Map{
testNetblock: {
"asn": mmdbtype.Uint32(64496),
"as_name": mmdbtype.String("Example Net"),
"country_code": mmdbtype.String("DE"),
},
})
f := openFile(t, path, newQueue())
answer := f.LookUp(netip.MustParsePrefix(testClient))
if answer.ASN != "" || answer.ASName != "" || answer.Country != "" {
t.Errorf("%s is placed %+v, want nowhere", testClient, answer)
}
}
func TestFileThatCannotBeReadIsAnError(t *testing.T) {
t.Parallel()
dir := t.TempDir()
missing := filepath.Join(dir, "missing.mmdb")
notDatabase := filepath.Join(dir, "not.mmdb")
writeFile(t, notDatabase, "not a lookup database\n")
for path, want := range map[string]string{
missing: "SWWAF_LOOKUP_DB_PATH cannot be read: open " + missing +
": no such file or directory",
notDatabase: "SWWAF_LOOKUP_DB_PATH " + notDatabase +
" is not a .mmdb file: error opening database: invalid MaxMind DB file",
} {
_, err := OpenFile(FileParams{
Path: path,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: newQueue(),
})
if err == nil || err.Error() != want {
t.Errorf("opening %s failed with %v, want %s", path, err, want)
}
}
}
// The tests below run readAfterChanges in a synctest bubble, where time is
// a clock of the test's own: time.Sleep moves it on at once, and
// synctest.Wait returns once readAfterChanges waits again, so that every
// reading due by then is done. The test sends the changes itself, as the
// watch of a directory cannot run in a bubble.
func TestReplacementCopiedOverTheFileInTwoPartsIsReadOnlyWhole(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "ipinfo_lite.mmdb")
writeDatabase(t, path, "DE")
queue := newQueue()
f := openFile(t, path, queue)
changes := watch(t, f)
other := filepath.Join(dir, "replacement.mmdb")
writeDatabase(t, other, "KP")
replacement, err := os.ReadFile(other) //nolint:gosec // a file the test wrote
if err != nil {
t.Fatalf("read %s: %v", other, err)
}
// The file in use is overwritten in place, and keeps giving what
// it gave. Its first part alone is not a .mmdb file.
file, err := os.Create(path) //nolint:gosec // a file the test wrote
if err != nil {
t.Fatalf("create %s: %v", path, err)
}
defer func() {
_ = file.Close()
}()
half := len(replacement) / 2
write(t, file, replacement[:half])
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
time.Sleep(quietTime - time.Nanosecond)
synctest.Wait()
wantCountry(t, f, "DE")
// The second part starts the wait again.
write(t, file, replacement[half:])
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
time.Sleep(quietTime - time.Nanosecond)
synctest.Wait()
wantCountry(t, f, "DE")
time.Sleep(time.Nanosecond)
synctest.Wait()
wantCountry(t, f, "KP")
if !f.LastRead().Equal(time.Now()) || f.ReadFailures() != 0 {
t.Errorf("read at %s, with %d failures; want read now, with none",
f.LastRead(), f.ReadFailures())
}
wantAlerts(t, queue)
})
}
func TestReplacementThatCannotBeReadLeavesTheFileInUseWithOneAlert(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
writeDatabase(t, path, "DE")
queue := newQueue()
f := openFile(t, path, queue)
read := f.LastRead()
changes := watch(t, f)
writeFile(t, path, "not a lookup database\n")
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
// Long after, the replacement has been read once.
time.Sleep(time.Hour)
synctest.Wait()
wantCountry(t, f, "DE")
if !f.LastRead().Equal(read) || f.ReadFailures() != 1 {
t.Errorf("read at %s, with %d failures; want read at %s, with one",
f.LastRead(), f.ReadFailures(), read)
}
wantAlerts(t, queue, alerts.Alert{
Time: read.Add(quietTime),
Event: alerts.EventFileError,
Reason: "the lookup database cannot be read, and the one read before stays in use",
Detail: map[string]any{
"file": path,
"error": "SWWAF_LOOKUP_DB_PATH " + path + " is not a .mmdb file: " +
"error opening database: invalid MaxMind DB file",
},
})
})
}
func TestChangeOfAnotherFileInTheDirectoryIsNoReplacement(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "ipinfo_lite.mmdb")
writeDatabase(t, path, "DE")
f := openFile(t, path, newQueue())
changes := watch(t, f)
// The wait that starts with the watch ends with a reading.
time.Sleep(quietTime)
synctest.Wait()
writeDatabase(t, path, "KP")
changes <- fsnotify.Event{Name: filepath.Join(dir, "other.mmdb"), Op: fsnotify.Create}
time.Sleep(quietTime)
synctest.Wait()
wantCountry(t, f, "DE")
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
time.Sleep(quietTime)
synctest.Wait()
wantCountry(t, f, "KP")
})
}
func TestReplacementSavedBeforeTheWatchStartsIsRead(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
writeDatabase(t, path, "DE")
f := openFile(t, path, newQueue())
// Saved after OpenFile read the file, and before its directory was
// watched, so that no change is seen for it.
writeDatabase(t, path, "KP")
watch(t, f)
time.Sleep(quietTime)
synctest.Wait()
wantCountry(t, f, "KP")
})
}
// newQueue returns a queue of alerts for a webhook that is never sent
// them, so that they wait in it for the test to look at.
func newQueue() *alerts.Queue {
return alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
Now: time.Now,
})
}
// writeDatabase writes a lookup database at path that places testNetblock
// in country, and no other address.
func writeDatabase(t *testing.T, path, country string) {
t.Helper()
lookuptest.Write(t, path, map[string]lookuptest.Network{
testNetblock: {ASN: "AS64496", ASName: "Example Net", Country: country},
})
}
// openFile opens the lookup database at path, which raises its alerts to
// queue.
func openFile(t *testing.T, path string, queue *alerts.Queue) *File {
t.Helper()
f, err := OpenFile(FileParams{
Path: path,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: queue,
})
if err != nil {
t.Fatalf("open %s: %v", path, err)
}
return f
}
// watch runs f's readAfterChanges until the test ends, and returns the
// channel that sends it changes.
func watch(t *testing.T, f *File) chan<- fsnotify.Event {
t.Helper()
changes := make(chan fsnotify.Event)
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
f.readAfterChanges(ctx, changes, nil)
close(stopped)
}()
t.Cleanup(func() {
stop()
<-stopped
})
return changes
}
// wantCountry checks the country f gives testClient.
func wantCountry(t *testing.T, f *File, want string) {
t.Helper()
got := f.LookUp(netip.MustParsePrefix(testClient)).Country
if got != want {
t.Errorf("%s is in %q, want %q", testClient, got, want)
}
}
// wantAlerts checks the alerts waiting in queue, and that it held none
// back.
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
t.Helper()
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != len(want) || (len(want) > 0 && !reflect.DeepEqual(waiting, want)) {
t.Errorf("alerts waiting %+v, want %+v", waiting, want)
}
if queue.Suppressed() != 0 {
t.Errorf("%d alerts held back, want none", queue.Suppressed())
}
}
// writeFile writes content to the file at path.
func writeFile(t *testing.T, path, content string) {
t.Helper()
err := os.WriteFile(path, []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", path, err)
}
}
// write writes data to the end of file.
func write(t *testing.T, file *os.File, data []byte) {
t.Helper()
_, err := file.Write(data)
if err != nil {
t.Fatalf("write: %v", err)
}
}
+140 -65
View File
@@ -1,7 +1,8 @@
// Package lookup looks up each client's country through the GeoJS web // Package lookup looks up each client's AS number and country, through
// service, and keeps the answers in memory, for at most 100,000 clients // the GeoJS web service or in the lookup database, the IPinfo Lite file
// and for 7 days each. The answers are written to lookups.json and read // SWWAF_LOOKUP_DB_PATH names. GeoJS's answers are kept in memory, for at
// from it by the state package. // most 100,000 clients and for 7 days each, and are written to
// lookups.json and read from it by the state package.
package lookup package lookup
import ( import (
@@ -14,17 +15,20 @@ import (
"net/http" "net/http"
"net/netip" "net/netip"
"slices" "slices"
"strconv"
"strings" "strings"
"sync" "sync"
"time" "time"
"github.com/hashicorp/golang-lru/v2/simplelru" "github.com/hashicorp/golang-lru/v2/simplelru"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
) )
// URL is GeoJS's country endpoint. Asked about several addresses at once, // URL is GeoJS's endpoint for an address's place and network. Asked about
// comma separated in its ip parameter, it answers with a list. // several addresses at once, comma separated in its ip parameter, it
const URL = "https://get.geojs.io/v1/ip/country.json" // answers with a list.
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.
@@ -39,9 +43,8 @@ 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
// timeout is how long a new client waits for its answer, and how long // unknownASN is the AS number GeoJS gives when it knows none.
// a request to GeoJS may take before it is abandoned. unknownASN = 64512
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.
@@ -61,6 +64,16 @@ 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.
@@ -68,16 +81,22 @@ type Params struct {
// Metrics count the requests to GeoJS, those that failed, and the // Metrics count the requests to GeoJS, those that failed, and the
// clients that go without an answer. // clients that go without an answer.
Metrics *metrics.Metrics Metrics *metrics.Metrics
// Alerts receive a source_failure alert each time GeoJS fails.
Alerts *alerts.Queue
} }
// GeoJS looks up clients' countries through GeoJS. At most one request // GeoJS looks up clients' AS numbers and countries through GeoJS. At most
// to GeoJS is under way at a time, and it asks about every client waiting, // one request to GeoJS is under way at a time, and it asks about every
// up to maxPerRequest. It is safe for concurrent use. // client waiting, 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
alerts *alerts.Queue
// httpClient follows no redirect, so that visitors' addresses go to // httpClient follows no redirect, so that visitors' addresses go to
// GeoJS alone: a redirect is a failure. // GeoJS alone: a redirect is a failure.
httpClient *http.Client httpClient *http.Client
@@ -95,11 +114,18 @@ type GeoJS struct {
retryAt time.Time retryAt time.Time
} }
// Answer is what GeoJS said about a client, as lookups.json holds it: its // Answer is what GeoJS or the lookup database said about a client: its AS
// country, "" when GeoJS cannot place it, when GeoJS said so, and when // number, such as AS64496, and the AS's name, both "" when the source knows
// the answer was last used. // no AS number for it; its country, "" when the source cannot place it;
// when the source said so; and, for GeoJS's answers, which lookups.json
// holds, when the answer was last used. The zero Answer is that of a
// client with no answer.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Answer struct { 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"`
@@ -124,9 +150,13 @@ 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,
alerts: params.Alerts,
httpClient: &http.Client{ httpClient: &http.Client{
CheckRedirect: func(*http.Request, []*http.Request) error { CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse return http.ErrUseLastResponse
@@ -137,23 +167,23 @@ func New(params Params) *GeoJS {
} }
} }
// Country returns the country GeoJS places client in, as a two-letter // LookUp returns the answer GeoJS gave about client, with its country as
// code in capitals, or "" when the country cannot be found: GeoJS cannot // a two-letter code in capitals, or the zero Answer when there is none
// place the client, or has not answered in time. An answer is kept for 7 // yet. An answer is kept for 7 days. Without one, the client is asked
// days. Without one, a client waits up to timeout for it, unless it has // about in the background, and, while Wait is set, the request waits up
// gone without one before; until GeoJS answers, the client is asked about // to Timeout for the answer, unless the client has gone without one
// again in the background. ctx is the context of the client's request, // before. ctx is the context of the client's request, and ends the wait
// and ends the wait when it ends. // 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 /64. // own address for IPv4, and an address in the same place for an IPv6 /64.
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 country return answer
} }
timer := time.NewTimer(timeout) timer := time.NewTimer(g.timeout)
defer timer.Stop() defer timer.Stop()
select { select {
@@ -165,7 +195,7 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
g.mu.Lock() g.mu.Lock()
defer g.mu.Unlock() defer g.mu.Unlock()
country, found := g.kept(client) answer, found := g.kept(client)
if !found { if !found {
g.metrics.GeoJSUnanswered.Inc() g.metrics.GeoJSUnanswered.Inc()
} }
@@ -175,7 +205,15 @@ func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
w.late = true w.late = true
} }
return country return answer
}
// 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
@@ -227,13 +265,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,
) (string, <-chan struct{}) { ) (Answer, <-chan struct{}) {
g.mu.Lock() g.mu.Lock()
defer g.mu.Unlock() defer g.mu.Unlock()
country, found := g.kept(client) answer, found := g.kept(client)
if found { if found {
return country, nil return answer, nil
} }
w, waiting := g.waiting[client] w, waiting := g.waiting[client]
@@ -244,10 +282,14 @@ 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 "", nil // too many clients wait already return Answer{}, nil // too many clients wait already
} }
if !g.asking { if !g.asking {
@@ -258,25 +300,25 @@ func (g *GeoJS) answerOrWait(
if w.late { if w.late {
g.metrics.GeoJSUnanswered.Inc() g.metrics.GeoJSUnanswered.Inc()
return "", nil return Answer{}, nil
} }
return "", w.asked return Answer{}, 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) (string, bool) { func (g *GeoJS) kept(client netip.Prefix) (Answer, 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 "", false return Answer{}, false
} }
kept.Used = now kept.Used = now
return kept.Country, true return *kept, 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
@@ -294,7 +336,8 @@ 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. // time, until none is left or GeoJS fails. Each answer kept is given to
// 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()
@@ -302,8 +345,16 @@ func (g *GeoJS) askAboutWaiting(ctx context.Context) {
return return
} }
countries, err := g.request(ctx, clients) given, err := g.request(ctx, clients)
if !g.keep(clients, countries, err) { kept, answered := g.keep(clients, given, err)
if g.answered != nil {
for _, answer := range kept {
g.answered(answer)
}
}
if !answered {
return return
} }
} }
@@ -335,32 +386,35 @@ func (g *GeoJS) nextClients() []netip.Prefix {
return clients return clients
} }
// keep notes how a request to GeoJS about clients ended, and reports // keep notes how a request to GeoJS about clients ended, given being the
// whether GeoJS answered about all of them. Each client whose address // answer for each address GeoJS's answer names. It returns the answers it
// GeoJS's answer names gets its answer, with no country when GeoJS gave // kept, and reports whether GeoJS answered about all of the clients. Each
// none. An answer that leaves an address out is a failure. After a // client whose address GeoJS's answer names gets its answer. An answer
// failure GeoJS is left alone for a while, and every client still waiting // that leaves an address out is a failure. After a failure GeoJS is left
// stops waiting and is asked about once GeoJS is asked again. // alone for a while, and every client still waiting stops waiting and is
// asked about once GeoJS is asked again.
func (g *GeoJS) keep( func (g *GeoJS) keep(
clients []netip.Prefix, countries map[netip.Addr]string, err error, clients []netip.Prefix, given map[netip.Addr]Answer, err error,
) bool { ) ([]Answer, 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 {
country, named := countries[client.Addr()] answer, named := given[client.Addr()]
if !named { if !named {
leftOut++ leftOut++
continue continue
} }
g.answers.Add(client, &Answer{ answer.Client, answer.Answered, answer.Used = client, now, now
Client: client, Country: country, Answered: now, Used: now, g.answers.Add(client, &answer)
}) kept = append(kept, answer)
close(g.waiting[client].asked) close(g.waiting[client].asked)
delete(g.waiting, client) delete(g.waiting, client)
} }
@@ -386,27 +440,37 @@ func (g *GeoJS) keep(
g.processLog.Warn("asking GeoJS failed", g.processLog.Warn("asking GeoJS failed",
"error", err.Error(), "asking_again_in", g.retryDelay.String()) "error", err.Error(), "asking_again_in", g.retryDelay.String())
g.alerts.Raise(alerts.Alert{
Event: alerts.EventSourceFailure,
Reason: "asking GeoJS failed",
Detail: map[string]any{
"source": "geojs", "error": err.Error(),
"asking_again_in": g.retryDelay.String(),
},
})
return false return kept, false
} }
g.retryDelay = 0 g.retryDelay = 0
return true return kept, true
} }
// request asks GeoJS about clients in one request, and returns the // request asks GeoJS about clients in one request, and returns the answer
// country it gave, in capitals, for each address its answer names. // for each address GeoJS's answer names: its AS number and the AS's name,
// 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]string, error) { ) (map[netip.Addr]Answer, 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, timeout) ctx, cancel := context.WithTimeout(ctx, g.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)
@@ -433,9 +497,12 @@ 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"`
Country string `json:"country"` ASN int64 `json:"asn"`
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)
@@ -443,14 +510,22 @@ func (g *GeoJS) request(
return nil, fmt.Errorf("read GeoJS's answer: %w", err) return nil, fmt.Errorf("read GeoJS's answer: %w", err)
} }
countries := make(map[netip.Addr]string, len(answers)) given := make(map[netip.Addr]Answer, 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 {
countries[addr] = strings.ToUpper(item.Country) continue
} }
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 countries, nil return given, nil
} }
+247 -16
View File
@@ -6,6 +6,8 @@ import (
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"net/netip" "net/netip"
"net/url"
"reflect"
"slices" "slices"
"strings" "strings"
"sync" "sync"
@@ -14,15 +16,20 @@ import (
"time" "time"
"github.com/prometheus/client_golang/prometheus/testutil" "github.com/prometheus/client_golang/prometheus/testutil"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
) )
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. // unplaced, and asNumber, kept as asn, and asName the AS it gives them.
germany = "DE" germany = "DE"
// unplaced is the address it cannot place. asNumber = 64496
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.
@@ -80,7 +87,7 @@ func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
var earlier sync.WaitGroup var earlier sync.WaitGroup
earlier.Go(func() { g.Country(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) }) earlier.Go(func() { g.LookUp(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) })
defer earlier.Wait() defer earlier.Wait()
waitForRequests(t, geojs, 1) waitForRequests(t, geojs, 1)
@@ -112,6 +119,58 @@ func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
}) })
} }
func TestRequestWaitsAsLongAsTheTimeoutSaysAndGeoJSIsAbandonedAfterIt(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
// A timeout longer than the default second, and a GeoJS that does
// not answer.
const longerTimeout = 3 * time.Second
m := metrics.New(1, "app")
g := lookup.New(lookup.Params{
URL: lookup.URL,
Timeout: longerTimeout,
Wait: true,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: m,
Alerts: alerts.New(alerts.Params{}),
})
g.SetTransport(&standIn{answers: hanging})
var (
request sync.WaitGroup
waited time.Duration
)
request.Go(func() {
began := time.Now()
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
waited = time.Since(began)
})
// A moment before the timeout runs out, GeoJS is still being asked:
// the request to it has not failed.
time.Sleep(longerTimeout - time.Millisecond)
synctest.Wait()
wantFailures(t, m, 0)
// As it runs out, the client's request goes on, and the request to
// GeoJS is abandoned, which counts as a failure.
request.Wait()
synctest.Wait()
if waited != longerTimeout {
t.Errorf("waited %s for the answer, want %s", waited, longerTimeout)
}
wantFailures(t, m, 1)
})
}
func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) { func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
t.Parallel() t.Parallel()
@@ -184,6 +243,96 @@ 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()
@@ -194,9 +343,12 @@ 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), Metrics: metrics.New(1, "app"),
Alerts: alerts.New(alerts.Params{}),
}) })
g.SetTransport(geojs) g.SetTransport(geojs)
@@ -211,6 +363,45 @@ func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
}) })
} }
func TestFailureRaisesASourceFailureAlertOncePerCooldown(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
geojs, clock, g, queue := startWithAlerts()
clients := newClients()
geojs.set(failing)
wantCountry(t, g, clients(), "")
want := alerts.Alert{
Time: clock.Now(),
Event: alerts.EventSourceFailure,
Reason: "asking GeoJS failed",
Detail: map[string]any{
"source": "geojs",
"error": "GeoJS answered 503 Service Unavailable",
"asking_again_in": "1s",
},
}
// The next failure, a second later, is a repeat within the
// cooldown.
clock.advance(time.Second)
wantCountry(t, g, clients(), "")
wantRequests(t, geojs, 2)
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || !reflect.DeepEqual(waiting[0], want) {
t.Errorf("alerts waiting %+v, want only %+v", waiting, want)
}
if queue.Suppressed() != 1 {
t.Errorf("%d alerts held back, want the repeat", queue.Suppressed())
}
})
}
func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) { func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) {
t.Parallel() t.Parallel()
@@ -355,12 +546,15 @@ 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) m := metrics.New(1, "app")
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,
Alerts: alerts.New(alerts.Params{}),
}) })
g.SetTransport(&standIn{answers: failing}) g.SetTransport(&standIn{answers: failing})
@@ -450,21 +644,24 @@ func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) {
} }
} }
list := make([]map[string]string, 0, len(addrs)) list := make([]map[string]any, 0, len(addrs))
for _, addr := range addrs { for _, addr := range addrs {
country := germany item := map[string]any{
"ip": addr, "asn": asNumber, "organization_name": asName,
"country_code": germany,
}
switch { switch {
case addr == unplaced: case addr == unplaced:
country = "" item = map[string]any{"ip": addr, "asn": 64512, "organization_name": "Unknown"}
case addr == leftOut && answers == answeringWithoutLeftOut: case addr == leftOut && answers == answeringWithoutLeftOut:
continue continue
case answers == answeringInLowerCase: case answers == answeringInLowerCase:
country = strings.ToLower(germany) item["country_code"] = strings.ToLower(germany)
} }
list = append(list, map[string]string{"ip": addr, "country": country}) list = append(list, item)
} }
var answer any = list var answer any = list
@@ -522,19 +719,43 @@ 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. // asking the stand-in by that clock, for which a request waits for its
// client's first answer.
func start() (*standIn, *testClock, *lookup.GeoJS) { func start() (*standIn, *testClock, *lookup.GeoJS) {
geojs, clock, g, _ := startWithAlerts()
return geojs, clock, g
}
// startWithAlerts is start, and returns the queue of the alerts GeoJS
// raises as well, for a webhook that is never sent them, with the default
// cooldown, by the same clock.
func startWithAlerts() (*standIn, *testClock, *lookup.GeoJS, *alerts.Queue) {
geojs := &standIn{} geojs := &standIn{}
clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)} clock := newClock()
queue := alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
Now: clock.Now,
})
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), Metrics: metrics.New(1, "app"),
Alerts: queue,
}) })
g.SetTransport(geojs) g.SetTransport(geojs)
return geojs, clock, g 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
@@ -553,7 +774,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.Country(t.Context(), client) got := g.LookUp(t.Context(), client).Country
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)
} }
@@ -598,6 +819,16 @@ func wantUnanswered(t *testing.T, m *metrics.Metrics, want float64) {
} }
} }
// wantFailures checks how many requests to GeoJS m counts as failed.
func wantFailures(t *testing.T, m *metrics.Metrics, want float64) {
t.Helper()
got := testutil.ToFloat64(m.GeoJSFailures)
if got != want {
t.Errorf("%v requests to GeoJS failed, want %v", got, want)
}
}
// waitForRequests waits until g has done all it can before time passes, // 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
@@ -0,0 +1,83 @@
// Package lookuptest writes lookup databases, IPinfo Lite files in their
// .mmdb form, for the tests of the packages that read them.
package lookuptest
import (
"bytes"
"net"
"os"
"testing"
"github.com/maxmind/mmdbwriter"
"github.com/maxmind/mmdbwriter/mmdbtype"
)
// fileMode is the mode of the files written: read and written by their
// owner alone.
const fileMode = 0o600
// Network is what a lookup database holds about a netblock, of the fields
// smallwebwaf reads: its AS number, such as AS64496, the AS's name, and
// its country, such as DE.
type Network struct {
ASN string
ASName string
Country string
}
// Write writes a lookup database at path that places each netblock in
// networks, such as 203.0.113.0/24, as its Network says, and no other
// address.
func Write(tb testing.TB, path string, networks map[string]Network) {
tb.Helper()
records := make(map[string]mmdbtype.Map, len(networks))
for netblock, network := range networks {
records[netblock] = mmdbtype.Map{
"asn": mmdbtype.String(network.ASN),
"as_name": mmdbtype.String(network.ASName),
"country_code": mmdbtype.String(network.Country),
}
}
WriteRecords(tb, path, records)
}
// WriteRecords writes a lookup database at path that holds each record in
// records for its netblock, and nothing for any other address.
func WriteRecords(tb testing.TB, path string, records map[string]mmdbtype.Map) {
tb.Helper()
tree, err := mmdbwriter.New(mmdbwriter.Options{
DatabaseType: "ipinfo_lite",
// The tests' clients are in the netblocks kept for documentation.
IncludeReservedNetworks: true,
})
if err != nil {
tb.Fatalf("new lookup database: %v", err)
}
for netblock, record := range records {
_, network, err := net.ParseCIDR(netblock)
if err != nil {
tb.Fatalf("netblock %q: %v", netblock, err)
}
err = tree.Insert(network, record)
if err != nil {
tb.Fatalf("insert %s: %v", netblock, err)
}
}
var database bytes.Buffer
_, err = tree.WriteTo(&database)
if err != nil {
tb.Fatalf("write the lookup database: %v", err)
}
err = os.WriteFile(path, database.Bytes(), fileMode)
if err != nil {
tb.Fatalf("write %s: %v", path, err)
}
}
+5 -2
View File
@@ -26,8 +26,11 @@ func TestSnapshotHoldsEachAnswerAndWhenItWasLastUsed(t *testing.T) {
wantCountry(t, g, placed, germany) wantCountry(t, g, placed, germany)
want := []lookup.Answer{ want := []lookup.Answer{
{Client: notPlaced, Country: "", Answered: asked, Used: asked}, {Client: notPlaced, 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
@@ -0,0 +1,155 @@
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
@@ -1,116 +0,0 @@
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
}
+104 -20
View File
@@ -6,11 +6,13 @@ 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"
"github.com/prometheus/client_golang/prometheus/collectors" "github.com/prometheus/client_golang/prometheus/collectors"
"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/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/remotelog" "sneak.berlin/go/smallwebwaf/internal/remotelog"
@@ -20,7 +22,8 @@ import (
// 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 *prometheus.Registry // registry gives every metric registered with it the label instance.
registry prometheus.Registerer
handler http.Handler handler http.Handler
inFlight prometheus.Gauge inFlight prometheus.Gauge
@@ -34,12 +37,13 @@ type Metrics struct {
offences *prometheus.CounterVec offences *prometheus.CounterVec
// ruleMatches are made by AddRules. // ruleMatches are made by AddRules.
ruleMatches *prometheus.CounterVec ruleMatches *prometheus.CounterVec
countries *countries 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 whose client counted // that failed. GeoJSUnanswered are the requests that needed their
// as coming from an unknown country because GeoJS had not answered // client's answer, for a setting that acts on it, and went on without
// about it in time. // it because GeoJS had not given it in time.
GeoJSRequests prometheus.Counter GeoJSRequests prometheus.Counter
GeoJSFailures prometheus.Counter GeoJSFailures prometheus.Counter
GeoJSUnanswered prometheus.Counter GeoJSUnanswered prometheus.Counter
@@ -53,14 +57,18 @@ 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 get series of their own // topN is how many countries and how many AS numbers get series of their
// (SWWAF_METRICS_TOP_N). // own (SWWAF_METRICS_TOP_N). Every metric carries instanceName
func New(topN int) *Metrics { // (SWWAF_INSTANCE_NAME) as its label instance.
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.NewRegistry(), registry: prometheus.WrapRegistererWith(
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.",
@@ -82,14 +90,16 @@ func New(topN int) *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, by its window.", "Requests that broke a rate limit or a byte limit, by its window and "+
[]string{"window"}), "its kind, requests or bytes.",
[]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.",
@@ -100,8 +110,8 @@ func New(topN int) *Metrics {
}), }),
GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{ GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_unanswered_total", Name: "smallwebwaf_geojs_unanswered_total",
Help: "Requests whose client counted as coming from an unknown " + Help: "Requests that needed their client's answer from GeoJS and " +
"country because GeoJS had not answered about it in time.", "went on without it, because GeoJS had not given 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),
@@ -118,16 +128,12 @@ func New(topN int) *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.rateLimitHits, m.sizeAndTimeLimitHits, m.offences, m.countries, m.asns,
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,
@@ -225,6 +231,72 @@ func (m *Metrics) AddRemoteLog(remote *remotelog.Sender) {
) )
} }
// AddLookupFile adds the metrics of the lookup database, read as the
// metrics are asked for: when the file in use was read, which lastRead
// returns, and the replacements of it that could not be read, which
// readFailures returns. The lookup package's File, which has both, cannot
// be named here: that package counts GeoJS's requests in these metrics.
func (m *Metrics) AddLookupFile(lastRead func() time.Time, readFailures func() int) {
m.registry.MustRegister(
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_lookup_database_last_read_timestamp_seconds",
Help: "When the lookup database in use was read, in seconds since 1970.",
}, func() float64 {
return float64(lastRead().Unix())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_lookup_database_read_failures_total",
Help: "Replacements of the lookup database that could not be read.",
}, func() float64 {
return float64(readFailures())
}),
)
}
// 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)
@@ -255,7 +327,15 @@ func (m *Metrics) RequestEnded(
} }
if line.LimitHit != "" { if line.LimitHit != "" {
m.rateLimitHits.WithLabelValues(line.LimitHit).Inc() // The log line names a byte limit's window with _bytes after it.
window, isBytes := strings.CutSuffix(line.LimitHit, "_bytes")
kind := ratelimit.KindRequests
if isBytes {
kind = ratelimit.KindBytes
}
m.rateLimitHits.WithLabelValues(window, kind).Inc()
} }
if limit != "" { if limit != "" {
@@ -267,7 +347,11 @@ func (m *Metrics) RequestEnded(
} }
if line.Country != "" { if line.Country != "" {
m.countries.add(line) m.countries.add(line.Country, line)
}
if line.ASN != "" {
m.asns.add(line.ASN, line)
} }
} }
+50 -7
View File
@@ -1,6 +1,7 @@
package proxy package proxy
import ( import (
"bytes"
"crypto/subtle" "crypto/subtle"
"encoding/json" "encoding/json"
"errors" "errors"
@@ -8,6 +9,7 @@ import (
"io" "io"
"net/http" "net/http"
"net/netip" "net/netip"
"os"
"strings" "strings"
"time" "time"
@@ -27,8 +29,13 @@ const banBodyMaxBytes = 4 << 10
const permanent = "permanent" const permanent = "permanent"
var ( var (
errNotBanToAdd = errors.New(
"the body is not a JSON object of netblock, duration and reason")
errNotNetblock = errors.New( errNotNetblock = errors.New(
"is not an address or a netblock, such as 203.0.113.9 or 203.0.113.0/24") "is not an address or a netblock, such as 203.0.113.9 or 203.0.113.0/24")
errMappedNetblock = errors.New(
"is IPv4-mapped: give the IPv4 netblock, such as 203.0.113.0/24")
errZone = errors.New("has a zone, which a netblock cannot have")
errNotDuration = errors.New( errNotDuration = errors.New(
"is not a duration above zero, such as 1h or 7d, or permanent") "is not a duration above zero, such as 1h or 7d, or permanent")
errNotAddress = errors.New("is not an address, such as 203.0.113.9") errNotAddress = errors.New("is not an address, such as 203.0.113.9")
@@ -111,6 +118,10 @@ type banToAdd struct {
// an admin, from now for the duration the body gives, with its reason, // an admin, from now for the duration the body gives, with its reason,
// and answers with that ban. // and answers with that ban.
func (rq *request) addBan() { func (rq *request) addBan() {
// The body must arrive within SWWAF_CLIENT_REQUEST_TIMEOUT, as any
// other request's must.
rq.stopReadingBody(rq.clientRequestDeadline())
toAdd, err := rq.readBanToAdd() toAdd, err := rq.readBanToAdd()
if refused := rq.refused.Load(); refused != nil { if refused := rq.refused.Load(); refused != nil {
rq.answer(*refused) // the body is over SWWAF_REQUEST_MAX_BYTES rq.answer(*refused) // the body is over SWWAF_REQUEST_MAX_BYTES
@@ -118,6 +129,16 @@ func (rq *request) addBan() {
return return
} }
if errors.Is(err, os.ErrDeadlineExceeded) {
rq.answer(refusal{
status: http.StatusRequestTimeout,
action: requestlog.ActionTimedOut,
limit: "SWWAF_CLIENT_REQUEST_TIMEOUT",
})
return
}
var ( var (
netblock netip.Prefix netblock netip.Prefix
expires time.Time expires time.Time
@@ -142,23 +163,33 @@ func (rq *request) addBan() {
rq.answerBans([]bans.Ban{ban}) rq.answerBans([]bans.Ban{ban})
} }
// readBanToAdd reads the body of POST BansPath, at most banBodyMaxBytes // readBanToAdd reads the body of POST BansPath: a JSON object with
// of it. // nothing but whitespace after it, in at most banBodyMaxBytes.
func (rq *request) readBanToAdd() (banToAdd, error) { func (rq *request) readBanToAdd() (banToAdd, error) {
var body io.ReadCloser = http.NoBody var body io.ReadCloser = http.NoBody
if rq.body != nil { if rq.body != nil {
body = rq.body body = rq.body
} }
data, err := io.ReadAll(http.MaxBytesReader(nil, body, banBodyMaxBytes))
if err != nil {
return banToAdd{}, fmt.Errorf("%w: %w", errNotBanToAdd, err)
}
var toAdd banToAdd var toAdd banToAdd
decoder := json.NewDecoder(http.MaxBytesReader(nil, body, banBodyMaxBytes)) decoder := json.NewDecoder(bytes.NewReader(data))
decoder.DisallowUnknownFields() decoder.DisallowUnknownFields()
err := decoder.Decode(&toAdd) err = decoder.Decode(&toAdd)
if err != nil { if err != nil {
return banToAdd{}, fmt.Errorf( return banToAdd{}, fmt.Errorf("%w: %w", errNotBanToAdd, err)
"the body is not a JSON object of netblock, duration and reason: %w", err) }
// Token returns io.EOF only when nothing but whitespace is left.
_, err = decoder.Token()
if !errors.Is(err, io.EOF) {
return banToAdd{}, fmt.Errorf("%w: more follows the object", errNotBanToAdd)
} }
return toAdd, nil return toAdd, nil
@@ -166,18 +197,30 @@ func (rq *request) readBanToAdd() (banToAdd, error) {
// banNetblock reads value, a netblock such as 203.0.113.0/24, or a // banNetblock reads value, a netblock such as 203.0.113.0/24, or a
// client's address, which stands for the netblock a ban on that client // client's address, which stands for the netblock a ban on that client
// covers. // covers. An IPv4-mapped netblock, such as ::ffff:203.0.113.0/120, is
// refused, since a client's address is looked up as IPv4 and a ban on it
// would refuse nothing, and so is a value with a zone.
func (h *handler) banNetblock(value string) (netip.Prefix, error) { func (h *handler) banNetblock(value string) (netip.Prefix, error) {
netblock, err := netip.ParsePrefix(value) netblock, err := netip.ParsePrefix(value)
if err == nil { if err == nil {
if netblock.Addr().Is4In6() {
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errMappedNetblock)
}
return netblock, nil return netblock, nil
} }
// ParsePrefix refuses a zone, but ParseAddr reads the /48 of
// 2001:db8::1%x/48 as part of the zone.
addr, err := netip.ParseAddr(value) addr, err := netip.ParseAddr(value)
if err != nil { if err != nil {
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errNotNetblock) return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errNotNetblock)
} }
if addr.Zone() != "" {
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errZone)
}
return h.netblock(addr), nil return h.netblock(addr), nil
} }
+61 -2
View File
@@ -173,7 +173,9 @@ func TestBanToAddGivesItsNetblockAndDuration(t *testing.T) {
{"198.51.100.7/16", "1h", "198.51.0.0/16", time.Hour}, {"198.51.100.7/16", "1h", "198.51.0.0/16", time.Hour},
{"2001:db8:6::/48", "1h", "2001:db8:6::/48", time.Hour}, {"2001:db8:6::/48", "1h", "2001:db8:6::/48", time.Hour},
} { } {
body := `{"netblock": "` + tc.netblock + `", "duration": "` + tc.duration + `"}` // Whitespace may follow the object.
body := `{"netblock": "` + tc.netblock + `", "duration": "` + tc.duration + `"}` +
"\r\n"
want := state.BanEntry{ want := state.BanEntry{
Netblock: netip.MustParsePrefix(tc.want), Start: start, Cause: bans.CauseAdmin, Netblock: netip.MustParsePrefix(tc.want), Start: start, Cause: bans.CauseAdmin,
} }
@@ -203,6 +205,29 @@ func TestBanToAddThatCannotBeReadIsRefused(t *testing.T) {
`{"netblock": "203.0.113", "duration": "1h"}`, `{"netblock": "203.0.113", "duration": "1h"}`,
`netblock "203.0.113" is not an address or a netblock`, `netblock "203.0.113" is not an address or a netblock`,
}, },
// A client's address is looked up as IPv4, so a ban on an
// IPv4-mapped netblock would refuse nothing.
{
`{"netblock": "::ffff:203.0.113.0/120", "duration": "1h"}`,
`netblock "::ffff:203.0.113.0/120" is IPv4-mapped`,
},
// Read as an address, its zone would be "x/48", and its ban on the
// /64 around it.
{
`{"netblock": "2001:db8::1%x/48", "duration": "1h"}`,
`netblock "2001:db8::1%x/48" has a zone`,
},
{
`{"netblock": "fe80::1%eth0", "duration": "1h"}`,
`netblock "fe80::1%eth0" has a zone`,
},
// Anything but whitespace after the object.
{
`{"netblock": "203.0.113.9", "duration": "1h"}` +
`{"netblock": "198.51.100.0/24", "duration": "1h"}`,
"more follows the object",
},
{`{"netblock": "203.0.113.9", "duration": "1h"} x`, "more follows the object"},
{`{"duration": "1h"}`, `netblock "" is not an address or a netblock`}, {`{"duration": "1h"}`, `netblock "" is not an address or a netblock`},
{`{"netblock": "203.0.113.9"}`, `duration "" is not a duration above zero`}, {`{"netblock": "203.0.113.9"}`, `duration "" is not a duration above zero`},
{ {
@@ -218,12 +243,16 @@ func TestBanToAddThatCannotBeReadIsRefused(t *testing.T) {
`duration "forever" is not a duration above zero, such as 1h or 7d, ` + `duration "forever" is not a duration above zero, such as 1h or 7d, ` +
`or permanent`, `or permanent`,
}, },
// Over the 4 KiB read of a body. // Over the 4 KiB read of a body, even when the object comes first.
{ {
`{"netblock": "203.0.113.9", "duration": "1h", "reason": "` + `{"netblock": "203.0.113.9", "duration": "1h", "reason": "` +
strings.Repeat("x", 4<<10) + `"}`, strings.Repeat("x", 4<<10) + `"}`,
"request body too large", "request body too large",
}, },
{
`{"netblock": "203.0.113.9", "duration": "1h"}` + strings.Repeat(" ", 4<<10),
"request body too large",
},
} { } {
got := s.admin(http.MethodPost, proxy.BansPath, tc.body, http.StatusBadRequest) got := s.admin(http.MethodPost, proxy.BansPath, tc.body, http.StatusBadRequest)
if !strings.Contains(string(got.body), tc.want) { if !strings.Contains(string(got.body), tc.want) {
@@ -257,6 +286,36 @@ func TestBanToAddOverTheRequestSizeLimitIsRefused(t *testing.T) {
} }
} }
func TestBanToAddSlowerThanTheClientRequestTimeoutIsRefused(t *testing.T) {
t.Parallel()
s, _, server := startWithClock(t, "", map[string]string{
adminToken: adminSecret,
metricsToken: token,
clientRequestTimeout: shortTimeoutSetting,
})
// The chunk announces 256 bytes and the rest of it never comes, so only
// the timeout ends the wait. A hold-up of the test process can only
// make the answer later, so the time is checked only for not being
// shorter than the timeout.
start := time.Now()
s.adminRequest(adminClient, adminBearer+"\r\nTransfer-Encoding: chunked",
http.MethodPost, proxy.BansPath, "100\r\n"+`{"netblock": "203.0.113.9", `,
http.StatusRequestTimeout, requestlog.ActionTimedOut)
if took := time.Since(start); took < shortTimeout {
t.Errorf("answered after %s, before the timeout of %s ran out", took, shortTimeout)
}
wantLimitHits(t, s.addr, clientRequestTimeout, 1)
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
}
func TestClientEndpointShowsTheClientAndItsBans(t *testing.T) { func TestClientEndpointShowsTheClientAndItsBans(t *testing.T) {
t.Parallel() t.Parallel()
+273
View File
@@ -0,0 +1,273 @@
package proxy_test
import (
"maps"
"net/http"
"net/netip"
"reflect"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL"
alertMaxPerHour = "SWWAF_ALERT_MAX_PER_HOUR"
// alertInstance is the instance every alert of these tests gives.
alertInstance = "fsn1app1/gitea"
)
func TestBanForABrokenLimitRaisesABanAlertWithItsNotes(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
rateLimitPerMinute: "1",
banScopeV4Prefix: "24",
})
start := clk.Now()
s.get(client, http.StatusOK, requestlog.ActionForward)
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
netblock := netip.MustParsePrefix("203.0.113.0/24")
ban := server.Ledger.Bans(netblock)[0]
// A request refused under the ban raises no other alert.
clk.advance(time.Minute)
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
wantAlerts(t, queue, banAlert(alerts.EventBan, start, client, bans.Ban{
Netblock: netblock, Cause: bans.CauseLimit,
Reason: "requests per minute over the limit of 1", Notes: ban.Notes,
}, requestlog.FormatTime(start.Add(time.Hour))))
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)
}
}
func TestAttackBanRaisesABanAlertThenAPermanentBanAlert(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithAlerts(t, map[string]string{
rulesDir: writeRules(t, testRules),
})
start := clk.Now()
netblock := netip.MustParsePrefix(client + "/32")
other := netip.MustParsePrefix(otherClient + "/32")
// The probe bans the client for seven days, and its next request makes
// the ban permanent. The request after that changes nothing.
s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
attackBan := server.Ledger.Bans(netblock)[0]
clk.advance(time.Minute)
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
permanentBan := server.Ledger.Bans(netblock)[0]
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
// Another client's probe after its first ban has run out without a
// request makes a permanent ban at once.
s.request(otherClient, "/.env", http.StatusForbidden, requestlog.ActionBanned)
clk.advance(7 * 24 * time.Hour)
s.request(otherClient, "/.env", http.StatusForbidden, requestlog.ActionBanned)
otherBans := server.Ledger.Bans(other)
wantAlerts(t, queue,
attackAlert(alerts.EventBan, start, client, attackBan,
requestlog.FormatTime(start.Add(7*24*time.Hour))),
attackAlert(alerts.EventPermanentBan, start.Add(time.Minute), client,
permanentBan, "permanent"),
attackAlert(alerts.EventBan, start.Add(time.Minute), otherClient, otherBans[0],
requestlog.FormatTime(start.Add(time.Minute+7*24*time.Hour))),
attackAlert(alerts.EventPermanentBan, start.Add(time.Minute+7*24*time.Hour),
otherClient, otherBans[1], "permanent"),
)
}
func TestObserveModeRaisesTheBanAlertsItWouldHave(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.
if held := server.Ledger.Snapshot(); len(held) != 1 || 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()
s, _, _, queue := startWithAlerts(t, map[string]string{
mode: observe,
rateLimitPerMinute: "2",
rulesDir: writeRules(t, testRules),
alertMaxPerHour: "2",
})
// The client's third request breaks the limit, and raises the first
// alert of the hour. Its fourth is within the cooldown.
for range 4 {
s.get(client, http.StatusOK, requestlog.ActionForward)
}
// The other client's first probe raises the second. Its second probe is
// 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
// never sent them, and returns the queue they wait in as well.
func startWithAlerts(
t *testing.T, env map[string]string,
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
t.Helper()
return startAppWithAlerts(t, func(http.ResponseWriter, *http.Request) {}, env)
}
// startAppWithAlerts is startWithAlerts in front of the app handler.
func startAppWithAlerts(
t *testing.T, handler http.HandlerFunc, env map[string]string,
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
t.Helper()
app := startApp(t, handler)
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
settings := map[string]string{
trustedProxies: trustLocalhost,
alertWebhookURL: "https://alerts.example/smallwebwaf",
instanceName: alertInstance,
}
maps.Copy(settings, env)
addr, out, server, queue := startProxyWithAlerts(t, app.URL, "", clk.Now, settings)
return &sender{t: t, addr: addr, out: out}, clk, server, queue
}
// banAlert returns the alert for event, raised by a request from client at
// the time raised, for ban, with its netblock, cause, reason and notes,
// which ends at expires, as the log line gives it.
func banAlert(
event string, raised time.Time, client string, ban bans.Ban, expires string,
) alerts.Alert {
return alerts.Alert{
Instance: alertInstance,
Time: raised,
Event: event,
Client: netip.MustParseAddr(client),
Netblock: ban.Netblock,
Reason: ban.Reason,
Detail: map[string]any{
"cause": ban.Cause, "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.
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
t.Helper()
got := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(got) != len(want) {
t.Fatalf("%d alerts wait, want %d: %+v", len(got), len(want), got)
}
for i := range want {
if !reflect.DeepEqual(got[i], want[i]) {
t.Errorf("alert %d is\n%+v\nwant\n%+v", i, got[i], want[i])
}
}
}
+191 -28
View File
@@ -4,7 +4,9 @@ import (
"net/netip" "net/netip"
"time" "time"
"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"
) )
@@ -16,81 +18,242 @@ 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. // 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,
// 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 check := rq.h.ledger.Check
if rq.h.config.Observe { if rq.h.config.Observe {
check = rq.h.ledger.Find // in observe mode the ban refuses nothing check = rq.h.ledger.Find // the ban refuses nothing, and stays as it is
} }
ban, banned := check(rq.client, now) ban, banned, madePermanent := check(rq.client, now)
if banned { if banned {
rq.line.BanExpires = banExpires(ban) rq.line.BanExpires = banExpires(ban)
} }
if madePermanent {
ban.Expires = time.Time{} // the ban made permanent, which Find leaves as it is
rq.alertBan(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 limit. In enforce mode such a request bans the // the client over a rate limit, as its limit percentage lowers it, which
// client's netblock, and sets the client's counters back to zero; in // breaks it.
// observe mode it does neither.
func (rq *request) limitBroken(now time.Time) bool { func (rq *request) limitBroken(now time.Time) bool {
group := clientGroup(rq.client) counts, hit, over := rq.h.limiter.Count(clientGroup(rq.client), now,
rq.limitPercent.percent)
counts, hit, over := rq.h.limiter.Count(group, now)
rq.line.Counts = counts rq.line.Counts = counts
if !over { if over {
return false rq.banForLimit(now, hit, rq.h.config.BanResponse)
} }
return over
}
// countBytes counts the request's bytes 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. The bytes are
// 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. 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
}
response, request := rq.out.bytes, rq.requestBytes()
if rq.upgraded != nil {
response += rq.upgraded.fromApp.Load()
request += rq.upgraded.toApp.Load()
}
var bytes int64
switch rq.h.config.BytesCount {
case "response":
bytes = response
case "request":
bytes = request
default: // both
bytes = response + request
}
now := rq.h.now()
counts, hit, over := rq.h.limiter.CountBytes(clientGroup(rq.client), now, bytes,
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)
}
}
// 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
if rq.h.config.Observe { netblock := rq.h.netblock(rq.client)
return true if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseLimit) {
return
} }
netblock := rq.h.netblock(rq.client) notes := bans.Notes{
ban := rq.h.ledger.BanForLimit(netblock, now, bans.Notes{ ASN: rq.line.ASN,
ASName: rq.line.ASName,
Country: rq.line.Country, Country: rq.line.Country,
Kind: hit.Kind,
Limit: hit.Limit, Limit: hit.Limit,
Window: hit.Window, Window: hit.Window,
Count: hit.Requests, Count: hit.Count,
Request: rq.noted(now), Request: rq.noted(now, status),
Requests: rq.netblockRequests(netblock), Requests: rq.netblockRequests(netblock),
}) }
rq.h.limiter.Reset(group)
percent := rq.limitPercent
if hit.Kind == ratelimit.KindBytes {
percent = rq.bytesPercent
}
notes.LimitPercent, notes.LimitPercentSetting = percent.logged()
if rq.h.config.Observe {
ban, wouldBan := rq.h.ledger.WouldBanForLimit(netblock, now, notes)
if wouldBan {
rq.alertBan(ban)
}
return
}
ban, made := rq.h.ledger.BanForLimit(netblock, now, notes)
rq.h.limiter.Reset(clientGroup(rq.client))
rq.line.BanExpires = banExpires(ban) rq.line.BanExpires = banExpires(ban)
return true if made {
rq.alertBan(ban)
}
} }
// 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. // attack, the match of rule, a ban rule. In observe mode it makes no ban,
// 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)
ban := rq.h.ledger.BanForAttack(netblock, now, bans.Notes{ if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseAttack) {
return
}
notes := bans.Notes{
ASN: rq.line.ASN,
ASName: rq.line.ASName,
Country: rq.line.Country, Country: rq.line.Country,
RuleID: rule.ID, RuleID: rule.ID,
Target: rule.Target, Target: rule.Target,
Request: rq.noted(now), Request: rq.noted(now, rq.h.config.BanResponse),
Requests: rq.netblockRequests(netblock), 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 {
rq.alertBan(ban)
}
} }
// noted is the request, refused at now with SWWAF_BAN_RESPONSE, as the // wouldAlertBan reports whether the alert for a ban on netblock for cause
// notes of the ban it makes keep it. // made at now would be sent. In observe mode the ban the request would
func (rq *request) noted(now time.Time) bans.Request { // 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
// 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
// observe mode, where ban is the ban that would have been made, or made
// permanent, mode, observe.
func (rq *request) alertBan(ban bans.Ban) {
event := alerts.EventBan
if ban.Permanent() {
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{
Event: event,
Client: rq.client,
Netblock: ban.Netblock,
ASN: ban.Notes.ASN,
ASName: ban.Notes.ASName,
Country: ban.Notes.Country,
Reason: ban.Reason,
Detail: detail,
})
}
// noted is the request, at now, with status, what the client was sent, or
// in observe mode would have been, as the notes of the ban it makes keep
// it.
func (rq *request) noted(now time.Time, status int) bans.Request {
return bans.Request{ 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: rq.h.config.BanResponse, Status: status,
UserAgent: rq.in.UserAgent(), UserAgent: rq.in.UserAgent(),
} }
} }
+6 -2
View File
@@ -281,7 +281,10 @@ 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,
@@ -362,8 +365,9 @@ 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' countries looked up at geojsURL, and a clock // X-Forwarded-For, clients' AS numbers and countries looked up at
// set to midnight, the start of a bucket in every window. // geojsURL, and a clock set to midnight, the start of a bucket in every
// 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) {
+96
View File
@@ -0,0 +1,96 @@
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, or
// SWWAF_UNKNOWN_LIMIT_PERCENT is below 100. 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
}
// limitPercentages returns a 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. 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
// SWWAF_COUNTRY_LIMIT_PERCENT gives its country, and, for a client
// without a country, SWWAF_UNKNOWN_LIMIT_PERCENT. For the byte limits,
// SWWAF_ASN_BYTES_PERCENT and SWWAF_COUNTRY_BYTES_PERCENT take the place
// of the first two for an AS number or a country they list.
func limitPercentages(
cfg *config.Config, asn, country string,
) (percentage, percentage) {
unknown := percentage{percent: whole}
if country == "" {
unknown = percentage{cfg.UnknownLimitPercent, "SWWAF_UNKNOWN_LIMIT_PERCENT"}
}
asnRequests := given(cfg.ASNLimitPercent, asn, "SWWAF_ASN_LIMIT_PERCENT")
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),
lowest(asnBytes, countryBytes, unknown)
}
// 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
}
+494
View File
@@ -0,0 +1,494 @@
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},
// 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 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. It returns the sender, the server and the queue of
// the alerts.
func startWithLookups(
t *testing.T, env map[string]string,
) (*sender, *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)
s, _, server, queue := startAppWithAlerts(t, readAndAnswer, settings)
return s, server, queue
}
// 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,6 +103,48 @@ 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 {
+583
View File
@@ -0,0 +1,583 @@
package proxy_test
import (
"bufio"
"io"
"net"
"net/http"
"net/netip"
"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 || 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
}
+6 -26
View File
@@ -1,31 +1,17 @@
package proxy package proxy
import ( import (
"context"
"net/netip"
"slices" "slices"
) )
// countryDenied reports whether the country lists refuse the request. // countryDenied reports whether the country lists refuse the request, by
// The client's country is looked up only while a list is set, and never // the client's country as it was looked up. A client without a country,
// for a client on a private, loopback or link-local address, which has // or whose country cannot be found, is refused only by
// no country. A client without a country, or whose country cannot be // SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES.
// found, is refused only by SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES. ctx is func (rq *request) countryDenied() bool {
// 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
@@ -33,9 +19,3 @@ func (rq *request) countryDenied(ctx context.Context) 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()
}
+99 -34
View File
@@ -56,16 +56,21 @@ 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)
for i, sent := range []struct{ client, country string }{ // The AS number GeoJS gives unplaced, 64512, counts as unknown.
{fromDE, "DE"}, {fromKP, "KP"}, {unplaced, ""}, for i, sent := range []struct{ client, asn, asName, country string }{
{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.Country != sent.country { if line.ASN != sent.asn || line.ASName != sent.asName ||
t.Errorf("log line has country %q, want %q", line.Country, sent.country) 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) {
@@ -130,7 +135,7 @@ func TestRequestRefusedByCountryIsNotCounted(t *testing.T) {
return return
} }
answer := []map[string]string{{"ip": r.URL.Query().Get("ip"), "country": "DE"}} answer := []geojsAnswer{{IP: r.URL.Query().Get("ip"), CountryCode: "DE"}}
err := json.NewEncoder(w).Encode(answer) err := json.NewEncoder(w).Encode(answer)
if err != nil { if err != nil {
@@ -174,20 +179,15 @@ func TestRequestRefusedByCountryIsNotCounted(t *testing.T) {
wantStatus(t, got, http.StatusOK) wantStatus(t, got, http.StatusOK)
} }
func TestCountryNotLookedUpWithoutAListOrForAPrivateAddress(t *testing.T) { func TestPrivateAddressIsNeverLookedUp(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 country list is set", nil, []string{fromKP, fromDE}}, {"no setting needs the lookup", nil},
{ {"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,7 +198,10 @@ func TestCountryNotLookedUpWithoutAListOrForAPrivateAddress(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)
for i, sent := range tc.clients { // "" sends no X-Forwarded-For: the client is 127.0.0.1.
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)
@@ -209,15 +212,26 @@ func TestCountryNotLookedUpWithoutAListOrForAPrivateAddress(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)
country, present := line.fields["country"] for _, field := range []string{"asn", "as_name", "country"} {
if !present || country != "" { value, present := line.fields[field]
t.Errorf("log line for %q has country %v, want an empty one", if !present || value != "" {
line.ClientIP, country) t.Errorf("log line for %q has %s %v, want an empty one",
line.ClientIP, field, value)
}
} }
} }
if len(asked()) != 0 { // GeoJS is asked about up to 200 waiting clients at once, so once it
t.Errorf("GeoJS was asked about %v, want nothing", asked()) // has been asked about fromDE, which comes last, it has been 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)
} }
}) })
} }
@@ -264,18 +278,31 @@ 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,
// and no other address. It returns its URL, and what returns the // each in an AS of its own, and no other address. It returns its URL, and
// addresses it has been asked about. // what returns the 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()
places := map[string]string{fromDE: "DE", fromKP: "KP"} geojsURL, asked, release := startHeldGeoJS(t)
release()
var asked struct { return geojsURL, asked
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) {
@@ -285,11 +312,11 @@ func startGeoJS(t *testing.T) (string, func() []string) {
asked.addrs = append(asked.addrs, addrs...) asked.addrs = append(asked.addrs, addrs...)
asked.mu.Unlock() asked.mu.Unlock()
answers := make([]map[string]string, 0, len(addrs)) <-released
answers := make([]geojsAnswer, 0, len(addrs))
for _, addr := range addrs { for _, addr := range addrs {
answers = append(answers, map[string]string{ answers = append(answers, answerAbout(addr))
"ip": addr, "country": places[addr],
})
} }
err := json.NewEncoder(w).Encode(answers) err := json.NewEncoder(w).Encode(answers)
@@ -299,10 +326,48 @@ func startGeoJS(t *testing.T) (string, func() []string) {
})) }))
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"}
} }
+6 -2
View File
@@ -24,7 +24,9 @@ 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 country is not looked up. // refused under that ban, for which the client is not looked up. GeoJS
// 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)
@@ -35,8 +37,10 @@ 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.Add(time.Second), LookedUp: start,
Requests: 4, Requests: 4,
Forwarded: 2, Forwarded: 2,
Refused: 2, Refused: 2,
+71
View File
@@ -0,0 +1,71 @@
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(clientGroup(rq.client))
} else {
rq.lookupAnswer = rq.h.geojs.LookUp(ctx, 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()
}
+377
View File
@@ -0,0 +1,377 @@
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{})
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)
}
}
+106 -45
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",status_class="2xx"}` forward := `{action="forward",instance="app",status_class="2xx"}`
notFound := `{action="admin",status_class="4xx"}` notFound := `{action="admin",instance="app",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,11 +159,13 @@ 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, "smallwebwaf_request_duration_seconds_count", 2) wantMetric(t, metrics,
wantMetric(t, metrics, "smallwebwaf_upstream_duration_seconds_count", 1) `smallwebwaf_request_duration_seconds_count{instance="app"}`, 2)
wantMetric(t, metrics, "smallwebwaf_requests_in_flight", 1) wantMetric(t, metrics,
metric(t, metrics, "go_goroutines") `smallwebwaf_upstream_duration_seconds_count{instance="app"}`, 1)
metric(t, metrics, "process_start_time_seconds") wantMetric(t, metrics, `smallwebwaf_requests_in_flight{instance="app"}`, 1)
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)
@@ -180,7 +182,7 @@ func TestMetricsCountTheTraffic(t *testing.T) {
}() }()
<-arrived <-arrived
wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 2) wantMetric(t, scrape(t, addr), `smallwebwaf_requests_in_flight{instance="app"}`, 2)
releaseApp() releaseApp()
err := <-ended err := <-ended
@@ -189,7 +191,7 @@ func TestMetricsCountTheTraffic(t *testing.T) {
} }
out.requestLines(t, 5) out.requestLines(t, 5)
wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 1) wantMetric(t, scrape(t, addr), `smallwebwaf_requests_in_flight{instance="app"}`, 1)
} }
func TestMetricsCountLimitsAndBans(t *testing.T) { func TestMetricsCountLimitsAndBans(t *testing.T) {
@@ -219,15 +221,16 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
metrics := s.scrape(scraper) metrics := s.scrape(scraper)
wantMetric(t, metrics, wantMetric(t, metrics,
`smallwebwaf_requests_total{action="denied",status_class="none"}`, 1) `smallwebwaf_requests_total{action="denied",instance="app",status_class="none"}`, 1)
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 1) wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+
wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 1) `kind="requests",window="minute"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 1) wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 1)
wantMetric(t, metrics, "smallwebwaf_active_bans", 1) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 1)
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 0) wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1)
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", 0) wantMetric(t, s.scrape(scraper), `smallwebwaf_active_bans{instance="app"}`, 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.
@@ -235,13 +238,14 @@ 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{window="minute"}`, 2) wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+
wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 2) `kind="requests",window="minute"}`, 2)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 2) wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 2)
wantMetric(t, metrics, "smallwebwaf_active_bans", 1) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 2)
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1) wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 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", 3) wantMetric(t, metrics, `smallwebwaf_tracked_clients{instance="app"}`, 3)
} }
func TestMetricsCountTheBansAnAdminMakes(t *testing.T) { func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
@@ -255,7 +259,7 @@ func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
rateLimitExemptNets: scraper, rateLimitExemptNets: scraper,
}) })
const admins = `smallwebwaf_bans_made_total{cause="admin"}` const admins = `smallwebwaf_bans_made_total{cause="admin",instance="app"}`
wantMetric(t, s.scrape(scraper), admins, 0) wantMetric(t, s.scrape(scraper), admins, 0)
@@ -287,7 +291,8 @@ func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
metricsTopN: "2", metricsTopN: "2",
deniedCountries: "kp", deniedCountries: "kp",
} }
addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, env) geojsURL, _ := startGeoJS(t)
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.
@@ -319,28 +324,82 @@ func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
metrics := scrape(t, addr) metrics := scrape(t, addr)
lines++ lines++
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3) wantMetric(t, metrics,
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`, 2) `smallwebwaf_country_requests_total{country="KP",instance="app"}`, 3)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 1) wantMetric(t, metrics,
wantMetric(t, metrics, `smallwebwaf_country_list_refusals_total{country="KP"}`, 3) `smallwebwaf_country_requests_total{country="DE",instance="app"}`, 2)
wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="KP"}`, 0) wantMetric(t, metrics,
wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`, 6) `smallwebwaf_country_requests_total{country="other",instance="app"}`, 1)
wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="KP"}`, wantMetric(t, metrics,
`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, `smallwebwaf_country_response_bytes_total{country="other"}`, wantMetric(t, metrics,
`smallwebwaf_country_response_bytes_total{country="other",instance="app"}`,
float64(len("hello"))) float64(len("hello")))
wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`) wantNoSeries(t, metrics,
`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, `smallwebwaf_country_requests_total{country="KP"}`, 3) wantMetric(t, metrics,
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`, 2) `smallwebwaf_country_requests_total{country="KP",instance="app"}`, 3)
wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 2) wantMetric(t, metrics,
wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`) `smallwebwaf_country_requests_total{country="FR",instance="app"}`, 2)
wantNoSeries(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`) wantMetric(t, metrics,
`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) {
@@ -370,16 +429,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") == 0 && for metric(t, metrics, `smallwebwaf_geojs_failures_total{instance="app"}`) == 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", 1) wantMetric(t, metrics, `smallwebwaf_geojs_requests_total{instance="app"}`, 1)
wantMetric(t, metrics, "smallwebwaf_geojs_failures_total", 1) wantMetric(t, metrics, `smallwebwaf_geojs_failures_total{instance="app"}`, 1)
wantMetric(t, metrics, "smallwebwaf_geojs_unanswered_total", 1) wantMetric(t, metrics, `smallwebwaf_geojs_unanswered_total{instance="app"}`, 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
@@ -422,8 +481,9 @@ 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 smallwebwaf_offences_total{kind="limit"}. It fails // their names, such as
// the test if there is no such series. // smallwebwaf_offences_total{instance="app",kind="limit"}. It fails the
// 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()
@@ -471,7 +531,8 @@ 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{limit="` + limit + `"}` series := `smallwebwaf_size_and_time_limit_hits_total{instance="app",limit="` +
limit + `"}`
metrics := scrape(t, addr) metrics := scrape(t, addr)
if hits == 0 { if hits == 0 {
+7 -5
View File
@@ -6,7 +6,6 @@ import (
"errors" "errors"
"io" "io"
"net/http" "net/http"
"os"
"reflect" "reflect"
"slices" "slices"
"strings" "strings"
@@ -123,10 +122,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()
hostname, _ := os.Hostname() bytes := float64(sent + received)
want := withTimings(line, requestlog.Line{ want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: hostname, Type: requestType, Time: line.Time, Instance: "app",
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),
@@ -134,7 +133,10 @@ 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{Minute: 1, Hour: 1, Day: 1}, Counts: ratelimit.Counts{
Minute: 1, Hour: 1, Day: 1,
MinuteBytes: bytes, HourBytes: bytes, DayBytes: bytes,
},
}) })
if !reflect.DeepEqual(line.Line, want) { 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)
@@ -318,7 +320,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), ProcessLog: requestlog.NewProcessLogger(io.Discard, cfg.InstanceName),
}) })
if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 || if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 ||
+52 -22
View File
@@ -11,6 +11,7 @@ import (
"strings" "strings"
"time" "time"
"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/config"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
@@ -53,9 +54,12 @@ 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' countries are looked up, normally // GeoJSURL is where clients' AS numbers and countries are looked up
// lookup.URL. GeoJS is asked only while a country list is set. // while SWWAF_LOOKUP_SOURCE is geojs, normally lookup.URL.
GeoJSURL string GeoJSURL 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.
@@ -63,17 +67,22 @@ type Params struct {
// Rules are the rule files' rules, which each request is checked // Rules are the rule files' rules, which each request is checked
// against. // against.
Rules *rules.Files Rules *rules.Files
// Alerts receive the alert for each ban the proxy makes or makes
// permanent, and for GeoJS failing.
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, and the metrics. // whose state the state files keep, the lookup database, nil unless
// SWWAF_LOOKUP_SOURCE is file, 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
Metrics *metrics.Metrics LookupFile *lookup.File
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
@@ -84,7 +93,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) m := metrics.New(params.Config.MetricsTopN, params.Config.InstanceName)
h := &handler{ h := &handler{
config: params.Config, config: params.Config,
requestLog: params.RequestLog, requestLog: params.RequestLog,
@@ -94,9 +103,12 @@ 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,
}), }),
ledger: bans.New(bans.Rules{ ledger: bans.New(bans.Rules{
LimitBanDuration: params.Config.LimitBanDuration, LimitBanDuration: params.Config.LimitBanDuration,
@@ -105,14 +117,24 @@ func New(params Params) *Server {
AttackBanDuration: params.Config.AttackBanDuration, AttackBanDuration: params.Config.AttackBanDuration,
MaxBans: params.Config.MaxBans, MaxBans: params.Config.MaxBans,
}), }),
geojs: lookup.New(lookup.Params{ lookupFile: params.LookupFile,
URL: params.GeoJSURL, rules: params.Rules,
Now: params.Now, alerts: params.Alerts,
ProcessLog: params.ProcessLog,
Metrics: m,
}),
rules: params.Rules,
} }
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)
@@ -129,10 +151,11 @@ 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,
Metrics: m, LookupFile: h.lookupFile,
Metrics: m,
} }
} }
@@ -149,7 +172,9 @@ type handler struct {
limiter *ratelimit.Limiter limiter *ratelimit.Limiter
ledger *bans.Ledger ledger *bans.Ledger
geojs *lookup.GeoJS geojs *lookup.GeoJS
lookupFile *lookup.File
rules *rules.Files rules *rules.Files
alerts *alerts.Queue
} }
// newTransport returns what carries requests to the app. It never goes // newTransport returns what carries requests to the app. It never goes
@@ -204,5 +229,10 @@ 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())
} }
+99 -31
View File
@@ -14,7 +14,9 @@ import (
"testing" "testing"
"time" "time"
"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"
@@ -65,6 +67,10 @@ const (
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"
@@ -202,8 +208,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' countries looked up at // startProxyWithGeoJS is startProxy with clients' AS numbers and
// geojsURL. // countries looked up at 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) {
@@ -216,43 +222,29 @@ 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. // env sets SWWAF_RULES_DIR, it is an empty directory, of no rules, and
// 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,
) (string, *output, *proxy.Server) { ) (string, *output, *proxy.Server) {
t.Helper() t.Helper()
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir()} addr, out, server, _ := startProxyWithAlerts(t, appURL, geojsURL, now, env)
maps.Copy(settings, env)
cfg, err := config.FromEnvironment(func(name string) (string, bool) { return addr, out, server
value, ok := settings[name] }
return value, ok // startProxyWithAlerts is startProxyWithClock, and returns the queue of
}) // the alerts the proxy raises as well, as newProxy makes them.
if err != nil { func startProxyWithAlerts(
t.Fatalf("settings %v: %v", settings, err) t *testing.T, appURL, geojsURL string, now func() time.Time,
} env map[string]string,
) (string, *output, *proxy.Server, *alerts.Queue) {
t.Helper()
out := &output{} server, out, alertQueue := newProxy(t, appURL, geojsURL, now, env)
processLog := requestlog.NewProcessLogger(out)
ruleFiles, err := rules.Load(rules.Params{
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
})
if err != nil {
t.Fatalf("rule files: %v", err)
}
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: out,
ProcessLog: processLog,
GeoJSURL: geojsURL,
Now: now,
Rules: ruleFiles,
})
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0") listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
if err != nil { if err != nil {
@@ -267,7 +259,83 @@ func startProxyWithClock(
_ = server.Close() _ = server.Close()
}) })
return listener.Addr().String(), out, server 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.
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)
cfg, err := config.FromEnvironment(func(name string) (string, bool) {
value, ok := settings[name]
return value, ok
})
if err != nil {
t.Fatalf("settings %v: %v", settings, err)
}
out := &output{}
processLog := requestlog.NewProcessLogger(out, cfg.InstanceName)
ruleFiles, err := rules.Load(rules.Params{
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
})
if err != nil {
t.Fatalf("rule files: %v", err)
}
alertQueue := alerts.New(alerts.Params{
WebhookURL: cfg.AlertWebhookURL,
Events: cfg.AlertEvents,
Cooldown: cfg.AlertCooldown,
MaxPerHour: cfg.AlertMaxPerHour,
Instance: cfg.InstanceName,
Now: now,
ProcessLog: processLog,
})
var lookupFile *lookup.File
if cfg.LookupSource == fileSource {
lookupFile, err = lookup.OpenFile(lookup.FileParams{
Path: cfg.LookupDBPath, Now: now, ProcessLog: processLog, Alerts: alertQueue,
})
if err != nil {
t.Fatalf("lookup database: %v", err)
}
}
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: out,
ProcessLog: processLog,
GeoJSURL: geojsURL,
LookupFile: lookupFile,
Now: now,
Rules: ruleFiles,
Alerts: alertQueue,
})
return server, out, 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,
+2 -1
View File
@@ -77,7 +77,8 @@ func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
const denied = "192.0.2.50" // in SWWAF_DENY_NETS const denied = "192.0.2.50" // in SWWAF_DENY_NETS
s, _, server := startWithClock(t, "", map[string]string{ geojsURL, _ := startGeoJS(t)
s, _, server := startWithClock(t, geojsURL, map[string]string{
rateLimitPerMinute: "1", rateLimitPerMinute: "1",
rateLimitExemptPaths: "/assets/,/favicon.ico", rateLimitExemptPaths: "/assets/,/favicon.ico",
denyNets: denied, denyNets: denied,
+91 -28
View File
@@ -3,6 +3,7 @@ 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"
@@ -15,6 +16,7 @@ import (
"sync/atomic" "sync/atomic"
"time" "time"
"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"
) )
@@ -48,7 +50,18 @@ type request struct {
client netip.Addr client netip.Addr
peer netip.Addr peer netip.Addr
peerTrusted bool peerTrusted bool
start time.Time // lookedUp is true once the client's AS number and country have been
// 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
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
@@ -59,6 +72,9 @@ 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
@@ -192,14 +208,17 @@ 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. For any other client, // client in SWWAF_ALLOW_NETS skips them, and is not looked up. For any
// SWWAF_DENY_NETS comes first, then a ban on its netblock, so that a // other client, SWWAF_DENY_NETS comes first, then a ban on its netblock,
// client either refuses is not looked up, and then the country lists; a // so that a client either refuses is not looked up, then the lookup of
// request any of them refuses is not counted for the rate limits. Then // its AS number and country, and then the country lists; a request any of
// come the rate limits, unless the client is in // them refuses is not counted for the rate limits. Then come the rate
// SWWAF_RATE_LIMIT_EXEMPT_NETS or the request's path is exempt under // 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) {
@@ -216,13 +235,21 @@ func (rq *request) checkClient(ctx context.Context) string {
return requestlog.ActionBanned return requestlog.ActionBanned
} }
if rq.countryDenied(ctx) { rq.lookUp(ctx)
if rq.countryDenied() {
return requestlog.ActionCountryDenied return requestlog.ActionCountryDenied
} }
exempt := isInside(rq.client, cfg.RateLimitExemptNets) || rq.counted = !isInside(rq.client, cfg.RateLimitExemptNets) &&
pathExempt(rq.in.URL, cfg.RateLimitExemptPaths) !pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
if !exempt && rq.limitBroken(now) { if rq.counted {
rq.limitPercent, rq.bytesPercent = limitPercentages(cfg, rq.line.ASN, rq.line.Country)
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
} }
@@ -287,7 +314,9 @@ 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. // the request's id set, without any X-Client-ASN or X-Client-Country the
// client sent, whatever SWWAF_ADD_LOOKUP_HEADERS says, and, while it is
// set, with the client's AS number and country in them.
func (rq *request) rewrite(pr *httputil.ProxyRequest) { 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
@@ -297,6 +326,12 @@ 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
@@ -307,11 +342,18 @@ 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. // connection it takes over, not through rq.out, and then copies
// what passes each way through res.Body, the connection to the app.
rq.stopTimers() rq.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
} }
@@ -414,10 +456,7 @@ 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
@@ -474,24 +513,48 @@ 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. // history, and then, for a client that was looked up, the lookup
// database's answer about it, or the answer from GeoJS kept about 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()
group := clientGroup(rq.client)
rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{ rq.h.limiter.AddToHistory(group, 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: requestBytes, RequestBytes: rq.requestBytes(),
ResponseBytes: rq.out.bytes, ResponseBytes: rq.out.bytes,
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit, BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
}) })
if !rq.lookedUp {
return
}
// The lookup database's answer was there at once.
if rq.h.config.LookupSource == "file" {
rq.h.addLookup(rq.lookupAnswer)
return
}
answer, kept := rq.h.geojs.Kept(group)
if kept {
rq.h.addLookup(answer)
}
}
// 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
+4 -1
View File
@@ -119,7 +119,10 @@ 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,
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1}, // Its 3 bytes in and 5 out, each way counted by default.
Counts: ratelimit.Counts{
Minute: 1, Hour: 1, Day: 1, MinuteBytes: 8, HourBytes: 8, DayBytes: 8,
},
}) })
if !reflect.DeepEqual(line.Line, want) { 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)
+8 -8
View File
@@ -197,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",rule_id="blocked"}`, 1) `smallwebwaf_rule_matches_total{action="block",instance="app",rule_id="blocked"}`, 1)
wantMetric(t, metrics, wantMetric(t, metrics,
`smallwebwaf_rule_matches_total{action="ban",rule_id="probe"}`, 1) `smallwebwaf_rule_matches_total{action="ban",instance="app",rule_id="probe"}`, 1)
wantMetric(t, metrics, "smallwebwaf_rules_loaded", 2) wantMetric(t, metrics, `smallwebwaf_rules_loaded{instance="app"}`, 2)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_requests_total{action="rule_blocked",`+
`smallwebwaf_requests_total{action="rule_blocked",status_class="4xx"}`, 1) `instance="app",status_class="4xx"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack"}`, 1) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack",instance="app"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 0) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 0)
wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1) wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 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
+4 -5
View File
@@ -10,8 +10,9 @@ 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. In enforce mode // and ActionBanned for a ban rule, or "" when none does. A ban rule bans
// a ban rule bans the client's netblock for a clear sign of attack. // the client's netblock for a clear sign of attack, or in observe mode
// raises the alert for the ban it would have made.
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,9 +30,7 @@ func (rq *request) checkRules(now time.Time) string {
case rules.ActionBlock: case rules.ActionBlock:
return requestlog.ActionRuleBlocked return requestlog.ActionRuleBlocked
case rules.ActionBan: case rules.ActionBan:
if !rq.h.config.Observe { rq.banForAttack(now, last)
rq.banForAttack(now, last)
}
return requestlog.ActionBanned return requestlog.ActionBanned
default: default:
+39 -4
View File
@@ -16,10 +16,10 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
start := midnight() start := midnight()
for i, r := range []ratelimit.Request{ for i, r := range []ratelimit.Request{
{Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100}, {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},
{Country: "FR", Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true}, {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,8 +33,6 @@ 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,
@@ -52,6 +50,43 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
} }
} }
func TestLookupReachesTheHistoryOfAClientInTheTable(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{})
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()
+200 -89
View File
@@ -1,9 +1,9 @@
// Package ratelimit keeps the table of clients: each client's requests // Package ratelimit keeps the table of clients: each client's requests
// counted over a minute, an hour and a day, as the "Counting method" // and bytes counted over a minute, an hour and a day, as the "Counting
// section of SPEC.md describes, which tell when a request takes the client // method" section of SPEC.md describes, which tell when a request takes
// over a rate limit, and each client's history since it was first seen. // the client over a rate limit or a byte limit, and each client's history
// At most 20,000 clients are kept, in memory, and written to clients.json // since it was first seen. At most 20,000 clients are kept, in memory, and
// and read from it by the state package. // written to clients.json and read from it by the state package.
package ratelimit package ratelimit
import ( import (
@@ -23,19 +23,30 @@ 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. Zero is no limit. // a day, and the most bytes. 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 against the limits, and keeps // Limiter counts each client's requests and bytes against the limits, and
// its history. It is safe for concurrent use. // keeps 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. // Client.buckets and Client.byteBuckets.
windows [3]window windows [3]window
mu sync.Mutex mu sync.Mutex
@@ -43,17 +54,23 @@ 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
// in each window, and its history. // of requests and of bytes in each window, and its history.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Client struct { 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"`
History History `json:"history"` MinuteBytes Buckets `json:"minute_bytes"`
HourBytes Buckets `json:"hour_bytes"`
DayBytes Buckets `json:"day_bytes"`
History History `json:"history"`
} }
// Buckets are a client's two buckets in one window: the requests in the // Buckets are a client's two buckets in one window: the requests, or the
// bucket under way, which began at Start, and in the bucket before it. // bytes, in the bucket under way, which began at Start, and in the bucket
// 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"`
@@ -66,8 +83,12 @@ 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"`
// Country is the client's country as it was last looked up, and // ASN, ASName and Country are the client's AS number, AS name and
// LookedUp when that was; both are empty while it never was. // country as last looked up, each empty when the lookup could not
// 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
@@ -97,14 +118,12 @@ type Responses struct {
// Offences are a client's offences, by kind. // Offences are a client's offences, by kind.
type Offences struct { type Offences struct {
// Limit is its requests that broke a rate limit. // Limit is its requests that broke a rate limit or a byte limit.
Limit int64 `json:"limit"` Limit int64 `json:"limit"`
} }
// 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
@@ -117,7 +136,8 @@ 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. // BrokeLimit is true for a request that broke a rate limit or a byte
// limit.
BrokeLimit bool BrokeLimit bool
} }
@@ -130,62 +150,74 @@ func New(limits Limits) *Limiter {
return &Limiter{ return &Limiter{
windows: [3]window{ windows: [3]window{
{name: "minute", length: time.Minute, limit: limits.PerMinute}, {
{name: "hour", length: time.Hour, limit: limits.PerHour}, name: "minute", length: time.Minute,
{name: "day", length: day, limit: limits.PerDay}, limit: limits.PerMinute, byteLimit: limits.BytesPerMinute,
},
{
name: "hour", length: time.Hour,
limit: limits.PerHour, byteLimit: limits.BytesPerHour,
},
{
name: "day", length: day,
limit: limits.PerDay, byteLimit: limits.BytesPerDay,
},
}, },
clients: clients, clients: clients,
} }
} }
// Hit is a request that takes a client over a rate limit. // Hit is a request that takes a client over a rate limit, or whose bytes
// take it over a byte limit.
type Hit struct { 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. // Limit is the window's limit, as the client's percentage of it.
Limit int64 Limit int64
// Requests is the client's requests counted in the window, this one // Count is the client's requests, or bytes, counted in the window,
// included. // this request's included.
Requests float64 Count float64
} }
// Counts are a client's requests in the minute, the hour and the day that // Counts are a client's requests and bytes in the minute, the hour and
// end at a request, that request included. // the day that end at a request, that request's included.
//
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
type Counts struct { 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 requests in each window. It // not it is refused, and returns the client's counts in each window. It
// reports whether the request takes the client over a limit, and the // reports whether the request takes the client over a rate limit, of
// window whose limit it goes over, the shortest if it is over several. // which the client gets the percentage percent, rounded down, and the hit:
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) { // the window whose limit it goes over, the shortest if it is over
l.mu.Lock() // several. A limit that is off stays off.
defer l.mu.Unlock() func (l *Limiter) Count(
client netip.Prefix, now time.Time, percent int64,
var ( ) (Counts, Hit, bool) {
requests [3]float64 return l.count(client, now, 1, 0, percent)
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 != ""
} }
// Reset sets client's counts in every window back to zero. Its history // CountBytes counts bytes, those of a request from client that has ended,
// keeps its totals. // at now, in every window, and returns the client's counts in each window.
// It reports whether the bytes take the client over a byte limit, of which
// the client gets the percentage percent, and the hit, as Count does.
func (l *Limiter) CountBytes(
client netip.Prefix, now time.Time, bytes, percent int64,
) (Counts, Hit, bool) {
return l.count(client, now, 0, bytes, percent)
}
// Reset sets client's counts 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()
@@ -193,6 +225,7 @@ 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{}
} }
} }
@@ -209,11 +242,6 @@ 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++
@@ -232,6 +260,25 @@ func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
} }
} }
// 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
// sent, as their histories count them. // sent, as their histories count them.
func (l *Limiter) Requests(netblock netip.Prefix) int64 { func (l *Limiter) Requests(netblock netip.Prefix) int64 {
@@ -312,12 +359,13 @@ 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, b := range c.buckets() { for i, w := range l.windows {
// The window that ends at now covers neither bucket once it for _, b := range []*Buckets{c.buckets()[i], c.byteBuckets()[i]} {
// begins after the bucket under way has ended. // The window that ends at now covers neither bucket once it
length := l.windows[i].length // begins after the bucket under way has ended.
if !now.Add(-length).Before(b.Start.Add(length)) { if !now.Add(-w.length).Before(b.Start.Add(w.length)) {
*b = Buckets{} *b = Buckets{}
}
} }
} }
@@ -325,6 +373,51 @@ 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 {
@@ -337,30 +430,48 @@ func (l *Limiter) get(client netip.Prefix) *Client {
return c return c
} }
// buckets returns c's buckets in the minute, the hour and the day. // buckets returns c's buckets of requests in the minute, the hour and the
// day.
func (c *Client) buckets() [3]*Buckets { func (c *Client) buckets() [3]*Buckets {
return [3]*Buckets{&c.Minute, &c.Hour, &c.Day} return [3]*Buckets{&c.Minute, &c.Hour, &c.Day}
} }
// window is a length of time over which requests are counted, and the // byteBuckets returns c's buckets of bytes in the minute, the hour and the
// most requests a client may make in it. // day.
type window struct { func (c *Client) byteBuckets() [3]*Buckets {
name string return [3]*Buckets{&c.MinuteBytes, &c.HourBytes, &c.DayBytes}
length time.Duration
limit int64
} }
// add counts a request at now in a window of length, and returns the // window is a length of time over which requests and bytes are counted,
// client's requests in the window that ends at now: those in the bucket // and the most requests and the most bytes a client may have in it.
// under way, and those in the bucket before it weighted by how much of type window struct {
// that bucket the window still covers. name string
length time.Duration
limit int64
byteLimit int64
}
// percentOf returns the percentage percent of limit, rounded down. It is
// written as limit's hundreds times percent, plus the rest's share, since
// limit*percent can overflow for a byte limit.
func percentOf(limit, percent int64) int64 {
const hundred = 100
return limit/hundred*percent + limit%hundred*percent/hundred
}
// add counts n requests, or n bytes, at now in a window of length, and
// returns the client's 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.
// //
// 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) float64 { func (b *Buckets) add(now time.Time, length time.Duration, n int64) float64 {
if now.Before(b.Start.Add(-time.Second)) { if now.Before(b.Start.Add(-time.Second)) {
*b = Buckets{} *b = Buckets{}
} }
@@ -377,7 +488,7 @@ func (b *Buckets) add(now time.Time, length time.Duration) float64 {
b.Current = 0 b.Current = 0
} }
b.Current++ b.Current += n
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)
+184 -6
View File
@@ -1,6 +1,7 @@
package ratelimit_test package ratelimit_test
import ( import (
"math"
"net/netip" "net/netip"
"testing" "testing"
"time" "time"
@@ -11,6 +12,10 @@ import (
// limit is the limit the tests set. // limit is the limit the tests set.
const limit = 3 const limit = 3
// 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"
@@ -62,22 +67,180 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
start := midnight() start := midnight()
for range limit { for range limit {
_, _, over := limiter.Count(client, start) _, _, over := limiter.Count(client, start, whole)
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) _, hit, over := limiter.Count(client, start, whole)
want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1} want := ratelimit.Hit{
Kind: ratelimit.KindRequests, Window: minute, Limit: limit, Count: limit + 1,
}
if !over || hit != want { 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})
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})
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)
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})
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{})
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})
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()
@@ -86,14 +249,14 @@ func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
start := midnight() start := midnight()
for range 3 { for range 3 {
limiter.Count(client, start) limiter.Count(client, start, whole)
} }
// A quarter into the next hour, the minute has only this request. The // 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)) counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4), whole)
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 {
@@ -261,9 +424,24 @@ func wantCount(
) { ) {
t.Helper() t.Helper()
_, hit, _ := limiter.Count(client, now) _, hit, _ := limiter.Count(client, now, whole)
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 -7
View File
@@ -16,7 +16,7 @@ func TestSnapshotListsTheClientsByAddress(t *testing.T) {
limiter := ratelimit.New(ratelimit.Limits{}) 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()) limiter.Count(netip.MustParsePrefix(want[i]), midnight(), whole)
} }
snapshot := limiter.Snapshot() snapshot := limiter.Snapshot()
@@ -63,7 +63,8 @@ func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
start := midnight() start := midnight()
limiter := ratelimit.New(ratelimit.Limits{}) limiter := ratelimit.New(ratelimit.Limits{})
limiter.Count(client, start) limiter.Count(client, start, whole)
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 {
@@ -76,19 +77,25 @@ func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
} }
// 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, which are emptied; the hour's and the day's stay, // minute's buckets, of requests and of bytes, which are emptied; the
// and so does the history. // hour's and the day's stay, 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 { if got.Minute.Current != 1 || got.MinuteBytes.Current != 5 {
t.Errorf("loaded just under two minutes on with minute buckets %+v", t.Errorf("loaded just under two minutes on with minute buckets %+v and %+v",
got.Minute) got.Minute, got.MinuteBytes)
} }
} }
+24 -9
View File
@@ -45,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. // broke a rate limit, or whose bytes broke a byte limit.
const OffenceLimit = "limit" const OffenceLimit = "limit"
// timeLayout is RFC 3339 with milliseconds. // timeLayout is RFC 3339 with milliseconds.
@@ -80,11 +80,14 @@ 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. // client is counted as. ASN, ASName and Country are the client's 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.
@@ -114,13 +117,25 @@ 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"`
// Counts are the client's requests as the rate limits counted them // LimitPercent and LimitPercentSetting are, for a request the rate
// with this one, for a request they counted. // limits counted whose client a biased threshold gives a percentage of
// the rate limits below 100, that percentage and the setting that gave
// it. BytesPercent and BytesPercentSetting are the same for the byte
// limits.
LimitPercent *int64 `json:"limit_percent,omitempty"`
LimitPercentSetting string `json:"limit_percent_setting,omitempty"`
BytesPercent *int64 `json:"bytes_percent,omitempty"`
BytesPercentSetting string `json:"bytes_percent_setting,omitempty"`
// Counts are, for a request the rate limits counted, the client's
// requests as they counted them with this one, and its bytes as the
// byte limits counted them, with this request's once it has ended if
// they count them.
Counts ratelimit.Counts `json:"counts,omitzero"` 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 rate limit the request went over: // LimitHit is the window whose limit the request went over, named as
// minute, hour or day. // Counts names its count: minute, hour or day for a rate limit, and
// minute_bytes, hour_bytes or day_bytes for a byte limit.
LimitHit string `json:"limit_hit,omitempty"` LimitHit string `json:"limit_hit,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"`
@@ -171,8 +186,8 @@ 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. // as a request line's, and instanceName, SWWAF_INSTANCE_NAME, as instance.
func NewProcessLogger(w io.Writer) *slog.Logger { func NewProcessLogger(w io.Writer, instanceName string) *slog.Logger {
handler := slog.NewJSONHandler(w, &slog.HandlerOptions{ handler := slog.NewJSONHandler(w, &slog.HandlerOptions{
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 {
@@ -183,5 +198,5 @@ func NewProcessLogger(w io.Writer) *slog.Logger {
}, },
}) })
return slog.New(handler).With("type", "process") return slog.New(handler).With("type", "process", "instance", instanceName)
} }
+5 -4
View File
@@ -65,12 +65,12 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
} }
} }
func TestProcessLinesAreMarkedProcess(t *testing.T) { func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
t.Parallel() t.Parallel()
var out bytes.Buffer var out bytes.Buffer
requestlog.NewProcessLogger(&out).Info("starting", "version", "v1") requestlog.NewProcessLogger(&out, "fsn1app1/gitea").Info("starting", "version", "v1")
var fields map[string]any var fields map[string]any
@@ -79,8 +79,9 @@ func TestProcessLinesAreMarkedProcess(t *testing.T) {
t.Fatalf("decode %q: %v", out.String(), err) t.Fatalf("decode %q: %v", out.String(), err)
} }
if fields["type"] != "process" || fields["msg"] != "starting" || if fields["type"] != "process" || fields["instance"] != "fsn1app1/gitea" ||
fields["level"] != "INFO" || fields["version"] != "v1" { fields["msg"] != "starting" || fields["level"] != "INFO" ||
fields["version"] != "v1" {
t.Errorf("process line %v", fields) t.Errorf("process line %v", fields)
} }
+28 -14
View File
@@ -22,6 +22,7 @@ import (
"github.com/fsnotify/fsnotify" "github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/config"
) )
@@ -100,6 +101,8 @@ type Params struct {
// ProcessLog receives how many rules were read, and the error in a // ProcessLog receives how many rules were read, and the error in a
// rule file edited while smallwebwaf runs. // rule file edited while smallwebwaf runs.
ProcessLog *slog.Logger ProcessLog *slog.Logger
// Alerts receive a file_error alert for that error.
Alerts *alerts.Queue
} }
// Files are the rule files of a running smallwebwaf, and the rules read // Files are the rule files of a running smallwebwaf, and the rules read
@@ -126,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
} }
@@ -223,13 +226,21 @@ 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. // logs the error that keeps the rules as they were, and raises a
// file_error alert for it, for the file it is in.
func (f *Files) readAgain() { func (f *Files) readAgain() {
rules, err := read(f.params.Dir) rules, path, err := read(f.params.Dir)
if err != nil { if err != nil {
f.params.ProcessLog.Error( const kept = "a rule file has an error, and the rules stay as they were"
"a rule file has an error, and the rules stay as they were",
"error", err.Error()) // 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": path, "error": err.Error()},
})
f.params.ProcessLog.Error(kept, "error", err.Error())
return return
} }
@@ -246,13 +257,14 @@ 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. A file whose name starts with a // files' names, and then of their lines, or an error, with the path of the
// dot, such as an editor's lock file .#50-app.rules, is not a rule file, // rule file it is in, or dir. A file whose name starts with a dot, such as
// as a shell's *.rules would not match it. // an editor's lock file .#50-app.rules, is not a rule file, as a shell's
func read(dir string) ([]Rule, error) { // *.rules would not match it.
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, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err) return nil, dir, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err)
} }
var rules []Rule var rules []Rule
@@ -266,13 +278,15 @@ func read(dir string) ([]Rule, error) {
continue continue
} }
rules, err = readFile(filepath.Join(dir, name), rules, places) path := filepath.Join(dir, name)
rules, err = readFile(path, rules, places)
if err != nil { if err != nil {
return nil, err return nil, path, 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
+32 -7
View File
@@ -7,12 +7,15 @@ import (
"maps" "maps"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"net/url"
"os" "os"
"path/filepath" "path/filepath"
"slices" "slices"
"strconv" "strconv"
"testing" "testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
) )
@@ -337,7 +340,7 @@ func TestEditsTakenInWhileRunning(t *testing.T) {
t.Parallel() t.Parallel()
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"}) dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
files, lines := watch(t, dir) files, lines, _ := watch(t, dir)
// matches reports whether path matches a rule. // matches reports whether path matches a rule.
matches := func(path string) bool { return len(files.Match(get(t, path))) == 1 } matches := func(path string) bool { return len(files.Match(get(t, path))) == 1 }
@@ -366,7 +369,7 @@ func TestBrokenEditKeepsTheRulesAsTheyWere(t *testing.T) {
t.Parallel() t.Parallel()
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"}) dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
files, lines := watch(t, dir) files, lines, queue := watch(t, dir)
// The edit's second line has an unknown action, so the rules stay as // The edit's second line has an unknown action, so the rules stay as
// they were, the first line's earlier version included. // they were, the first line's earlier version included.
@@ -380,13 +383,27 @@ 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.
wantFileError := func() {
t.Helper()
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || waiting[0].Event != alerts.EventFileError ||
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)
}
}
wantFileError()
wantMatched(t, files, get(t, "/first"), "first") wantMatched(t, files, get(t, "/first"), "first")
wantMatched(t, files, get(t, "/second")) wantMatched(t, files, get(t, "/second"))
// Once mended, the file is read again. // Once mended, the file is read again, and raises no alert.
save(t, dir, firstFile, "first path block ^/edited\nsecond path ban ^/second\n") save(t, dir, firstFile, "first path block ^/edited\nsecond path ban ^/second\n")
lines.waitUntil(t, func() bool { return len(files.Match(get(t, "/second"))) == 1 }) lines.waitUntil(t, func() bool { return len(files.Match(get(t, "/second"))) == 1 })
wantMatched(t, files, get(t, "/edited"), "first") wantMatched(t, files, get(t, "/edited"), "first")
wantFileError()
} }
func TestDefaultFileBansProbesAtTheSiteRootAlone(t *testing.T) { func TestDefaultFileBansProbesAtTheSiteRootAlone(t *testing.T) {
@@ -504,7 +521,8 @@ func save(t *testing.T, dir, name, content string) {
} }
// newParams returns Params for the rule files in dir, switched on, with // newParams returns Params for the rule files in dir, switched on, with
// the process log in the processLog returned. // the process log in the processLog returned, and the alerts waiting in a
// queue for a webhook that is never sent them.
func newParams(dir string) (rules.Params, processLog) { func newParams(dir string) (rules.Params, processLog) {
lines := make(processLog, maxLogLines) lines := make(processLog, maxLogLines)
@@ -512,6 +530,12 @@ func newParams(dir string) (rules.Params, processLog) {
Dir: dir, Dir: dir,
Enabled: true, Enabled: true,
ProcessLog: slog.New(slog.NewJSONHandler(lines, nil)), ProcessLog: slog.New(slog.NewJSONHandler(lines, nil)),
Alerts: alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
Now: time.Now,
}),
}, lines }, lines
} }
@@ -531,8 +555,9 @@ func load(t *testing.T, files ruleFiles) *rules.Files {
} }
// watch loads the rules in dir, runs their Watch until the test ends, and // watch loads the rules in dir, runs their Watch until the test ends, and
// waits until it watches the directory. // waits until it watches the directory. It returns the alerts' queue as
func watch(t *testing.T, dir string) (*rules.Files, processLog) { // well.
func watch(t *testing.T, dir string) (*rules.Files, processLog, *alerts.Queue) {
t.Helper() t.Helper()
params, lines := newParams(dir) params, lines := newParams(dir)
@@ -557,7 +582,7 @@ func watch(t *testing.T, dir string) (*rules.Files, processLog) {
lines.waitFor(t, watching) lines.waitFor(t, watching)
return files, lines return files, lines, params.Alerts
} }
// wantRefused checks that loading the rule files in dir fails with the // wantRefused checks that loading the rule files in dir fails with the
+3
View File
@@ -13,6 +13,8 @@ import (
"time" "time"
"github.com/fsnotify/fsnotify" "github.com/fsnotify/fsnotify"
"sneak.berlin/go/smallwebwaf/internal/alerts"
) )
// The tests below run readAfterChanges in a synctest bubble, where time is // The tests below run readAfterChanges in a synctest bubble, where time is
@@ -91,6 +93,7 @@ func load(t *testing.T, dir string) *Files {
files, err := Load(Params{ files, err := Load(Params{
Dir: dir, Enabled: true, ProcessLog: slog.New(slog.DiscardHandler), Dir: dir, Enabled: true, ProcessLog: slog.New(slog.DiscardHandler),
Alerts: alerts.New(alerts.Params{}),
}) })
if err != nil { if err != nil {
t.Fatalf("load: %v", err) t.Fatalf("load: %v", err)
+140 -53
View File
@@ -1,6 +1,7 @@
// Package smallwebwaf runs the smallwebwaf process: it reads the settings, // Package smallwebwaf runs the smallwebwaf process: it reads the settings,
// the rule files and the state files, serves requests until it is told to // the rule files, the lookup database and the state files, serves requests
// stop, and then stops in an orderly way, writing the state files. // until it is told to stop, and then stops in an orderly way, writing the
// state files.
package smallwebwaf package smallwebwaf
import ( import (
@@ -15,6 +16,7 @@ import (
"syscall" "syscall"
"time" "time"
"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/lookup"
"sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/proxy"
@@ -63,11 +65,12 @@ func Main(version string) int {
}) })
} }
// Run reads the settings, the rule files and the state files, then serves // Run reads the settings, the rule files, the lookup database and the
// requests until ctx is done. It returns the process's exit status, 1 // state files, then serves requests until ctx is done. It returns the
// when smallwebwaf cannot start. // process's exit status, 1 when smallwebwaf cannot start.
func Run(ctx context.Context, params Params) int { func Run(ctx context.Context, params Params) int {
processLog := requestlog.NewProcessLogger(params.Stdout) processLog := requestlog.NewProcessLogger(params.Stdout,
config.InstanceName(params.LookupEnv))
cfg, err := config.FromEnvironment(params.LookupEnv) cfg, err := config.FromEnvironment(params.LookupEnv)
if err != nil { if err != nil {
@@ -85,16 +88,22 @@ 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)
stopSending := startSending(ctx, remote, processLog) stopSending := startSending(ctx, remote, processLog)
defer stopSending() defer stopSending()
} }
// The state files and the alerts give times in UTC.
now := func() time.Time { return time.Now().UTC() }
alertQueue := newAlertQueue(cfg, now, processLog)
ruleFiles, err := rules.Load(rules.Params{ ruleFiles, err := rules.Load(rules.Params{
Dir: cfg.RulesDir, Dir: cfg.RulesDir,
Enabled: cfg.RulesEnabled, Enabled: cfg.RulesEnabled,
ProcessLog: processLog, ProcessLog: processLog,
Alerts: alertQueue,
}) })
if err != nil { if err != nil {
processLog.Error("cannot use the rule files", "error", err.Error()) processLog.Error("cannot use the rule files", "error", err.Error())
@@ -102,32 +111,18 @@ func Run(ctx context.Context, params Params) int {
return 1 return 1
} }
// The state files give times in UTC. server, err := newServer(cfg, stdout, processLog, now, ruleFiles, alertQueue)
now := func() time.Time { return time.Now().UTC() } if err != nil {
processLog.Error("cannot use the lookup database", "error", err.Error())
return 1
}
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: stdout,
ProcessLog: processLog,
GeoJSURL: lookup.URL,
Now: now,
Rules: ruleFiles,
})
if remote != nil { if remote != nil {
server.Metrics.AddRemoteLog(remote) server.Metrics.AddRemoteLog(remote)
} }
files, err := state.Load(state.Params{ files, err := loadStateFiles(cfg, server, alertQueue, now, processLog)
Dir: cfg.StateDir,
WriteDelay: cfg.StateWriteDelay,
CounterInterval: cfg.StateCounterInterval,
Ledger: server.Ledger,
Limiter: server.Limiter,
GeoJS: server.GeoJS,
Now: now,
ProcessLog: processLog,
Metrics: server.Metrics,
})
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())
@@ -147,7 +142,89 @@ func Run(ctx context.Context, params Params) int {
"address", listener.Addr().String(), "address", listener.Addr().String(),
"settings", cfg) "settings", cfg)
return serve(ctx, server.Server, listener, files, ruleFiles, processLog) return serve(ctx, server, listener, files, ruleFiles, alertQueue, processLog)
}
// newServer returns the server smallwebwaf runs, with the metrics of the
// alerts, after reading the lookup database while SWWAF_LOOKUP_SOURCE is
// file. A lookup database that cannot be read is an error.
func newServer(
cfg *config.Config, stdout io.Writer, processLog *slog.Logger,
now func() time.Time, ruleFiles *rules.Files, alertQueue *alerts.Queue,
) (*proxy.Server, error) {
var lookupFile *lookup.File
if cfg.LookupSource == "file" {
var err error
lookupFile, err = lookup.OpenFile(lookup.FileParams{
Path: cfg.LookupDBPath,
Now: now,
ProcessLog: processLog,
Alerts: alertQueue,
})
if err != nil {
return nil, err
}
}
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: stdout,
ProcessLog: processLog,
GeoJSURL: lookup.URL,
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
// alertQueue, as state.Load does.
func loadStateFiles(
cfg *config.Config, server *proxy.Server, alertQueue *alerts.Queue,
now func() time.Time, processLog *slog.Logger,
) (*state.Files, error) {
return state.Load(state.Params{
Dir: cfg.StateDir,
WriteDelay: cfg.StateWriteDelay,
CounterInterval: cfg.StateCounterInterval,
Ledger: server.Ledger,
Limiter: server.Limiter,
GeoJS: server.GeoJS,
Alerts: alertQueue,
Now: now,
ProcessLog: processLog,
Metrics: server.Metrics,
})
}
// newAlertQueue returns the queue of the alerts to the webhook, Slack and
// ntfy, with the settings for them.
func newAlertQueue(
cfg *config.Config, now func() time.Time, processLog *slog.Logger,
) *alerts.Queue {
return alerts.New(alerts.Params{
WebhookURL: cfg.AlertWebhookURL,
WebhookHeaders: cfg.AlertWebhookHeaders,
SlackURL: cfg.AlertSlackWebhookURL,
NtfyURL: cfg.AlertNtfyURL,
NtfyToken: cfg.AlertNtfyToken,
Events: cfg.AlertEvents,
Cooldown: cfg.AlertCooldown,
MaxPerHour: cfg.AlertMaxPerHour,
Instance: cfg.InstanceName,
Now: now,
ProcessLog: processLog,
})
} }
// newRemoteLogSender returns a sender of the log lines to // newRemoteLogSender returns a sender of the log lines to
@@ -188,12 +265,15 @@ 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, and reads the rule files again // due, takes in an admin's edits of them, reads the rule files again as
// as they change, until ctx is done. Then it gives the requests in // they change, and the lookup database when it is replaced, and sends the
// progress shutdownTimeout to finish, and writes every state file. // alerts, until ctx is done. Then it gives the requests in progress
// shutdownTimeout to finish, and writes every state file, alerts.json with
// the alerts still waiting.
func serve( func serve(
ctx context.Context, server *http.Server, listener net.Listener, ctx context.Context, server *proxy.Server, listener net.Listener,
files *state.Files, ruleFiles *rules.Files, processLog *slog.Logger, files *state.Files, ruleFiles *rules.Files, alertQueue *alerts.Queue,
processLog *slog.Logger,
) int { ) int {
served := make(chan error, 1) served := make(chan error, 1)
@@ -204,24 +284,15 @@ func serve(
writing, stopWriting := context.WithCancel(ctx) writing, stopWriting := context.WithCancel(ctx)
defer stopWriting() defer stopWriting()
written := make(chan struct{}) written := inBackground(func() { files.Run(writing) })
watched := make(chan struct{}) watched := inBackground(func() { files.Watch(writing) })
rulesWatched := make(chan struct{}) rulesWatched := inBackground(func() { ruleFiles.Watch(writing) })
lookupFileWatched := inBackground(func() {
go func() { if server.LookupFile != nil {
files.Run(writing) server.LookupFile.Watch(writing)
close(written) }
}() })
alertsSent := inBackground(func() { alertQueue.Run(writing) })
go func() {
files.Watch(writing)
close(watched)
}()
go func() {
ruleFiles.Watch(writing)
close(rulesWatched)
}()
select { select {
case err := <-served: case err := <-served:
@@ -253,7 +324,8 @@ func serve(
} }
// Run and Watch have ended, so nothing else reads or writes the // Run and Watch have ended, so nothing else reads or writes the
// files. Every request has ended too, but for two kinds // files, and no alert is being sent, so that alerts.json keeps every
// alert not yet sent. Every request has ended too, but for two kinds
// that Go's server does not wait for: one cut off because Shutdown // that Go's server does not wait for: one cut off because Shutdown
// timed out, and one whose connection switched protocols, such as a // timed out, and one whose connection switched protocols, such as a
// WebSocket. Such a request adds to its client's history only as it // WebSocket. Such a request adds to its client's history only as it
@@ -262,6 +334,8 @@ func serve(
<-written <-written
<-watched <-watched
<-rulesWatched <-rulesWatched
<-lookupFileWatched
<-alertsSent
err = files.WriteAll() err = files.WriteAll()
if err != nil { if err != nil {
@@ -274,3 +348,16 @@ func serve(
return 0 return 0
} }
// inBackground runs task on a goroutine of its own, and returns a channel
// that is closed once task has returned.
func inBackground(task func()) <-chan struct{} {
done := make(chan struct{})
go func() {
task()
close(done)
}()
return done
}
+534 -10
View File
@@ -14,9 +14,11 @@ import (
"strconv" "strconv"
"strings" "strings"
"sync" "sync"
"sync/atomic"
"testing" "testing"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
"sneak.berlin/go/smallwebwaf/internal/smallwebwaf" "sneak.berlin/go/smallwebwaf/internal/smallwebwaf"
) )
@@ -37,11 +39,20 @@ const (
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL" stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
rulesDir = "SWWAF_RULES_DIR" rulesDir = "SWWAF_RULES_DIR"
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name lookupSource = "SWWAF_LOOKUP_SOURCE"
lookupDBPath = "SWWAF_LOOKUP_DB_PATH"
instanceName = "SWWAF_INSTANCE_NAME"
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
// 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.
@@ -98,12 +109,16 @@ 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. // returns its exit status. SWWAF_LOOKUP_SOURCE is off unless env sets it,
// 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
}, },
@@ -116,7 +131,10 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
out := &output{} out := &output{}
status := run(t.Context(), map[string]string{"SWWAF_REQUEST_MAX_BYTES": "lots"}, out) status := run(t.Context(), map[string]string{
"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)
} }
@@ -124,7 +142,8 @@ 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["level"] != "ERROR" || if line["type"] != "process" || line["instance"] != instance ||
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)
} }
@@ -135,7 +154,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, "SWWAF_METRICS_TOKEN"} { for _, name := range []string{adminToken, metricsToken} {
t.Run(name, func(t *testing.T) { t.Run(name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -229,6 +248,46 @@ func TestServesUntilToldToStop(t *testing.T) {
out.line(t, "msg", "stopped") out.line(t, "msg", "stopped")
} }
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()
@@ -441,6 +500,107 @@ 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.
"SWWAF_RATE_LIMIT_EXEMPT_NETS": 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 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()
@@ -516,7 +676,8 @@ 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",
"SWWAF_METRICS_TOKEN": token, metricsToken: token,
instanceName: instance,
} }
out := runUntilStopped(t, env, func(url string) { out := runUntilStopped(t, env, func(url string) {
@@ -526,16 +687,17 @@ 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 0", `smallwebwaf_remote_log_lines_sent_total{instance="fsn1app1/gitea"} 0`,
"smallwebwaf_remote_log_buffer_depth 1", `smallwebwaf_remote_log_buffer_depth{instance="fsn1app1/gitea"} 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)
} }
} }
if strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total 0\n") || const dropped = "\nsmallwebwaf_remote_log_lines_dropped_total" +
!strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total ") { `{instance="fsn1app1/gitea"} `
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)
} }
@@ -545,6 +707,170 @@ 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) {
t.Parallel()
webhook := startWebhook(t)
rules := t.TempDir()
err := os.WriteFile(filepath.Join(rules, "50-app.rules"),
[]byte(`probe path ban ^/\.env$`+"\n"), 0o600)
if err != nil {
t.Fatalf("write the rule file: %v", err)
}
dir := t.TempDir()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: dir,
rulesDir: rules,
"SWWAF_ALERT_WEBHOOK_URL": webhook.url,
"SWWAF_ALERT_WEBHOOK_HEADERS": "Authorization:Bearer " + adminSecret,
instanceName: instance,
}
runUntilStopped(t, env, func(url string) {
// The probe bans the client, and the webhook is sent the alert.
wantRefused(t, url+".env")
post := webhook.waitFor(t, "ban", true)
if post.alert["client"] != localhost || post.alert["netblock"] != localhost+"/32" ||
post.authorization != "Bearer "+adminSecret {
t.Errorf("the webhook was sent %v, with Authorization %q", post.alert,
post.authorization)
}
// The webhook fails, so the alert for the ban made permanent by the
// client's next request waits.
webhook.failing.Store(true)
wantRefused(t, url)
webhook.waitFor(t, "permanent_ban", false)
})
// alerts.json keeps it as smallwebwaf stops, and once started again,
// smallwebwaf sends it.
var file struct {
Waiting map[string][]struct {
Event string `json:"event"`
} `json:"waiting"`
}
path := filepath.Join(dir, "alerts.json")
data, err := os.ReadFile(path) //nolint:gosec // a file in the test's directory
if err == nil {
err = json.Unmarshal(data, &file)
}
waiting := file.Waiting["webhook"]
if err != nil || len(waiting) != 1 || waiting[0].Event != "permanent_ban" {
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
// ban does not cover.
const token = "0123456789abcdef0123456789abcdef"
webhook.failing.Store(false)
env["SWWAF_ALLOW_NETS"] = localhost
env[metricsToken] = token
runUntilStopped(t, env, func(url string) {
webhook.waitFor(t, "permanent_ban", true)
const ofWebhook = `{destination="webhook",instance="fsn1app1/gitea"}`
metrics := metricsWith(t, url+"_smallwebwaf/metrics", token,
"\nsmallwebwaf_alerts_sent_total"+ofWebhook+" 1\n")
for _, series := range []string{"failed", "suppressed", "dropped"} {
zero := "\nsmallwebwaf_alerts_" + series + "_total" + ofWebhook + " 0\n"
if !strings.Contains(metrics, zero) {
t.Errorf("no %q in the metrics:\n%s", zero, metrics)
}
}
})
}
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) {
@@ -806,6 +1132,70 @@ 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) {
@@ -917,6 +1307,140 @@ 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
// it is sent, and answers 204, or 503 while failing.
type webhook struct {
url string
failing atomic.Bool
mu sync.Mutex
posts []webhookPost
}
// webhookPost is an alert the webhook was sent, with the Authorization
// header sent with it, and whether the webhook took it.
type webhookPost struct {
alert map[string]any
authorization string
answered bool
}
// startWebhook starts a webhook that takes every alert.
func startWebhook(t *testing.T) *webhook {
t.Helper()
w := &webhook{}
server := httptest.NewServer(http.HandlerFunc(
func(rw http.ResponseWriter, r *http.Request) {
var alert map[string]any
_ = json.NewDecoder(r.Body).Decode(&alert)
failing := w.failing.Load()
w.mu.Lock()
w.posts = append(w.posts, webhookPost{
alert: alert, authorization: r.Header.Get("Authorization"),
answered: !failing,
})
w.mu.Unlock()
if failing {
rw.WriteHeader(http.StatusServiceUnavailable)
return
}
rw.WriteHeader(http.StatusNoContent)
}))
t.Cleanup(server.Close)
w.url = server.URL + "/alerts"
return w
}
// waitFor waits until the webhook has been sent an alert for event that
// it took, or, unless answered, failed, and returns it. It waits as long
// as that takes, so that a slow test process cannot fail the test.
func (w *webhook) waitFor(t *testing.T, event string, answered bool) webhookPost {
t.Helper()
for {
w.mu.Lock()
for _, post := range w.posts {
if post.alert["event"] == event && post.answered == answered {
w.mu.Unlock()
return post
}
}
w.mu.Unlock()
time.Sleep(pollInterval)
}
}
// statusFrom returns the status a request to url from the client at // statusFrom returns the status a request to url from the client at
// from, as X-Forwarded-For names it, is answered with. // from, as X-Forwarded-For names it, is answered with.
func statusFrom(t *testing.T, url, from string) int { func statusFrom(t *testing.T, url, from string) int {
+152 -30
View File
@@ -1,11 +1,13 @@
// 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, and lookups.json GeoJS's answers. Load reads them at start, // history, lookups.json GeoJS's answers, and alerts.json the cooldowns,
// Watch takes in an admin's edit of one while smallwebwaf runs, and Run // the hour under way and the alerts waiting for each destination. Load
// and WriteAll write them. The disk is read and written outside the // reads them at start, Watch takes in an admin's edit of one while
// parts' locks, which are held only to take a snapshot or to put in what // smallwebwaf runs, and Run and WriteAll write them. The disk is read and
// a file holds, so that no request waits on the disk. // 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 (
@@ -17,13 +19,16 @@ 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/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"
@@ -42,13 +47,18 @@ const (
bansJSON = "bans.json" bansJSON = "bans.json"
clientsJSON = "clients.json" clientsJSON = "clients.json"
lookupsJSON = "lookups.json" lookupsJSON = "lookups.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")
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.
@@ -60,10 +70,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 and GeoJS hold the state. // Ledger, Limiter, GeoJS and Alerts hold the state. Alerts also
// receive a file_error alert for an edit set aside, and for a write
// that fails while smallwebwaf runs.
Ledger *bans.Ledger Ledger *bans.Ledger
Limiter *ratelimit.Limiter Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS GeoJS *lookup.GeoJS
Alerts *alerts.Queue
// 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
@@ -121,6 +134,14 @@ type lookupsFile struct {
Lookups []lookup.Answer `json:"lookups"` Lookups []lookup.Answer `json:"lookups"`
} }
// alertsFile is alerts.json, indented for an admin to read and edit.
type alertsFile struct {
Version int `json:"version"`
Cooldowns []alerts.Cooldown `json:"cooldowns"`
Hour alerts.Hour `json:"hour"`
Waiting map[string][]alerts.Alert `json:"waiting"`
}
// 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
// check refuses the first entry without a field it needs, which would // check refuses the first entry without a field it needs, which would
// otherwise be read as something the entry does not say. data is the // otherwise be read as something the entry does not say. data is the
@@ -147,23 +168,25 @@ 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)
alertsRead, alertsErr := f.read(alertsJSON)
err = errors.Join(bansErr, clientsErr, lookupsErr) 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,
"alerts_waiting", alertsRead)
return f, nil return f, nil
} }
// Run writes bans.json WriteDelay after a ban is made, with every ban // Run writes bans.json WriteDelay after a ban is made, with every ban
// made in between, and every file every CounterInterval, until ctx is // made in between, and every file every CounterInterval, until ctx is
// done. A write that fails is logged, and the file is written again at // done. A write that fails is logged, raised as a file_error alert, and
// its next write. Each write takes in an admin's edit of its file first, // the file is written again at its next write. Each write takes in an
// as writeFile describes. // admin's edit of its file first, as writeFile describes.
func (f *Files) Run(ctx context.Context) { func (f *Files) Run(ctx context.Context) {
interval := time.NewTicker(f.params.CounterInterval) interval := time.NewTicker(f.params.CounterInterval)
defer interval.Stop() defer interval.Stop()
@@ -181,9 +204,11 @@ func (f *Files) Run(ctx context.Context) {
case <-bansDue: case <-bansDue:
bansDue = nil bansDue = nil
f.logFailure(f.writeFile(bansJSON)) f.logFailure(bansJSON, f.writeFile(bansJSON))
case <-interval.C: case <-interval.C:
f.logFailure(f.WriteAll()) for _, name := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} {
f.logFailure(name, f.writeFile(name))
}
} }
} }
} }
@@ -192,7 +217,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(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
@@ -227,7 +252,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: case bansJSON, clientsJSON, lookupsJSON, alertsJSON:
f.fileChanged(name) f.fileChanged(name)
} }
case err = <-watcher.Errors: case err = <-watcher.Errors:
@@ -237,11 +262,22 @@ func (f *Files) Watch(ctx context.Context) {
} }
} }
// logFailure logs a write that failed. // logFailure logs a write of the state file name that failed, and raises
func (f *Files) logFailure(err error) { // a file_error alert for it.
func (f *Files) logFailure(name string, err error) {
if err != nil { if err != nil {
f.params.ProcessLog.Error("writing the state files failed", const failed = "writing the state files failed"
"error", err.Error())
// 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: failed,
Detail: map[string]any{
"file": filepath.Join(f.params.Dir, name), "error": err.Error(),
},
})
f.params.ProcessLog.Error(failed, "error", err.Error())
} }
} }
@@ -362,6 +398,32 @@ 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 alertsJSON:
// 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,
})
for _, waiting := range file.Waiting {
entries += len(waiting)
}
} }
f.sums[name] = sha256.Sum256(data) f.sums[name] = sha256.Sum256(data)
@@ -413,8 +475,9 @@ func (f *Files) writeFile(name string) error {
// setAside renames the state file name, an edit that does not parse with // setAside renames the state file name, an edit that does not parse with
// parseErr, to name.bad, for the admin to mend, and logs it with where in // parseErr, to name.bad, for the admin to mend, and logs it with where in
// the file the error is. If the rename fails, the edit is left as it is, // the file the error is, and raises a file_error alert for it. If the
// and the error returned is parseErr joined with the rename's. // rename fails, the edit is left as it is, and the error returned is
// parseErr joined with the rename's.
func (f *Files) setAside(name string, parseErr error) error { func (f *Files) setAside(name string, parseErr error) error {
path := filepath.Join(f.params.Dir, name) path := filepath.Join(f.params.Dir, name)
@@ -423,8 +486,16 @@ func (f *Files) setAside(name string, parseErr error) error {
return errors.Join(parseErr, err) return errors.Join(parseErr, err)
} }
f.params.ProcessLog.Error("set aside an edit of a state file that does not parse", const setAside = "set aside an edit of a state file that does not parse"
"file", path+".bad", "error", parseErr.Error())
// 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: setAside,
Detail: map[string]any{"file": path + ".bad", "error": parseErr.Error()},
})
f.params.ProcessLog.Error(setAside, "file", path+".bad", "error", parseErr.Error())
f.params.Metrics.StateFileEditSetAside(name) f.params.Metrics.StateFileEditSetAside(name)
return nil return nil
@@ -445,8 +516,21 @@ func (f *Files) encode(name string) ([]byte, error) {
return append(data, '\n'), nil return append(data, '\n'), nil
case clientsJSON: case clientsJSON:
return encodeOnePerLine("clients", f.params.Limiter.Snapshot()) return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
default: // lookups.json case lookupsJSON:
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot()) return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
default: // alerts.json
held := f.params.Alerts.Snapshot()
file := alertsFile{
Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour,
Waiting: held.Waiting,
}
data, err := json.MarshalIndent(file, "", " ")
if err != nil {
return nil, err
}
return append(data, '\n'), nil
} }
} }
@@ -531,8 +615,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 in a window but no start, which would drop // requests, or with requests or bytes in a window but no start, which
// them and give the client a fresh allowance. // would drop 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 {
@@ -544,6 +628,12 @@ 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")
} }
} }
@@ -581,8 +671,40 @@ func (f *lookupsFile) check(data []byte) error {
return nil return nil
} }
// countsWithoutStart reports whether b holds requests but no start, which // check refuses a cooldown without its event or when its alert was sent,
// places them in time. // which would hold back no repeat, alerts waiting for a destination with
// another name than webhook, slack or ntfy, most likely misspelt, and an
// alert waiting without its event or its time.
func (f *alertsFile) check([]byte) error {
for i, cooldown := range f.Cooldowns {
switch {
case cooldown.Event == "":
return fmt.Errorf("cooldowns %w", missing(i, "event"))
case cooldown.Sent.IsZero():
return fmt.Errorf("cooldowns %w", missing(i, "sent"))
}
}
for _, destination := range slices.Sorted(maps.Keys(f.Waiting)) {
if !slices.Contains(alerts.Destinations(), destination) {
return fmt.Errorf("waiting %q %w", destination, errDestination)
}
for i, alert := range f.Waiting[destination] {
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 nil
}
// 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)
} }
+352 -32
View File
@@ -10,8 +10,10 @@ import (
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"net/netip" "net/netip"
"net/url"
"os" "os"
"path/filepath" "path/filepath"
"reflect"
"slices" "slices"
"strconv" "strconv"
"strings" "strings"
@@ -19,6 +21,7 @@ import (
"testing/synctest" "testing/synctest"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"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"
@@ -31,6 +34,10 @@ const (
bansJSON = "bans.json" bansJSON = "bans.json"
clientsJSON = "clients.json" clientsJSON = "clients.json"
lookupsJSON = "lookups.json" lookupsJSON = "lookups.json"
alertsJSON = "alerts.json"
// The AS number and AS name the tests' clients are looked up in.
asn = "AS64496"
asName = "Example Net"
// What the process log says once Watch watches the directory, and as // What the process log says once Watch watches the directory, and as
// it takes in an edit. // it takes in an edit.
watching = "watching the state files for edits" watching = "watching the state files for edits"
@@ -38,6 +45,9 @@ const (
// maxLogLines is how many lines of the process log wait for a test to // maxLogLines is how many lines of the process log wait for a test to
// read them. // read them.
maxLogLines = 64 maxLogLines = 64
// whole is the percentage of each limit a client gets when nothing
// lowers its limits.
whole = 100
) )
// permanentBansJSON is bans.json holding permanentBan. // permanentBansJSON is bans.json holding permanentBan.
@@ -51,6 +61,8 @@ const permanentBansJSON = `{
"cause": "admin", "cause": "admin",
"reason": "scrapes every commit", "reason": "scrapes every commit",
"notes": { "notes": {
"asn": "AS64496",
"as_name": "Example Net",
"country": "DE", "country": "DE",
"limit": 1000, "limit": 1000,
"window": "minute", "window": "minute",
@@ -85,6 +97,69 @@ const liftedBansJSON = `{"version": 1, "bans": [{"netblock": "203.0.113.9/32", `
`"start": "2026-10-06T00:00:00Z", "expires": "2026-10-06T01:00:00Z", ` + `"start": "2026-10-06T00:00:00Z", "expires": "2026-10-06T01:00:00Z", ` +
`"cause": "limit", "lifted": "2026-10-06T00:10:00Z"}]}` `"cause": "limit", "lifted": "2026-10-06T00:10:00Z"}]}`
// filledAlertsJSON is alerts.json holding the alerts of fill.
const filledAlertsJSON = `{
"version": 1,
"cooldowns": [
{
"event": "file_error",
"netblock": "",
"file": "/var/lib/smallwebwaf/bans.json",
"sent": "2026-10-06T00:00:00Z",
"suppressed_repeats": 0
},
{
"event": "ban",
"netblock": "203.0.113.9/32",
"sent": "2026-10-06T00:00:00Z",
"suppressed_repeats": 1
}
],
"hour": {
"start": "2026-10-06T00:00:00Z",
"sent": 2,
"held_back": {
"source_failure": 1
}
},
"waiting": {
"webhook": [
{
"instance": "fsn1app1/gitea",
"time": "2026-10-06T00:00:00Z",
"event": "ban",
"client": "203.0.113.9",
"netblock": "203.0.113.9/32",
"asn": "",
"as_name": "",
"country": "DE",
"reason": "requests per minute over the limit of 1",
"detail": {
"cause": "limit"
},
"suppressed_repeats": 0
},
{
"instance": "fsn1app1/gitea",
"time": "2026-10-06T00:00:00Z",
"event": "file_error",
"client": "",
"netblock": "",
"asn": "",
"as_name": "",
"country": "",
"reason": "writing the state files failed",
"detail": {
"error": "no space left on device",
"file": "/var/lib/smallwebwaf/bans.json"
},
"suppressed_repeats": 0
}
]
}
}
`
func TestFilesWrittenAndReadBack(t *testing.T) { func TestFilesWrittenAndReadBack(t *testing.T) {
t.Parallel() t.Parallel()
@@ -111,13 +186,69 @@ func TestFilesWrittenAndReadBack(t *testing.T) {
wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot()) wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot())
wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot()) wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot())
if got, want := after.Alerts.Snapshot(), before.Alerts.Snapshot(); !reflect.DeepEqual(
got, want) {
t.Errorf("%s read back\n%+v\nwant\n%+v", alertsJSON, got, want)
}
// Each one-per-line file lists its entries by client, and nothing // Each one-per-line file lists its entries by client, and nothing
// but the three files is left in the directory. // but the four files is left in the directory.
wantEntries(t, filepath.Join(dir, clientsJSON), "clients", wantEntries(t, filepath.Join(dir, clientsJSON), "clients",
"192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64") "192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64")
wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups", wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups",
"192.0.2.1/32", "203.0.113.9/32") "192.0.2.1/32", "203.0.113.9/32")
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
}
func TestAlertsJSONIsIndentedWithTheCooldownsTheHourAndTheAlertsWaiting(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
fill(params)
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
got := readFile(t, filepath.Join(dir, alertsJSON))
if got != filledAlertsJSON {
t.Errorf("alerts.json\n%s\nwant\n%s", got, filledAlertsJSON)
}
}
func TestSourceFailureCooldownKeptInAlertsJSONAcrossARestart(t *testing.T) {
t.Parallel()
dir := t.TempDir()
before := newParams(dir)
files := load(t, before)
failure := alerts.Alert{
Event: alerts.EventSourceFailure, Reason: "asking GeoJS failed",
Detail: map[string]any{"source": "geojs"},
}
before.Alerts.Raise(failure)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
// After the restart, the cooldown read back holds back a repeat for the
// same source.
after := newParams(dir)
load(t, after)
after.Alerts.Raise(failure)
waiting := after.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || after.Alerts.Suppressed() != 1 {
t.Errorf("%d alerts wait and %d are held back, want the one read back and 1",
len(waiting), after.Alerts.Suppressed())
}
} }
func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) { func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) {
@@ -146,8 +277,10 @@ func TestMissingFilesAreEmptyState(t *testing.T) {
params := newParams(t.TempDir()) params := newParams(t.TempDir())
load(t, params) load(t, params)
held := params.Alerts.Snapshot()
if len(params.Ledger.Snapshot()) != 0 || len(params.Limiter.Snapshot()) != 0 || if len(params.Ledger.Snapshot()) != 0 || len(params.Limiter.Snapshot()) != 0 ||
len(params.GeoJS.Snapshot()) != 0 { len(params.GeoJS.Snapshot()) != 0 || len(held.Cooldowns) != 0 ||
len(held.Waiting[alerts.DestinationWebhook]) != 0 || held.Hour.Sent != 0 {
t.Error("state from no files") t.Error("state from no files")
} }
} }
@@ -189,6 +322,16 @@ func TestFileThatDoesNotParseStopsTheStart(t *testing.T) {
`{"version": 1, "bans": [{"netblock": "203.0.113.300/32"}]}`, `{"version": 1, "bans": [{"netblock": "203.0.113.300/32"}]}`,
`: netip.ParsePrefix("203.0.113.300/32")`, `: netip.ParsePrefix("203.0.113.300/32")`,
}, },
{
"an unknown field of an alert waiting", alertsJSON,
`{"version": 1, "waiting": {"webhook": [{"event": "ban", "evnet": "ban"}]}}`,
`: json: unknown field "evnet"`,
},
{
"alerts waiting for an unknown destination", alertsJSON,
`{"version": 1, "waiting": {"webhook": [], "slak": []}}`,
`: waiting "slak" is not webhook, slack or ntfy`,
},
} { } {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -252,6 +395,12 @@ func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
`"hour": {"current": 3}}]}`, `"hour": {"current": 3}}]}`,
`: entry 1 has no "hour.start"`, `: entry 1 has no "hour.start"`,
}, },
{
"a client with bytes in a window without its start", clientsJSON,
`{"version": 1, "clients": [{"client": "203.0.113.9/32", ` +
`"minute_bytes": {"previous": 5120}}]}`,
`: entry 1 has no "minute_bytes.start"`,
},
{ {
"an answer without a client", lookupsJSON, "an answer without a client", lookupsJSON,
`{"version": 1, "lookups": [{"country": "DE", ` + answer + `}]}`, `{"version": 1, "lookups": [{"country": "DE", ` + answer + `}]}`,
@@ -278,6 +427,60 @@ func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
} }
} }
func TestAlertsJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name, content string
// want is what the error says after the file's path.
want string
}{
{
"a cooldown without its event",
`{"version": 1, "cooldowns": [{"sent": "2026-10-06T00:00:00Z"}]}`,
`: cooldowns entry 1 has no "event"`,
},
{
"a cooldown without when it was sent",
`{"version": 1, "cooldowns": [{"event": "ban"}]}`,
`: cooldowns entry 1 has no "sent"`,
},
{
"an alert waiting without its event",
`{"version": 1, "waiting": {"ntfy": [{"time": "2026-10-06T00:00:00Z"}]}}`,
`: waiting ntfy entry 1 has no "event"`,
},
{
"an alert waiting without its time",
`{"version": 1, "waiting": {"slack": [{"event": "ban"}]}}`,
`: waiting slack entry 1 has no "time"`,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
wantRefused(t, alertsJSON, tc.content, tc.want)
})
}
}
func TestAlertsJSONWithWaitingAsAListStopsTheStartSayingWhatToChange(t *testing.T) {
t.Parallel()
// alerts.json as it was written before alerts went to Slack and ntfy too,
// with no alert waiting, or one.
for _, waiting := range []string{
`[]`,
`[{"event": "ban", "time": "2026-10-06T00:00:00Z"}]`,
} {
wantRefused(t, alertsJSON, `{"version": 1, "cooldowns": [], `+
`"hour": {"start": "2026-10-06T00:00:00Z", "sent": 0, "held_back": {}}, `+
`"waiting": `+waiting+`}`,
`: waiting is a list, but now lists the alerts by destination: put the `+
`list under "webhook", as "waiting": {"webhook": [...]}, or remove the file`)
}
}
func TestBanWithAnotherCauseStopsTheStart(t *testing.T) { func TestBanWithAnotherCauseStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
@@ -294,7 +497,7 @@ func TestBanWithAnotherCauseStopsTheStart(t *testing.T) {
func TestUnknownVersionStopsTheStart(t *testing.T) { func TestUnknownVersionStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
for _, file := range []string{bansJSON, clientsJSON, lookupsJSON} { for _, file := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} {
for _, content := range []string{`{"version": 2}`, `{}`} { for _, content := range []string{`{"version": 2}`, `{}`} {
t.Run(file+" "+content, func(t *testing.T) { t.Run(file+" "+content, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -344,12 +547,12 @@ func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) {
// A second ban, made while the first waits to be written, puts the // A second ban, made while the first waits to be written, puts the
// write off no further, and is written with it. // write off no further, and is written with it.
first := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), first, _ := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"),
midnight(), bans.Notes{}) midnight(), bans.Notes{})
time.Sleep(5 * time.Second) time.Sleep(5 * time.Second)
second := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"), second, _ := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"),
midnight(), bans.Notes{}) midnight(), bans.Notes{})
time.Sleep(5*time.Second - time.Nanosecond) time.Sleep(5*time.Second - time.Nanosecond)
@@ -396,8 +599,8 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
time.Sleep(time.Nanosecond) time.Sleep(time.Nanosecond)
synctest.Wait() synctest.Wait()
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
removeFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) removeFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
} }
}) })
} }
@@ -440,6 +643,57 @@ func TestEditJustBeforeAScheduledWriteSurvivesIt(t *testing.T) {
}) })
} }
func TestWriteThatFailsWhileRunningRaisesAFileErrorAlertOncePerCooldown(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
params := newParams(dir)
params.CounterInterval = time.Minute
run(t, load(t, params).Run)
// A directory in the way of clients.json's temporary file fails each
// of its writes, while the other files are written. It holds a file,
// so that the write cannot remove it.
err := os.Mkdir(filepath.Join(dir, clientsJSON+".tmp"), 0o700)
if err == nil {
err = os.WriteFile(filepath.Join(dir, clientsJSON+".tmp", "kept"), nil, 0o600)
}
if err != nil {
t.Fatalf("put a directory in the way: %v", err)
}
time.Sleep(time.Minute)
synctest.Wait()
waiting := params.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 {
t.Fatalf("%d alerts wait, want 1", len(waiting))
}
message, _ := waiting[0].Detail["error"].(string)
if waiting[0].Event != alerts.EventFileError ||
waiting[0].Reason != "writing the state files failed" ||
waiting[0].Detail["file"] != filepath.Join(dir, clientsJSON) ||
!strings.Contains(message, clientsJSON+".tmp") {
t.Fatalf("alerts waiting %+v, want a file_error alert for clients.json, "+
"naming its temporary file", waiting)
}
// The next write fails too, within the cooldown, which holds it back.
time.Sleep(time.Minute)
synctest.Wait()
waiting = params.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || params.Alerts.Suppressed() != 1 {
t.Errorf("%d alerts wait and %d are held back, want 1 and 1",
len(waiting), params.Alerts.Suppressed())
}
})
}
func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) { func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
t.Parallel() t.Parallel()
@@ -462,7 +716,7 @@ func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) {
params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(), params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(),
bans.Notes{}) bans.Notes{})
params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight()) params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight(), whole)
err = files.WriteAll() err = files.WriteAll()
if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") { if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") {
@@ -496,8 +750,8 @@ func TestWritesAreCountedInTheMetrics(t *testing.T) {
} }
const ( const (
ofBans = `{file="bans.json"}` ofBans = `{file="bans.json",instance="app"}`
ofClients = `{file="clients.json"}` ofClients = `{file="clients.json",instance="app"}`
) )
got := scrape(t, params) got := scrape(t, params)
@@ -565,7 +819,7 @@ func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) {
t.Errorf("bans.json is now %v (%v), want the socket", info, err) t.Errorf("bans.json is now %v (%v), want the socket", info, err)
} }
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
wantWriteFailed(t, params, bansJSON) wantWriteFailed(t, params, bansJSON)
} }
@@ -637,6 +891,27 @@ func TestEditOfEachFileTakenIn(t *testing.T) {
wantTakenIn(t, lines, dir, lookupsJSON) wantTakenIn(t, lines, dir, lookupsJSON)
wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(), wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(),
[]lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}}) []lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}})
// A netblock with bits past its length is read as the netblock it is
// in.
edit(t, dir, alertsJSON, `{"version": 1, "cooldowns": [{"event": "ban", `+
`"netblock": "198.51.100.9/24", "sent": "2026-10-06T00:00:00Z"}], `+
`"waiting": {"webhook": [{"event": "file_error", "time": "2026-10-06T00:00:00Z"}]}}`)
wantTakenIn(t, lines, dir, alertsJSON)
want := alerts.State{
Cooldowns: []alerts.Cooldown{{
Event: alerts.EventBan, Netblock: netip.MustParsePrefix("198.51.100.0/24"),
Sent: midnight(),
}},
Hour: alerts.Hour{HeldBack: map[string]int{}},
Waiting: map[string][]alerts.Alert{
alerts.DestinationWebhook: {{Event: alerts.EventFileError, Time: midnight()}},
},
}
if got := params.Alerts.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("%s taken in as\n%+v\nwant\n%+v", alertsJSON, got, want)
}
} }
func TestOwnWritesAreNotTakenIn(t *testing.T) { func TestOwnWritesAreNotTakenIn(t *testing.T) {
@@ -708,7 +983,7 @@ func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) {
`"start": "2026-10-06T00:00:00Z", "expires": null}]}`) `"start": "2026-10-06T00:00:00Z", "expires": null}]}`)
wantTakenIn(t, lines, dir, bansJSON) wantTakenIn(t, lines, dir, bansJSON)
_, banned := params.Ledger.Check(client, midnight()) _, banned, _ := params.Ledger.Check(client, midnight())
if !banned { if !banned {
t.Error("the ban added to bans.json does not refuse") t.Error("the ban added to bans.json does not refuse")
} }
@@ -717,7 +992,7 @@ func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) {
edit(t, dir, bansJSON, `{"version": 1, "bans": []}`) edit(t, dir, bansJSON, `{"version": 1, "bans": []}`)
wantTakenIn(t, lines, dir, bansJSON) wantTakenIn(t, lines, dir, bansJSON)
_, banned = params.Ledger.Check(client, midnight()) _, banned, _ = params.Ledger.Check(client, midnight())
if banned { if banned {
t.Error("the ban removed from bans.json still refuses") t.Error("the ban removed from bans.json still refuses")
} }
@@ -807,7 +1082,7 @@ func TestBanLiftedByAnEditWhileRunning(t *testing.T) {
netblock := netip.MustParsePrefix(liftedClient + "/32") netblock := netip.MustParsePrefix(liftedClient + "/32")
params.Ledger.BanForLimit(netblock, midnight(), bans.Notes{}) params.Ledger.BanForLimit(netblock, midnight(), bans.Notes{})
_, banned := params.Ledger.Find(netblock.Addr(), afterLifting()) _, banned, _ := params.Ledger.Find(netblock.Addr(), afterLifting())
if !banned { if !banned {
t.Fatal("the ban does not refuse before it is lifted") t.Fatal("the ban does not refuse before it is lifted")
} }
@@ -844,7 +1119,7 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
edit(t, dir, bansJSON, broken) edit(t, dir, bansJSON, broken)
edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`) edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`)
wantTakenIn(t, lines, dir, clientsJSON) wantTakenIn(t, lines, dir, clientsJSON)
wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON)
// The next write sets it aside, logged with where the error is, and // The next write sets it aside, logged with where the error is, and
// writes bans.json again from what smallwebwaf still holds. // writes bans.json again from what smallwebwaf still holds.
@@ -861,7 +1136,14 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
t.Errorf("set aside with %v", line) t.Errorf("set aside with %v", line)
} }
wantFiles(t, dir, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON) // It is raised as a file_error alert, with the same file and error.
waiting := params.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || waiting[0].Event != alerts.EventFileError ||
waiting[0].Detail["file"] != path+".bad" || waiting[0].Detail["error"] != message {
t.Errorf("alerts waiting %+v, want a file_error alert for %s", waiting, path+".bad")
}
wantFiles(t, dir, alertsJSON, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON)
if got := readFile(t, path+".bad"); got != broken { if got := readFile(t, path+".bad"); got != broken {
t.Errorf("bans.json.bad holds\n%s\nwant the edit", got) t.Errorf("bans.json.bad holds\n%s\nwant the edit", got)
@@ -894,7 +1176,7 @@ func TestEditsTakenInAreCountedInTheMetrics(t *testing.T) {
wantTakenIn(t, lines, dir, bansJSON) wantTakenIn(t, lines, dir, bansJSON)
wantMetric(t, scrape(t, params), wantMetric(t, scrape(t, params),
`smallwebwaf_state_file_edits_taken_in_total{file="bans.json"}`, 2) `smallwebwaf_state_file_edits_taken_in_total{file="bans.json",instance="app"}`, 2)
} }
func TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) { func TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) {
@@ -959,7 +1241,7 @@ func TestEditsSetAsideAreCountedInTheMetrics(t *testing.T) {
} }
wantMetric(t, scrape(t, params), wantMetric(t, scrape(t, params),
`smallwebwaf_state_file_edits_set_aside_total{file="bans.json"}`, 1) `smallwebwaf_state_file_edits_set_aside_total{file="bans.json",instance="app"}`, 1)
} }
func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) { func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) {
@@ -990,10 +1272,11 @@ func midnight() time.Time {
} }
// newParams returns Params for the state files in dir, with parts that // newParams returns Params for the state files in dir, with parts that
// hold nothing yet. GeoJS is never asked. // hold nothing yet. GeoJS is never asked, and the alerts, at most two an
// hour, are never sent.
func newParams(dir string) state.Params { func newParams(dir string) state.Params {
discard := slog.New(slog.DiscardHandler) discard := slog.New(slog.DiscardHandler)
m := metrics.New(1) m := metrics.New(1, "app")
return state.Params{ return state.Params{
Dir: dir, Dir: dir,
@@ -1010,6 +1293,14 @@ func newParams(dir string) state.Params {
GeoJS: lookup.New(lookup.Params{ GeoJS: lookup.New(lookup.Params{
Now: midnight, ProcessLog: discard, Metrics: m, Now: midnight, ProcessLog: discard, Metrics: m,
}), }),
Alerts: alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
MaxPerHour: 2,
Instance: "fsn1app1/gitea",
Now: midnight,
}),
Now: midnight, Now: midnight,
ProcessLog: discard, ProcessLog: discard,
Metrics: m, Metrics: m,
@@ -1017,32 +1308,59 @@ func newParams(dir string) state.Params {
} }
// fill puts a permanent ban an admin made, a ban for a broken limit and // fill puts a permanent ban an admin made, a ban for a broken limit and
// one for a clear sign of attack, clients with counts and histories, and // one for a clear sign of attack, clients with counts and histories,
// GeoJS answers into the parts of params. // GeoJS answers, and alerts, as filledAlertsJSON holds them, into the
// parts of params.
func fill(params state.Params) { func fill(params state.Params) {
now := midnight() now := midnight()
client := netip.MustParsePrefix("203.0.113.9/32") client := netip.MustParsePrefix("203.0.113.9/32")
params.Ledger.Load([]bans.Ban{permanentBan()}) params.Ledger.Load([]bans.Ban{permanentBan()})
params.Ledger.BanForLimit(client, now, bans.Notes{Country: "DE", Limit: 1}) params.Ledger.BanForLimit(client, now, bans.Notes{
ASN: asn, ASName: asName, Country: "DE", Limit: 1,
})
params.Ledger.BanForAttack(netip.MustParsePrefix("192.0.2.1/32"), now, params.Ledger.BanForAttack(netip.MustParsePrefix("192.0.2.1/32"), now,
bans.Notes{RuleID: "env-file", Target: "path"}) bans.Notes{RuleID: "env-file", Target: "path"})
for _, c := range []string{"2001:db8::/64", "203.0.113.9/32", "192.0.2.1/32"} { for _, c := range []string{"2001:db8::/64", "203.0.113.9/32", "192.0.2.1/32"} {
params.Limiter.Count(netip.MustParsePrefix(c), now) params.Limiter.Count(netip.MustParsePrefix(c), now, whole)
} }
params.Limiter.CountBytes(client, now, 8, whole)
params.Limiter.AddToHistory(client, now, ratelimit.Request{ params.Limiter.AddToHistory(client, now, ratelimit.Request{
Country: "DE", Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5, Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5,
}) })
params.Limiter.AddLookup(client, now.Add(-time.Hour), asn, asName, "DE")
params.GeoJS.Load([]lookup.Answer{ params.GeoJS.Load([]lookup.Answer{
{Client: client, Country: "DE", Answered: now.Add(-time.Hour), Used: now}, {
Client: client, ASN: asn, ASName: asName, Country: "DE",
Answered: now.Add(-time.Hour), Used: now,
},
{ {
Client: netip.MustParsePrefix("192.0.2.1/32"), Client: netip.MustParsePrefix("192.0.2.1/32"),
Answered: now.Add(-time.Hour), Used: now.Add(-time.Minute), Answered: now.Add(-time.Hour), Used: now.Add(-time.Minute),
}, },
}) })
// An alert waiting, a repeat of it the cooldown holds back, another
// alert waiting, and one past the two an hour, for the hour's summary.
ban := alerts.Alert{
Event: alerts.EventBan, Client: client.Addr(), Netblock: client, Country: "DE",
Reason: "requests per minute over the limit of 1",
Detail: map[string]any{"cause": "limit"},
}
params.Alerts.Raise(ban)
params.Alerts.Raise(ban)
params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError, Reason: "writing the state files failed",
Detail: map[string]any{
"file": "/var/lib/smallwebwaf/bans.json", "error": "no space left on device",
},
})
params.Alerts.Raise(alerts.Alert{
Event: alerts.EventSourceFailure, Reason: "asking GeoJS failed",
})
} }
// permanentBan is the ban permanentBansJSON holds. // permanentBan is the ban permanentBansJSON holds.
@@ -1053,6 +1371,8 @@ func permanentBan() bans.Ban {
Cause: bans.CauseAdmin, Cause: bans.CauseAdmin,
Reason: "scrapes every commit", Reason: "scrapes every commit",
Notes: bans.Notes{ Notes: bans.Notes{
ASN: asn,
ASName: asName,
Country: "DE", Country: "DE",
Limit: 1000, Limit: 1000,
Window: "minute", Window: "minute",
@@ -1098,13 +1418,13 @@ func wantLiftedBanKept(
netblock := netip.MustParsePrefix(liftedClient + "/32") netblock := netip.MustParsePrefix(liftedClient + "/32")
_, banned := ledger.Check(netblock.Addr(), afterLifting()) _, banned, _ := ledger.Check(netblock.Addr(), afterLifting())
if banned { if banned {
t.Error("the lifted ban refuses") t.Error("the lifted ban refuses")
} }
// Were the lifted ban counted, the next would last three hours. // Were the lifted ban counted, the next would last three hours.
ban := ledger.BanForLimit(netblock, afterLifting(), bans.Notes{}) ban, _ := ledger.BanForLimit(netblock, afterLifting(), bans.Notes{})
if ban.Expires.Sub(ban.Start) != time.Hour { if ban.Expires.Sub(ban.Start) != time.Hour {
t.Errorf("the next ban lasts %s, want 1h", ban.Expires.Sub(ban.Start)) t.Errorf("the next ban lasts %s, want 1h", ban.Expires.Sub(ban.Start))
} }
@@ -1350,8 +1670,8 @@ func scrape(t *testing.T, params state.Params) string {
} }
// metric returns the value of series in text, the metrics, such as // metric returns the value of series in text, the metrics, such as
// smallwebwaf_state_file_writes_total{file="bans.json"}, or fails the test // smallwebwaf_state_file_writes_total{file="bans.json",instance="app"}, or
// if there is no such series. // fails the test if there is no such series.
func metric(t *testing.T, text, series string) float64 { func metric(t *testing.T, text, series string) float64 {
t.Helper() t.Helper()
@@ -1380,7 +1700,7 @@ func wantWriteFailed(t *testing.T, params state.Params, name string) {
t.Helper() t.Helper()
got := scrape(t, params) got := scrape(t, params)
file := `{file="` + name + `"}` file := `{file="` + name + `",instance="app"}`
wantMetric(t, got, "smallwebwaf_state_file_writes_total"+file, 1) wantMetric(t, got, "smallwebwaf_state_file_writes_total"+file, 1)
wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+file, 1) wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+file, 1)
+7 -5
View File
@@ -7,9 +7,10 @@
# 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, the volume and both images are removed however the script # containers run with SWWAF_LOOKUP_SOURCE=off, so that no address is sent
# ends. Building the app needs network access, for nixpkgs' binary cache. # to GeoJS. The containers, the volume and both images are removed however
# script/check does not run this. # the script ends. Building the app needs network access, for nixpkgs'
# 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)"
@@ -63,12 +64,13 @@ 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 and a rate limit of one request a minute, and wait until it is # volume, a rate limit of one request a minute and no client looked up,
# healthy. # and wait until it is 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)"
Executable
+19
View File
@@ -0,0 +1,19 @@
#!/bin/sh
# script/tidy: write go.mod and go.sum as `go mod tidy` writes them, which
# the test phase of the Dockerfile checks. This builds the Dockerfile's
# tidy-files stage, which holds the two files alone, and --output writes
# them into the working tree. The build makes no image, so it has no tag.
# --no-cache because `go mod tidy` asks the module proxy, whose answers a
# cached layer would repeat.
set -eu
ROOT="$(cd "$(dirname "$0")/.." && pwd -P)"
main() {
cd "$ROOT"
docker build --no-cache \
--target tidy-files \
--output "type=local,dest=$ROOT" .
}
main "$@"