1 Commits
Author SHA1 Message Date
clawbot 5bf7802404 Admin endpoints for bans and clients on the single listener (closes #27)
check / check (push) Successful in 4m1s
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-06 22:57:12 +00:00
71 changed files with 1168 additions and 16266 deletions
-4
View File
@@ -61,10 +61,6 @@ linters:
desc: >- desc: >-
Test-support code belongs in test files and in packages whose Test-support code belongs in test files and in packages whose
directory name ends in test, not in the shipped binary. directory name ends in test, not in the shipped binary.
- pkg: sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest
desc: >-
Test-support code belongs in test files and in packages whose
directory name ends in test, not in the shipped binary.
# Only decisions already recorded in the Go package defaults are # Only decisions already recorded in the Go package defaults are
# listed here. Every entry matches the module path exactly. # listed here. Every entry matches the module path exactly.
gomodguard_v2: gomodguard_v2:
+1 -25
View File
@@ -29,12 +29,6 @@ RUN go mod download
COPY . . COPY . .
# go.mod and go.sum must be as `go mod tidy` writes them, which is what
# `make tidy` does. Checked before the tests, which a missing go.sum line
# fails with a message that does not name `make tidy`.
RUN go mod tidy -diff || \
{ echo "go.mod or go.sum is not tidy: run make tidy" >&2; exit 1; }
# Go's build cache is kept on a tmpfs, out of the image: nothing uses it # Go's build cache is kept on a tmpfs, out of the image: nothing uses it
# after this step, and writing it into the image takes seconds. # after this step, and writing it into the image takes seconds.
RUN --mount=type=tmpfs,target=/root/.cache/go-build \ RUN --mount=type=tmpfs,target=/root/.cache/go-build \
@@ -42,25 +36,7 @@ RUN --mount=type=tmpfs,target=/root/.cache/go-build \
{ echo "--- Rerunning with -v for details ---"; \ { echo "--- Rerunning with -v for details ---"; \
go test -timeout 90s -race -v ./...; exit 1; } go test -timeout 90s -race -v ./...; exit 1; }
# Tidy stage: `go mod tidy` in the test phase's Go, so that the files it # Build stage. Nothing is wanted from the two phases above; the copies
# writes pass the test phase's check. Nothing else depends on it, so only
# script/tidy, which names the stage after it, builds it.
#
# golang 1.27.1-trixie, 2026-09-19
FROM golang@sha256:3b77fc618ec235a1ab412de7737f120dd507c57e8d87de4cbb7994fb94275ed5 AS tidy
WORKDIR /src
COPY . .
RUN go mod tidy
# go.mod and go.sum alone, which script/tidy writes into the working tree.
FROM scratch AS tidy-files
COPY --from=tidy /src/go.mod /src/go.sum /
# Build stage. Nothing is wanted from the lint and test phases; the copies
# are what make BuildKit build them first, so the image, which needs this # are what make BuildKit build them first, so the image, which needs this
# stage, cannot be produced unless lint and test passed. # stage, cannot be produced unless lint and test passed.
# #
+3 -7
View File
@@ -1,10 +1,9 @@
.PHONY: bootstrap setup test lint fmt fmt-check tidy check docker hooks build \ .PHONY: bootstrap setup test lint fmt fmt-check check docker hooks build run \
run example-app example-app
# Makefile targets are thin shims; the implementations live in script/ # Makefile targets are thin shims; the implementations live in script/
# per the scripts-to-rule-them-all pattern (see the Entrypoints section # per the scripts-to-rule-them-all pattern (see the Entrypoints section
# of README.md). tidy writes go.mod and go.sum as `go mod tidy` does, # of README.md). build and run are for working on the code by hand;
# which test checks. build and run are for working on the code by hand;
# example-app checks the image with an app built on it. # example-app checks the image with an app built on it.
bootstrap: bootstrap:
@@ -25,9 +24,6 @@ fmt:
fmt-check: fmt-check:
@script/fmt-check @script/fmt-check
tidy:
@script/tidy
check: check:
@script/check @script/check
+247 -982
View File
File diff suppressed because it is too large Load Diff
+1 -4
View File
@@ -5,8 +5,6 @@ go 1.26.0
require ( require (
github.com/fsnotify/fsnotify v1.10.1 github.com/fsnotify/fsnotify v1.10.1
github.com/hashicorp/golang-lru/v2 v2.0.7 github.com/hashicorp/golang-lru/v2 v2.0.7
github.com/maxmind/mmdbwriter v1.2.0
github.com/oschwald/maxminddb-golang/v2 v2.7.0
github.com/prometheus/client_golang v1.24.1 github.com/prometheus/client_golang v1.24.1
) )
@@ -18,7 +16,6 @@ require (
github.com/prometheus/client_model v0.6.2 // indirect github.com/prometheus/client_model v0.6.2 // indirect
github.com/prometheus/common v0.70.1 // indirect github.com/prometheus/common v0.70.1 // indirect
github.com/prometheus/procfs v0.21.1 // indirect github.com/prometheus/procfs v0.21.1 // indirect
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba // indirect golang.org/x/sys v0.47.0 // indirect
golang.org/x/sys v0.48.0 // indirect
google.golang.org/protobuf v1.36.11 // indirect google.golang.org/protobuf v1.36.11 // indirect
) )
+10 -12
View File
@@ -2,6 +2,8 @@ github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM=
github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho= github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho=
github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo= github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
@@ -12,12 +14,10 @@ github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJn
github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ=
github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc= github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc=
github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw= github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw=
github.com/maxmind/mmdbwriter v1.2.0 h1:hyvDopImmgvle3aR8AaddxXnT0iQH2KWJX3vNfkwzYM=
github.com/maxmind/mmdbwriter v1.2.0/go.mod h1:EQmKHhk2y9DRVvyNxwCLKC5FrkXZLx4snc5OlLY5XLE=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA=
github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ=
github.com/oschwald/maxminddb-golang/v2 v2.7.0 h1:ZcAr3GYc2LYC8aec2mCMX9+QOF0EolH3jDFKRV/Z1+U= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/oschwald/maxminddb-golang/v2 v2.7.0/go.mod h1:DuKJLbbug6TXC0yJXgs1MWifvXHmudRWzMobMIUu04g= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU= github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU=
github.com/prometheus/client_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE= github.com/prometheus/client_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE=
github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk= github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk=
@@ -26,17 +26,15 @@ github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi
github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc= github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc=
github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI= github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI=
github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY= github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY=
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ= go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ=
go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ= go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ=
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBseWJUpBw5I82+2U4M=
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
-929
View File
@@ -1,929 +0,0 @@
// Package alerts sends alerts on bans, on traffic over an anomaly
// threshold, on a source that fails and on a file with an error to each
// destination set: to the webhook
// SWWAF_ALERT_WEBHOOK_URL names, each as one JSON object, as the "Alert
// webhook schema" section of SPEC.md describes, to the Slack incoming
// webhook SWWAF_ALERT_SLACK_WEBHOOK_URL names, as a message, and to the
// 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"
// EventAnomaly is a count of requests or bytes over an anomaly
// threshold.
EventAnomaly = "anomaly"
// EventWAFBlock comes with the Core Rule Set; nothing raises it yet.
EventWAFBlock = "waf_block"
// EventReputationHit is a request whose client a blocklist or a DNSBL
// zone lists.
EventReputationHit = "reputation_hit"
// EventSourceFailure is GeoJS failing or refusing smallwebwaf, a fetch
// of a list failing, or a query to a DNSBL zone failing or refused.
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 sent as an hour ends: of the alerts held
// back in it past SWWAF_ALERT_MAX_PER_HOUR, and of the repeats held
// back by the cooldowns dropped as it ends, which no alert let through
// has given. It is sent with SWWAF_ALERT_MAX_PER_HOUR off too, for
// those repeats. SWWAF_ALERT_EVENTS does not name it.
EventSummary = "summary"
)
// 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", for a source_failure, its
// "source", and for an anomaly, its "scope", with the "asn" or the
// "name" of some scopes, which the cooldown tells repeats by.
Reason string `json:"reason"`
Detail map[string]any `json:"detail"`
// SuppressedRepeats is how many repeats of the alert the cooldown
// held back since the last one let through. For a summary, it is how
// many the cooldowns dropped as the hour ended had held back that no
// alert let through gave.
SuppressedRepeats int `json:"suppressed_repeats"`
}
// Cooldown is, for an event on a netblock, about a file or a source, or
// for an anomaly in a scope, when the last alert let through was raised,
// and how many repeats the cooldown has held back since, as alerts.json
// holds it.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Cooldown struct {
Event string `json:"event"`
Netblock netip.Prefix `json:"netblock"`
File string `json:"file,omitempty"`
Source string `json:"source,omitempty"`
Scope string `json:"scope,omitempty"`
ASN string `json:"asn,omitempty"`
Name string `json:"name,omitempty"`
Sent time.Time `json:"sent"`
SuppressedRepeats int `json:"suppressed_repeats"`
}
// 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, source or scope.
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, or in the same
// scope with the same AS number or name, as its detail names them. Each
// is empty for an alert without one.
type cooldownKey struct {
event string
netblock netip.Prefix
file string
source string
scope string
asn string
name string
}
// cooldownKeyOf returns what makes another alert a repeat of alert.
func cooldownKeyOf(alert *Alert) cooldownKey {
file, _ := alert.Detail["file"].(string)
source, _ := alert.Detail["source"].(string)
scope, _ := alert.Detail["scope"].(string)
asn, _ := alert.Detail["asn"].(string)
name, _ := alert.Detail["name"].(string)
return cooldownKey{alert.Event, alert.Netblock, file, source, scope, asn, name}
}
// 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. The next one let through gives that count, unless an hour of
// the clock ends first after the cooldown has run out: the cooldown is
// then dropped, and that hour's summary gives the count. Past MaxPerHour
// alerts let through in the hour under way, an alert is held back for
// that hour's summary instead, which is sent once the hour has ended; it
// starts no cooldown. Raise never waits: an alert let through joins the
// queue of each destination, from which Run sends it, and with queueSize
// alerts waiting for a destination, the oldest is dropped.
func (q *Queue) Raise(alert Alert) {
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, alert.Event) {
return
}
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, source, scope, AS
// number and name.
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),
cmp.Compare(a.Scope, b.Scope), cmp.Compare(a.ASN, b.ASN),
cmp.Compare(a.Name, b.Name))
})
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,
cooldown.Scope, cooldown.ASN, cooldown.Name,
}
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,
Scope: key.scope, ASN: key.asn, Name: key.name, Sent: now,
}
}
// endHour ends the hour under way, if now is past it. It drops the
// cooldowns that have run out, whatever repeats they held back, so that
// they do not pile up, and queues that hour's summary when alerts were
// held back in it past MaxPerHour, or when a cooldown dropped had held
// back repeats, which no alert let through has given: the summary gives
// them.
func (q *Queue) endHour(now time.Time) {
start := now.Truncate(time.Hour)
if !start.After(q.hour.Start) {
return
}
repeats := 0
for key, cooldown := range q.cooldowns {
if now.Sub(cooldown.Sent) >= q.params.Cooldown {
repeats += cooldown.SuppressedRepeats
delete(q.cooldowns, key)
}
}
heldBack := 0
for _, count := range q.hour.HeldBack {
heldBack += count
}
var reasons []string
if heldBack > 0 {
reasons = append(reasons, fmt.Sprintf("%d alerts held back in the hour from %s, "+
"past the %d an hour SWWAF_ALERT_MAX_PER_HOUR allows", heldBack,
q.hour.Start.Format(time.RFC3339), q.params.MaxPerHour))
}
if repeats > 0 {
reasons = append(reasons, fmt.Sprintf("%d repeats held back by "+
"SWWAF_ALERT_COOLDOWN that no later alert gives", repeats))
}
if len(reasons) > 0 {
q.queue(&Alert{
Instance: q.params.Instance,
Time: now,
Event: EventSummary,
Reason: strings.Join(reasons, "; "),
Detail: map[string]any{
"hour": q.hour.Start, "count": heldBack, "events": q.hour.HeldBack,
},
SuppressedRepeats: repeats,
})
}
q.hour = Hour{Start: start, HeldBack: map[string]int{}}
}
// 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
@@ -1,16 +0,0 @@
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
}
}
}
-406
View File
@@ -1,406 +0,0 @@
// Package anomaly counts requests and bytes over a minute and an hour, per
// client, per surrounding netblock, per AS number, for the whole service
// and per named netblock, and raises an anomaly alert for a count over its
// threshold, as "Anomaly thresholds" under "Configuration surface" in
// SPEC.md describes. It refuses and bans nothing. At most 20,000 counters
// are kept, in memory, and written to alerts.json and read from it by the
// state package.
package anomaly
import (
"cmp"
"fmt"
"net/netip"
"slices"
"sync"
"time"
"github.com/hashicorp/golang-lru/v2/simplelru"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
// maxCounters is how many counters are kept. Past it, the counter counted
// least recently is dropped, and starts afresh if it is counted again.
const maxCounters = 20000
// The scopes, what a counter counts, as the settings, alerts.json and the
// alerts name them.
const (
// ScopeClient is one client: an IPv4 address, or an IPv6 /64.
ScopeClient = "client"
// ScopeNet is the netblock around a client, SWWAF_ANOMALY_NET_V4_PREFIX
// or SWWAF_ANOMALY_NET_V6_PREFIX long.
ScopeNet = "net"
// ScopeASN is an AS number.
ScopeASN = "asn"
// ScopeTotal is the whole service.
ScopeTotal = "total"
// ScopeWatch is a named netblock of SWWAF_WATCH_NETS.
ScopeWatch = "watch"
)
// Scopes returns every scope.
func Scopes() []string {
return []string{ScopeClient, ScopeNet, ScopeASN, ScopeTotal, ScopeWatch}
}
// The windows a counter counts in, as the alerts name them.
const (
minute = "minute"
hour = "hour"
)
// Thresholds are the most requests and the most bytes a scope may have
// counted in a minute and in an hour before an alert is raised. Zero is
// off.
type Thresholds struct {
RequestsPerMinute int64
RequestsPerHour int64
BytesPerMinute int64
BytesPerHour int64
}
// NamedNetblock is a netblock SWWAF_WATCH_NETS names.
type NamedNetblock struct {
Name string
Netblock netip.Prefix
}
// Params are what New needs.
type Params struct {
// The thresholds of each scope: SWWAF_ANOMALY_CLIENT_*,
// SWWAF_ANOMALY_NET_*, SWWAF_ANOMALY_ASN_*, SWWAF_ANOMALY_TOTAL_* and
// SWWAF_WATCH_*.
Client, Net, ASN, Total, Watch Thresholds
// NetV4Prefix and NetV6Prefix are the lengths of the netblock around a
// client (SWWAF_ANOMALY_NET_V4_PREFIX and SWWAF_ANOMALY_NET_V6_PREFIX).
NetV4Prefix, NetV6Prefix int
// NamedNetblocks are SWWAF_WATCH_NETS.
NamedNetblocks []NamedNetblock
// Alerts receive the anomaly alerts.
Alerts *alerts.Queue
}
// Counter is one scope's counts, as alerts.json holds them: the scope,
// with the netblock, the AS number or the name that tells it from the
// others in that scope, and its two buckets of requests and of bytes in
// the minute and in the hour. A bucket whose threshold is off counts
// nothing, and is left out.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Counter struct {
Scope string `json:"scope"`
Netblock netip.Prefix `json:"netblock,omitzero"`
ASN string `json:"asn,omitempty"`
Name string `json:"name,omitempty"`
Minute ratelimit.Buckets `json:"minute,omitzero"`
Hour ratelimit.Buckets `json:"hour,omitzero"`
MinuteBytes ratelimit.Buckets `json:"minute_bytes,omitzero"`
HourBytes ratelimit.Buckets `json:"hour_bytes,omitzero"`
}
// Request is a request that has ended, as the counters count it.
type Request struct {
// Client is the client's address, and ClientGroup the client it is
// counted as: its IPv4 address, or its IPv6 /64.
Client netip.Addr
ClientGroup netip.Prefix
// ASN, ASName and Country are the client's as looked up, each "" when
// unknown.
ASN, ASName, Country string
// Bytes are the request's bytes, as SWWAF_BYTES_COUNT counts them.
Bytes int64
}
// Counters counts each request in the scopes it is in. It is safe for
// concurrent use.
type Counters struct {
params Params
mu sync.Mutex
counters *simplelru.LRU[key, *Counter]
}
// key is what tells a counter from the others: its scope, with its
// netblock, AS number or name.
type key struct {
scope string
netblock netip.Prefix
asn string
name string
}
// New returns Counters for params, with nothing counted yet.
func New(params Params) *Counters {
counters, err := simplelru.NewLRU[key, *Counter](maxCounters, nil)
if err != nil {
panic(err) // NewLRU fails only for a size below one
}
return &Counters{params: params, counters: counters}
}
// Count counts r, a request that has ended, at now, in each scope it is
// in whose thresholds are not all off: its client, the netblock around
// it, its AS number once known, the whole service, and each named
// netblock it is in. Only the counts whose threshold is set are counted.
// For each scope whose count is over a threshold, it raises an anomaly
// alert, for the first such count in the order requests and bytes in the
// minute, then in the hour; the alert queue's cooldown holds back the
// repeats. Nothing is refused or banned.
func (c *Counters) Count(now time.Time, r Request) {
var raised []alerts.Alert
c.mu.Lock()
for _, scope := range c.scopesOf(r) {
counter, found := c.counters.Get(scope.key)
if !found {
counter = scope.key.counter()
c.counters.Add(scope.key, counter)
}
over, passed := counter.add(now, r.Bytes, scope.thresholds)
if passed {
raised = append(raised, alertFor(r, scope.key, over))
}
}
c.mu.Unlock()
for _, alert := range raised {
c.params.Alerts.Raise(alert)
}
}
// Snapshot returns every counter, sorted by scope, then by netblock, AS
// number and name, as alerts.json lists them.
func (c *Counters) Snapshot() []Counter {
c.mu.Lock()
counters := make([]Counter, 0, c.counters.Len())
for _, counter := range c.counters.Values() {
counters = append(counters, *counter)
}
c.mu.Unlock()
slices.SortFunc(counters, func(a, b Counter) int {
return cmp.Or(cmp.Compare(a.Scope, b.Scope), a.Netblock.Compare(b.Netblock),
cmp.Compare(a.ASN, b.ASN), cmp.Compare(a.Name, b.Name))
})
return counters
}
// Load puts counters, read from alerts.json, in place of those held, in
// the order they were last counted, as the starts of their buckets tell,
// so that the one counted least recently is dropped first. Each netblock
// is masked to its length, so that 203.0.113.9/24 is 203.0.113.0/24.
// Buckets whose time has passed at now are emptied, and a counter left
// with every bucket empty is dropped.
func (c *Counters) Load(counters []Counter, now time.Time) {
counters = slices.Clone(counters)
slices.SortStableFunc(counters, func(a, b Counter) int {
return a.lastStart().Compare(b.lastStart())
})
c.mu.Lock()
defer c.mu.Unlock()
c.counters.Purge()
for _, counter := range counters {
counter.Netblock = counter.Netblock.Masked()
empty := true
for _, count := range counter.counts() {
if count.buckets.Passed(now, count.length) {
*count.buckets = ratelimit.Buckets{}
}
empty = empty && *count.buckets == ratelimit.Buckets{}
}
if !empty {
c.counters.Add(counter.key(), &counter)
}
}
}
// scope is a scope a request is counted in, and its thresholds.
type scope struct {
key key
thresholds Thresholds
}
// scopesOf returns the scopes r is in whose thresholds are not all off.
func (c *Counters) scopesOf(r Request) []scope {
p := c.params
client := r.Client.Unmap()
all := []scope{
{key{scope: ScopeClient, netblock: r.ClientGroup}, p.Client},
{key{scope: ScopeNet, netblock: c.netAround(client)}, p.Net},
{key{scope: ScopeTotal}, p.Total},
}
if r.ASN != "" {
all = append(all, scope{key{scope: ScopeASN, asn: r.ASN}, p.ASN})
}
for _, named := range p.NamedNetblocks {
if named.Netblock.Contains(client) {
all = append(all, scope{
key{scope: ScopeWatch, netblock: named.Netblock, name: named.Name}, p.Watch,
})
}
}
return slices.DeleteFunc(all, func(s scope) bool {
return s.thresholds == Thresholds{}
})
}
// netAround returns the netblock around client that ScopeNet counts it
// in: NetV4Prefix or NetV6Prefix long.
func (c *Counters) netAround(client netip.Addr) netip.Prefix {
length := c.params.NetV6Prefix
if client.Is4() {
length = c.params.NetV4Prefix
}
return netip.PrefixFrom(client, length).Masked()
}
// overThreshold is a count over its threshold: what it counts, requests or
// bytes, its window, the count and the threshold.
type overThreshold struct {
kind, window string
count float64
threshold int64
}
// add counts a request of bytes at now in each of c's counts whose
// threshold, in thresholds, is set, and returns the first count over its
// threshold, and whether there is one.
func (c *Counter) add(
now time.Time, bytes int64, thresholds Thresholds,
) (overThreshold, bool) {
// In the order of counts.
inOrder := [4]int64{
thresholds.RequestsPerMinute, thresholds.BytesPerMinute,
thresholds.RequestsPerHour, thresholds.BytesPerHour,
}
var (
first overThreshold
passed bool
)
for i, count := range c.counts() {
threshold := inOrder[i]
if threshold == 0 {
continue
}
n := int64(1)
if count.kind == ratelimit.KindBytes {
n = bytes
}
counted := count.buckets.Add(now, count.length, n)
if !passed && counted > float64(threshold) {
first = overThreshold{count.kind, count.window, counted, threshold}
passed = true
}
}
return first, passed
}
// bucketCount is one of a counter's four counts: requests or bytes, in a
// window of length, and the buckets they are counted in.
type bucketCount struct {
kind, window string
length time.Duration
buckets *ratelimit.Buckets
}
// counts returns c's counts: requests and bytes in the minute, then in
// the hour.
func (c *Counter) counts() [4]bucketCount {
return [4]bucketCount{
{ratelimit.KindRequests, minute, time.Minute, &c.Minute},
{ratelimit.KindBytes, minute, time.Minute, &c.MinuteBytes},
{ratelimit.KindRequests, hour, time.Hour, &c.Hour},
{ratelimit.KindBytes, hour, time.Hour, &c.HourBytes},
}
}
// lastStart returns the start of c's latest bucket, which tells, to the
// minute or to the hour, when c was last counted.
func (c *Counter) lastStart() time.Time {
var latest time.Time
for _, count := range c.counts() {
if count.buckets.Start.After(latest) {
latest = count.buckets.Start
}
}
return latest
}
// key returns what tells c from the other counters.
func (c *Counter) key() key {
return key{scope: c.Scope, netblock: c.Netblock, asn: c.ASN, name: c.Name}
}
// counter returns a counter for k, with nothing counted yet.
func (k key) counter() *Counter {
return &Counter{Scope: k.scope, Netblock: k.netblock, ASN: k.asn, Name: k.name}
}
// alertFor returns the anomaly alert for o, a count over its threshold in
// the scope k, which r took over it. It gives r's client, with its AS
// number, AS name and country, and the netblock counted, of a client, the
// netblock around it or a named netblock. Its detail gives the scope, the
// AS number or the name of a scope that has one, the window, what is
// counted, the count and the threshold.
func alertFor(r Request, k key, o overThreshold) alerts.Alert {
detail := map[string]any{
"scope": k.scope, "window": o.window, "kind": o.kind, "count": o.count,
"threshold": o.threshold,
}
var counted string
switch k.scope {
case ScopeClient:
counted = "the client " + k.netblock.String()
case ScopeNet:
counted = "the netblock " + k.netblock.String()
case ScopeASN:
counted = k.asn
detail["asn"] = k.asn
case ScopeTotal:
counted = "the whole service"
default: // watch
counted = "the named netblock " + k.name + ", " + k.netblock.String()
detail["name"] = k.name
}
return alerts.Alert{
Event: alerts.EventAnomaly,
Client: r.Client,
Netblock: k.netblock,
ASN: r.ASN,
ASName: r.ASName,
Country: r.Country,
Reason: fmt.Sprintf("%s per %s of %s over the threshold of %d", o.kind, o.window,
counted, o.threshold),
Detail: detail,
}
}
-238
View File
@@ -1,238 +0,0 @@
package anomaly_test
import (
"encoding/json"
"fmt"
"net/netip"
"net/url"
"reflect"
"slices"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
)
// maxCounters is how many counters are kept.
const maxCounters = 20000
func TestEachScopeHasACooldownOfItsOwn(t *testing.T) {
t.Parallel()
queue := newQueue()
office := netip.MustParsePrefix("203.0.113.0/24")
overAtTheSecond := anomaly.Thresholds{RequestsPerMinute: 1}
counters := anomaly.New(anomaly.Params{
Client: overAtTheSecond, Net: overAtTheSecond, ASN: overAtTheSecond,
Total: overAtTheSecond, Watch: overAtTheSecond,
// The netblock around a client is the client's own, and two names
// name one netblock.
NetV4Prefix: 32,
NamedNetblocks: []anomaly.NamedNetblock{
{Name: "office", Netblock: office}, {Name: "hq", Netblock: office},
},
Alerts: queue,
})
// The first client's second request is over the threshold in the six
// scopes it is in. The other client's two are both over it in the whole
// service and in each named netblock, three repeats each, and its
// second is over it in the scopes of its own, its client, its netblock
// and its AS number, which are no repeats.
for _, r := range []anomaly.Request{
{Client: netip.MustParseAddr("203.0.113.9"), ASN: "AS64496"},
{Client: netip.MustParseAddr("203.0.113.10"), ASN: "AS64511"},
} {
r.ClientGroup = netip.PrefixFrom(r.Client, 32)
for range 2 {
counters.Count(midnight(), r)
}
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 9 || queue.Suppressed() != 6 {
t.Fatalf("%d alerts wait and %d are held back, want 9 and 6: %+v",
len(waiting), queue.Suppressed(), waiting)
}
// alerts.json keeps each scope's cooldown: each alert raised again
// after a restart is a repeat.
data, err := json.Marshal(queue.Snapshot())
if err != nil {
t.Fatalf("encode: %v", err)
}
var read alerts.State
err = json.Unmarshal(data, &read)
if err != nil {
t.Fatalf("decode: %v", err)
}
after := newQueue()
after.Load(read)
for _, alert := range read.Waiting[alerts.DestinationWebhook] {
after.Raise(alert)
}
if after.Suppressed() != 9 {
t.Errorf("after loading, %d alerts are held back, want 9", after.Suppressed())
}
}
func TestKeepsAtMost20000CountersDroppingTheLeastRecentlyCounted(t *testing.T) {
t.Parallel()
counters := newCounters(anomaly.Params{
Client: anomaly.Thresholds{RequestsPerMinute: 1000},
})
for i := range maxCounters {
counters.Count(midnight(), request(i))
}
// Counted again, the first client is the most recently counted, and
// the second is dropped for a new one.
counters.Count(midnight(), request(0))
counters.Count(midnight(), request(maxCounters))
got := counters.Snapshot()
if len(got) != maxCounters || !holds(got, 0) || holds(got, 1) ||
!holds(got, maxCounters) {
t.Errorf("%d counters, holding the first client %v, the second %v and the "+
"new one %v, want %d, the first and the new one", len(got), holds(got, 0),
holds(got, 1), holds(got, maxCounters), maxCounters)
}
}
func TestLoadEmptiesBucketsWhoseTimeHasPassedAndDropsEmptyCounters(t *testing.T) {
t.Parallel()
counters := newCounters(anomaly.Params{
Net: anomaly.Thresholds{RequestsPerMinute: 1000, RequestsPerHour: 1000},
Total: anomaly.Thresholds{RequestsPerMinute: 1000},
NetV4Prefix: 24,
})
halfAnHourOn := midnight().Add(30 * time.Minute)
// Half an hour on, the hour's buckets count still, and the minute's
// do not.
counters.Load([]anomaly.Counter{
{
Scope: anomaly.ScopeNet,
Netblock: netip.MustParsePrefix("203.0.113.9/24"),
Minute: ratelimit.Buckets{Start: midnight(), Current: 5},
Hour: ratelimit.Buckets{Start: midnight(), Current: 7},
},
{
Scope: anomaly.ScopeTotal,
Minute: ratelimit.Buckets{Start: midnight(), Current: 1},
},
}, halfAnHourOn)
// The whole service's counter, left empty, is dropped, and the
// netblock read is masked to its length.
netblock := anomaly.Counter{
Scope: anomaly.ScopeNet,
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
Hour: ratelimit.Buckets{Start: midnight(), Current: 7},
}
if got, want := counters.Snapshot(), []anomaly.Counter{netblock}; !reflect.DeepEqual(
got, want) {
t.Errorf("counters read\n%+v\nwant\n%+v", got, want)
}
// A request from the netblock is counted with the requests read.
counters.Count(halfAnHourOn, anomaly.Request{
Client: netip.MustParseAddr("203.0.113.9"),
ClientGroup: netip.MustParsePrefix("203.0.113.9/32"),
})
netblock.Minute = ratelimit.Buckets{Start: halfAnHourOn, Current: 1}
netblock.Hour.Current = 8
want := []anomaly.Counter{netblock, {
Scope: anomaly.ScopeTotal,
Minute: ratelimit.Buckets{Start: halfAnHourOn, Current: 1},
}}
if got := counters.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("counters after a request\n%+v\nwant\n%+v", got, want)
}
}
func TestLoadDropsTheLeastRecentlyCountedFirst(t *testing.T) {
t.Parallel()
counters := newCounters(anomaly.Params{
Client: anomaly.Thresholds{RequestsPerMinute: 1000},
})
now := midnight().Add(time.Minute)
// The second half of the file was counted in the minute before the
// first half.
read := make([]anomaly.Counter, 0, maxCounters)
for i := range maxCounters {
start := now
if i >= maxCounters/2 {
start = midnight()
}
read = append(read, anomaly.Counter{
Scope: anomaly.ScopeClient, Netblock: request(i).ClientGroup,
Minute: ratelimit.Buckets{Start: start, Current: 1},
})
}
counters.Load(read, now)
counters.Count(now, request(maxCounters))
got := counters.Snapshot()
if !holds(got, 0) || holds(got, maxCounters/2) {
t.Errorf("holding the first client of the file %v, and the first counted in "+
"the minute before %v, want only the first", holds(got, 0),
holds(got, maxCounters/2))
}
}
// midnight is the time of the tests' requests.
func midnight() time.Time {
return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
}
// newCounters returns Counters for params, whose alerts go nowhere.
func newCounters(params anomaly.Params) *anomaly.Counters {
params.Alerts = alerts.New(alerts.Params{})
return anomaly.New(params)
}
// newQueue returns a queue of alerts to a webhook, with the default
// cooldown, which keeps them waiting, since it is never run.
func newQueue() *alerts.Queue {
return alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
MaxPerHour: 60,
Now: midnight,
})
}
// request returns a request from client number i, an address in
// 10.0.0.0/8.
func request(i int) anomaly.Request {
client := netip.MustParseAddr(fmt.Sprintf("10.%d.%d.%d", i>>16, i>>8&255, i&255))
return anomaly.Request{Client: client, ClientGroup: netip.PrefixFrom(client, 32)}
}
// holds reports whether counters hold the counter of client number i.
func holds(counters []anomaly.Counter, i int) bool {
return slices.ContainsFunc(counters, func(counter anomaly.Counter) bool {
return counter.Netblock == request(i).ClientGroup
})
}
+11 -14
View File
@@ -65,16 +65,13 @@ func TestReasonOfTheBansSmallwebwafMakes(t *testing.T) {
ledger := bans.New(defaultRules()) ledger := bans.New(defaultRules())
limit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(), limit := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
bans.Notes{Kind: "requests", Limit: 1000, Window: "minute"}) bans.Notes{Limit: 1000, Window: "minute"})
byteLimit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.3/32"), attack := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
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 {
@@ -104,12 +101,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",
@@ -137,7 +134,7 @@ func TestLiftedBanForAnAttackRefusesNothingAndMakesNoBanLonger(t *testing.T) {
now := midnight().Add(2 * time.Hour) now := midnight().Add(2 * time.Hour)
_, banned, _ := ledger.Find(netblock.Addr(), now) _, banned := ledger.Find(netblock.Addr(), now)
if banned { if banned {
t.Error("the lifted ban refuses") t.Error("the lifted ban refuses")
} }
@@ -148,7 +145,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)
} }
@@ -158,7 +155,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"),
@@ -223,7 +220,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)
@@ -264,11 +261,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")
} }
+42 -139
View File
@@ -1,8 +1,8 @@
// Package bans is the ban ledger: the bans smallwebwaf makes on the // Package bans is the ban ledger: the bans smallwebwaf makes on the
// netblocks of clients that break a rate limit or a byte limit or show a // netblocks of clients that break a rate limit or show a clear sign of
// clear sign of attack, and those an admin makes, with their notes, as // attack, and those an admin makes, with their notes, as the "Bans"
// the "Bans" section of SPEC.md describes. The bans are kept in memory, // section of SPEC.md describes. The bans are kept in memory, and written
// and written to bans.json and read from it by the state package. // to bans.json and read from it by the state package.
package bans package bans
import ( import (
@@ -91,35 +91,22 @@ func (b Ban) ActiveAt(now time.Time) bool {
// //
//nolint:tagliatelle // the state files use snake_case, as the request log does //nolint:tagliatelle // the state files use snake_case, as the request log does
type Notes struct { type Notes struct {
// ASN, ASName and Country are the client's AS number, AS name and // Country is the client's country, when it was looked up.
// country, when they were looked up: when the request that caused the
// ban was made, or when GeoJS answered about the client afterwards.
ASN string `json:"asn"`
ASName string `json:"as_name"`
Country string `json:"country"` Country string `json:"country"`
// Kind, Limit, Window and Count are, for a ban for a broken limit, // Limit, Window and Count are, for a ban for a broken limit, the limit
// what the limit was on, "requests" for a rate limit or "bytes" for a // that was broken, its window, "minute", "hour" or "day", and the
// byte limit, the limit that was broken, its window, "minute", "hour" // count reached: the client's requests in the window, the one that
// or "day", and the count reached: the client's requests, or bytes, in // broke the limit included. These are the requests that counted
// the window, those of the request that broke the limit included. // toward the ban, and the window is the time over which they came.
// These are what counted toward the ban, and the window is the time
// over which they came.
Kind string `json:"kind,omitempty"`
Limit int64 `json:"limit,omitempty"` Limit int64 `json:"limit,omitempty"`
Window string `json:"window,omitempty"` Window string `json:"window,omitempty"`
Count float64 `json:"count,omitempty"` Count float64 `json:"count,omitempty"`
// LimitPercent and LimitPercentSetting are, for a ban for a limit a
// biased threshold lowered, the client's percentage of that kind of
// limit, of which Limit is the result, and the setting that gave it.
// Both are left out for a limit that was not lowered.
LimitPercent *int64 `json:"limit_percent,omitempty"`
LimitPercentSetting string `json:"limit_percent_setting,omitempty"`
// RuleID and Target are, for a ban for a clear sign of attack, the id // RuleID and Target are, for a ban for a clear sign of attack, the id
// of the rule file rule that matched, and its target. // of the rule file rule that matched, and its target.
RuleID string `json:"rule_id,omitempty"` RuleID string `json:"rule_id,omitempty"`
Target string `json:"target,omitempty"` Target string `json:"target,omitempty"`
// Request is the request that broke the limit, or whose bytes broke // Request is the request that broke the limit, or that was the clear
// it, or that was the clear sign of attack. // 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
@@ -207,43 +194,39 @@ 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.
// The last result reports whether the request made the ban permanent. func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool) {
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, false return Ban{}, false
} }
ban.Notes.Requests++ ban.Notes.Requests++
ban.Notes.Refused++ ban.Notes.Refused++
madePermanent := ban.Cause == CauseAttack && !ban.Permanent() if ban.Cause == CauseAttack && !ban.Permanent() {
if madePermanent {
ban.Expires = time.Time{} ban.Expires = time.Time{}
l.markChanged() l.markChanged()
} }
return *ban, true, madePermanent return *ban, true
} }
// Find is Check without counting the request among those the ban // Find is Check without counting the request among those the ban
// refused, and without making the ban permanent: in observe mode a ban // refused: in observe mode a ban refuses nothing.
// refuses nothing. The last result reports whether Check would have made func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool) {
// the ban permanent.
func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool, bool) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
ban := l.active(client, now) ban := l.active(client, now)
if ban == nil { if ban == nil {
return Ban{}, false, false return Ban{}, false
} }
return *ban, true, ban.Cause == CauseAttack && !ban.Permanent() return *ban, true
} }
// activeBan returns the ban in bans, a netblock's bans oldest first, that // activeBan returns the ban in bans, a netblock's bans oldest first, that
@@ -261,80 +244,28 @@ 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, and true. A first ban lasts LimitBanDuration. A ban // returns the ban. A first ban lasts LimitBanDuration. A ban made within
// made within LimitBanRepeatWindow after the netblock's ban that ended // LimitBanRepeatWindow after the netblock's ban that ended last, other
// last, other than one for a clear sign of attack or a lifted one, lasts // than one for a clear sign of attack or a lifted one, lasts repeatFactor
// repeatFactor times as long as that one. A ban that would be longer // times as long as that one. A ban that would be longer than
// than MaxBanDuration is permanent instead. If a ban on netblock is still // 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 with false, and no other is made. The ledger fills in the // returned and no other is made. The ledger fills in the notes' Refused
// notes' Refused and EarlierBans itself, and gives the ban the reason // and EarlierBans itself, and gives the ban the reason "requests per
// "<Kind> per <Window> over the limit of <Limit>", from the notes, such // <Window> over the limit of <Limit>", from the notes.
// as "requests per minute over the limit of 1000". func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes) Ban {
func (l *Ledger) BanForLimit( reason := fmt.Sprintf("requests per %s over the limit of %d",
netblock netip.Prefix, now time.Time, notes Notes, notes.Window, notes.Limit)
) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, true)
}
// WouldBanForLimit returns what BanForLimit would, without making the ban: return l.ban(netblock, now, CauseLimit, reason, notes)
// what observe mode would have done.
func (l *Ledger) WouldBanForLimit(
netblock netip.Prefix, now time.Time, notes Notes,
) (Ban, bool) {
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, false)
} }
// BanForAttack bans netblock at now for a clear sign of attack, with // BanForAttack bans netblock at now for a clear sign of attack, with
// notes, and returns the ban, and whether it made it, as BanForLimit // notes, and returns the ban, as BanForLimit does. A first ban lasts
// does. A first ban lasts AttackBanDuration; once the netblock has had // AttackBanDuration; once the netblock has had one that was not lifted,
// one that was not lifted, the next is permanent. Its reason is "matched // the next is permanent. Its reason is "matched the rule <RuleID>".
// the rule <RuleID>". func (l *Ledger) BanForAttack(netblock netip.Prefix, now time.Time, notes Notes) Ban {
func (l *Ledger) BanForAttack( return l.ban(netblock, now, CauseAttack, "matched the rule "+notes.RuleID, notes)
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
@@ -426,28 +357,6 @@ func (l *Ledger) Bans(netblock netip.Prefix) []Ban {
return slices.Clone(*bans) return slices.Clone(*bans)
} }
// AddLookup gives the notes of netblock's bans that have no AS number, AS
// name or country yet those of a client in it, as the lookup answered
// about it. It is not a request from netblock, and leaves when it was last
// seen unchanged. It does not have bans.json written at once: the notes
// are written with its next write, as the counts in them are.
func (l *Ledger) AddLookup(netblock netip.Prefix, asn, asName, country string) {
l.mu.Lock()
defer l.mu.Unlock()
bans, found := l.netblocks.Peek(netblock)
if !found {
return
}
for i := range *bans {
notes := &(*bans)[i].Notes
if notes.ASN == "" && notes.ASName == "" && notes.Country == "" {
notes.ASN, notes.ASName, notes.Country = asn, asName, country
}
}
}
// Made returns how many bans for cause have been made since the start: // Made returns how many bans for cause have been made since the start:
// for CauseLimit and CauseAttack, by the ledger; for CauseAdmin, by an // for CauseLimit and CauseAttack, by the ledger; for CauseAdmin, by an
// admin, with BanForAdmin or in an edit of bans.json, as LoadEdit counts // admin, with BanForAdmin or in an edit of bans.json, as LoadEdit counts
@@ -574,12 +483,10 @@ func (l *Ledger) holds(netblock netip.Prefix, start time.Time) bool {
} }
// ban bans netblock at now for cause, with reason and notes, as // ban bans netblock at now for cause, with reason and notes, as
// BanForLimit and BanForAttack describe, and returns the ban, and whether // BanForLimit and BanForAttack describe, and returns the ban.
// 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, keep bool, netblock netip.Prefix, now time.Time, cause, reason string, notes Notes,
) (Ban, bool) { ) Ban {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
@@ -590,7 +497,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, false return *active
} }
held = *bans held = *bans
@@ -606,15 +513,11 @@ 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, true return ban
} }
// earlierBans returns how many bans a netblock with the bans held, oldest // earlierBans returns how many bans a netblock with the bans held, oldest
+44 -179
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,22 +123,12 @@ 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, made := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) first := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
if !made { again := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
t.Error("the first ban was not made")
}
again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{}) if again != first || len(ledger.Bans(netblock)) != 1 {
t.Errorf("a limit broken during a ban gave %+v and %d bans, want %+v and 1",
if made || again != first || len(ledger.Bans(netblock)) != 1 { again, len(ledger.Bans(netblock)), first)
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)
} }
} }
@@ -147,21 +137,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")
} }
@@ -179,14 +169,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")
} }
@@ -208,7 +198,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{})
@@ -243,8 +233,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 ||
@@ -261,7 +251,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 {
@@ -272,31 +262,23 @@ func TestRequestDuringAnAttackBanMakesItPermanent(t *testing.T) {
wantChanged(t, ledger, true) wantChanged(t, ledger, true)
// In observe mode the ban refuses nothing, and stays as it is, while // In observe mode the ban refuses nothing, and stays as it is.
// Find tells that the request would have made it permanent. got, _ := ledger.Find(netblock.Addr(), midnight().Add(time.Hour))
got, _, wouldMakePermanent := ledger.Find(netblock.Addr(), midnight().Add(time.Hour)) if got.Permanent() {
if got.Permanent() || ledger.Bans(netblock)[0].Permanent() || !wouldMakePermanent { t.Fatal("a request found under the ban made it permanent")
t.Fatalf("a request found under the ban left it %+v, would have made it "+
"permanent %t, want it as it was, and true", got, wouldMakePermanent)
} }
wantChanged(t, ledger, false) // A request it refuses makes it permanent, and bans.json due.
got, _ = ledger.Check(netblock.Addr(), midnight().Add(time.Hour))
// A request it refuses makes it permanent, says so, and makes if !got.Permanent() || !ledger.Bans(netblock)[0].Permanent() {
// bans.json due. t.Fatalf("after a request during the ban, it is %+v, want it permanent", got)
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)
// The next request finds it permanent already. _, banned := ledger.Check(netblock.Addr(), midnight().Add(100*365*day))
_, banned, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(100*365*day)) if !banned {
if !banned || madePermanent { t.Error("the permanent ban ended")
t.Errorf("a later request is banned %t, and made the ban permanent %t, "+
"want banned by the permanent ban", banned, madePermanent)
} }
} }
@@ -307,8 +289,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",
@@ -317,14 +299,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, "+
@@ -332,84 +314,6 @@ func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
} }
} }
func TestWouldBanGivesTheBanWithoutMakingIt(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
wantChanged(t, ledger, true)
// While the first ban lasts, none would be made.
during, would := ledger.WouldBanForAttack(netblock, midnight(), bans.Notes{})
if would || 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()
@@ -418,16 +322,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, _, madePermanent := ledger.Check(netblock.Addr(), limit.Start) got, _ := ledger.Check(netblock.Addr(), limit.Start)
if got.Permanent() || madePermanent { if got.Permanent() {
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")
} }
} }
@@ -442,7 +346,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{
@@ -454,45 +358,6 @@ func TestRequestTextsAreCutTo256Bytes(t *testing.T) {
} }
} }
func TestLookupFillsTheNotesOfTheNetblocksBansWithoutOne(t *testing.T) {
t.Parallel()
ledger := bans.New(defaultRules())
netblock := netip.MustParsePrefix("203.0.113.9/32")
other := netip.MustParsePrefix("198.51.100.7/32")
// A ban made with the client's lookup, one made before it came, after
// the first ended, and one on another netblock.
ledger.BanForLimit(netblock, midnight(), bans.Notes{
ASN: "AS64497", ASName: "Other Net", Country: "FR",
})
ledger.BanForLimit(netblock, midnight().Add(time.Hour), bans.Notes{})
ledger.BanForLimit(other, midnight(), bans.Notes{})
ledger.AddLookup(netblock, "AS64496", "Example Net", "DE")
held := ledger.Bans(netblock)
if len(held) != 2 {
t.Fatalf("%s has %d bans, want 2", netblock, len(held))
}
for i, want := range []bans.Notes{
{ASN: "AS64497", ASName: "Other Net", Country: "FR"},
{ASN: "AS64496", ASName: "Example Net", Country: "DE"},
} {
got := held[i].Notes
if got.ASN != want.ASN || got.ASName != want.ASName || got.Country != want.Country {
t.Errorf("ban %d's notes give %q, %q and %q, want %q, %q and %q", i+1,
got.ASN, got.ASName, got.Country, want.ASN, want.ASName, want.Country)
}
}
if notes := ledger.Bans(other)[0].Notes; notes.ASN != "" || notes.Country != "" {
t.Errorf("the ban on %s has %q and %q, want neither",
other, notes.ASN, notes.Country)
}
}
// defaultRules are the rules at the settings' defaults. // defaultRules are the rules at the settings' defaults.
func defaultRules() bans.Rules { func defaultRules() bans.Rules {
return bans.Rules{ return bans.Rules{
+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}
+32 -896
View File
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
-233
View File
@@ -1,233 +0,0 @@
package lookup
import (
"context"
"fmt"
"log/slog"
"net/netip"
"os"
"path/filepath"
"sync"
"time"
"github.com/fsnotify/fsnotify"
"github.com/oschwald/maxminddb-golang/v2"
"sneak.berlin/go/smallwebwaf/internal/alerts"
)
// quietTime is how long the lookup database must go without a change
// before it is read again, so that a file still being copied in is read
// only once whole.
const quietTime = 2 * time.Second
// FileParams are what OpenFile needs.
type FileParams struct {
// Path is the lookup database, the IPinfo Lite file in its .mmdb form
// (SWWAF_LOOKUP_DB_PATH).
Path string
// Now tells the time, normally time.Now.
Now func() time.Time
// ProcessLog receives each reading of the file, and why a replacement
// of it cannot be read.
ProcessLog *slog.Logger
// Alerts receive a file_error alert for each replacement that cannot
// be read.
Alerts *alerts.Queue
}
// File looks up clients' AS numbers and countries in the lookup database,
// held in memory, and reads it again when it is replaced. It is safe for
// concurrent use.
type File struct {
params FileParams
mu sync.Mutex
// reader is the database in use, and lastRead when it was read.
// readFailures are the replacements that could not be read.
reader *maxminddb.Reader
lastRead time.Time
readFailures int
}
// record is what the lookup database holds about a network, of the fields
// smallwebwaf reads.
type record struct {
ASN string `maxminddb:"asn"`
ASName string `maxminddb:"as_name"`
CountryCode string `maxminddb:"country_code"`
}
// OpenFile reads the lookup database. A file that cannot be read, or that
// is not a .mmdb file, is an error.
func OpenFile(params FileParams) (*File, error) {
reader, err := read(params.Path)
if err != nil {
return nil, err
}
f := &File{params: params}
f.use(reader)
return f, nil
}
// LookUp returns what the lookup database says about client: its AS
// number, such as AS64496, the AS's name, and its country, such as DE,
// each "" when the database does not give it, as for an address missing
// from it. The database is asked about the client's first address, as
// GeoJS is.
func (f *File) LookUp(client netip.Prefix) Answer {
f.mu.Lock()
reader := f.reader
f.mu.Unlock()
var found record
// A record that cannot be decoded places the client nowhere, as a
// missing one does.
err := reader.Lookup(client.Addr()).Decode(&found)
if err != nil {
found = record{}
}
return Answer{
Client: client,
ASN: found.ASN,
ASName: found.ASName,
Country: found.CountryCode,
Answered: f.params.Now(),
}
}
// LastRead returns when the lookup database in use was read.
func (f *File) LastRead() time.Time {
f.mu.Lock()
defer f.mu.Unlock()
return f.lastRead
}
// ReadFailures returns how many replacements of the lookup database could
// not be read.
func (f *File) ReadFailures() int {
f.mu.Lock()
defer f.mu.Unlock()
return f.readFailures
}
// Watch watches the directory of the lookup database until ctx is done,
// and reads the file again once it has gone without a change for
// quietTime, after it is replaced, written or removed, and after Watch
// starts watching. If the directory cannot be watched, that is logged, and
// the database read at start stays in use.
func (f *File) Watch(ctx context.Context) {
watcher, err := fsnotify.NewWatcher()
if err == nil {
defer func() {
_ = watcher.Close()
}()
err = watcher.Add(filepath.Dir(f.params.Path))
}
if err != nil {
f.params.ProcessLog.Error("cannot watch the lookup database for replacements",
"error", err.Error())
return
}
f.params.ProcessLog.Info("watching the lookup database for replacements",
"file", f.params.Path)
f.readAfterChanges(ctx, watcher.Events, watcher.Errors)
}
// readAfterChanges reads the lookup database again once quietTime has
// passed without a change to it from events, until ctx is done, and logs
// the errors from errs. A change to another file in its directory does not
// count. The wait starts at once, as if for a change, so that a file
// replaced after OpenFile read it, and before its directory was watched,
// is read too.
func (f *File) readAfterChanges(
ctx context.Context, events <-chan fsnotify.Event, errs <-chan error,
) {
path := filepath.Clean(f.params.Path)
quiet := time.NewTimer(quietTime)
defer quiet.Stop()
for {
select {
case <-ctx.Done():
return
case event := <-events:
if filepath.Clean(event.Name) == path {
quiet.Reset(quietTime)
}
case <-quiet.C:
f.readAgain()
case err := <-errs:
f.params.ProcessLog.Warn("watching the lookup database failed",
"error", err.Error())
}
}
}
// readAgain reads the lookup database again, in place of the one in use,
// or, if it cannot be read, counts that, raises a file_error alert for it
// and logs it, and the one in use stays in use.
func (f *File) readAgain() {
reader, err := read(f.params.Path)
if err != nil {
const kept = "the lookup database cannot be read, " +
"and the one read before stays in use"
f.mu.Lock()
f.readFailures++
f.mu.Unlock()
// Raised before it is logged, so that the alert is there once the
// log line is.
f.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventFileError,
Reason: kept,
Detail: map[string]any{"file": f.params.Path, "error": err.Error()},
})
f.params.ProcessLog.Error(kept, "error", err.Error())
return
}
f.use(reader)
}
// use puts reader in use, in place of the database read before, and logs
// that the file was read.
func (f *File) use(reader *maxminddb.Reader) {
f.mu.Lock()
f.reader = reader
f.lastRead = f.params.Now()
f.mu.Unlock()
f.params.ProcessLog.Info("read the lookup database", "file", f.params.Path)
}
// read reads the lookup database at path. The whole file is read into
// memory, rather than mapped into it as the reader can, so that a file
// overwritten in place cannot change, or end, under a lookup.
func read(path string) (*maxminddb.Reader, error) {
data, err := os.ReadFile(path) //nolint:gosec // the file the admin names
if err != nil {
return nil, fmt.Errorf("SWWAF_LOOKUP_DB_PATH cannot be read: %w", err)
}
reader, err := maxminddb.OpenBytes(data)
if err != nil {
return nil, fmt.Errorf("SWWAF_LOOKUP_DB_PATH %s is not a .mmdb file: %w", path, err)
}
return reader, nil
}
-379
View File
@@ -1,379 +0,0 @@
package lookup
import (
"context"
"log/slog"
"net/netip"
"net/url"
"os"
"path/filepath"
"reflect"
"testing"
"testing/synctest"
"time"
"github.com/fsnotify/fsnotify"
"github.com/maxmind/mmdbwriter/mmdbtype"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
)
// testNetblock is the netblock the tests' lookup databases place, and
// testClient a client in it.
const (
testNetblock = "203.0.113.0/24"
testClient = "203.0.113.9/32"
)
func TestFilePlacesClientsAndCountsAnAddressMissingFromItAsUnknown(t *testing.T) {
t.Parallel()
germany := lookuptest.Network{ASN: "AS64496", ASName: "Example Net", Country: "DE"}
northKorea := lookuptest.Network{ASN: "AS64511", ASName: "Other Net", Country: "KP"}
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
lookuptest.Write(t, path, map[string]lookuptest.Network{
testNetblock: germany,
"2001:db8::/32": northKorea,
})
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
f, err := OpenFile(FileParams{
Path: path,
Now: func() time.Time { return now },
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: newQueue(),
})
if err != nil {
t.Fatalf("open %s: %v", path, err)
}
for client, want := range map[string]lookuptest.Network{
testClient: germany,
// An IPv6 client is its /64.
"2001:db8:1:2::/64": northKorea,
"198.51.100.7/32": {},
} {
prefix := netip.MustParsePrefix(client)
got := f.LookUp(prefix)
if got != (Answer{
Client: prefix, ASN: want.ASN, ASName: want.ASName, Country: want.Country,
Answered: now,
}) {
t.Errorf("%s has the answer %+v, want %+v, answered %s", client, got, want, now)
}
}
}
func TestRecordThatCannotBeReadPlacesTheClientNowhere(t *testing.T) {
t.Parallel()
// The AS number is a number, where a string belongs. The writer writes
// a record's fields in the order of their names, so as_name is read
// before the AS number fails.
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
lookuptest.WriteRecords(t, path, map[string]mmdbtype.Map{
testNetblock: {
"asn": mmdbtype.Uint32(64496),
"as_name": mmdbtype.String("Example Net"),
"country_code": mmdbtype.String("DE"),
},
})
f := openFile(t, path, newQueue())
answer := f.LookUp(netip.MustParsePrefix(testClient))
if answer.ASN != "" || answer.ASName != "" || answer.Country != "" {
t.Errorf("%s is placed %+v, want nowhere", testClient, answer)
}
}
func TestFileThatCannotBeReadIsAnError(t *testing.T) {
t.Parallel()
dir := t.TempDir()
missing := filepath.Join(dir, "missing.mmdb")
notDatabase := filepath.Join(dir, "not.mmdb")
writeFile(t, notDatabase, "not a lookup database\n")
for path, want := range map[string]string{
missing: "SWWAF_LOOKUP_DB_PATH cannot be read: open " + missing +
": no such file or directory",
notDatabase: "SWWAF_LOOKUP_DB_PATH " + notDatabase +
" is not a .mmdb file: error opening database: invalid MaxMind DB file",
} {
_, err := OpenFile(FileParams{
Path: path,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: newQueue(),
})
if err == nil || err.Error() != want {
t.Errorf("opening %s failed with %v, want %s", path, err, want)
}
}
}
// The tests below run readAfterChanges in a synctest bubble, where time is
// a clock of the test's own: time.Sleep moves it on at once, and
// synctest.Wait returns once readAfterChanges waits again, so that every
// reading due by then is done. The test sends the changes itself, as the
// watch of a directory cannot run in a bubble.
func TestReplacementCopiedOverTheFileInTwoPartsIsReadOnlyWhole(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "ipinfo_lite.mmdb")
writeDatabase(t, path, "DE")
queue := newQueue()
f := openFile(t, path, queue)
changes := watch(t, f)
other := filepath.Join(dir, "replacement.mmdb")
writeDatabase(t, other, "KP")
replacement, err := os.ReadFile(other) //nolint:gosec // a file the test wrote
if err != nil {
t.Fatalf("read %s: %v", other, err)
}
// The file in use is overwritten in place, and keeps giving what
// it gave. Its first part alone is not a .mmdb file.
file, err := os.Create(path) //nolint:gosec // a file the test wrote
if err != nil {
t.Fatalf("create %s: %v", path, err)
}
defer func() {
_ = file.Close()
}()
half := len(replacement) / 2
write(t, file, replacement[:half])
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
time.Sleep(quietTime - time.Nanosecond)
synctest.Wait()
wantCountry(t, f, "DE")
// The second part starts the wait again.
write(t, file, replacement[half:])
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
time.Sleep(quietTime - time.Nanosecond)
synctest.Wait()
wantCountry(t, f, "DE")
time.Sleep(time.Nanosecond)
synctest.Wait()
wantCountry(t, f, "KP")
if !f.LastRead().Equal(time.Now()) || f.ReadFailures() != 0 {
t.Errorf("read at %s, with %d failures; want read now, with none",
f.LastRead(), f.ReadFailures())
}
wantAlerts(t, queue)
})
}
func TestReplacementThatCannotBeReadLeavesTheFileInUseWithOneAlert(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
writeDatabase(t, path, "DE")
queue := newQueue()
f := openFile(t, path, queue)
read := f.LastRead()
changes := watch(t, f)
writeFile(t, path, "not a lookup database\n")
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
// Long after, the replacement has been read once.
time.Sleep(time.Hour)
synctest.Wait()
wantCountry(t, f, "DE")
if !f.LastRead().Equal(read) || f.ReadFailures() != 1 {
t.Errorf("read at %s, with %d failures; want read at %s, with one",
f.LastRead(), f.ReadFailures(), read)
}
wantAlerts(t, queue, alerts.Alert{
Time: read.Add(quietTime),
Event: alerts.EventFileError,
Reason: "the lookup database cannot be read, and the one read before stays in use",
Detail: map[string]any{
"file": path,
"error": "SWWAF_LOOKUP_DB_PATH " + path + " is not a .mmdb file: " +
"error opening database: invalid MaxMind DB file",
},
})
})
}
func TestChangeOfAnotherFileInTheDirectoryIsNoReplacement(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "ipinfo_lite.mmdb")
writeDatabase(t, path, "DE")
f := openFile(t, path, newQueue())
changes := watch(t, f)
// The wait that starts with the watch ends with a reading.
time.Sleep(quietTime)
synctest.Wait()
writeDatabase(t, path, "KP")
changes <- fsnotify.Event{Name: filepath.Join(dir, "other.mmdb"), Op: fsnotify.Create}
time.Sleep(quietTime)
synctest.Wait()
wantCountry(t, f, "DE")
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
time.Sleep(quietTime)
synctest.Wait()
wantCountry(t, f, "KP")
})
}
func TestReplacementSavedBeforeTheWatchStartsIsRead(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
writeDatabase(t, path, "DE")
f := openFile(t, path, newQueue())
// Saved after OpenFile read the file, and before its directory was
// watched, so that no change is seen for it.
writeDatabase(t, path, "KP")
watch(t, f)
time.Sleep(quietTime)
synctest.Wait()
wantCountry(t, f, "KP")
})
}
// newQueue returns a queue of alerts for a webhook that is never sent
// them, so that they wait in it for the test to look at.
func newQueue() *alerts.Queue {
return alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
Now: time.Now,
})
}
// writeDatabase writes a lookup database at path that places testNetblock
// in country, and no other address.
func writeDatabase(t *testing.T, path, country string) {
t.Helper()
lookuptest.Write(t, path, map[string]lookuptest.Network{
testNetblock: {ASN: "AS64496", ASName: "Example Net", Country: country},
})
}
// openFile opens the lookup database at path, which raises its alerts to
// queue.
func openFile(t *testing.T, path string, queue *alerts.Queue) *File {
t.Helper()
f, err := OpenFile(FileParams{
Path: path,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: queue,
})
if err != nil {
t.Fatalf("open %s: %v", path, err)
}
return f
}
// watch runs f's readAfterChanges until the test ends, and returns the
// channel that sends it changes.
func watch(t *testing.T, f *File) chan<- fsnotify.Event {
t.Helper()
changes := make(chan fsnotify.Event)
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
f.readAfterChanges(ctx, changes, nil)
close(stopped)
}()
t.Cleanup(func() {
stop()
<-stopped
})
return changes
}
// wantCountry checks the country f gives testClient.
func wantCountry(t *testing.T, f *File, want string) {
t.Helper()
got := f.LookUp(netip.MustParsePrefix(testClient)).Country
if got != want {
t.Errorf("%s is in %q, want %q", testClient, got, want)
}
}
// wantAlerts checks the alerts waiting in queue, and that it held none
// back.
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
t.Helper()
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != len(want) || (len(want) > 0 && !reflect.DeepEqual(waiting, want)) {
t.Errorf("alerts waiting %+v, want %+v", waiting, want)
}
if queue.Suppressed() != 0 {
t.Errorf("%d alerts held back, want none", queue.Suppressed())
}
}
// writeFile writes content to the file at path.
func writeFile(t *testing.T, path, content string) {
t.Helper()
err := os.WriteFile(path, []byte(content), 0o600)
if err != nil {
t.Fatalf("write %s: %v", path, err)
}
}
// write writes data to the end of file.
func write(t *testing.T, file *os.File, data []byte) {
t.Helper()
_, err := file.Write(data)
if err != nil {
t.Fatalf("write: %v", err)
}
}
+65 -140
View File
@@ -1,8 +1,7 @@
// Package lookup looks up each client's AS number and country, through // Package lookup looks up each client's country through the GeoJS web
// the GeoJS web service or in the lookup database, the IPinfo Lite file // service, and keeps the answers in memory, for at most 100,000 clients
// SWWAF_LOOKUP_DB_PATH names. GeoJS's answers are kept in memory, for at // and for 7 days each. The answers are written to lookups.json and read
// most 100,000 clients and for 7 days each, and are written to // from it by the state package.
// lookups.json and read from it by the state package.
package lookup package lookup
import ( import (
@@ -15,20 +14,17 @@ 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 endpoint for an address's place and network. Asked about // URL is GeoJS's country endpoint. Asked about several addresses at once,
// several addresses at once, comma separated in its ip parameter, it // comma separated in its ip parameter, it answers with a list.
// answers with a list. const URL = "https://get.geojs.io/v1/ip/country.json"
const URL = "https://get.geojs.io/v1/ip/geo.json"
const ( const (
// keepFor is how long an answer is used instead of asking GeoJS again. // keepFor is how long an answer is used instead of asking GeoJS again.
@@ -43,8 +39,9 @@ const (
maxWaiting = 10000 maxWaiting = 10000
// maxPerRequest is how many addresses one request to GeoJS asks about. // maxPerRequest is how many addresses one request to GeoJS asks about.
maxPerRequest = 200 maxPerRequest = 200
// unknownASN is the AS number GeoJS gives when it knows none. // timeout is how long a new client waits for its answer, and how long
unknownASN = 64512 // a request to GeoJS may take before it is abandoned.
timeout = time.Second
// After a failure GeoJS is not asked again for a second, and for // After a failure GeoJS is not asked again for a second, and for
// retryDelayFactor times as long after each further failure in a row, // retryDelayFactor times as long after each further failure in a row,
// up to five minutes. // up to five minutes.
@@ -64,16 +61,6 @@ var (
type Params struct { type Params struct {
// URL is where GeoJS is asked, normally URL. // URL is where GeoJS is asked, normally URL.
URL string URL string
// Timeout is how long a request waits for its client's first answer,
// and how long a request to GeoJS may take before it is abandoned
// (SWWAF_LOOKUP_TIMEOUT).
Timeout time.Duration
// Wait is true when a setting needs each request's answer before the
// request goes on. Otherwise no request waits for one.
Wait bool
// Answered, unless nil, is given each answer GeoJS gives, once it is
// kept.
Answered func(Answer)
// Now tells the time, normally time.Now. // Now tells the time, normally time.Now.
Now func() time.Time Now func() time.Time
// ProcessLog receives GeoJS's failures. // ProcessLog receives GeoJS's failures.
@@ -81,22 +68,16 @@ 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' AS numbers and countries through GeoJS. At most // GeoJS looks up clients' countries through GeoJS. At most one request
// one request to GeoJS is under way at a time, and it asks about every // to GeoJS is under way at a time, and it asks about every client waiting,
// client waiting, up to maxPerRequest. It is safe for concurrent use. // up to maxPerRequest. It is safe for concurrent use.
type GeoJS struct { type GeoJS struct {
url string url string
timeout time.Duration
wait bool
answered func(Answer)
now func() time.Time now func() time.Time
processLog *slog.Logger processLog *slog.Logger
metrics *metrics.Metrics metrics *metrics.Metrics
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
@@ -114,18 +95,11 @@ type GeoJS struct {
retryAt time.Time retryAt time.Time
} }
// Answer is what GeoJS or the lookup database said about a client: its AS // Answer is what GeoJS said about a client, as lookups.json holds it: its
// number, such as AS64496, and the AS's name, both "" when the source knows // country, "" when GeoJS cannot place it, when GeoJS said so, and when
// no AS number for it; its country, "" when the source cannot place it; // the answer was last used.
// when the source said so; and, for GeoJS's answers, which lookups.json
// holds, when the answer was last used. The zero Answer is that of a
// client with no answer.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Answer struct { type Answer struct {
Client netip.Prefix `json:"client"` Client netip.Prefix `json:"client"`
ASN string `json:"asn"`
ASName string `json:"as_name"`
Country string `json:"country"` Country string `json:"country"`
Answered time.Time `json:"answered"` Answered time.Time `json:"answered"`
Used time.Time `json:"used"` Used time.Time `json:"used"`
@@ -150,13 +124,9 @@ 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
@@ -167,23 +137,23 @@ func New(params Params) *GeoJS {
} }
} }
// LookUp returns the answer GeoJS gave about client, with its country as // Country returns the country GeoJS places client in, as a two-letter
// a two-letter code in capitals, or the zero Answer when there is none // code in capitals, or "" when the country cannot be found: GeoJS cannot
// yet. An answer is kept for 7 days. Without one, the client is asked // place the client, or has not answered in time. An answer is kept for 7
// about in the background, and, while Wait is set, the request waits up // days. Without one, a client waits up to timeout for it, unless it has
// to Timeout for the answer, unless the client has gone without one // gone without one before; until GeoJS answers, the client is asked about
// before. ctx is the context of the client's request, and ends the wait // again in the background. ctx is the context of the client's request,
// when it ends. // and ends the wait when it ends.
// //
// GeoJS is asked about the client's first address, which is the client's // GeoJS is asked about the client's first address, which is the client's
// own address for IPv4, and an address in the same place for an IPv6 /64. // own address for IPv4, and an address in the same place for an IPv6 /64.
func (g *GeoJS) LookUp(ctx context.Context, client netip.Prefix) Answer { func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string {
answer, asked := g.answerOrWait(ctx, client) country, asked := g.answerOrWait(ctx, client)
if asked == nil { if asked == nil {
return answer return country
} }
timer := time.NewTimer(g.timeout) timer := time.NewTimer(timeout)
defer timer.Stop() defer timer.Stop()
select { select {
@@ -195,7 +165,7 @@ func (g *GeoJS) LookUp(ctx context.Context, client netip.Prefix) Answer {
g.mu.Lock() g.mu.Lock()
defer g.mu.Unlock() defer g.mu.Unlock()
answer, found := g.kept(client) country, found := g.kept(client)
if !found { if !found {
g.metrics.GeoJSUnanswered.Inc() g.metrics.GeoJSUnanswered.Inc()
} }
@@ -205,15 +175,7 @@ func (g *GeoJS) LookUp(ctx context.Context, client netip.Prefix) Answer {
w.late = true w.late = true
} }
return answer return country
}
// Kept returns client's answer, if one is kept, without asking GeoJS.
func (g *GeoJS) Kept(client netip.Prefix) (Answer, bool) {
g.mu.Lock()
defer g.mu.Unlock()
return g.kept(client)
} }
// Snapshot returns every answer kept, sorted by client, as lookups.json // Snapshot returns every answer kept, sorted by client, as lookups.json
@@ -265,13 +227,13 @@ func (g *GeoJS) Load(answers []Answer) {
// nil when there is nothing to wait for. // nil when there is nothing to wait for.
func (g *GeoJS) answerOrWait( func (g *GeoJS) answerOrWait(
ctx context.Context, client netip.Prefix, ctx context.Context, client netip.Prefix,
) (Answer, <-chan struct{}) { ) (string, <-chan struct{}) {
g.mu.Lock() g.mu.Lock()
defer g.mu.Unlock() defer g.mu.Unlock()
answer, found := g.kept(client) country, found := g.kept(client)
if found { if found {
return answer, nil return country, nil
} }
w, waiting := g.waiting[client] w, waiting := g.waiting[client]
@@ -282,14 +244,10 @@ func (g *GeoJS) answerOrWait(
g.ask(ctx) g.ask(ctx)
if !g.wait {
return Answer{}, nil // the answer is not needed before the request goes on
}
if w == nil { if w == nil {
g.metrics.GeoJSUnanswered.Inc() g.metrics.GeoJSUnanswered.Inc()
return Answer{}, nil // too many clients wait already return "", nil // too many clients wait already
} }
if !g.asking { if !g.asking {
@@ -300,25 +258,25 @@ func (g *GeoJS) answerOrWait(
if w.late { if w.late {
g.metrics.GeoJSUnanswered.Inc() g.metrics.GeoJSUnanswered.Inc()
return Answer{}, nil return "", nil
} }
return Answer{}, w.asked return "", w.asked
} }
// kept returns client's answer, if GeoJS gave it less than keepFor ago, // kept returns client's answer, if GeoJS gave it less than keepFor ago,
// and notes that it was used. // and notes that it was used.
func (g *GeoJS) kept(client netip.Prefix) (Answer, bool) { func (g *GeoJS) kept(client netip.Prefix) (string, bool) {
now := g.now() now := g.now()
kept, found := g.answers.Get(client) kept, found := g.answers.Get(client)
if !found || now.Sub(kept.Answered) >= keepFor { if !found || now.Sub(kept.Answered) >= keepFor {
return Answer{}, false return "", false
} }
kept.Used = now kept.Used = now
return *kept, true return kept.Country, true
} }
// ask starts asking GeoJS about the waiting clients, unless a request to // ask starts asking GeoJS about the waiting clients, unless a request to
@@ -336,8 +294,7 @@ func (g *GeoJS) ask(ctx context.Context) {
} }
// askAboutWaiting asks GeoJS about the waiting clients, one request at a // askAboutWaiting asks GeoJS about the waiting clients, one request at a
// time, until none is left or GeoJS fails. Each answer kept is given to // time, until none is left or GeoJS fails.
// Answered, outside the lock, since Answered takes locks of its own.
func (g *GeoJS) askAboutWaiting(ctx context.Context) { func (g *GeoJS) askAboutWaiting(ctx context.Context) {
for { for {
clients := g.nextClients() clients := g.nextClients()
@@ -345,16 +302,8 @@ func (g *GeoJS) askAboutWaiting(ctx context.Context) {
return return
} }
given, err := g.request(ctx, clients) countries, err := g.request(ctx, clients)
kept, answered := g.keep(clients, given, err) if !g.keep(clients, countries, err) {
if g.answered != nil {
for _, answer := range kept {
g.answered(answer)
}
}
if !answered {
return return
} }
} }
@@ -386,35 +335,32 @@ func (g *GeoJS) nextClients() []netip.Prefix {
return clients return clients
} }
// keep notes how a request to GeoJS about clients ended, given being the // keep notes how a request to GeoJS about clients ended, and reports
// answer for each address GeoJS's answer names. It returns the answers it // whether GeoJS answered about all of them. Each client whose address
// kept, and reports whether GeoJS answered about all of the clients. Each // GeoJS's answer names gets its answer, with no country when GeoJS gave
// client whose address GeoJS's answer names gets its answer. An answer // none. An answer that leaves an address out is a failure. After a
// that leaves an address out is a failure. After a failure GeoJS is left // failure GeoJS is left alone for a while, and every client still waiting
// alone for a while, and every client still waiting stops waiting and is // stops waiting and is asked about once GeoJS is asked again.
// asked about once GeoJS is asked again.
func (g *GeoJS) keep( func (g *GeoJS) keep(
clients []netip.Prefix, given map[netip.Addr]Answer, err error, clients []netip.Prefix, countries map[netip.Addr]string, err error,
) ([]Answer, bool) { ) bool {
g.mu.Lock() g.mu.Lock()
defer g.mu.Unlock() defer g.mu.Unlock()
now := g.now() now := g.now()
kept := make([]Answer, 0, len(clients))
leftOut := 0 leftOut := 0
for _, client := range clients { for _, client := range clients {
answer, named := given[client.Addr()] country, named := countries[client.Addr()]
if !named { if !named {
leftOut++ leftOut++
continue continue
} }
answer.Client, answer.Answered, answer.Used = client, now, now g.answers.Add(client, &Answer{
g.answers.Add(client, &answer) Client: client, Country: country, Answered: now, Used: now,
kept = append(kept, answer) })
close(g.waiting[client].asked) close(g.waiting[client].asked)
delete(g.waiting, client) delete(g.waiting, client)
} }
@@ -440,37 +386,27 @@ 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 kept, false return false
} }
g.retryDelay = 0 g.retryDelay = 0
return kept, true return true
} }
// request asks GeoJS about clients in one request, and returns the answer // request asks GeoJS about clients in one request, and returns the
// for each address GeoJS's answer names: its AS number and the AS's name, // country it gave, in capitals, for each address its answer names.
// both "" for the AS number 64512, which GeoJS gives when it knows none,
// and its country, in capitals.
func (g *GeoJS) request( func (g *GeoJS) request(
ctx context.Context, clients []netip.Prefix, ctx context.Context, clients []netip.Prefix,
) (map[netip.Addr]Answer, error) { ) (map[netip.Addr]string, error) {
addrs := make([]string, 0, len(clients)) addrs := make([]string, 0, len(clients))
for _, client := range clients { for _, client := range clients {
addrs = append(addrs, client.Addr().String()) addrs = append(addrs, client.Addr().String())
} }
ctx, cancel := context.WithTimeout(ctx, g.timeout) ctx, cancel := context.WithTimeout(ctx, timeout)
defer cancel() defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, g.url, http.NoBody) req, err := http.NewRequestWithContext(ctx, http.MethodGet, g.url, http.NoBody)
@@ -497,12 +433,9 @@ func (g *GeoJS) request(
return nil, fmt.Errorf("%w %s", errStatus, res.Status) return nil, fmt.Errorf("%w %s", errStatus, res.Status)
} }
//nolint:tagliatelle // GeoJS's own names
var answers []struct { var answers []struct {
IP string `json:"ip"` IP string `json:"ip"`
ASN int64 `json:"asn"` Country string `json:"country"`
ASName string `json:"organization_name"`
CountryCode string `json:"country_code"`
} }
err = json.NewDecoder(io.LimitReader(res.Body, maxResponseBytes)).Decode(&answers) err = json.NewDecoder(io.LimitReader(res.Body, maxResponseBytes)).Decode(&answers)
@@ -510,22 +443,14 @@ func (g *GeoJS) request(
return nil, fmt.Errorf("read GeoJS's answer: %w", err) return nil, fmt.Errorf("read GeoJS's answer: %w", err)
} }
given := make(map[netip.Addr]Answer, len(answers)) countries := make(map[netip.Addr]string, len(answers))
for _, item := range answers { for _, item := range answers {
addr, err := netip.ParseAddr(item.IP) addr, err := netip.ParseAddr(item.IP)
if err != nil { if err == nil {
continue countries[addr] = strings.ToUpper(item.Country)
} }
answer := Answer{Country: strings.ToUpper(item.CountryCode)}
if item.ASN != 0 && item.ASN != unknownASN {
answer.ASN = "AS" + strconv.FormatInt(item.ASN, 10)
answer.ASName = item.ASName
}
given[addr] = answer
} }
return given, nil return countries, nil
} }
+16 -247
View File
@@ -6,8 +6,6 @@ import (
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"net/netip" "net/netip"
"net/url"
"reflect"
"slices" "slices"
"strings" "strings"
"sync" "sync"
@@ -16,20 +14,15 @@ 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, and asNumber, kept as asn, and asName the AS it gives them. // unplaced.
germany = "DE" germany = "DE"
asNumber = 64496 // unplaced is the address it cannot place.
asn = "AS64496"
asName = "Example Net"
// unplaced is the address it cannot place, for which it gives the AS
// number 64512 and the AS name Unknown, as GeoJS does.
unplaced = "192.0.2.1" unplaced = "192.0.2.1"
// leftOut is the address it leaves out of its answer when // leftOut is the address it leaves out of its answer when
// answeringWithoutLeftOut. // answeringWithoutLeftOut.
@@ -87,7 +80,7 @@ func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
var earlier sync.WaitGroup var earlier sync.WaitGroup
earlier.Go(func() { g.LookUp(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) }) earlier.Go(func() { g.Country(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) })
defer earlier.Wait() defer earlier.Wait()
waitForRequests(t, geojs, 1) waitForRequests(t, geojs, 1)
@@ -119,58 +112,6 @@ func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) {
}) })
} }
func TestRequestWaitsAsLongAsTheTimeoutSaysAndGeoJSIsAbandonedAfterIt(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
// A timeout longer than the default second, and a GeoJS that does
// not answer.
const longerTimeout = 3 * time.Second
m := metrics.New(1, "app")
g := lookup.New(lookup.Params{
URL: lookup.URL,
Timeout: longerTimeout,
Wait: true,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: m,
Alerts: alerts.New(alerts.Params{}),
})
g.SetTransport(&standIn{answers: hanging})
var (
request sync.WaitGroup
waited time.Duration
)
request.Go(func() {
began := time.Now()
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
waited = time.Since(began)
})
// A moment before the timeout runs out, GeoJS is still being asked:
// the request to it has not failed.
time.Sleep(longerTimeout - time.Millisecond)
synctest.Wait()
wantFailures(t, m, 0)
// As it runs out, the client's request goes on, and the request to
// GeoJS is abandoned, which counts as a failure.
request.Wait()
synctest.Wait()
if waited != longerTimeout {
t.Errorf("waited %s for the answer, want %s", waited, longerTimeout)
}
wantFailures(t, m, 1)
})
}
func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) { func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) {
t.Parallel() t.Parallel()
@@ -243,96 +184,6 @@ func TestCountryIsKeptInCapitals(t *testing.T) {
}) })
} }
func TestAnswerHoldsTheASNumberTheASNameAndTheCountry(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
_, clock, g := start()
placed := netip.MustParsePrefix("203.0.113.9/32")
notPlaced := netip.MustParsePrefix(unplaced + "/32")
now := clock.Now()
// For the client it cannot place, GeoJS gives the AS number 64512
// and the AS name Unknown, which count as unknown.
for client, want := range map[netip.Prefix]lookup.Answer{
placed: {
Client: placed, ASN: asn, ASName: asName, Country: germany,
Answered: now, Used: now,
},
notPlaced: {Client: notPlaced, Answered: now, Used: now},
} {
got := g.LookUp(t.Context(), client)
if got != want {
t.Errorf("answer for %s\n%+v\nwant\n%+v", client, got, want)
}
}
})
}
func TestWithoutWaitTheRequestGoesOnAtOnceAndTheAnswerIsGivenWhenItComes(
t *testing.T,
) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
var (
mu sync.Mutex
given []lookup.Answer
)
geojs := &standIn{answers: answeringSlowly}
clock := newClock()
m := metrics.New(1, "app")
g := lookup.New(lookup.Params{
URL: lookup.URL,
Timeout: timeout,
Answered: func(answer lookup.Answer) {
mu.Lock()
defer mu.Unlock()
given = append(given, answer)
},
Now: clock.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Metrics: m,
Alerts: alerts.New(alerts.Params{}),
})
g.SetTransport(geojs)
client := netip.MustParsePrefix("203.0.113.9/32")
// The request goes on at once, without an answer, and GeoJS is asked
// about the client, which it answers most of a second later.
began := time.Now()
got := g.LookUp(t.Context(), client)
if took := time.Since(began); took != 0 || got != (lookup.Answer{}) {
t.Errorf("waited %s for %+v, want no wait and no answer", took, got)
}
waitForRequests(t, geojs, 1)
wantAsked(t, geojs, 0, "203.0.113.9")
time.Sleep(timeout)
synctest.Wait()
now := clock.Now()
want := lookup.Answer{
Client: client, ASN: asn, ASName: asName, Country: germany,
Answered: now, Used: now,
}
mu.Lock()
if !slices.Equal(given, []lookup.Answer{want}) {
t.Errorf("answers given %+v, want only %+v", given, want)
}
mu.Unlock()
wantCountry(t, g, client, germany)
wantUnanswered(t, m, 0)
})
}
func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) { func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
t.Parallel() t.Parallel()
@@ -343,12 +194,9 @@ func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) {
geojs := &standIn{answers: hanging} geojs := &standIn{answers: hanging}
g := lookup.New(lookup.Params{ g := lookup.New(lookup.Params{
URL: lookup.URL, URL: lookup.URL,
Timeout: timeout,
Wait: true,
Now: time.Now, Now: time.Now,
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)), ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
Metrics: metrics.New(1, "app"), Metrics: metrics.New(1),
Alerts: alerts.New(alerts.Params{}),
}) })
g.SetTransport(geojs) g.SetTransport(geojs)
@@ -363,45 +211,6 @@ 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()
@@ -546,15 +355,12 @@ func TestClientsWithoutAnAnswerAreCounted(t *testing.T) {
t.Parallel() t.Parallel()
synctest.Test(t, func(t *testing.T) { synctest.Test(t, func(t *testing.T) {
m := metrics.New(1, "app") m := metrics.New(1)
g := lookup.New(lookup.Params{ g := lookup.New(lookup.Params{
URL: lookup.URL, URL: lookup.URL,
Timeout: timeout,
Wait: true,
Now: time.Now, Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler), ProcessLog: slog.New(slog.DiscardHandler),
Metrics: m, Metrics: m,
Alerts: alerts.New(alerts.Params{}),
}) })
g.SetTransport(&standIn{answers: failing}) g.SetTransport(&standIn{answers: failing})
@@ -644,24 +450,21 @@ func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) {
} }
} }
list := make([]map[string]any, 0, len(addrs)) list := make([]map[string]string, 0, len(addrs))
for _, addr := range addrs { for _, addr := range addrs {
item := map[string]any{ country := germany
"ip": addr, "asn": asNumber, "organization_name": asName,
"country_code": germany,
}
switch { switch {
case addr == unplaced: case addr == unplaced:
item = map[string]any{"ip": addr, "asn": 64512, "organization_name": "Unknown"} country = ""
case addr == leftOut && answers == answeringWithoutLeftOut: case addr == leftOut && answers == answeringWithoutLeftOut:
continue continue
case answers == answeringInLowerCase: case answers == answeringInLowerCase:
item["country_code"] = strings.ToLower(germany) country = strings.ToLower(germany)
} }
list = append(list, item) list = append(list, map[string]string{"ip": addr, "country": country})
} }
var answer any = list var answer any = list
@@ -719,43 +522,19 @@ func (c *testClock) advance(d time.Duration) {
} }
// start returns a stand-in for GeoJS that answers, a clock, and a GeoJS // start returns a stand-in for GeoJS that answers, a clock, and a GeoJS
// asking the stand-in by that clock, for which a request waits for its // asking the stand-in by that clock.
// client's first answer.
func start() (*standIn, *testClock, *lookup.GeoJS) { func start() (*standIn, *testClock, *lookup.GeoJS) {
geojs, clock, g, _ := startWithAlerts()
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 := newClock() clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
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, "app"), Metrics: metrics.New(1),
Alerts: queue,
}) })
g.SetTransport(geojs) g.SetTransport(geojs)
return geojs, clock, g, queue return geojs, clock, g
}
// 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
@@ -774,7 +553,7 @@ func newClients() func() netip.Prefix {
func wantCountry(t *testing.T, g *lookup.GeoJS, client netip.Prefix, want string) { func wantCountry(t *testing.T, g *lookup.GeoJS, client netip.Prefix, want string) {
t.Helper() t.Helper()
got := g.LookUp(t.Context(), client).Country got := g.Country(t.Context(), client)
if got != want { if got != want {
t.Errorf("%s is in %q, want %q", client, got, want) t.Errorf("%s is in %q, want %q", client, got, want)
} }
@@ -819,16 +598,6 @@ func wantUnanswered(t *testing.T, m *metrics.Metrics, want float64) {
} }
} }
// wantFailures checks how many requests to GeoJS m counts as failed.
func wantFailures(t *testing.T, m *metrics.Metrics, want float64) {
t.Helper()
got := testutil.ToFloat64(m.GeoJSFailures)
if got != want {
t.Errorf("%v requests to GeoJS failed, want %v", got, want)
}
}
// waitForRequests waits until g has done all it can before time passes, // waitForRequests waits until g has done all it can before time passes,
// checks that GeoJS has had count requests, and returns the addresses each // checks that GeoJS has had count requests, and returns the addresses each
// asked about. // asked about.
-83
View File
@@ -1,83 +0,0 @@
// Package lookuptest writes lookup databases, IPinfo Lite files in their
// .mmdb form, for the tests of the packages that read them.
package lookuptest
import (
"bytes"
"net"
"os"
"testing"
"github.com/maxmind/mmdbwriter"
"github.com/maxmind/mmdbwriter/mmdbtype"
)
// fileMode is the mode of the files written: read and written by their
// owner alone.
const fileMode = 0o600
// Network is what a lookup database holds about a netblock, of the fields
// smallwebwaf reads: its AS number, such as AS64496, the AS's name, and
// its country, such as DE.
type Network struct {
ASN string
ASName string
Country string
}
// Write writes a lookup database at path that places each netblock in
// networks, such as 203.0.113.0/24, as its Network says, and no other
// address.
func Write(tb testing.TB, path string, networks map[string]Network) {
tb.Helper()
records := make(map[string]mmdbtype.Map, len(networks))
for netblock, network := range networks {
records[netblock] = mmdbtype.Map{
"asn": mmdbtype.String(network.ASN),
"as_name": mmdbtype.String(network.ASName),
"country_code": mmdbtype.String(network.Country),
}
}
WriteRecords(tb, path, records)
}
// WriteRecords writes a lookup database at path that holds each record in
// records for its netblock, and nothing for any other address.
func WriteRecords(tb testing.TB, path string, records map[string]mmdbtype.Map) {
tb.Helper()
tree, err := mmdbwriter.New(mmdbwriter.Options{
DatabaseType: "ipinfo_lite",
// The tests' clients are in the netblocks kept for documentation.
IncludeReservedNetworks: true,
})
if err != nil {
tb.Fatalf("new lookup database: %v", err)
}
for netblock, record := range records {
_, network, err := net.ParseCIDR(netblock)
if err != nil {
tb.Fatalf("netblock %q: %v", netblock, err)
}
err = tree.Insert(network, record)
if err != nil {
tb.Fatalf("insert %s: %v", netblock, err)
}
}
var database bytes.Buffer
_, err = tree.WriteTo(&database)
if err != nil {
tb.Fatalf("write the lookup database: %v", err)
}
err = os.WriteFile(path, database.Bytes(), fileMode)
if err != nil {
tb.Fatalf("write %s: %v", path, err)
}
}
+2 -5
View File
@@ -26,11 +26,8 @@ func TestSnapshotHoldsEachAnswerAndWhenItWasLastUsed(t *testing.T) {
wantCountry(t, g, placed, germany) wantCountry(t, g, placed, germany)
want := []lookup.Answer{ want := []lookup.Answer{
{Client: notPlaced, Answered: asked, Used: asked}, {Client: notPlaced, Country: "", Answered: asked, Used: asked},
{ {Client: placed, Country: germany, Answered: asked, Used: asked.Add(time.Hour)},
Client: placed, ASN: asn, ASName: asName, Country: germany,
Answered: asked, Used: asked.Add(time.Hour),
},
} }
if got := g.Snapshot(); !slices.Equal(got, want) { if got := g.Snapshot(); !slices.Equal(got, want) {
t.Errorf("snapshot\n%+v\nwant\n%+v", got, want) t.Errorf("snapshot\n%+v\nwant\n%+v", got, want)
-155
View File
@@ -1,155 +0,0 @@
package metrics
import (
"sync"
"github.com/prometheus/client_golang/prometheus"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// other is the label value under which the countries or AS numbers
// outside the busiest are counted.
const other = "other"
// busiest are the metrics by one thing the lookup finds of the client,
// its country or its AS number, for requests whose client's is known.
// The topN busiest countries or AS numbers, by their requests since the
// start, have series of their own, and the others are counted under
// other, so that there are never more than topN + 1 series. One that
// drops out of the busiest loses its series, and its next requests are
// counted under other; one that becomes one of them gets a series that
// counts from then on. Each series therefore only ever goes up.
type busiest struct {
topN int
requests *prometheus.CounterVec
requestBytes *prometheus.CounterVec
responseBytes *prometheus.CounterVec
// refused are, by country, the requests the country lists refused; nil
// by AS number.
refused *prometheus.CounterVec
mu sync.Mutex
// seen is each country's or AS number's requests since the start, by
// which they are ranked.
seen map[string]int64
// top are those with series of their own.
top map[string]bool
}
// newCountries returns the metrics by country, with series of their own
// for the topN busiest countries.
func newCountries(topN int) *busiest {
countries := newBusiest(topN, "country", "the client's country")
countries.refused = counterVec("smallwebwaf_country_list_refusals_total",
"Requests the country lists refused, by the client's country.",
[]string{"country"})
return countries
}
// newASNs returns the metrics by AS number, with series of their own for
// the topN busiest AS numbers.
func newASNs(topN int) *busiest {
return newBusiest(topN, "asn", "the client's AS number")
}
// newBusiest returns the metrics by label, which is described as
// description, with series of their own for the topN busiest values.
func newBusiest(topN int, label, description string) *busiest {
by := []string{label}
return &busiest{
topN: topN,
requests: counterVec("smallwebwaf_"+label+"_requests_total",
"Requests, by "+description+".", by),
requestBytes: counterVec("smallwebwaf_"+label+"_request_bytes_total",
"Request body bytes, by "+description+".", by),
responseBytes: counterVec("smallwebwaf_"+label+"_response_bytes_total",
"Response body bytes, by "+description+".", by),
seen: map[string]int64{},
top: map[string]bool{},
}
}
// Describe and Collect make the metrics a prometheus.Collector, so that
// they are registered together.
func (b *busiest) Describe(ch chan<- *prometheus.Desc) {
for _, vec := range b.vecs() {
vec.Describe(ch)
}
}
// Collect is the other half of prometheus.Collector, with Describe.
func (b *busiest) Collect(ch chan<- prometheus.Metric) {
for _, vec := range b.vecs() {
vec.Collect(ch)
}
}
// vecs returns the metrics: by AS number, those of requests and bytes; by
// country, the refusals by the country lists as well.
func (b *busiest) vecs() []*prometheus.CounterVec {
vecs := []*prometheus.CounterVec{b.requests, b.requestBytes, b.responseBytes}
if b.refused != nil {
vecs = append(vecs, b.refused)
}
return vecs
}
// add counts a request from its log line, whose client's country or AS
// number, value, is known.
func (b *busiest) add(value string, line *requestlog.Line) {
b.mu.Lock()
defer b.mu.Unlock()
b.seen[value]++
label := b.label(value)
b.requests.WithLabelValues(label).Inc()
b.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes))
b.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes))
if b.refused != nil && line.Action == requestlog.ActionCountryDenied {
b.refused.WithLabelValues(label).Inc()
}
}
// label returns the label a request from value is counted under: value
// while it is one of the busiest, other while it is not. A value busier
// than the least busy of them takes its place, and that one's series are
// dropped.
func (b *busiest) label(value string) string {
if b.top[value] {
return value
}
if len(b.top) < b.topN {
b.top[value] = true
return value
}
least := ""
for top := range b.top {
if least == "" || b.seen[top] < b.seen[least] {
least = top
}
}
if b.seen[value] <= b.seen[least] {
return other
}
delete(b.top, least)
for _, vec := range b.vecs() {
vec.DeleteLabelValues(least)
}
b.top[value] = true
return value
}
+116
View File
@@ -0,0 +1,116 @@
package metrics
import (
"sync"
"github.com/prometheus/client_golang/prometheus"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// other is the label under which the countries outside the busiest are
// counted.
const other = "other"
// countries are the metrics by the client's country, for requests whose
// client's country is known. The topN busiest countries, by their requests
// since the start, have series of their own, and the others are counted
// under other, so that there are never more than topN + 1 series. A
// country that drops out of the busiest loses its series, and its next
// requests are counted under other; one that becomes one of them gets a
// series that counts from then on. Each series therefore only ever goes
// up.
type countries struct {
topN int
requests *prometheus.CounterVec
requestBytes *prometheus.CounterVec
responseBytes *prometheus.CounterVec
// refused are the requests the country lists refused.
refused *prometheus.CounterVec
mu sync.Mutex
// seen is each country's requests since the start, by which the
// countries are ranked. GeoJS gives two-letter codes, so it holds at
// most a few hundred.
seen map[string]int64
// top are the countries with series of their own.
top map[string]bool
}
// newCountries returns the metrics by country, with series of their own
// for the topN busiest countries.
func newCountries(topN int) *countries {
byCountry := []string{"country"}
return &countries{
topN: topN,
requests: counterVec("smallwebwaf_country_requests_total",
"Requests, by the client's country.", byCountry),
requestBytes: counterVec("smallwebwaf_country_request_bytes_total",
"Request body bytes, by the client's country.", byCountry),
responseBytes: counterVec("smallwebwaf_country_response_bytes_total",
"Response body bytes, by the client's country.", byCountry),
refused: counterVec("smallwebwaf_country_list_refusals_total",
"Requests the country lists refused, by the client's country.",
byCountry),
seen: map[string]int64{},
top: map[string]bool{},
}
}
// add counts a request from its log line, whose country is known.
func (c *countries) add(line *requestlog.Line) {
c.mu.Lock()
defer c.mu.Unlock()
c.seen[line.Country]++
label := c.label(line.Country)
c.requests.WithLabelValues(label).Inc()
c.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes))
c.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes))
if line.Action == requestlog.ActionCountryDenied {
c.refused.WithLabelValues(label).Inc()
}
}
// label returns the label a request from country is counted under: the
// country while it is one of the busiest, other while it is not. A
// country busier than the least busy of them takes its place, and that
// country's series are dropped.
func (c *countries) label(country string) string {
if c.top[country] {
return country
}
if len(c.top) < c.topN {
c.top[country] = true
return country
}
least := ""
for top := range c.top {
if least == "" || c.seen[top] < c.seen[least] {
least = top
}
}
if c.seen[country] <= c.seen[least] {
return other
}
delete(c.top, least)
for _, vec := range []*prometheus.CounterVec{
c.requests, c.requestBytes, c.responseBytes, c.refused,
} {
vec.DeleteLabelValues(least)
}
c.top[country] = true
return country
}
+22 -184
View File
@@ -6,26 +6,21 @@ 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/config"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/remotelog" "sneak.berlin/go/smallwebwaf/internal/remotelog"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
) )
// Metrics are smallwebwaf's metrics. They are safe for concurrent use. // Metrics are smallwebwaf's metrics. They are safe for concurrent use.
type Metrics struct { type Metrics struct {
// registry gives every metric registered with it the label instance. registry *prometheus.Registry
registry prometheus.Registerer
handler http.Handler handler http.Handler
inFlight prometheus.Gauge inFlight prometheus.Gauge
@@ -37,17 +32,14 @@ type Metrics struct {
rateLimitHits *prometheus.CounterVec rateLimitHits *prometheus.CounterVec
sizeAndTimeLimitHits *prometheus.CounterVec sizeAndTimeLimitHits *prometheus.CounterVec
offences *prometheus.CounterVec offences *prometheus.CounterVec
// ruleMatches are made by AddRules, and reputationHits by // ruleMatches are made by AddRules.
// AddReputation. ruleMatches *prometheus.CounterVec
ruleMatches *prometheus.CounterVec countries *countries
reputationHits *prometheus.CounterVec
countries *busiest
asns *busiest
// GeoJSRequests are the requests to GeoJS, and GeoJSFailures those // GeoJSRequests are the requests to GeoJS, and GeoJSFailures those
// that failed. GeoJSUnanswered are the requests that needed their // that failed. GeoJSUnanswered are the requests whose client counted
// client's answer, for a setting that acts on it, and went on without // as coming from an unknown country because GeoJS had not answered
// it because GeoJS had not given it in time. // about it in time.
GeoJSRequests prometheus.Counter GeoJSRequests prometheus.Counter
GeoJSFailures prometheus.Counter GeoJSFailures prometheus.Counter
GeoJSUnanswered prometheus.Counter GeoJSUnanswered prometheus.Counter
@@ -61,18 +53,14 @@ type Metrics struct {
} }
// New returns the metrics, with the Go runtime's and the process's own. // New returns the metrics, with the Go runtime's and the process's own.
// topN is how many countries and how many AS numbers get series of their // topN is how many countries get series of their own
// own (SWWAF_METRICS_TOP_N). Every metric carries instanceName // (SWWAF_METRICS_TOP_N).
// (SWWAF_INSTANCE_NAME) as its label instance. func New(topN int) *Metrics {
func New(topN int, instanceName string) *Metrics {
byStatus := []string{"status_class", "action"} byStatus := []string{"status_class", "action"}
byFile := []string{"file"} byFile := []string{"file"}
registry := prometheus.NewRegistry()
m := &Metrics{ m := &Metrics{
registry: prometheus.WrapRegistererWith( registry: prometheus.NewRegistry(),
prometheus.Labels{"instance": instanceName}, registry),
handler: promhttp.HandlerFor(registry, promhttp.HandlerOpts{}),
inFlight: prometheus.NewGauge(prometheus.GaugeOpts{ inFlight: prometheus.NewGauge(prometheus.GaugeOpts{
Name: "smallwebwaf_requests_in_flight", Name: "smallwebwaf_requests_in_flight",
Help: "Requests under way.", Help: "Requests under way.",
@@ -94,16 +82,14 @@ func New(topN int, instanceName string) *Metrics {
Help: "How long requests passed to the app took, from then to their end.", Help: "How long requests passed to the app took, from then to their end.",
}), }),
rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total", rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total",
"Requests that broke a rate limit or a byte limit, by its window and "+ "Requests that broke a rate limit, by its window.",
"its kind, requests or bytes.", []string{"window"}),
[]string{"window", "kind"}),
sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total", sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total",
"Requests that passed a size or time limit, by its setting.", "Requests that passed a size or time limit, by its setting.",
[]string{"limit"}), []string{"limit"}),
offences: counterVec("smallwebwaf_offences_total", offences: counterVec("smallwebwaf_offences_total",
"Offences, by kind.", []string{"kind"}), "Offences, by kind.", []string{"kind"}),
countries: newCountries(topN), countries: newCountries(topN),
asns: newASNs(topN),
GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{ GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_requests_total", Name: "smallwebwaf_geojs_requests_total",
Help: "Requests to GeoJS.", Help: "Requests to GeoJS.",
@@ -114,8 +100,8 @@ func New(topN int, instanceName string) *Metrics {
}), }),
GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{ GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{
Name: "smallwebwaf_geojs_unanswered_total", Name: "smallwebwaf_geojs_unanswered_total",
Help: "Requests that needed their client's answer from GeoJS and " + Help: "Requests whose client counted as coming from an unknown " +
"went on without it, because GeoJS had not given it in time.", "country because GeoJS had not answered about it in time.",
}), }),
stateFileWrites: counterVec("smallwebwaf_state_file_writes_total", stateFileWrites: counterVec("smallwebwaf_state_file_writes_total",
"Writes of each state file.", byFile), "Writes of each state file.", byFile),
@@ -132,12 +118,16 @@ func New(topN int, instanceName string) *Metrics {
byFile), byFile),
} }
m.handler = promhttp.HandlerFor(m.registry, promhttp.HandlerOpts{})
m.registry.MustRegister( m.registry.MustRegister(
collectors.NewGoCollector(), collectors.NewGoCollector(),
collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}), collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}),
m.inFlight, m.requests, m.requestBytes, m.responseBytes, m.inFlight, m.requests, m.requestBytes, m.responseBytes,
m.requestDuration, m.upstreamDuration, m.requestDuration, m.upstreamDuration,
m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences, m.countries, m.asns, m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences,
m.countries.requests, m.countries.requestBytes, m.countries.responseBytes,
m.countries.refused,
m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered, m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered,
m.stateFileWrites, m.stateFileWriteFailures, m.stateFileWrites, m.stateFileWriteFailures,
m.stateFileLastWrite, m.stateFileSize, m.stateFileLastWrite, m.stateFileSize,
@@ -235,146 +225,6 @@ 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())
}),
)
}
// AddReputation adds the metrics of the lists fetched from URLs and of the
// DNSBL zones, by source, each list's URL or each zone, its key masked as
// config.MaskZoneKey masks it: the requests whose client a blocklist or a
// zone's verdict lists, which ReputationHit counts, and, read from lists
// and dnsbl as the metrics are asked for, for a list, the fetches that
// failed and when the copy in use was fetched, and for a zone, the queries
// made and those that failed. It is called once, before ReputationHit.
func (m *Metrics) AddReputation(lists *reputation.Lists, dnsbl *reputation.DNSBL) {
const (
sourceLabel = "source"
failuresHelp = "Fetches of the list, or queries to the DNSBL zone, that failed."
)
m.reputationHits = counterVec("smallwebwaf_reputation_hits_total",
"Requests whose client a blocklist or a DNSBL zone lists, by the "+
"blocklist's URL or the zone.",
[]string{sourceLabel})
m.registry.MustRegister(m.reputationHits)
for _, zone := range dnsbl.Zones() {
source := prometheus.Labels{sourceLabel: config.MaskZoneKey(zone)}
m.registry.MustRegister(
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_reputation_queries_total",
Help: "Queries to the DNSBL zone.",
ConstLabels: source,
}, func() float64 {
return float64(dnsbl.Queries(zone))
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_reputation_failures_total",
Help: failuresHelp,
ConstLabels: source,
}, func() float64 {
return float64(dnsbl.Failures(zone))
}),
)
}
for _, listURL := range lists.URLs() {
source := prometheus.Labels{sourceLabel: listURL}
m.registry.MustRegister(
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_reputation_failures_total",
Help: failuresHelp,
ConstLabels: source,
}, func() float64 {
return float64(lists.Failures(listURL))
}),
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
Name: "smallwebwaf_reputation_last_fetch_timestamp_seconds",
Help: "When the copy of the list in use was fetched, in seconds since " +
"1970, or 0 while there is none.",
ConstLabels: source,
}, func() float64 {
fetched := lists.Fetched(listURL)
if fetched.IsZero() {
return 0
}
return float64(fetched.Unix())
}),
)
}
}
// ReputationHit counts a request whose client source lists: a blocklist,
// by its URL, or a DNSBL zone, its key masked.
func (m *Metrics) ReputationHit(source string) {
m.reputationHits.WithLabelValues(source).Inc()
}
// AddAlerts adds the metrics of the alerts sent to each destination set,
// read from queue as the metrics are asked for, by destination: the
// alerts sent, the requests to the destination that failed, the alerts
// held back, which are the same for every destination, and those
// dropped. With no destination set, it adds none.
func (m *Metrics) AddAlerts(queue *alerts.Queue) {
for _, name := range queue.DestinationsSet() {
destination := prometheus.Labels{"destination": name}
m.registry.MustRegister(
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_sent_total",
Help: "Alerts the destination took.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Counts(name).Sent)
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_failed_total",
Help: "Requests to the destination that failed.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Counts(name).Failed)
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_suppressed_total",
Help: "Alerts held back: repeats within SWWAF_ALERT_COOLDOWN, and " +
"alerts past SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Suppressed())
}),
prometheus.NewCounterFunc(prometheus.CounterOpts{
Name: "smallwebwaf_alerts_dropped_total",
Help: "Alerts dropped, the oldest first, from a full queue, and " +
"alerts given up as the destination refused them.",
ConstLabels: destination,
}, func() float64 {
return float64(queue.Counts(name).Dropped)
}),
)
}
}
// ServeHTTP answers with the metrics in the Prometheus text format. // ServeHTTP answers with the metrics in the Prometheus text format.
func (m *Metrics) ServeHTTP(w http.ResponseWriter, r *http.Request) { func (m *Metrics) ServeHTTP(w http.ResponseWriter, r *http.Request) {
m.handler.ServeHTTP(w, r) m.handler.ServeHTTP(w, r)
@@ -405,15 +255,7 @@ func (m *Metrics) RequestEnded(
} }
if line.LimitHit != "" { if line.LimitHit != "" {
// The log line names a byte limit's window with _bytes after it. m.rateLimitHits.WithLabelValues(line.LimitHit).Inc()
window, isBytes := strings.CutSuffix(line.LimitHit, "_bytes")
kind := ratelimit.KindRequests
if isBytes {
kind = ratelimit.KindBytes
}
m.rateLimitHits.WithLabelValues(window, kind).Inc()
} }
if limit != "" { if limit != "" {
@@ -425,11 +267,7 @@ func (m *Metrics) RequestEnded(
} }
if line.Country != "" { if line.Country != "" {
m.countries.add(line.Country, line) m.countries.add(line)
}
if line.ASN != "" {
m.asns.add(line.ASN, line)
} }
} }
-273
View File
@@ -1,273 +0,0 @@
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])
}
}
}
-30
View File
@@ -1,30 +0,0 @@
package proxy
import (
"net/netip"
"testing"
"sneak.berlin/go/smallwebwaf/internal/config"
)
func TestWithEveryAnomalyThresholdOffARequestIsNotCounted(t *testing.T) {
t.Parallel()
// A request from a client looked up through GeoJS, with every anomaly
// threshold off. Its handler has neither GeoJS's answers nor the
// anomaly counters, nor a clock, and the request no response: reading
// any of them to count the request panics.
rq := &request{
h: &handler{config: &config.Config{LookupSource: "geojs"}},
client: netip.MustParseAddr("203.0.113.9"),
lookedUp: true,
}
defer func() {
if r := recover(); r != nil {
t.Errorf("counting the request did work, with every threshold off: %v", r)
}
}()
rq.countAnomalies()
}
-377
View File
@@ -1,377 +0,0 @@
package proxy_test
import (
"maps"
"net/http"
"net/netip"
"reflect"
"slices"
"strconv"
"sync/atomic"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The anomaly thresholds: the prefix of a scope followed by the end of a
// count.
const (
anomalyClient = "SWWAF_ANOMALY_CLIENT_"
anomalyNet = "SWWAF_ANOMALY_NET_"
anomalyASN = "SWWAF_ANOMALY_ASN_"
anomalyTotal = "SWWAF_ANOMALY_TOTAL_"
anomalyWatch = "SWWAF_WATCH_"
requestsPerMinute = "REQUESTS_PER_MINUTE"
requestsPerHour = "REQUESTS_PER_HOUR"
bytesPerMinute = "BYTES_PER_MINUTE"
bytesPerHour = "BYTES_PER_HOUR"
)
// The other anomaly settings.
const (
anomalyNetV4Prefix = "SWWAF_ANOMALY_NET_V4_PREFIX"
anomalyNetV6Prefix = "SWWAF_ANOMALY_NET_V6_PREFIX"
watchNets = "SWWAF_WATCH_NETS"
)
const (
// clientsNet is the netblock around client at the default length, and
// office a named netblock of the same.
clientsNet = "203.0.113.0/24"
office = "office=" + clientsNet
// aLot is a threshold no test reaches.
aLot = "1000"
// hour is the window an alert names for a threshold per hour.
hour = "hour"
)
func TestEachScopeAndWindowOverItsThresholdAlertsOncePerCooldown(t *testing.T) {
t.Parallel()
for _, scope := range []struct {
prefix, scope string
// netblock is the alert's, and counted what its reason names. extra
// is what its detail gives besides what every anomaly alert's does.
netblock netip.Prefix
counted string
extra map[string]any
}{
{
anomalyClient, anomaly.ScopeClient, netip.MustParsePrefix(client + "/32"),
"the client " + client + "/32", nil,
},
{
anomalyNet, anomaly.ScopeNet, netip.MustParsePrefix(clientsNet),
"the netblock " + clientsNet, nil,
},
{anomalyASN, anomaly.ScopeASN, netip.Prefix{}, asnDE, map[string]any{"asn": asnDE}},
{anomalyTotal, anomaly.ScopeTotal, netip.Prefix{}, "the whole service", nil},
{
anomalyWatch, anomaly.ScopeWatch, netip.MustParsePrefix(clientsNet),
"the named netblock office, " + clientsNet, map[string]any{"name": "office"},
},
} {
for _, threshold := range []struct {
end, kind, window string
// value is the threshold, which the third upload of 100 bytes
// takes the count over, to count.
value int64
count float64
}{
{requestsPerMinute, ratelimit.KindRequests, minute, 2, 3},
{requestsPerHour, ratelimit.KindRequests, hour, 2, 3},
{bytesPerMinute, ratelimit.KindBytes, minute, 250, 300},
{bytesPerHour, ratelimit.KindBytes, hour, 250, 300},
} {
setting := scope.prefix + threshold.end
value := strconv.FormatInt(threshold.value, 10)
t.Run(setting, func(t *testing.T) {
t.Parallel()
s, clk, server, queue := startWithLookupsAndClock(t, map[string]string{
setting: value, watchNets: office,
})
start := clk.Now()
// The third upload takes the count over the threshold, and the
// fourth, within the cooldown, is held back. Each is passed to
// the app.
for range 4 {
s.uploadFrom(client)
}
detail := map[string]any{
"scope": scope.scope, "window": threshold.window, "kind": threshold.kind,
"count": threshold.count, "threshold": threshold.value,
}
maps.Copy(detail, scope.extra)
wantAlerts(t, queue, alerts.Alert{
Instance: alertInstance,
Time: start,
Event: alerts.EventAnomaly,
Client: netip.MustParseAddr(client),
Netblock: scope.netblock,
ASN: asnDE,
ASName: asNameDE,
Country: "DE",
Reason: threshold.kind + " per " + threshold.window + " of " +
scope.counted + " over the threshold of " + value,
Detail: detail,
})
wantAlertedAgainOnceTheCooldownHasRunOut(t, s, clk, queue)
if held := server.Ledger.Snapshot(); len(held) != 0 {
t.Errorf("the ledger holds %+v, want no ban", held)
}
})
}
}
}
// wantAlertedAgainOnceTheCooldownHasRunOut checks that, once the cooldown
// has run out after a first alert, which held back one repeat, the next
// count over the threshold, at the latest three uploads from client on,
// raises another alert, giving that repeat.
func wantAlertedAgainOnceTheCooldownHasRunOut(
t *testing.T, s *sender, clk *clock, queue *alerts.Queue,
) {
t.Helper()
clk.advance(15 * time.Minute)
for range 3 {
s.uploadFrom(client)
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 2 || !waiting[1].Time.Equal(clk.Now()) ||
waiting[1].SuppressedRepeats != 1 {
t.Errorf("alerts wait %+v, want the first and another, with 1 repeat", waiting)
}
}
func TestEveryRequestIsCountedWhateverIsDoneWithIt(t *testing.T) {
t.Parallel()
const (
allowed = "192.0.2.7" // in SWWAF_ALLOW_NETS
exempt = "192.0.2.10" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
denied = "192.0.2.20" // in SWWAF_DENY_NETS
)
s, _, _, queue := startAppWithAlerts(t, readAndAnswer, map[string]string{
anomalyClient + requestsPerMinute: "2",
allowNets: allowed,
rateLimitExemptNets: exempt,
rateLimitExemptPaths: "/static/",
denyNets: denied,
})
// The third request of each takes its client's count over the threshold
// of 2.
for _, sent := range []struct {
from, path string
status int
action string
}{
{allowed, "/", http.StatusOK, requestlog.ActionForward},
{exempt, "/", http.StatusOK, requestlog.ActionForward},
{client, "/static/app.js", http.StatusOK, requestlog.ActionForward},
{denied, "/", http.StatusForbidden, requestlog.ActionDenied},
} {
for range 3 {
s.request(sent.from, sent.path, sent.status, sent.action)
}
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
got := make([]string, 0, len(waiting))
for _, alert := range waiting {
got = append(got, alert.Client.String())
}
if want := []string{allowed, exempt, client, denied}; !slices.Equal(got, want) {
t.Errorf("alerts for the clients %v, want %v", got, want)
}
}
func TestThresholdsOffCountNothingAndAlertNothing(t *testing.T) {
t.Parallel()
// With every threshold off, nothing is counted.
s, server, queue := startWithLookups(t, map[string]string{watchNets: office})
for range 5 {
s.uploadFrom(client)
}
if counters := server.Anomalies.Snapshot(); len(counters) != 0 {
t.Errorf("counters %+v, want none", counters)
}
wantAlerts(t, queue)
// With one set, its count alone is counted, in its scope alone.
s, clk, server, queue := startWithLookupsAndClock(t, map[string]string{
anomalyNet + requestsPerMinute: aLot, watchNets: office,
})
for range 5 {
s.uploadFrom(client)
}
want := []anomaly.Counter{{
Scope: anomaly.ScopeNet,
Netblock: netip.MustParsePrefix(clientsNet),
Minute: ratelimit.Buckets{Start: clk.Now(), Current: 5},
}}
if got := server.Anomalies.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("counters\n%+v\nwant\n%+v", got, want)
}
wantAlerts(t, queue)
}
func TestNetblockAroundAClientIsAsLongAsTheSettingsSay(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
// Each client of sent sends one request, and want gives the
// netblocks they are counted in, each with its requests.
sent []string
want map[string]int64
}{
{
"by default", nil,
[]string{client, "203.0.113.200", "192.0.2.7", ipv6Client, "2001:db8:0:ffff::1"},
map[string]int64{clientsNet: 2, "192.0.2.0/24": 1, "2001:db8::/48": 2},
},
{
"as set", map[string]string{anomalyNetV4Prefix: "16", anomalyNetV6Prefix: "32"},
[]string{client, "203.0.200.1", ipv6Client, "2001:db8:ffff::1"},
map[string]int64{"203.0.0.0/16": 2, "2001:db8::/32": 2},
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{anomalyNet + requestsPerMinute: aLot}
maps.Copy(env, tc.env)
s, _, server, _ := startAppWithAlerts(t, readAndAnswer, env)
for _, from := range tc.sent {
s.get(from, http.StatusOK, requestlog.ActionForward)
}
got := map[string]int64{}
for _, counter := range server.Anomalies.Snapshot() {
got[counter.Netblock.String()] = counter.Minute.Current
}
if !maps.Equal(got, tc.want) {
t.Errorf("requests by netblock %v, want %v", got, tc.want)
}
})
}
}
func TestClientIsCountedForItsASNumberOnceTheLookupGivesOne(t *testing.T) {
t.Parallel()
s, server, _ := startWithLookups(t, map[string]string{
anomalyASN + requestsPerMinute: aLot,
})
// The lookup database does not hold unplaced.
for _, from := range []string{fromDE, fromDE, fromKP, noCountry, unplaced} {
s.uploadFrom(from)
}
got := map[string]int64{}
for _, counter := range server.Anomalies.Snapshot() {
got[counter.ASN] = counter.Minute.Current
}
if want := map[string]int64{asnDE: 2, asnKP: 1, "AS64500": 1}; !maps.Equal(got, want) {
t.Errorf("requests by AS number %v, want %v", got, want)
}
}
func TestRequestCountsForTheASNumberGeoJSGivesBeforeItEnds(t *testing.T) {
t.Parallel()
// The stand-in for GeoJS answers only once released, which the app
// does as it answers the request, and then waits until the answer is
// kept.
geojsURL, _, release := startHeldGeoJS(t)
var server atomic.Pointer[proxy.Server]
app := startApp(t, func(http.ResponseWriter, *http.Request) {
release()
waitUntil(func() bool {
_, kept := server.Load().GeoJS.Kept(netip.MustParsePrefix(fromDE + "/32"))
return kept
})
})
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
addr, out, started := startProxyWithClock(t, app.URL, geojsURL, clk.Now,
map[string]string{
trustedProxies: trustLocalhost,
lookupTimeout: "1h",
anomalyASN + requestsPerMinute: aLot,
})
server.Store(started)
// The request went on without the answer, and is counted for the AS
// number it gives.
s := &sender{t: t, addr: addr, out: out}
if line := s.get(fromDE, http.StatusOK, requestlog.ActionForward); line.ASN != "" {
t.Errorf("log line has AS number %q, want none: the request waited", line.ASN)
}
want := []anomaly.Counter{{
Scope: anomaly.ScopeASN, ASN: asnDE,
Minute: ratelimit.Buckets{Start: clk.Now(), Current: 1},
}}
if got := started.Anomalies.Snapshot(); !reflect.DeepEqual(got, want) {
t.Errorf("counters\n%+v\nwant\n%+v", got, want)
}
}
func TestEachNamedNetblockCountsTheClientsInIt(t *testing.T) {
t.Parallel()
s, _, server, _ := startAppWithAlerts(t, readAndAnswer, map[string]string{
anomalyWatch + requestsPerMinute: aLot,
watchNets: office + ",wide=203.0.0.0/16,other=198.51.100.0/25",
})
// client is in office and in wide.
for _, from := range []string{client, "203.0.200.1", "192.0.2.7"} {
s.get(from, http.StatusOK, requestlog.ActionForward)
}
got := map[string]int64{}
for _, counter := range server.Anomalies.Snapshot() {
got[counter.Name] = counter.Minute.Current
}
if want := map[string]int64{"office": 1, "wide": 2}; !maps.Equal(got, want) {
t.Errorf("requests by named netblock %v, want %v", got, want)
}
}
+28 -193
View File
@@ -4,9 +4,7 @@ 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"
) )
@@ -18,244 +16,81 @@ func (rq *request) banResponse(action string) *refusal {
} }
// banned reports whether a ban on a netblock the client is in covers the // banned reports whether a ban on a netblock the client is in covers the
// request at now, and notes for the log line when that ban ends. A // request at now, and notes for the log line when that ban ends.
// 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 // the ban refuses nothing, and stays as it is check = rq.h.ledger.Find // in observe mode the ban refuses nothing
} }
ban, banned, madePermanent := check(rq.client, now) ban, banned := 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 rate limit, as its limit percentage lowers it, which // the client over a limit. In enforce mode such a request bans the
// breaks it. // client's netblock, and sets the client's counters back to zero; in
// observe mode it does neither.
func (rq *request) limitBroken(now time.Time) bool { func (rq *request) limitBroken(now time.Time) bool {
counts, hit, over := rq.h.limiter.Count(clientGroup(rq.client), now, group := clientGroup(rq.client)
rq.limitPercent.percent)
counts, hit, over := rq.h.limiter.Count(group, now)
rq.line.Counts = counts rq.line.Counts = counts
if over { if !over {
rq.banForLimit(now, hit, rq.h.config.BanResponse) return false
} }
return over
}
// countBytes counts the request's bytes, as countedBytes gives them, for
// the byte limits, once its response has ended, and notes the client's
// byte totals for the log line; its requests stay there as the rate limits
// counted them. Only a request passed to the app has them counted, and
// only one the rate limits counted; in observe mode, not one that enforce
// mode would have refused. Bytes that take the client over a byte limit,
// as its limit percentage for the byte limits lowers it, break it; the
// response was passed on whole.
func (rq *request) countBytes() {
if !rq.counted || rq.line.WouldAction != "" {
return
}
now := rq.h.now()
counts, hit, over := rq.h.limiter.CountBytes(clientGroup(rq.client), now,
rq.countedBytes(), rq.bytesPercent.percent)
rq.line.Counts.MinuteBytes = counts.MinuteBytes
rq.line.Counts.HourBytes = counts.HourBytes
rq.line.Counts.DayBytes = counts.DayBytes
if over {
rq.banForLimit(now, hit, rq.out.status)
}
}
// countedBytes returns the request's bytes, once it has ended, as the
// byte limits and the anomaly thresholds count them: the response's body
// bytes, the request's, or both, as SWWAF_BYTES_COUNT says. For an
// upgraded connection, such as a WebSocket, which has closed by then, what
// it carried from the app counts with the response's and what it carried
// from the client with the request's.
func (rq *request) countedBytes() int64 {
response, request := rq.out.bytes, rq.requestBytes()
if rq.upgraded != nil {
response += rq.upgraded.fromApp.Load()
request += rq.upgraded.toApp.Load()
}
switch rq.h.config.BytesCount {
case "response":
return response
case "request":
return request
default: // both
return response + request
}
}
// banForLimit bans the client's netblock at now for a broken limit, the
// one hit names, and notes the offence for the log line. status is what
// the client was sent, or is sent: SWWAF_BAN_RESPONSE for a request over
// a rate limit, the app's answer for one whose bytes broke a byte limit.
// The ban's notes give the client's limit percentage for that kind of
// limit. The ban sets the client's counters back to zero. In observe mode
// it makes no ban and sets nothing back, and raises the alert for the ban
// it would have made, if that alert would be sent.
func (rq *request) banForLimit(now time.Time, hit ratelimit.Hit, status int) {
rq.line.LimitHit = hit.Window rq.line.LimitHit = hit.Window
if hit.Kind == ratelimit.KindBytes {
rq.line.LimitHit += "_bytes" // as counts names the byte totals
}
rq.line.Offence = requestlog.OffenceLimit rq.line.Offence = requestlog.OffenceLimit
netblock := rq.h.netblock(rq.client) if rq.h.config.Observe {
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseLimit) { return true
return
} }
notes := bans.Notes{ netblock := rq.h.netblock(rq.client)
ASN: rq.line.ASN, ban := rq.h.ledger.BanForLimit(netblock, now, bans.Notes{
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.Count, Count: hit.Requests,
Request: rq.noted(now, status), Request: rq.noted(now),
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)
if made { return true
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. In observe mode it makes no ban, // attack, the match of rule, a ban rule.
// and raises the alert for the ban it would have made, if that alert
// would be sent.
func (rq *request) banForAttack(now time.Time, rule rules.Rule) { func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
netblock := rq.h.netblock(rq.client) netblock := rq.h.netblock(rq.client)
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseAttack) { ban := rq.h.ledger.BanForAttack(netblock, now, bans.Notes{
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, rq.h.config.BanResponse), Request: rq.noted(now),
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)
if made {
rq.alertBan(ban)
}
}
// wouldAlertBan reports whether the alert for a ban on netblock for cause
// made at now would be sent. In observe mode the ban the request would
// have made is worked out only then, at most once per
// SWWAF_ALERT_COOLDOWN and never with no webhook set: its notes count the
// netblock's requests, which can mean going through every client.
func (rq *request) wouldAlertBan(
netblock netip.Prefix, now time.Time, cause string,
) bool {
event := alerts.EventBan
if rq.h.ledger.WouldBePermanent(netblock, now, cause) {
event = alerts.EventPermanentBan
}
return rq.h.alerts.WouldSend(event, netblock)
}
// alertBan raises the alert for ban, which the request made, or made
// 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,
}) })
rq.line.BanExpires = banExpires(ban)
} }
// noted is the request, at now, with status, what the client was sent, or // noted is the request, refused at now with SWWAF_BAN_RESPONSE, as the
// in observe mode would have been, as the notes of the ban it makes keep // notes of the ban it makes keep it.
// it. func (rq *request) noted(now time.Time) bans.Request {
func (rq *request) noted(now time.Time, status int) bans.Request {
return bans.Request{ return bans.Request{
Time: now, Time: now,
Method: rq.in.Method, Method: rq.in.Method,
Host: rq.in.Host, Host: rq.in.Host,
Path: rq.in.URL.RequestURI(), Path: rq.in.URL.RequestURI(),
Status: status, Status: rq.h.config.BanResponse,
UserAgent: rq.in.UserAgent(), UserAgent: rq.in.UserAgent(),
} }
} }
+2 -6
View File
@@ -281,10 +281,7 @@ func TestBanNotes(t *testing.T) {
Cause: bans.CauseLimit, Cause: bans.CauseLimit,
Reason: "requests per minute over the limit of 1", Reason: "requests per minute over the limit of 1",
Notes: bans.Notes{ Notes: bans.Notes{
ASN: asnDE,
ASName: asNameDE,
Country: "DE", Country: "DE",
Kind: "requests",
Limit: 1, Limit: 1,
Window: minute, Window: minute,
Count: 2, Count: 2,
@@ -365,9 +362,8 @@ func (c *clock) advance(d time.Duration) {
// startWithClock starts smallwebwaf in front of an app that answers 200, // startWithClock starts smallwebwaf in front of an app that answers 200,
// with the settings in env on top of trusting localhost's // with the settings in env on top of trusting localhost's
// X-Forwarded-For, clients' AS numbers and countries looked up at // X-Forwarded-For, clients' countries looked up at geojsURL, and a clock
// geojsURL, and a clock set to midnight, the start of a bucket in every // set to midnight, the start of a bucket in every window.
// window.
func startWithClock( func startWithClock(
t *testing.T, geojsURL string, env map[string]string, t *testing.T, geojsURL string, env map[string]string,
) (*sender, *clock, *proxy.Server) { ) (*sender, *clock, *proxy.Server) {
-118
View File
@@ -1,118 +0,0 @@
package proxy
import (
"sneak.berlin/go/smallwebwaf/internal/config"
)
// whole is the percentage of each limit a client gets when no biased
// threshold lowers its limits.
const whole = 100
// percentage is a client's limit percentage for the rate limits or for
// the byte limits, as the biased thresholds give it, and the setting that
// gave it: "" with whole when none lowers that kind of limit.
type percentage struct {
percent int64
setting string
}
// biasedThresholdsSet reports whether a biased threshold can lower a
// client's limits: one of its lists is not empty,
// SWWAF_UNKNOWN_LIMIT_PERCENT is below 100, or SWWAF_ASN_LIMIT_PERCENT_URL
// is set. The client's lookup is then needed before its request goes on.
func biasedThresholdsSet(cfg *config.Config) bool {
return len(cfg.ASNLimitPercent) > 0 || len(cfg.CountryLimitPercent) > 0 ||
len(cfg.ASNBytesPercent) > 0 || len(cfg.CountryBytesPercent) > 0 ||
cfg.UnknownLimitPercent < whole || cfg.ASNLimitPercentURL != ""
}
// limitPercentages returns the client's limit percentages, for the rate
// limits and for the byte limits, by its AS number and country as looked
// up, each "" when unknown, and the blocklists and DNSBL zones that list
// it. Each is the lowest of those the settings give it, the first of them
// in the order below when several are lowest: the percentage
// SWWAF_ASN_LIMIT_PERCENT gives its AS number, the one the file
// SWWAF_ASN_LIMIT_PERCENT_URL names gives it, the one
// SWWAF_COUNTRY_LIMIT_PERCENT gives its country, for a client without a
// country, SWWAF_UNKNOWN_LIMIT_PERCENT, for a client a blocklist lists,
// the percentage of SWWAF_BLOCKLIST_ACTION while it is limit, and for a
// client a DNSBL zone's verdict lists, the percentage of
// SWWAF_REPUTATION_ACTION while it is limit. For the byte limits,
// SWWAF_ASN_BYTES_PERCENT and SWWAF_COUNTRY_BYTES_PERCENT take the place
// of the first three for an AS number or a country they list.
func (rq *request) limitPercentages() (percentage, percentage) {
cfg := rq.h.config
asn, country := rq.line.ASN, rq.line.Country
unknown := percentage{percent: whole}
if country == "" {
unknown = percentage{cfg.UnknownLimitPercent, "SWWAF_UNKNOWN_LIMIT_PERCENT"}
}
fetched := percentage{percent: whole}
if percent, listed := rq.h.lists.ASNLimitPercent(asn); listed {
fetched = percentage{percent, "SWWAF_ASN_LIMIT_PERCENT_URL"}
}
blocklisted := percentage{percent: whole}
if rq.blocklisted && cfg.BlocklistAction == "limit" {
blocklisted = percentage{cfg.BlocklistLimitPercent, "SWWAF_BLOCKLIST_ACTION"}
}
dnsblListed := percentage{percent: whole}
if rq.dnsblListed && cfg.ReputationAction == "limit" {
dnsblListed = percentage{cfg.ReputationLimitPercent, "SWWAF_REPUTATION_ACTION"}
}
asnRequests := lowest(given(cfg.ASNLimitPercent, asn, "SWWAF_ASN_LIMIT_PERCENT"),
fetched)
countryRequests := given(cfg.CountryLimitPercent, country,
"SWWAF_COUNTRY_LIMIT_PERCENT")
asnBytes, countryBytes := asnRequests, countryRequests
if _, listed := cfg.ASNBytesPercent[asn]; listed {
asnBytes = given(cfg.ASNBytesPercent, asn, "SWWAF_ASN_BYTES_PERCENT")
}
if _, listed := cfg.CountryBytesPercent[country]; listed {
countryBytes = given(cfg.CountryBytesPercent, country, "SWWAF_COUNTRY_BYTES_PERCENT")
}
return lowest(asnRequests, countryRequests, unknown, blocklisted, dnsblListed),
lowest(asnBytes, countryBytes, unknown, blocklisted, dnsblListed)
}
// given returns the percentage percents, the setting named setting, gives
// code, an AS number or a country, or whole when it does not list code.
func given(percents map[string]int64, code, setting string) percentage {
percent, listed := percents[code]
if !listed {
return percentage{percent: whole}
}
return percentage{percent, setting}
}
// lowest returns the lowest of percentages below whole, the first of them
// when several are lowest, or whole when none is below it.
func lowest(percentages ...percentage) percentage {
low := percentage{percent: whole}
for _, p := range percentages {
if p.percent < low.percent {
low = p
}
}
return low
}
// logged returns p as the log line and the notes of a ban give it: its
// percent and setting, or nil and "" for whole, which they leave out.
func (p percentage) logged() (*int64, string) {
if p.percent == whole {
return nil, ""
}
return &p.percent, p.setting
}
-504
View File
@@ -1,504 +0,0 @@
package proxy_test
import (
"fmt"
"io"
"maps"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"path/filepath"
"strings"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The biased thresholds.
const (
asnLimitPercent = "SWWAF_ASN_LIMIT_PERCENT"
countryLimitPercent = "SWWAF_COUNTRY_LIMIT_PERCENT"
asnBytesPercent = "SWWAF_ASN_BYTES_PERCENT"
countryBytesPercent = "SWWAF_COUNTRY_BYTES_PERCENT"
unknownLimitPercent = "SWWAF_UNKNOWN_LIMIT_PERCENT"
)
const (
// asnDEHalf and countryDEHalf give fromDE's AS number and its country
// half of every limit, and asnDEQuarter gives its AS number a quarter.
asnDEHalf = asnDE + ":50"
asnDEQuarter = asnDE + ":25"
countryDEHalf = "de:50"
// noCountry is in an AS of its own, AS64500, and in no country.
noCountry = "192.0.2.80"
// fourAMinute is the rate limit these tests set: half of it is 2
// requests a minute, a quarter of it 1.
fourAMinute = "4"
// twoUploads is the byte limit these tests set: 199 bytes, which an
// upload, a request with a body and its answer, 100 bytes, is within,
// and half of which, 99 bytes, it is over.
twoUploads = "199"
// none is how percentText gives a percentage left out.
none = "none"
)
func TestEachBiasedThresholdLowersTheRateLimits(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
setting, value, from string
}{
{asnLimitPercent, asnDEHalf, fromDE},
{countryLimitPercent, countryDEHalf, fromDE},
{unknownLimitPercent, "50", unplaced},
} {
t.Run(tc.setting, func(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{
rateLimitPerMinute: fourAMinute, tc.setting: tc.value,
})
// Half of 4 requests a minute: the third breaks the limit.
for _, sent := range []struct {
status int
action string
}{
{http.StatusOK, requestlog.ActionForward},
{http.StatusOK, requestlog.ActionForward},
{http.StatusForbidden, requestlog.ActionRateLimited},
} {
line := s.get(tc.from, sent.status, sent.action)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
"50 from "+tc.setting)
}
// fromKP, which no setting lists, has the whole limit.
for range 3 {
line := s.get(fromKP, http.StatusOK, requestlog.ActionForward)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
none)
}
})
}
}
func TestEachBiasedThresholdLowersTheByteLimits(t *testing.T) {
t.Parallel()
// The AS numbers and countries are given in either case.
for _, tc := range []struct {
setting, value, from string
}{
{asnLimitPercent, asnDEHalf, fromDE},
{countryLimitPercent, "DE:50", fromDE},
{unknownLimitPercent, "50", unplaced},
{asnBytesPercent, "as64496:50", fromDE},
{countryBytesPercent, countryDEHalf, fromDE},
} {
t.Run(tc.setting, func(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{
bytesLimitPerMinute: twoUploads, tc.setting: tc.value,
})
// The upload's 100 bytes are over half of 199, 99.
line := s.uploadFrom(tc.from)
if line.LimitHit != minuteBytes {
t.Errorf("log line has limit_hit %q, want %s", line.LimitHit, minuteBytes)
}
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
"50 from "+tc.setting)
// fromKP, which no setting lists, has the whole limit.
line = s.uploadFrom(fromKP)
if line.LimitHit != "" {
t.Errorf("log line for %s has limit_hit %q, want none", fromKP, line.LimitHit)
}
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting, none)
})
}
}
func TestBytesPercentSettingsTakeThePlaceOfTheOthersForByteLimits(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
// limitPercent and bytesPercent are the log line's, as percentText
// gives them, and limitHit is its limit_hit.
limitPercent, bytesPercent, limitHit string
}{
{
"lowering the byte limits alone",
map[string]string{asnBytesPercent: asnDEHalf},
none, "50 from " + asnBytesPercent, minuteBytes,
},
{
"lowering the byte limits alone, by country",
map[string]string{countryBytesPercent: countryDEHalf},
none, "50 from " + countryBytesPercent, minuteBytes,
},
{
"raising the byte limits back",
map[string]string{asnLimitPercent: asnDEHalf, asnBytesPercent: asnDE + ":100"},
"50 from " + asnLimitPercent, none, "",
},
{
"raising the byte limits back, by country",
map[string]string{countryLimitPercent: countryDEHalf, countryBytesPercent: "de:100"},
"50 from " + countryLimitPercent, none, "",
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{bytesLimitPerMinute: twoUploads}
maps.Copy(env, tc.env)
s, _, _ := startWithLookups(t, env)
// The upload's 100 bytes are over 99, half of 199, and within 199.
line := s.uploadFrom(fromDE)
if line.LimitHit != tc.limitHit {
t.Errorf("log line has limit_hit %q, want %q", line.LimitHit, tc.limitHit)
}
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
tc.limitPercent)
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
tc.bytesPercent)
})
}
}
func TestZeroPercentIsAZeroAllowance(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{asnLimitPercent: asnDE + ":0"})
// The first request breaks the limit, and bans the client; the log line
// gives the 0.
line := s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited)
if line.fields["limit_percent"] != float64(0) ||
line.fields["limit_percent_setting"] != asnLimitPercent {
t.Errorf("log line has limit_percent %v from %v, want 0 from %s",
line.fields["limit_percent"], line.fields["limit_percent_setting"],
asnLimitPercent)
}
s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned)
}
func TestLowestPercentageApplies(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
from string
// want is the log line's limit_percent, as percentText gives it.
want string
}{
{
"the country's",
map[string]string{asnLimitPercent: asnDEHalf, countryLimitPercent: "de:25"},
fromDE, "25 from " + countryLimitPercent,
},
{
"the AS number's",
map[string]string{asnLimitPercent: asnDEQuarter, countryLimitPercent: countryDEHalf},
fromDE, "25 from " + asnLimitPercent,
},
{
"the AS number's, the first of two alike",
map[string]string{asnLimitPercent: asnDEQuarter, countryLimitPercent: "de:25"},
fromDE, "25 from " + asnLimitPercent,
},
{
"that for a client without a country",
map[string]string{asnLimitPercent: "AS64500:50", unknownLimitPercent: "25"},
noCountry, "25 from " + unknownLimitPercent,
},
{
// SWWAF_UNKNOWN_LIMIT_PERCENT is left at its default, 100.
"the AS number's, for a client without a country",
map[string]string{asnLimitPercent: "AS64500:25"},
noCountry, "25 from " + asnLimitPercent,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{rateLimitPerMinute: fourAMinute}
maps.Copy(env, tc.env)
s, _, _ := startWithLookups(t, env)
// A quarter of 4 requests a minute: the second breaks the limit.
s.get(tc.from, http.StatusOK, requestlog.ActionForward)
line := s.get(tc.from, http.StatusForbidden, requestlog.ActionRateLimited)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
tc.want)
})
}
}
func TestUnknownLimitPercentGivesEveryClientWithoutACountryItsPercentage(t *testing.T) {
t.Parallel()
s, _, _ := startWithLookups(t, map[string]string{
rateLimitPerMinute: fourAMinute, unknownLimitPercent: "50",
})
// One the lookup database does not hold, and one on a private address,
// which is never looked up: the third request of each breaks half of 4.
for _, from := range []string{unplaced, "10.0.0.8"} {
s.get(from, http.StatusOK, requestlog.ActionForward)
s.get(from, http.StatusOK, requestlog.ActionForward)
s.get(from, http.StatusForbidden, requestlog.ActionRateLimited)
}
// One in a country has the whole limit.
for range 3 {
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
}
}
func TestClientWithoutAnAnswerInTimeHasTheUnknownLimitPercent(t *testing.T) {
t.Parallel()
// In a synctest bubble, as TestRequestWaitsAsLongAsTheLookupTimeoutSays
// says, with a GeoJS that never answers.
synctest.Test(t, func(t *testing.T) {
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
time.Now, map[string]string{unknownLimitPercent: "0"})
// Once the second the request waits for its answer is up, the client
// counts as without a country, and its zero allowance refuses the
// request before it reaches the app.
serveFromDE(t, server, http.MethodGet, http.NoBody)
line := out.requestLine(t)
wantLine(t, line, http.StatusForbidden, requestlog.ActionRateLimited)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
"0 from "+unknownLimitPercent)
})
}
func TestRequestWaitsForItsLookupWhileABiasedThresholdIsSet(t *testing.T) {
t.Parallel()
const timeout = 3 * time.Second
for _, tc := range []struct {
setting, value string
waits bool
}{
{asnLimitPercent, asnDEHalf, true},
{countryLimitPercent, countryDEHalf, true},
{asnBytesPercent, asnDEHalf, true},
{countryBytesPercent, countryDEHalf, true},
{unknownLimitPercent, "99", true},
{asnLimitPercentURL, asnURL, true},
// At 100, its default, it lowers no limit.
{unknownLimitPercent, "100", false},
} {
t.Run(tc.setting+"="+tc.value, func(t *testing.T) {
t.Parallel()
// In a synctest bubble, as TestRequestWaitsAsLongAsTheLookupTimeoutSays
// says, with a GeoJS that never answers.
synctest.Test(t, func(t *testing.T) {
// The request's body is over SWWAF_REQUEST_MAX_BYTES, so that it
// is refused after the checks, and never reaches the app.
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
time.Now, map[string]string{
lookupTimeout: timeout.String(), requestMaxBytes: "1",
tc.setting: tc.value,
})
began := time.Now()
serveFromDE(t, server, http.MethodPost, strings.NewReader("ab"))
want := time.Duration(0)
if tc.waits {
want = timeout
}
if waited := time.Since(began); waited != want {
t.Errorf("the request waited %s for its answer, want %s", waited, want)
}
wantLine(t, out.requestLine(t), http.StatusRequestEntityTooLarge,
requestlog.ActionTooLarge)
// The bubble's clock stops once this function returns, so the
// request to GeoJS, which a request that did not wait leaves
// under way, has to be abandoned before then.
time.Sleep(timeout)
})
})
}
}
func TestBanForALoweredLimitGivesThePercentageInItsNotesAndItsAlert(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
// before is how many uploads come before the one that breaks a
// limit, which is answered with status and logged with action.
before int
status int
action string
// reason and want are the ban's reason, and its notes' limit
// percentage, as percentText gives it.
reason, want string
}{
{
// A quarter of 12 requests a minute is 3: the fourth breaks it.
"a rate limit",
map[string]string{rateLimitPerMinute: "12", asnLimitPercent: asnDEQuarter},
3, http.StatusForbidden, requestlog.ActionRateLimited,
"requests per minute over the limit of 3", "25 from " + asnLimitPercent,
},
{
// The byte limits' percentage, not the rate limits'.
"a byte limit",
map[string]string{
bytesLimitPerMinute: twoUploads, asnLimitPercent: asnDEQuarter,
asnBytesPercent: asnDEHalf,
},
0, http.StatusOK, requestlog.ActionForward,
"bytes per minute over the limit of 99", "50 from " + asnBytesPercent,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
s, server, queue := startWithLookups(t, tc.env)
for range tc.before {
s.uploadFrom(fromDE)
}
s.requestWithBody(http.MethodPost, fromDE, "/", uploadHeader, uploadBody,
tc.status, tc.action)
held := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32"))
if len(held) != 1 {
t.Fatalf("bans %+v, want one", held)
}
notes := held[0].Notes
if held[0].Reason != tc.reason {
t.Errorf("the ban's reason is %q, want %q", held[0].Reason, tc.reason)
}
wantPercent(t, "the notes' limit_percent", notes.LimitPercent,
notes.LimitPercentSetting, tc.want)
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 {
t.Fatalf("%d alerts wait, want the ban's alone: %+v", len(waiting), waiting)
}
alerted, _ := waiting[0].Detail["notes"].(bans.Notes)
wantPercent(t, "the alert's notes' limit_percent", alerted.LimitPercent,
alerted.LimitPercentSetting, tc.want)
})
}
}
// startWithLookups is startWithLookupsAndClock for a test that needs no
// clock.
func startWithLookups(
t *testing.T, env map[string]string,
) (*sender, *proxy.Server, *alerts.Queue) {
t.Helper()
s, _, server, queue := startWithLookupsAndClock(t, env)
return s, server, queue
}
// startWithLookupsAndClock is startAppWithAlerts in front of
// readAndAnswer, with the settings in env on top of clients looked up in a
// lookup database, which places fromDE and fromKP in the AS numbers and
// countries the stand-in for GeoJS gives them, noCountry in AS64500 and no
// country, and no other address.
func startWithLookupsAndClock(
t *testing.T, env map[string]string,
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
t.Helper()
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
lookuptest.Write(t, path, map[string]lookuptest.Network{
fromDE + "/32": {ASN: asnDE, ASName: asNameDE, Country: "DE"},
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
noCountry + "/32": {ASN: "AS64500", ASName: "Nowhere Net"},
})
settings := map[string]string{lookupSource: fileSource, lookupDBPath: path}
maps.Copy(settings, env)
return startAppWithAlerts(t, readAndAnswer, settings)
}
// uploadFrom is upload from the client at from.
func (s *sender) uploadFrom(from string) logLine {
s.t.Helper()
line, _ := s.requestWithBody(http.MethodPost, from, "/", uploadHeader, uploadBody,
http.StatusOK, requestlog.ActionForward)
return line
}
// serveFromDE hands a request from fromDE with method and body straight to
// server's handler, without the network, and returns once it is answered.
func serveFromDE(t *testing.T, server *proxy.Server, method string, body io.Reader) {
t.Helper()
req := httptest.NewRequestWithContext(t.Context(), method, "/", body)
req.RemoteAddr = net.JoinHostPort(fromDE, "1234")
server.Handler.ServeHTTP(httptest.NewRecorder(), req)
}
// wantPercent checks a limit percentage that a log line or a ban's notes
// give, what, and the setting that gave it, against want, as percentText
// gives them.
func wantPercent(t *testing.T, what string, percent *int64, setting, want string) {
t.Helper()
if got := percentText(percent, setting); got != want {
t.Errorf("%s is %s, want %s", what, got, want)
}
}
// percentText gives a limit percentage and the setting that gave it as
// text, such as "50 from SWWAF_ASN_LIMIT_PERCENT", or none when both are
// left out.
func percentText(percent *int64, setting string) string {
switch {
case percent == nil && setting == "":
return none
case percent == nil:
return "none from " + setting
default:
return fmt.Sprintf("%d from %s", *percent, setting)
}
}
-42
View File
@@ -103,48 +103,6 @@ func (b *responseBody) Close() error {
return b.body.Close() return b.body.Close()
} }
// upgradedConn is the connection to the app once the app has switched
// protocols, as for a WebSocket. ReverseProxy writes to it what the client
// sends and reads from it what the app sends, on goroutines of its own,
// until the connection closes; it counts the bytes each way, for the byte
// limits.
type upgradedConn struct {
io.ReadWriteCloser
// fromApp is how many bytes the app has sent, and toApp how many the
// client has.
fromApp atomic.Int64
toApp atomic.Int64
}
// Read reads what the app sends.
func (c *upgradedConn) Read(p []byte) (int, error) {
n, err := c.ReadWriteCloser.Read(p)
c.fromApp.Add(int64(n))
return n, err
}
// Write sends the app what the client sent.
func (c *upgradedConn) Write(p []byte) (int, error) {
n, err := c.ReadWriteCloser.Write(p)
c.toApp.Add(int64(n))
return n, err
}
// CloseWrite tells the app that the client sends no more, while what the
// app sends still passes. ReverseProxy calls it once the client has
// stopped sending, and closes the connection there if it is not supported.
func (c *upgradedConn) CloseWrite() error {
conn, ok := c.ReadWriteCloser.(interface{ CloseWrite() error })
if !ok {
return http.ErrNotSupported
}
return conn.CloseWrite()
}
// limitBody returns body, cut off with an *http.MaxBytesError after // limitBody returns body, cut off with an *http.MaxBytesError after
// maxBytes, or unchanged if maxBytes is zero, which is off. // maxBytes, or unchanged if maxBytes is zero, which is off.
func limitBody(body io.ReadCloser, maxBytes int64) io.ReadCloser { func limitBody(body io.ReadCloser, maxBytes int64) io.ReadCloser {
-583
View File
@@ -1,583 +0,0 @@
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
}
+26 -6
View File
@@ -1,17 +1,31 @@
package proxy package proxy
import ( import (
"context"
"net/netip"
"slices" "slices"
) )
// countryDenied reports whether the country lists refuse the request, by // countryDenied reports whether the country lists refuse the request.
// the client's country as it was looked up. A client without a country, // The client's country is looked up only while a list is set, and never
// or whose country cannot be found, is refused only by // for a client on a private, loopback or link-local address, which has
// SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES. // no country. A client without a country, or whose country cannot be
func (rq *request) countryDenied() bool { // found, is refused only by SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES. ctx is
// the request's own context.
func (rq *request) countryDenied(ctx context.Context) bool {
denied := rq.h.config.DeniedCountries denied := rq.h.config.DeniedCountries
allowed := rq.h.config.ExclusivelyAllowedCountries allowed := rq.h.config.ExclusivelyAllowedCountries
country := rq.line.Country
if len(denied) == 0 && len(allowed) == 0 {
return false
}
var country string
if hasCountry(rq.client) {
country = rq.h.geojs.Country(ctx, clientGroup(rq.client))
}
rq.line.Country = country
if slices.Contains(denied, country) { if slices.Contains(denied, country) {
return true return true
@@ -19,3 +33,9 @@ func (rq *request) countryDenied() bool {
return len(allowed) > 0 && !slices.Contains(allowed, country) return len(allowed) > 0 && !slices.Contains(allowed, country)
} }
// hasCountry reports whether addr can be placed in a country: private,
// loopback and link-local addresses cannot.
func hasCountry(addr netip.Addr) bool {
return !addr.IsPrivate() && !addr.IsLoopback() && !addr.IsLinkLocalUnicast()
}
+34 -99
View File
@@ -56,21 +56,16 @@ func TestCountryLists(t *testing.T) {
maps.Copy(env, tc.env) maps.Copy(env, tc.env)
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env) addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env)
// The AS number GeoJS gives unplaced, 64512, counts as unknown. for i, sent := range []struct{ client, country string }{
for i, sent := range []struct{ client, asn, asName, country string }{ {fromDE, "DE"}, {fromKP, "KP"}, {unplaced, ""},
{fromDE, asnDE, asNameDE, "DE"}, {fromKP, asnKP, asNameKP, "KP"},
{unplaced, "", "", ""},
} { } {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, sent.client) req.Header.Set(forwardedFor, sent.client)
got := do(t, req) got := do(t, req)
line := out.requestLines(t, i+1)[i] line := out.requestLines(t, i+1)[i]
if line.ASN != sent.asn || line.ASName != sent.asName || if line.Country != sent.country {
line.Country != sent.country { t.Errorf("log line has country %q, want %q", line.Country, sent.country)
t.Errorf("log line has %q, %q and %q, want %q, %q and %q",
line.ASN, line.ASName, line.Country,
sent.asn, sent.asName, sent.country)
} }
if slices.Contains(tc.refused, sent.client) { if slices.Contains(tc.refused, sent.client) {
@@ -135,7 +130,7 @@ func TestRequestRefusedByCountryIsNotCounted(t *testing.T) {
return return
} }
answer := []geojsAnswer{{IP: r.URL.Query().Get("ip"), CountryCode: "DE"}} answer := []map[string]string{{"ip": r.URL.Query().Get("ip"), "country": "DE"}}
err := json.NewEncoder(w).Encode(answer) err := json.NewEncoder(w).Encode(answer)
if err != nil { if err != nil {
@@ -179,15 +174,20 @@ func TestRequestRefusedByCountryIsNotCounted(t *testing.T) {
wantStatus(t, got, http.StatusOK) wantStatus(t, got, http.StatusOK)
} }
func TestPrivateAddressIsNeverLookedUp(t *testing.T) { func TestCountryNotLookedUpWithoutAListOrForAPrivateAddress(t *testing.T) {
t.Parallel() t.Parallel()
for _, tc := range []struct { for _, tc := range []struct {
name string name string
env map[string]string env map[string]string
clients []string // "" sends no X-Forwarded-For: the client is 127.0.0.1
}{ }{
{"no setting needs the lookup", nil}, {"no country list is set", nil, []string{fromKP, fromDE}},
{"a country list is set", map[string]string{deniedCountries: "kp"}}, {
"private, loopback and link-local addresses",
map[string]string{deniedCountries: "kp"},
[]string{"10.0.0.5", "192.168.1.9", "fd00::5", "", "169.254.0.9", "fe80::9"},
},
} { } {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -198,10 +198,7 @@ func TestPrivateAddressIsNeverLookedUp(t *testing.T) {
maps.Copy(env, tc.env) maps.Copy(env, tc.env)
addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env) addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env)
// "" sends no X-Forwarded-For: the client is 127.0.0.1. for i, sent := range tc.clients {
for i, sent := range []string{
"10.0.0.5", "192.168.1.9", "fd00::5", "", "169.254.0.9", "fe80::9",
} {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
if sent != "" { if sent != "" {
req.Header.Set(forwardedFor, sent) req.Header.Set(forwardedFor, sent)
@@ -212,26 +209,15 @@ func TestPrivateAddressIsNeverLookedUp(t *testing.T) {
line := out.requestLines(t, i+1)[i] line := out.requestLines(t, i+1)[i]
wantLine(t, line, http.StatusOK, requestlog.ActionForward) wantLine(t, line, http.StatusOK, requestlog.ActionForward)
for _, field := range []string{"asn", "as_name", "country"} { country, present := line.fields["country"]
value, present := line.fields[field] if !present || country != "" {
if !present || value != "" { t.Errorf("log line for %q has country %v, want an empty one",
t.Errorf("log line for %q has %s %v, want an empty one", line.ClientIP, country)
line.ClientIP, field, value)
}
} }
} }
// GeoJS is asked about up to 200 waiting clients at once, so once it if len(asked()) != 0 {
// has been asked about fromDE, which comes last, it has been asked t.Errorf("GeoJS was asked about %v, want nothing", asked())
// about every client before it that waited for an answer.
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, fromDE)
wantStatus(t, do(t, req), http.StatusOK)
waitUntil(func() bool { return slices.Contains(asked(), fromDE) })
if got := asked(); !slices.Equal(got, []string{fromDE}) {
t.Errorf("GeoJS was asked about %v, want %s alone", got, fromDE)
} }
}) })
} }
@@ -278,31 +264,18 @@ func TestExclusiveListRefusesAPrivateAddressUnlessAllowed(t *testing.T) {
} }
} }
// startGeoJS starts a stand-in for GeoJS, which places fromDE and fromKP, // startGeoJS starts a stand-in for GeoJS, which places fromDE and fromKP
// each in an AS of its own, and no other address. It returns its URL, and // and no other address. It returns its URL, and what returns the
// what returns the addresses it has been asked about. // addresses it has been asked about.
func startGeoJS(t *testing.T) (string, func() []string) { func startGeoJS(t *testing.T) (string, func() []string) {
t.Helper() t.Helper()
geojsURL, asked, release := startHeldGeoJS(t) places := map[string]string{fromDE: "DE", fromKP: "KP"}
release()
return geojsURL, asked var asked struct {
} mu sync.Mutex
addrs []string
// startHeldGeoJS is startGeoJS for a stand-in that answers nothing until }
// release is called. Each request to it waits until then.
func startHeldGeoJS(t *testing.T) (string, func() []string, func()) {
t.Helper()
var (
asked struct {
mu sync.Mutex
addrs []string
}
released = make(chan struct{})
once sync.Once
)
geojs := httptest.NewServer(http.HandlerFunc( geojs := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) { func(w http.ResponseWriter, r *http.Request) {
@@ -312,11 +285,11 @@ func startHeldGeoJS(t *testing.T) (string, func() []string, func()) {
asked.addrs = append(asked.addrs, addrs...) asked.addrs = append(asked.addrs, addrs...)
asked.mu.Unlock() asked.mu.Unlock()
<-released answers := make([]map[string]string, 0, len(addrs))
answers := make([]geojsAnswer, 0, len(addrs))
for _, addr := range addrs { for _, addr := range addrs {
answers = append(answers, answerAbout(addr)) answers = append(answers, map[string]string{
"ip": addr, "country": places[addr],
})
} }
err := json.NewEncoder(w).Encode(answers) err := json.NewEncoder(w).Encode(answers)
@@ -326,48 +299,10 @@ func startHeldGeoJS(t *testing.T) (string, func() []string, func()) {
})) }))
t.Cleanup(geojs.Close) t.Cleanup(geojs.Close)
release := func() { once.Do(func() { close(released) }) }
// Run before geojs.Close, which waits for every request to be answered.
t.Cleanup(release)
return geojs.URL, func() []string { return geojs.URL, func() []string {
asked.mu.Lock() asked.mu.Lock()
defer asked.mu.Unlock() defer asked.mu.Unlock()
return slices.Clone(asked.addrs) return slices.Clone(asked.addrs)
}, release
}
// The AS numbers and names the stand-in for GeoJS gives fromDE and
// fromKP, as they are logged.
const (
asnDE = "AS64496"
asNameDE = "Example Net"
asnKP = "AS64511"
asNameKP = "Other Net"
)
// geojsAnswer is an answer of GeoJS about one address, with the fields
// smallwebwaf reads.
//
//nolint:tagliatelle // GeoJS's own names
type geojsAnswer struct {
IP string `json:"ip"`
ASN int `json:"asn"`
ASName string `json:"organization_name"`
CountryCode string `json:"country_code,omitempty"`
}
// answerAbout is what the stand-in for GeoJS answers about addr: for an
// address it cannot place, the AS number 64512 and the AS name Unknown
// with no country, as GeoJS does.
func answerAbout(addr string) geojsAnswer {
switch addr {
case fromDE:
return geojsAnswer{IP: addr, ASN: 64496, ASName: asNameDE, CountryCode: "DE"}
case fromKP:
return geojsAnswer{IP: addr, ASN: 64511, ASName: asNameKP, CountryCode: "KP"}
} }
return geojsAnswer{IP: addr, ASN: 64512, ASName: "Unknown"}
} }
+2 -6
View File
@@ -24,9 +24,7 @@ func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
start := clk.Now() start := clk.Now()
// Two let through, one over the limit, which bans the client, and one // Two let through, one over the limit, which bans the client, and one
// refused under that ban, for which the client is not looked up. GeoJS // refused under that ban, for which the country is not looked up.
// answers about the client at its first request, and its later ones
// use that answer.
s.get(fromDE, http.StatusOK, requestlog.ActionForward) s.get(fromDE, http.StatusOK, requestlog.ActionForward)
clk.advance(time.Second) clk.advance(time.Second)
s.get(fromDE, http.StatusOK, requestlog.ActionForward) s.get(fromDE, http.StatusOK, requestlog.ActionForward)
@@ -37,10 +35,8 @@ func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) {
want := ratelimit.History{ want := ratelimit.History{
FirstSeen: start, FirstSeen: start,
LastSeen: start.Add(2 * time.Second), LastSeen: start.Add(2 * time.Second),
ASN: asnDE,
ASName: asNameDE,
Country: "DE", Country: "DE",
LookedUp: start, LookedUp: start.Add(time.Second),
Requests: 4, Requests: 4,
Forwarded: 2, Forwarded: 2,
Refused: 2, Refused: 2,
-71
View File
@@ -1,71 +0,0 @@
package proxy
import (
"context"
"net/http"
"net/netip"
"sneak.berlin/go/smallwebwaf/internal/lookup"
)
// The headers in which the app is passed the client's AS number and
// country while SWWAF_ADD_LOOKUP_HEADERS is set. Go writes every header
// name in this form, as it sends it and as it receives it, so X-Client-ASN
// arrives as X-Client-Asn, and Del removes a client's own whatever their
// case; header names are not case-sensitive.
const (
asnHeader = "X-Client-Asn"
countryHeader = "X-Client-Country"
)
// lookUp looks up the client's AS number and country, in the lookup
// database or through GeoJS, and notes them for the log line, unless
// SWWAF_LOOKUP_SOURCE is off or the client is on a private, loopback or
// link-local address, which no lookup can place. The lookup database
// answers at once. With GeoJS, while a setting needs the answer, such as a
// country list or a biased threshold, a new client's request waits for it.
// ctx is the request's own context.
func (rq *request) lookUp(ctx context.Context) {
if rq.h.config.LookupSource == "off" || !canBePlaced(rq.client) {
return
}
if rq.h.config.LookupSource == "file" {
rq.lookupAnswer = rq.h.lookupFile.LookUp(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
@@ -1,377 +0,0 @@
package proxy_test
import (
"net"
"net/http"
"net/http/httptest"
"net/netip"
"path/filepath"
"slices"
"sync"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// asnAndCountry is what a lookup gives a client: its AS number, AS name
// and country.
type asnAndCountry struct{ asn, asName, country string }
// fileSource is the SWWAF_LOOKUP_SOURCE that looks clients up in the
// lookup database.
const fileSource = "file"
func TestEveryClientIsLookedUpWithoutWaitingWhileNoSettingNeedsIt(t *testing.T) {
t.Parallel()
// The stand-in for GeoJS answers only once released. A request that
// waited for it would wait an hour, and get no answer within
// waitLimit.
geojsURL, asked, release := startHeldGeoJS(t)
s, _, server := startWithClock(t, geojsURL, map[string]string{
lookupTimeout: "1h",
rateLimitPerMinute: "1",
})
// fromDE's second request breaks the limit and bans it, and fromKP
// comes too. None waits for GeoJS.
for _, line := range []logLine{
s.get(fromDE, http.StatusOK, requestlog.ActionForward),
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited),
s.get(fromKP, http.StatusOK, requestlog.ActionForward),
} {
got := asnAndCountry{line.ASN, line.ASName, line.Country}
if got != (asnAndCountry{}) {
t.Errorf("log line has %+v before GeoJS answered, want nothing", got)
}
}
// Once GeoJS answers, each answer reaches the client's history, and
// fromDE's reaches the notes of its ban.
release()
netblock := netip.MustParsePrefix(fromDE + "/32")
waitUntil(func() bool {
return historyOf(t, server, fromDE).ASN != "" &&
historyOf(t, server, fromKP).ASN != "" &&
server.Ledger.Bans(netblock)[0].Notes.ASN != ""
})
de := asnAndCountry{asnDE, asNameDE, "DE"}
for addr, want := range map[string]asnAndCountry{
fromDE: de, fromKP: {asnKP, asNameKP, "KP"},
} {
h := historyOf(t, server, addr)
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != want {
t.Errorf("%s's history has %+v, want %+v", addr, got, want)
}
}
notes := server.Ledger.Bans(netblock)[0].Notes
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != de {
t.Errorf("the ban's notes have %+v, want %+v", got, de)
}
// GeoJS was asked about each client once, fromKP after fromDE, whose
// request was under way when fromKP came.
if got := asked(); !slices.Equal(got, []string{fromDE, fromKP}) {
t.Errorf("GeoJS was asked about %v, want %s and %s", got, fromDE, fromKP)
}
}
func TestASNumberAndNameInTheLogLineTheHistoryTheBanNotesAndTheAlert(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
addr, out, server, queue := startProxyWithAlerts(t, app.URL, geojsURL, clk.Now,
map[string]string{
trustedProxies: trustLocalhost,
alertWebhookURL: "https://alerts.example/smallwebwaf",
rateLimitPerMinute: "1",
})
s := &sender{t: t, addr: addr, out: out}
// The answer is kept before the requests, so GeoJS is not asked, and
// gives no answer of its own.
netblock := netip.MustParsePrefix(fromDE + "/32")
server.GeoJS.Load([]lookup.Answer{{
Client: netblock, ASN: asnDE, ASName: asNameDE, Country: "DE",
Answered: clk.Now(), Used: clk.Now(),
}})
want := asnAndCountry{asnDE, asNameDE, "DE"}
for _, line := range []logLine{
s.get(fromDE, http.StatusOK, requestlog.ActionForward),
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited),
} {
if got := (asnAndCountry{line.ASN, line.ASName, line.Country}); got != want {
t.Errorf("log line has %+v, want %+v", got, want)
}
}
h := historyOf(t, server, fromDE)
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != want ||
!h.LookedUp.Equal(clk.Now()) {
t.Errorf("history has %+v, looked up at %s; want %+v, at %s",
got, h.LookedUp, want, clk.Now())
}
notes := server.Ledger.Bans(netblock)[0].Notes
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != want {
t.Errorf("the ban's notes have %+v, want %+v", got, want)
}
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 {
t.Fatalf("alerts waiting %+v, want the ban's alone", waiting)
}
alert := waiting[0]
if got := (asnAndCountry{alert.ASN, alert.ASName, alert.Country}); got != want {
t.Errorf("the ban's alert has %+v, want %+v", got, want)
}
}
func TestLookupSourceOffLooksNoClientUp(t *testing.T) {
t.Parallel()
geojsURL, asked := startGeoJS(t)
s, clk, server := startWithClock(t, geojsURL, map[string]string{lookupSource: "off"})
// Even an answer kept from before is not used.
server.GeoJS.Load([]lookup.Answer{{
Client: netip.MustParsePrefix(fromDE + "/32"), ASN: asnDE, ASName: asNameDE,
Country: "DE", Answered: clk.Now(), Used: clk.Now(),
}})
for _, from := range []string{fromDE, fromKP} {
line := s.get(from, http.StatusOK, requestlog.ActionForward)
got := asnAndCountry{line.ASN, line.ASName, line.Country}
if got != (asnAndCountry{}) {
t.Errorf("log line has %+v, want nothing", got)
}
}
if h := historyOf(t, server, fromDE); h.ASN != "" || !h.LookedUp.IsZero() {
t.Errorf("history has %q, looked up at %s, want no lookup", h.ASN, h.LookedUp)
}
if len(asked()) != 0 {
t.Errorf("GeoJS was asked about %v, want nothing", asked())
}
}
func TestClientsAreLookedUpInTheLookupDatabaseAndGeoJSIsNotAsked(t *testing.T) {
t.Parallel()
geojsURL, asked := startGeoJS(t)
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
lookuptest.Write(t, path, map[string]lookuptest.Network{
fromDE + "/32": {ASN: asnDE, ASName: asNameDE, Country: "DE"},
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
})
s, clk, server := startWithClock(t, geojsURL, map[string]string{
lookupSource: fileSource,
lookupDBPath: path,
allowedCountries: "DE",
rateLimitPerMinute: "1",
})
// fromDE's second request breaks the limit and bans it. The list
// refuses fromKP, and unplaced, which the file does not hold.
de := asnAndCountry{asnDE, asNameDE, "DE"}
for _, tc := range []struct {
line logLine
want asnAndCountry
}{
{s.get(fromDE, http.StatusOK, requestlog.ActionForward), de},
{s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited), de},
{
s.get(fromKP, http.StatusForbidden, requestlog.ActionCountryDenied),
asnAndCountry{asnKP, asNameKP, "KP"},
},
{
s.get(unplaced, http.StatusForbidden, requestlog.ActionCountryDenied),
asnAndCountry{},
},
} {
got := asnAndCountry{tc.line.ASN, tc.line.ASName, tc.line.Country}
if got != tc.want {
t.Errorf("log line has %+v, want %+v", got, tc.want)
}
}
h := historyOf(t, server, fromDE)
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != de ||
!h.LookedUp.Equal(clk.Now()) {
t.Errorf("history has %+v, looked up at %s; want %+v, at %s",
got, h.LookedUp, de, clk.Now())
}
notes := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32"))[0].Notes
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != de {
t.Errorf("the ban's notes have %+v, want %+v", got, de)
}
if len(asked()) != 0 {
t.Errorf("GeoJS was asked about %v, want nothing", asked())
}
}
func TestLookupHeadersArePassedToTheAppAndTheClientsOwnRemoved(t *testing.T) {
t.Parallel()
var (
mu sync.Mutex
got [][2][]string // each request's X-Client-ASN and X-Client-Country
)
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
mu.Lock()
defer mu.Unlock()
got = append(got, [2][]string{
r.Header.Values("X-Client-Asn"), r.Header.Values("X-Client-Country"),
})
})
geojsURL, _ := startGeoJS(t)
addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{
trustedProxies: trustLocalhost,
addLookupHeaders: "true",
})
s := &sender{t: t, addr: addr, out: out}
// Each client sends headers of its own. fromDE's first request waits
// for its answer, which the app is passed; unplaced has none to pass,
// and a client on a private address is not looked up.
for _, from := range []string{fromDE, unplaced, "10.0.0.8"} {
s.requestWithHeader(from, "/", clientsOwnLookupHeaders,
http.StatusOK, requestlog.ActionForward)
}
mu.Lock()
defer mu.Unlock()
want := [][2][]string{{{asnDE}, {"DE"}}, {nil, nil}, {nil, nil}}
if !slices.EqualFunc(got, want, func(a, b [2][]string) bool {
return slices.Equal(a[0], b[0]) && slices.Equal(a[1], b[1])
}) {
t.Errorf("the app was passed %v, want %v", got, want)
}
}
func TestClientsOwnLookupHeadersAreRemovedWhileTheSettingIsOff(t *testing.T) {
t.Parallel()
var (
mu sync.Mutex
asn, country []string
)
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
mu.Lock()
defer mu.Unlock()
asn, country = r.Header.Values("X-Client-Asn"), r.Header.Values("X-Client-Country")
})
geojsURL, _ := startGeoJS(t)
addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{
trustedProxies: trustLocalhost,
})
s := &sender{t: t, addr: addr, out: out}
s.requestWithHeader(fromDE, "/", clientsOwnLookupHeaders,
http.StatusOK, requestlog.ActionForward)
mu.Lock()
defer mu.Unlock()
if asn != nil || country != nil {
t.Errorf("the app was passed X-Client-ASN %v and X-Client-Country %v, want neither",
asn, country)
}
}
func TestRequestWaitsAsLongAsTheLookupTimeoutSays(t *testing.T) {
t.Parallel()
// The test runs in a synctest bubble, where the time package runs on a
// clock of the test's own: the wait lasts exactly as long as it should,
// however slowly the test process runs. Nothing in it may wait on the
// network, which would keep that clock from moving on: the request is
// handed to the proxy's handler, and GeoJS is one that never answers.
synctest.Test(t, func(t *testing.T) {
// Not the default second. The exclusive list needs the answer, and
// the app is never reached.
const timeout = 3 * time.Second
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
time.Now, map[string]string{
lookupTimeout: timeout.String(),
allowedCountries: "DE",
})
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/",
http.NoBody)
req.RemoteAddr = net.JoinHostPort(fromDE, "1234")
began := time.Now()
server.Handler.ServeHTTP(httptest.NewRecorder(), req)
if waited := time.Since(began); waited != timeout {
t.Errorf("the request waited %s for its answer, want %s", waited, timeout)
}
// Without an answer, the client is in no country the list allows.
wantLine(t, out.requestLine(t), http.StatusForbidden,
requestlog.ActionCountryDenied)
})
}
// unansweredGeoJSURL is where a GeoJS that never answers is asked: a
// request to it waits, without the network, until it is abandoned.
// TestMain registers it with Go's default transport, through which GeoJS
// is asked.
const unansweredGeoJSURL = "unanswered://geojs/v1/ip/geo.json"
func TestMain(m *testing.M) {
transport, _ := http.DefaultTransport.(*http.Transport)
transport.RegisterProtocol("unanswered", unansweredGeoJS{})
m.Run()
}
// unansweredGeoJS is the GeoJS at unansweredGeoJSURL.
type unansweredGeoJS struct{}
// RoundTrip waits until req is abandoned.
func (unansweredGeoJS) RoundTrip(req *http.Request) (*http.Response, error) {
<-req.Context().Done()
return nil, req.Context().Err()
}
// clientsOwnLookupHeaders are the X-Client-ASN and X-Client-Country a
// client sends of its own, each twice, in two cases.
const clientsOwnLookupHeaders = "X-Client-ASN: AS1\r\nx-client-asn: AS2\r\n" +
"X-CLIENT-COUNTRY: KP\r\nx-client-country: CN"
// waitUntil waits until done reports true, for at most waitLimit.
func waitUntil(done func() bool) {
deadline := time.Now().Add(waitLimit)
for !done() && time.Now().Before(deadline) {
time.Sleep(pollInterval)
}
}
+45 -106
View File
@@ -148,8 +148,8 @@ func TestMetricsCountTheTraffic(t *testing.T) {
wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound) wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound)
out.requestLines(t, 2) out.requestLines(t, 2)
forward := `{action="forward",instance="app",status_class="2xx"}` forward := `{action="forward",status_class="2xx"}`
notFound := `{action="admin",instance="app",status_class="4xx"}` notFound := `{action="admin",status_class="4xx"}`
// The request for the metrics is itself under way. // The request for the metrics is itself under way.
metrics := scrape(t, addr) metrics := scrape(t, addr)
@@ -159,13 +159,11 @@ func TestMetricsCountTheTraffic(t *testing.T) {
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+forward, 5) wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+forward, 5)
wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+notFound, wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+notFound,
float64(len("Not Found\n"))) float64(len("Not Found\n")))
wantMetric(t, metrics, wantMetric(t, metrics, "smallwebwaf_request_duration_seconds_count", 2)
`smallwebwaf_request_duration_seconds_count{instance="app"}`, 2) wantMetric(t, metrics, "smallwebwaf_upstream_duration_seconds_count", 1)
wantMetric(t, metrics, wantMetric(t, metrics, "smallwebwaf_requests_in_flight", 1)
`smallwebwaf_upstream_duration_seconds_count{instance="app"}`, 1) metric(t, metrics, "go_goroutines")
wantMetric(t, metrics, `smallwebwaf_requests_in_flight{instance="app"}`, 1) metric(t, metrics, "process_start_time_seconds")
metric(t, metrics, `go_goroutines{instance="app"}`)
metric(t, metrics, `process_start_time_seconds{instance="app"}`)
// A request the app holds is under way until it ends. // A request the app holds is under way until it ends.
httpClient := newClient(t) httpClient := newClient(t)
@@ -182,7 +180,7 @@ func TestMetricsCountTheTraffic(t *testing.T) {
}() }()
<-arrived <-arrived
wantMetric(t, scrape(t, addr), `smallwebwaf_requests_in_flight{instance="app"}`, 2) wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 2)
releaseApp() releaseApp()
err := <-ended err := <-ended
@@ -191,7 +189,7 @@ func TestMetricsCountTheTraffic(t *testing.T) {
} }
out.requestLines(t, 5) out.requestLines(t, 5)
wantMetric(t, scrape(t, addr), `smallwebwaf_requests_in_flight{instance="app"}`, 1) wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 1)
} }
func TestMetricsCountLimitsAndBans(t *testing.T) { func TestMetricsCountLimitsAndBans(t *testing.T) {
@@ -221,16 +219,15 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
metrics := s.scrape(scraper) metrics := s.scrape(scraper)
wantMetric(t, metrics, wantMetric(t, metrics,
`smallwebwaf_requests_total{action="denied",instance="app",status_class="none"}`, 1) `smallwebwaf_requests_total{action="denied",status_class="none"}`, 1)
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+ wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 1)
`kind="requests",window="minute"}`, 1) wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 1)
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 1) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_permanent_bans", 0)
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 0)
clk.advance(time.Hour) clk.advance(time.Hour)
wantMetric(t, s.scrape(scraper), `smallwebwaf_active_bans{instance="app"}`, 0) wantMetric(t, s.scrape(scraper), "smallwebwaf_active_bans", 0)
// A limit broken again right after would ban for three hours, longer // A limit broken again right after would ban for three hours, longer
// than SWWAF_MAX_BAN_DURATION, so the ban is permanent. // than SWWAF_MAX_BAN_DURATION, so the ban is permanent.
@@ -238,14 +235,13 @@ func TestMetricsCountLimitsAndBans(t *testing.T) {
s.get(client, 0, requestlog.ActionRateLimited) s.get(client, 0, requestlog.ActionRateLimited)
metrics = s.scrape(scraper) metrics = s.scrape(scraper)
wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{instance="app",`+ wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 2)
`kind="requests",window="minute"}`, 2) wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 2)
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 2) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 2)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 2) wantMetric(t, metrics, "smallwebwaf_active_bans", 1)
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 1)
// denied, client, and the scraper as of its earlier requests. // denied, client, and the scraper as of its earlier requests.
wantMetric(t, metrics, `smallwebwaf_tracked_clients{instance="app"}`, 3) wantMetric(t, metrics, "smallwebwaf_tracked_clients", 3)
} }
func TestMetricsCountTheBansAnAdminMakes(t *testing.T) { func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
@@ -259,7 +255,7 @@ func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
rateLimitExemptNets: scraper, rateLimitExemptNets: scraper,
}) })
const admins = `smallwebwaf_bans_made_total{cause="admin",instance="app"}` const admins = `smallwebwaf_bans_made_total{cause="admin"}`
wantMetric(t, s.scrape(scraper), admins, 0) wantMetric(t, s.scrape(scraper), admins, 0)
@@ -291,8 +287,7 @@ func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
metricsTopN: "2", metricsTopN: "2",
deniedCountries: "kp", deniedCountries: "kp",
} }
geojsURL, _ := startGeoJS(t) addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, env)
addr, out, server := startProxyWithClock(t, app.URL, geojsURL, time.Now, env)
// The answers are kept before the requests, so that none waits for // The answers are kept before the requests, so that none waits for
// GeoJS. // GeoJS.
@@ -324,82 +319,28 @@ func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
metrics := scrape(t, addr) metrics := scrape(t, addr)
lines++ lines++
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
`smallwebwaf_country_requests_total{country="KP",instance="app"}`, 3) wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`, 2)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 1)
`smallwebwaf_country_requests_total{country="DE",instance="app"}`, 2) wantMetric(t, metrics, `smallwebwaf_country_list_refusals_total{country="KP"}`, 3)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="KP"}`, 0)
`smallwebwaf_country_requests_total{country="other",instance="app"}`, 1) wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`, 6)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="KP"}`,
`smallwebwaf_country_list_refusals_total{country="KP",instance="app"}`, 3)
wantMetric(t, metrics,
`smallwebwaf_country_request_bytes_total{country="KP",instance="app"}`, 0)
wantMetric(t, metrics,
`smallwebwaf_country_request_bytes_total{country="DE",instance="app"}`, 6)
wantMetric(t, metrics,
`smallwebwaf_country_response_bytes_total{country="KP",instance="app"}`,
float64(3*len("Forbidden\n"))) float64(3*len("Forbidden\n")))
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="other"}`,
`smallwebwaf_country_response_bytes_total{country="other",instance="app"}`,
float64(len("hello"))) float64(len("hello")))
wantNoSeries(t, metrics, wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`)
`smallwebwaf_country_requests_total{country="FR",instance="app"}`)
// Once FR is busier than DE, it takes DE's place: its series counts // Once FR is busier than DE, it takes DE's place: its series counts
// from then on, and DE's is gone. // from then on, and DE's is gone.
send(fromFR, 3, http.StatusOK) send(fromFR, 3, http.StatusOK)
metrics = scrape(t, addr) metrics = scrape(t, addr)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3)
`smallwebwaf_country_requests_total{country="KP",instance="app"}`, 3) wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`, 2)
wantMetric(t, metrics, wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 2)
`smallwebwaf_country_requests_total{country="FR",instance="app"}`, 2) wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`)
wantMetric(t, metrics, wantNoSeries(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`)
`smallwebwaf_country_requests_total{country="other",instance="app"}`, 2)
wantNoSeries(t, metrics,
`smallwebwaf_country_requests_total{country="DE",instance="app"}`)
wantNoSeries(t, metrics,
`smallwebwaf_country_request_bytes_total{country="DE",instance="app"}`)
}
func TestMetricsByASNumberKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
t.Parallel()
geojsURL, _ := startGeoJS(t)
s, clk, server := startWithClock(t, geojsURL, map[string]string{
metricsToken: token,
metricsTopN: "1",
})
// The answers are kept before the requests, so that GeoJS gives none
// of its own. Each client is in an AS of its own.
answer := func(addr, asn string) lookup.Answer {
return lookup.Answer{
Client: netip.MustParsePrefix(addr + "/32"), ASN: asn,
Answered: clk.Now(), Used: clk.Now(),
}
}
server.GeoJS.Load([]lookup.Answer{
answer(fromDE, "AS64501"), answer(fromKP, "AS64502"),
})
// With one AS number of its own, the other is counted as other. The
// metrics are asked for from a private address, which has no AS number.
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
s.get(fromKP, http.StatusOK, requestlog.ActionForward)
metrics := s.scrape("10.0.0.9")
wantMetric(t, metrics,
`smallwebwaf_asn_requests_total{asn="AS64501",instance="app"}`, 2)
wantMetric(t, metrics,
`smallwebwaf_asn_requests_total{asn="other",instance="app"}`, 1)
wantMetric(t, metrics,
`smallwebwaf_asn_request_bytes_total{asn="AS64501",instance="app"}`, 0)
wantMetric(t, metrics,
`smallwebwaf_asn_response_bytes_total{asn="other",instance="app"}`, 0)
wantNoSeries(t, metrics,
`smallwebwaf_asn_requests_total{asn="AS64502",instance="app"}`)
} }
func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) { func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) {
@@ -429,16 +370,16 @@ func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) {
deadline := time.Now().Add(waitLimit) deadline := time.Now().Add(waitLimit)
metrics := scrape(t, addr) metrics := scrape(t, addr)
for metric(t, metrics, `smallwebwaf_geojs_failures_total{instance="app"}`) == 0 && for metric(t, metrics, "smallwebwaf_geojs_failures_total") == 0 &&
time.Now().Before(deadline) { time.Now().Before(deadline) {
time.Sleep(pollInterval) time.Sleep(pollInterval)
metrics = scrape(t, addr) metrics = scrape(t, addr)
} }
wantMetric(t, metrics, `smallwebwaf_geojs_requests_total{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_geojs_requests_total", 1)
wantMetric(t, metrics, `smallwebwaf_geojs_failures_total{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_geojs_failures_total", 1)
wantMetric(t, metrics, `smallwebwaf_geojs_unanswered_total{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_geojs_unanswered_total", 1)
} }
// keptAnswer returns GeoJS's answer that the client at addr is in // keptAnswer returns GeoJS's answer that the client at addr is in
@@ -481,9 +422,8 @@ func (s *sender) scrape(from string) string {
// metric returns the value of series in metrics, which are in the // metric returns the value of series in metrics, which are in the
// Prometheus text format. series is a name and its labels in the order of // Prometheus text format. series is a name and its labels in the order of
// their names, such as // their names, such as smallwebwaf_offences_total{kind="limit"}. It fails
// smallwebwaf_offences_total{instance="app",kind="limit"}. It fails the // the test if there is no such series.
// test if there is no such series.
func metric(t *testing.T, metrics, series string) float64 { func metric(t *testing.T, metrics, series string) float64 {
t.Helper() t.Helper()
@@ -531,8 +471,7 @@ func wantNoSeries(t *testing.T, metrics, series string) {
func wantLimitHits(t *testing.T, addr, limit string, hits int) { func wantLimitHits(t *testing.T, addr, limit string, hits int) {
t.Helper() t.Helper()
series := `smallwebwaf_size_and_time_limit_hits_total{instance="app",limit="` + series := `smallwebwaf_size_and_time_limit_hits_total{limit="` + limit + `"}`
limit + `"}`
metrics := scrape(t, addr) metrics := scrape(t, addr)
if hits == 0 { if hits == 0 {
+5 -7
View File
@@ -6,6 +6,7 @@ import (
"errors" "errors"
"io" "io"
"net/http" "net/http"
"os"
"reflect" "reflect"
"slices" "slices"
"strings" "strings"
@@ -122,10 +123,10 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) { func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
t.Helper() t.Helper()
bytes := float64(sent + received) hostname, _ := os.Hostname()
want := withTimings(line, requestlog.Line{ want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: "app", Type: requestType, Time: line.Time, Instance: hostname,
ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host, ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
Path: rawPath, Query: rawQuery, Protocol: protocol, Path: rawPath, Query: rawQuery, Protocol: protocol,
Status: http.StatusTeapot, RequestBytes: int64(sent), Status: http.StatusTeapot, RequestBytes: int64(sent),
@@ -133,10 +134,7 @@ func wantRequestFields(t *testing.T, line logLine, host string, sent, received i
RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32", RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8", ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward, UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
Counts: ratelimit.Counts{ Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
Minute: 1, Hour: 1, Day: 1,
MinuteBytes: bytes, HourBytes: bytes, DayBytes: bytes,
},
}) })
if !reflect.DeepEqual(line.Line, want) { if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want) t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
@@ -320,7 +318,7 @@ func TestServerHasTheDefaultLimits(t *testing.T) {
server := proxy.New(proxy.Params{ server := proxy.New(proxy.Params{
Config: cfg, Config: cfg,
RequestLog: io.Discard, RequestLog: io.Discard,
ProcessLog: requestlog.NewProcessLogger(io.Discard, cfg.InstanceName), ProcessLog: requestlog.NewProcessLogger(io.Discard),
}) })
if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 || if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 ||
+21 -100
View File
@@ -11,14 +11,11 @@ import (
"strings" "strings"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/config" "sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
"sneak.berlin/go/smallwebwaf/internal/rules" "sneak.berlin/go/smallwebwaf/internal/rules"
) )
@@ -56,12 +53,9 @@ type Params struct {
RequestLog io.Writer RequestLog io.Writer
// ProcessLog receives the process's own messages. // ProcessLog receives the process's own messages.
ProcessLog *slog.Logger ProcessLog *slog.Logger
// GeoJSURL is where clients' AS numbers and countries are looked up // GeoJSURL is where clients' countries are looked up, normally
// while SWWAF_LOOKUP_SOURCE is geojs, normally lookup.URL. // lookup.URL. GeoJS is asked only while a country list is set.
GeoJSURL string GeoJSURL string
// 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.
@@ -69,28 +63,17 @@ 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, for each count over an anomaly threshold, for each request
// whose client a blocklist or a DNSBL zone lists, and for GeoJS failing,
// a fetch of a list failing or a query to a DNSBL zone 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, the lookup database, nil unless // whose state the state files keep, and the metrics.
// SWWAF_LOOKUP_SOURCE is file, the lists fetched from URLs, which its Run
// fetches, the DNSBL zones' verdicts, and the metrics.
type Server struct { type Server struct {
*http.Server *http.Server
Ledger *bans.Ledger Ledger *bans.Ledger
Limiter *ratelimit.Limiter Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS GeoJS *lookup.GeoJS
Anomalies *anomaly.Counters Metrics *metrics.Metrics
LookupFile *lookup.File
Lists *reputation.Lists
DNSBL *reputation.DNSBL
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
@@ -101,8 +84,7 @@ type Server struct {
// applies the timeouts and size limits from then on. // applies the timeouts and size limits from then on.
func New(params Params) *Server { func New(params Params) *Server {
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn) errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
m := metrics.New(params.Config.MetricsTopN, params.Config.InstanceName) m := metrics.New(params.Config.MetricsTopN)
lists, dnsbl := newReputation(params)
h := &handler{ h := &handler{
config: params.Config, config: params.Config,
requestLog: params.RequestLog, requestLog: params.RequestLog,
@@ -112,12 +94,9 @@ 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,
@@ -126,40 +105,16 @@ func New(params Params) *Server {
AttackBanDuration: params.Config.AttackBanDuration, AttackBanDuration: params.Config.AttackBanDuration,
MaxBans: params.Config.MaxBans, MaxBans: params.Config.MaxBans,
}), }),
anomalies: anomaly.New(anomaly.Params{ geojs: lookup.New(lookup.Params{
Client: params.Config.AnomalyClient, URL: params.GeoJSURL,
Net: params.Config.AnomalyNet, Now: params.Now,
ASN: params.Config.AnomalyASN, ProcessLog: params.ProcessLog,
Total: params.Config.AnomalyTotal, Metrics: m,
Watch: params.Config.AnomalyWatch,
NetV4Prefix: params.Config.AnomalyNetV4Prefix,
NetV6Prefix: params.Config.AnomalyNetV6Prefix,
NamedNetblocks: params.Config.WatchNets,
Alerts: params.Alerts,
}), }),
lookupFile: params.LookupFile, rules: params.Rules,
lists: lists,
dnsbl: dnsbl,
rules: params.Rules,
alerts: params.Alerts,
} }
h.geojs = lookup.New(lookup.Params{
URL: params.GeoJSURL,
Timeout: params.Config.LookupTimeout,
// The country lists, the headers and the biased thresholds act on
// the answer before the request goes on.
Wait: len(params.Config.DeniedCountries) > 0 ||
len(params.Config.ExclusivelyAllowedCountries) > 0 ||
params.Config.AddLookupHeaders || biasedThresholdsSet(params.Config),
Answered: h.addLookup,
Now: params.Now,
ProcessLog: params.ProcessLog,
Metrics: m,
Alerts: params.Alerts,
})
m.AddBansAndClients(h.ledger, h.limiter, params.Now) m.AddBansAndClients(h.ledger, h.limiter, params.Now)
m.AddRules(params.Rules) m.AddRules(params.Rules)
m.AddReputation(h.lists, h.dnsbl)
return &Server{ return &Server{
Server: &http.Server{ Server: &http.Server{
@@ -174,36 +129,13 @@ func New(params Params) *Server {
MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10), MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
ErrorLog: errorLog, ErrorLog: errorLog,
}, },
Ledger: h.ledger, Ledger: h.ledger,
Limiter: h.limiter, Limiter: h.limiter,
GeoJS: h.geojs, GeoJS: h.geojs,
Anomalies: h.anomalies, Metrics: m,
LookupFile: h.lookupFile,
Lists: h.lists,
DNSBL: h.dnsbl,
Metrics: m,
} }
} }
// newReputation returns the lists fetched from URLs and the DNSBL zones'
// verdicts, as the settings in params name them, with none fetched or
// asked for yet.
func newReputation(params Params) (*reputation.Lists, *reputation.DNSBL) {
cfg := params.Config
lists := reputation.New(reputation.Params{
BlocklistURLs: cfg.BlocklistURLs, Refresh: cfg.BlocklistRefresh,
ASNLimitPercentURL: cfg.ASNLimitPercentURL, Now: params.Now,
ProcessLog: params.ProcessLog, Alerts: params.Alerts,
})
dnsbl := reputation.NewDNSBL(reputation.DNSBLParams{
Zones: cfg.DNSBLZones, Resolver: cfg.DNSBLResolver, CacheTTL: cfg.ReputationCacheTTL,
Timeout: cfg.ReputationTimeout, Now: params.Now, ProcessLog: params.ProcessLog,
Alerts: params.Alerts,
})
return lists, dnsbl
}
// handler is the proxy. It holds what every request shares; what belongs // handler is the proxy. It holds what every request shares; what belongs
// to one request is in a request. // to one request is in a request.
type handler struct { type handler struct {
@@ -217,12 +149,7 @@ type handler struct {
limiter *ratelimit.Limiter limiter *ratelimit.Limiter
ledger *bans.Ledger ledger *bans.Ledger
geojs *lookup.GeoJS geojs *lookup.GeoJS
anomalies *anomaly.Counters
lookupFile *lookup.File
lists *reputation.Lists
dnsbl *reputation.DNSBL
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
@@ -259,7 +186,6 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
// Once the request has ended, before its log line is written. // Once the request has ended, before its log line is written.
defer rq.addToHistory() defer rq.addToHistory()
defer rq.countAnomalies()
refused := rq.check(r.Context()) refused := rq.check(r.Context())
rq.checked = time.Now() rq.checked = time.Now()
@@ -278,10 +204,5 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return return
} }
// Once the response has ended, before the request is added to its
// client's history. Deferred, since ReverseProxy panics to end a
// response it cannot finish.
defer rq.countBytes()
rq.forward(r.Context()) rq.forward(r.Context())
} }
+31 -99
View File
@@ -14,9 +14,7 @@ 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"
@@ -67,10 +65,6 @@ 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"
@@ -208,8 +202,8 @@ func startProxy(t *testing.T, appURL string, env map[string]string) (string, *ou
return startProxyWithGeoJS(t, appURL, "", env) return startProxyWithGeoJS(t, appURL, "", env)
} }
// startProxyWithGeoJS is startProxy with clients' AS numbers and // startProxyWithGeoJS is startProxy with clients' countries looked up at
// countries looked up at geojsURL. // geojsURL.
func startProxyWithGeoJS( func startProxyWithGeoJS(
t *testing.T, appURL, geojsURL string, env map[string]string, t *testing.T, appURL, geojsURL string, env map[string]string,
) (string, *output) { ) (string, *output) {
@@ -222,29 +216,43 @@ func startProxyWithGeoJS(
// startProxyWithClock is startProxyWithGeoJS with requests counted and // startProxyWithClock is startProxyWithGeoJS with requests counted and
// bans made by the time now tells, and returns the server as well. Unless // bans made by the time now tells, and returns the server as well. Unless
// env sets SWWAF_RULES_DIR, it is an empty directory, of no rules, and // env sets SWWAF_RULES_DIR, it is an empty directory, of no rules.
// unless it sets SWWAF_INSTANCE_NAME, that is app, the label instance of
// every metric.
func startProxyWithClock( func startProxyWithClock(
t *testing.T, appURL, geojsURL string, now func() time.Time, t *testing.T, appURL, geojsURL string, now func() time.Time,
env map[string]string, env map[string]string,
) (string, *output, *proxy.Server) { ) (string, *output, *proxy.Server) {
t.Helper() t.Helper()
addr, out, server, _ := startProxyWithAlerts(t, appURL, geojsURL, now, env) settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir()}
maps.Copy(settings, env)
return addr, out, server cfg, err := config.FromEnvironment(func(name string) (string, bool) {
} value, ok := settings[name]
// startProxyWithAlerts is startProxyWithClock, and returns the queue of return value, ok
// the alerts the proxy raises as well, as newProxy makes them. })
func startProxyWithAlerts( if err != nil {
t *testing.T, appURL, geojsURL string, now func() time.Time, t.Fatalf("settings %v: %v", settings, err)
env map[string]string, }
) (string, *output, *proxy.Server, *alerts.Queue) {
t.Helper()
server, out, alertQueue := newProxy(t, appURL, geojsURL, now, env) out := &output{}
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 {
@@ -259,83 +267,7 @@ func startProxyWithAlerts(
_ = server.Close() _ = server.Close()
}) })
return listener.Addr().String(), out, server, alertQueue return listener.Addr().String(), out, server
}
// 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,
+1 -2
View File
@@ -77,8 +77,7 @@ 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
geojsURL, _ := startGeoJS(t) s, _, server := startWithClock(t, "", map[string]string{
s, _, server := startWithClock(t, geojsURL, map[string]string{
rateLimitPerMinute: "1", rateLimitPerMinute: "1",
rateLimitExemptPaths: "/assets/,/favicon.ico", rateLimitExemptPaths: "/assets/,/favicon.ico",
denyNets: denied, denyNets: denied,
-56
View File
@@ -1,56 +0,0 @@
package proxy
import (
"context"
"sneak.berlin/go/smallwebwaf/internal/alerts"
)
// blocklistDenied notes the blocklists that list the client, as
// noteListed does, and reports whether SWWAF_BLOCKLIST_ACTION, being deny,
// refuses the request. Being limit, it lowers the client's limits instead
// (see limitPercentages), and being log, it does nothing more.
func (rq *request) blocklistDenied() bool {
listedBy := rq.h.lists.ListedBy(rq.client)
rq.blocklisted = len(listedBy) > 0
rq.noteListed(listedBy, "listed by a blocklist")
return rq.blocklisted && rq.h.config.BlocklistAction == "deny"
}
// dnsblDenied notes the DNSBL zones whose verdict lists the client, as
// noteListed does, and reports whether SWWAF_REPUTATION_ACTION, being
// deny, refuses the request. Being limit, it lowers the client's limits
// instead (see limitPercentages), and being log, it does nothing more. A
// zone without a verdict on the client is asked about it in the
// background, and the request does not wait for the answer. ctx is the
// request's own context.
func (rq *request) dnsblDenied(ctx context.Context) bool {
listedBy := rq.h.dnsbl.ListedBy(ctx, rq.client)
rq.dnsblListed = len(listedBy) > 0
rq.noteListed(listedBy, "listed by a DNSBL zone")
return rq.dnsblListed && rq.h.config.ReputationAction == "deny"
}
// noteListed adds sources, the URLs of the blocklists or the DNSBL zones,
// their keys masked, that list the client, to the log line's reputation,
// counts each of them in the metrics, and raises a reputation_hit alert,
// with reason, for each.
func (rq *request) noteListed(sources []string, reason string) {
rq.line.Reputation = append(rq.line.Reputation, sources...)
for _, source := range sources {
rq.h.metrics.ReputationHit(source)
rq.h.alerts.Raise(alerts.Alert{
Event: alerts.EventReputationHit,
Client: rq.client,
Netblock: clientGroup(rq.client),
ASN: rq.line.ASN,
ASName: rq.line.ASName,
Country: rq.line.Country,
Reason: reason,
Detail: map[string]any{"source": source},
})
}
}
-660
View File
@@ -1,660 +0,0 @@
package proxy_test
import (
"maps"
"net/http"
"net/netip"
"slices"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// The reputation settings.
const (
blocklistURLs = "SWWAF_BLOCKLIST_URLS"
blocklistAction = "SWWAF_BLOCKLIST_ACTION"
asnLimitPercentURL = "SWWAF_ASN_LIMIT_PERCENT_URL"
)
// The actions of SWWAF_BLOCKLIST_ACTION and SWWAF_REPUTATION_ACTION:
// limitHalf gives a listed client half of every limit, and limitQuarter a
// quarter.
const (
actionDeny = "deny"
actionLog = "log"
limitHalf = "limit:50"
limitQuarter = "limit:25"
)
// The lists these tests name, which are never fetched: each test puts in
// the copies it needs, as reputation.json would at start.
const (
dropURL = "https://lists.example/drop.txt"
torURL = "https://lists.example/tor.txt"
asnURL = "https://lists.example/asn.txt"
)
func TestEachBlocklistActionForAListedAddressAndAListedNetblock(t *testing.T) {
t.Parallel()
forward, denied := requestlog.ActionForward, requestlog.ActionDenied
for _, tc := range []struct {
action string
// statuses and actions are those of a listed client's three
// requests, and percent their limit_percent, as percentText gives it.
statuses []int
actions []string
percent string
}{
{
actionDeny, []int{http.StatusForbidden, http.StatusForbidden, http.StatusForbidden},
[]string{denied, denied, denied}, none,
},
{
// Half of 4 requests a minute: the third breaks the limit.
limitHalf, []int{http.StatusOK, http.StatusOK, http.StatusForbidden},
[]string{forward, forward, requestlog.ActionRateLimited},
"50 from " + blocklistAction,
},
{
actionLog, []int{http.StatusOK, http.StatusOK, http.StatusOK},
[]string{forward, forward, forward}, none,
},
} {
t.Run(tc.action, func(t *testing.T) {
t.Parallel()
s, server, _ := startWithLookups(t, map[string]string{
rateLimitPerMinute: fourAMinute, blocklistURLs: dropURL,
blocklistAction: tc.action,
})
// fromDE is listed as an address, and fromKP in a netblock.
loadLists(t, server, map[string][]string{
dropURL: {"; DROP", fromDE, "198.51.100.0/24 ; SBL1"},
})
for _, from := range []string{fromDE, fromKP} {
for i := range 3 {
line := s.get(from, tc.statuses[i], tc.actions[i])
wantReputation(t, line, dropURL)
wantPercent(t, "limit_percent", line.LimitPercent,
line.LimitPercentSetting, tc.percent)
// A request refused for the list is not counted.
counted := line.fields["counts"] != nil
if counted != (tc.actions[i] != denied) {
t.Errorf("request from %s counted %t, logged %s", from, counted,
tc.actions[i])
}
}
}
// A client no list lists has the whole limit.
for range 3 {
line := s.get(unplaced, http.StatusOK, forward)
wantReputation(t, line)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
none)
}
// A refusal for the list makes no ban.
if held := server.Ledger.Snapshot(); tc.action == actionDeny && len(held) != 0 {
t.Errorf("bans %+v, want none", held)
}
})
}
}
func TestBlocklistsComeAfterTheCountryListsAndSkipAllowNets(t *testing.T) {
t.Parallel()
s, server, queue := startWithLookups(t, map[string]string{
blocklistURLs: dropURL, deniedCountries: "kp", allowNets: fromDE,
})
loadLists(t, server, map[string][]string{dropURL: {fromDE, fromKP}})
// fromKP's country refuses it before the list is looked at, and fromDE,
// in SWWAF_ALLOW_NETS, is not checked at all: neither is noted, nor
// alerted.
wantReputation(t, s.get(fromKP, http.StatusForbidden, requestlog.ActionCountryDenied))
wantReputation(t, s.get(fromDE, http.StatusOK, requestlog.ActionForward))
wantAlerts(t, queue)
}
func TestObserveModeForwardsAClientABlocklistDeniesAndAlertsIt(t *testing.T) {
t.Parallel()
s, server, queue := startWithLookups(t, map[string]string{
blocklistURLs: dropURL, mode: observe,
})
loadLists(t, server, map[string][]string{dropURL: {fromDE}})
line := s.get(fromDE, http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionDenied)
wantReputation(t, line, dropURL)
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || waiting[0].Event != alerts.EventReputationHit {
t.Errorf("alerts waiting %+v, want a reputation_hit alert", waiting)
}
}
func TestBlocklistLimitTakesPartInTheLowestPercentageOfEveryLimit(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
action, asnPercent string
// want is the upload's limit_percent and bytes_percent, as
// percentText gives them, and limitHit its limit_hit.
want, limitHit string
}{
{limitHalf, asnDEQuarter, "25 from " + asnLimitPercent, minuteBytes},
{limitQuarter, asnDEHalf, "25 from " + blocklistAction, minuteBytes},
// The AS number's, the first of two alike.
{limitQuarter, asnDEQuarter, "25 from " + asnLimitPercent, minuteBytes},
{actionLog, asnDE + ":100", none, ""},
} {
t.Run(tc.action+" "+tc.asnPercent, func(t *testing.T) {
t.Parallel()
s, server, _ := startWithLookups(t, map[string]string{
bytesLimitPerMinute: twoUploads, blocklistURLs: dropURL,
blocklistAction: tc.action, asnLimitPercent: tc.asnPercent,
})
loadLists(t, server, map[string][]string{dropURL: {fromDE}})
// The upload's 100 bytes are over 49, a quarter of 199, and 99,
// half of it, and within 199.
line := s.uploadFrom(fromDE)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
tc.want)
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
tc.want)
if line.LimitHit != tc.limitHit {
t.Errorf("log line has limit_hit %q, want %q", line.LimitHit, tc.limitHit)
}
})
}
}
func TestASNLimitPercentFileCountsAsTheSettingDoesTheLowerWinning(t *testing.T) {
t.Parallel()
const (
fromURL = "25 from " + asnLimitPercentURL
fromSetting = "25 from " + asnLimitPercent
)
for _, tc := range []struct {
name string
env map[string]string
file string
// limitPercent and bytesPercent are the upload's, as percentText
// gives them.
limitPercent, bytesPercent string
}{
{"the file's alone", nil, asnDEQuarter, fromURL, fromURL},
{
"the file's, lower than the setting's",
map[string]string{asnLimitPercent: asnDEHalf}, asnDEQuarter, fromURL, fromURL,
},
{
"the setting's, lower than the file's",
map[string]string{asnLimitPercent: asnDEQuarter}, asnDEHalf,
fromSetting, fromSetting,
},
{
"the setting's, the first of two alike",
map[string]string{asnLimitPercent: asnDEQuarter}, asnDEQuarter,
fromSetting, fromSetting,
},
{"none, for an AS number the file does not list", nil, asnKP + ":25", none, none},
{
"SWWAF_ASN_BYTES_PERCENT's in place of the file's for the byte limits",
map[string]string{asnBytesPercent: asnDE + ":100"}, asnDEQuarter, fromURL, none,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
env := map[string]string{asnLimitPercentURL: asnURL}
maps.Copy(env, tc.env)
s, server, _ := startWithLookups(t, env)
loadLists(t, server, map[string][]string{asnURL: {"# by AS number", tc.file}})
line := s.uploadFrom(fromDE)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
tc.limitPercent)
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
tc.bytesPercent)
})
}
}
func TestEachBlocklistThatListsAClientRaisesAnAlertOncePerCooldownAndIsCounted(
t *testing.T,
) {
t.Parallel()
const emptyURL = "https://lists.example/empty.txt"
s, clk, server, queue := startWithLookupsAndClock(t, map[string]string{
blocklistURLs: dropURL + "," + torURL + "," + emptyURL,
blocklistAction: actionLog,
metricsToken: token,
})
loadLists(t, server, map[string][]string{dropURL: {fromDE}, torURL: {fromDE}})
// The second request's alerts are repeats, which the cooldown holds
// back.
for range 2 {
wantReputation(t, s.get(fromDE, http.StatusOK, requestlog.ActionForward),
dropURL, torURL)
}
hit := func(source string) alerts.Alert {
return alerts.Alert{
Instance: alertInstance,
Time: clk.Now(),
Event: alerts.EventReputationHit,
Client: netip.MustParseAddr(fromDE),
Netblock: netip.MustParsePrefix(fromDE + "/32"),
ASN: asnDE,
ASName: asNameDE,
Country: "DE",
Reason: "listed by a blocklist",
Detail: map[string]any{"source": source},
}
}
wantAlerts(t, queue, hit(dropURL), hit(torURL))
if queue.Suppressed() != 2 {
t.Errorf("%d alerts held back, want the second request's 2", queue.Suppressed())
}
// Each list's hits, none of its fetches failed, and when its copy was
// fetched, 0 for the one without.
metrics := s.scrape(unplaced)
fetched := float64(listsFetched().Unix())
for listURL, want := range map[string]struct{ hits, fetched float64 }{
dropURL: {2, fetched}, torURL: {2, fetched}, emptyURL: {0, 0},
} {
labels := `{instance="` + alertInstance + `",source="` + listURL + `"}`
if want.hits == 0 {
wantNoSeries(t, metrics, "smallwebwaf_reputation_hits_total"+labels)
} else {
wantMetric(t, metrics, "smallwebwaf_reputation_hits_total"+labels, want.hits)
}
wantMetric(t, metrics, "smallwebwaf_reputation_failures_total"+labels, 0)
wantMetric(t, metrics, "smallwebwaf_reputation_last_fetch_timestamp_seconds"+labels,
want.fetched)
}
}
// The DNSBL settings.
const (
dnsblZones = "SWWAF_DNSBL_ZONES"
dnsblResolver = "SWWAF_DNSBL_RESOLVER"
reputationAction = "SWWAF_REPUTATION_ACTION"
)
// The DNSBL zones these tests name, which are never asked about the
// clients the tests send requests from: each test puts in the verdicts it
// needs, as reputation.json would at start. A query a test does start is
// sent to noResolver, where nothing listens, so that none leaves the host.
const (
dnsblZone = "dnsbl.example"
otherZone = "other.example"
noResolver = "127.0.0.1:9"
)
func TestEachReputationActionForAClientADNSBLZoneLists(t *testing.T) {
t.Parallel()
forward, denied := requestlog.ActionForward, requestlog.ActionDenied
for _, tc := range []struct {
action string
// statuses and actions are those of a listed client's three
// requests, and percent their limit_percent, as percentText gives it.
statuses []int
actions []string
percent string
}{
{
actionDeny, []int{http.StatusForbidden, http.StatusForbidden, http.StatusForbidden},
[]string{denied, denied, denied}, none,
},
{
// Half of 4 requests a minute: the third breaks the limit.
limitHalf, []int{http.StatusOK, http.StatusOK, http.StatusForbidden},
[]string{forward, forward, requestlog.ActionRateLimited},
"50 from " + reputationAction,
},
{
actionLog, []int{http.StatusOK, http.StatusOK, http.StatusOK},
[]string{forward, forward, forward}, none,
},
} {
t.Run(tc.action, func(t *testing.T) {
t.Parallel()
s, server, _ := startWithLookups(t, map[string]string{
rateLimitPerMinute: fourAMinute, dnsblZones: dnsblZone + "," + otherZone,
dnsblResolver: noResolver, reputationAction: tc.action,
})
listedBy := map[string][]string{
fromDE: {dnsblZone, otherZone}, fromKP: {otherZone}, unplaced: nil,
}
loadVerdicts(server, listedBy)
for _, from := range []string{fromDE, fromKP} {
for i := range 3 {
line := s.get(from, tc.statuses[i], tc.actions[i])
// In the order SWWAF_DNSBL_ZONES names them.
wantReputation(t, line, listedBy[from]...)
wantPercent(t, "limit_percent", line.LimitPercent,
line.LimitPercentSetting, tc.percent)
// A request refused for the verdict is not counted.
counted := line.fields["counts"] != nil
if counted != (tc.actions[i] != denied) {
t.Errorf("request from %s counted %t, logged %s", from, counted,
tc.actions[i])
}
}
}
// A client no zone lists has the whole limit.
for range 3 {
line := s.get(unplaced, http.StatusOK, forward)
wantReputation(t, line)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
none)
}
// A refusal for the verdict makes no ban, and every client had its
// verdicts, so no zone was asked.
if held := server.Ledger.Snapshot(); tc.action == actionDeny && len(held) != 0 {
t.Errorf("bans %+v, want none", held)
}
if queries := server.DNSBL.Queries(dnsblZone); queries != 0 {
t.Errorf("%d queries, want none", queries)
}
})
}
}
func TestDNSBLZonesComeAfterTheBlocklistsAndSkipAllowNets(t *testing.T) {
t.Parallel()
s, server, queue := startWithLookups(t, map[string]string{
blocklistURLs: dropURL, dnsblZones: dnsblZone, dnsblResolver: noResolver,
reputationAction: actionDeny, allowNets: fromDE,
})
loadLists(t, server, map[string][]string{dropURL: {fromKP}})
loadVerdicts(server, map[string][]string{fromKP: {dnsblZone}, fromDE: {dnsblZone}})
// The blocklist refuses fromKP before its verdict is looked at, and
// fromDE, in SWWAF_ALLOW_NETS, is not checked at all: neither is noted
// for the zone, nor alerted, nor asked about.
wantReputation(t, s.get(fromKP, http.StatusForbidden, requestlog.ActionDenied),
dropURL)
wantReputation(t, s.get(fromDE, http.StatusOK, requestlog.ActionForward))
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || waiting[0].Detail["source"] != dropURL {
t.Errorf("alerts waiting %+v, want the blocklist's reputation_hit alone", waiting)
}
if queries := server.DNSBL.Queries(dnsblZone); queries != 0 {
t.Errorf("%d queries, want none", queries)
}
}
func TestObserveModeForwardsAClientADNSBLZoneDeniesAndAlertsIt(t *testing.T) {
t.Parallel()
s, server, queue := startWithLookups(t, map[string]string{
dnsblZones: dnsblZone, dnsblResolver: noResolver, reputationAction: actionDeny,
mode: observe,
})
loadVerdicts(server, map[string][]string{fromDE: {dnsblZone}})
line := s.get(fromDE, http.StatusOK, requestlog.ActionForward)
wantWouldAction(t, line, requestlog.ActionDenied)
wantReputation(t, line, dnsblZone)
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || waiting[0].Event != alerts.EventReputationHit {
t.Errorf("alerts waiting %+v, want a reputation_hit alert", waiting)
}
}
func TestReputationLimitTakesPartInTheLowestPercentageOfEveryLimit(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
blocklistAction, reputationAction string
// want is the request's limit_percent and bytes_percent, as
// percentText gives them.
want string
}{
{limitHalf, limitQuarter, "25 from " + reputationAction},
{limitQuarter, limitHalf, "25 from " + blocklistAction},
// The blocklist's, the first of two alike.
{limitQuarter, limitQuarter, "25 from " + blocklistAction},
{actionLog, limitQuarter, "25 from " + reputationAction},
{actionLog, actionLog, none},
} {
t.Run(tc.blocklistAction+" "+tc.reputationAction, func(t *testing.T) {
t.Parallel()
s, server, _ := startWithLookups(t, map[string]string{
blocklistURLs: dropURL, blocklistAction: tc.blocklistAction,
dnsblZones: dnsblZone, dnsblResolver: noResolver,
reputationAction: tc.reputationAction,
})
loadLists(t, server, map[string][]string{dropURL: {fromDE}})
loadVerdicts(server, map[string][]string{fromDE: {dnsblZone}})
// Named by the blocklist, then by the zone.
line := s.get(fromDE, http.StatusOK, requestlog.ActionForward)
wantReputation(t, line, dropURL, dnsblZone)
wantPercent(t, "limit_percent", line.LimitPercent, line.LimitPercentSetting,
tc.want)
wantPercent(t, "bytes_percent", line.BytesPercent, line.BytesPercentSetting,
tc.want)
})
}
}
func TestEachZoneThatListsAClientRaisesAnAlertOncePerCooldownAndIsCounted(
t *testing.T,
) {
t.Parallel()
s, clk, server, queue := startWithLookupsAndClock(t, map[string]string{
dnsblZones: dnsblZone + "," + otherZone, dnsblResolver: noResolver,
reputationAction: actionLog, metricsToken: token,
})
loadVerdicts(server, map[string][]string{
fromDE: {dnsblZone, otherZone}, unplaced: nil,
})
// The second request's alerts are repeats, which the cooldown holds
// back.
for range 2 {
wantReputation(t, s.get(fromDE, http.StatusOK, requestlog.ActionForward),
dnsblZone, otherZone)
}
hit := func(zone string) alerts.Alert {
return alerts.Alert{
Instance: alertInstance,
Time: clk.Now(),
Event: alerts.EventReputationHit,
Client: netip.MustParseAddr(fromDE),
Netblock: netip.MustParsePrefix(fromDE + "/32"),
ASN: asnDE,
ASName: asNameDE,
Country: "DE",
Reason: "listed by a DNSBL zone",
Detail: map[string]any{"source": zone},
}
}
wantAlerts(t, queue, hit(dnsblZone), hit(otherZone))
if queue.Suppressed() != 2 {
t.Errorf("%d alerts held back, want the second request's 2", queue.Suppressed())
}
// Each zone's hits, and its queries and their failures, none, since
// every client had its verdicts.
metrics := s.scrape(unplaced)
for _, zone := range []string{dnsblZone, otherZone} {
labels := `{instance="` + alertInstance + `",source="` + zone + `"}`
wantMetric(t, metrics, "smallwebwaf_reputation_hits_total"+labels, 2)
wantMetric(t, metrics, "smallwebwaf_reputation_queries_total"+labels, 0)
wantMetric(t, metrics, "smallwebwaf_reputation_failures_total"+labels, 0)
}
}
func TestZoneKeyIsMaskedInTheLogTheAlertAndTheMetrics(t *testing.T) {
t.Parallel()
const (
key = "abcdefghijklmnopqrstuvwxyz"
keyed = key + ".xbl.dq.spamhaus.net"
masked = "********.xbl.dq.spamhaus.net"
)
s, server, queue := startWithLookups(t, map[string]string{
dnsblZones: keyed, dnsblResolver: noResolver, reputationAction: actionLog,
metricsToken: token,
})
loadVerdicts(server, map[string][]string{fromDE: {keyed}, unplaced: nil})
wantReputation(t, s.get(fromDE, http.StatusOK, requestlog.ActionForward), masked)
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || waiting[0].Detail["source"] != masked {
t.Errorf("alerts waiting %+v, want a reputation_hit alert from %s", waiting,
masked)
}
metrics := s.scrape(unplaced)
wantMetric(t, metrics, `smallwebwaf_reputation_hits_total{instance="`+
alertInstance+`",source="`+masked+`"}`, 1)
for name, shown := range map[string]string{
"the log": s.out.text(), "the metrics": metrics,
} {
if strings.Contains(shown, key) {
t.Errorf("%s shows the key:\n%s", name, shown)
}
}
}
func TestRequestFromAClientWithoutAVerdictHasTheZoneAskedAboutIt(t *testing.T) {
t.Parallel()
s, server, _ := startWithLookups(t, map[string]string{
dnsblZones: dnsblZone + "," + otherZone, dnsblResolver: noResolver,
})
server.DNSBL.Load([]reputation.Verdict{{
Zone: otherZone, Client: netip.MustParseAddr(fromDE), Listed: true,
Fetched: verdictsFetched(),
}})
// The verdict of the other zone is used, and dnsbl.example, which has
// none, is asked about the client in the background, once: the second
// request finds the query under way, or the zone left alone after it
// failed, since nothing answers at noResolver.
for range 2 {
wantReputation(t, s.get(fromDE, http.StatusOK, requestlog.ActionForward), otherZone)
}
if queries := server.DNSBL.Queries(dnsblZone); queries != 1 {
t.Errorf("%d queries to %s, want 1", queries, dnsblZone)
}
if queries := server.DNSBL.Queries(otherZone); queries != 0 {
t.Errorf("%d queries to %s, want none", queries, otherZone)
}
}
// listsFetched is when loadLists has the copies fetched.
func listsFetched() time.Time {
return time.Date(2026, 10, 5, 0, 0, 0, 0, time.UTC)
}
// verdictsFetched is when loadVerdicts has the verdicts fetched: half a
// day before the time the tests' clock is set to, so that they are in use
// until half a day later.
func verdictsFetched() time.Time {
return time.Date(2026, 10, 5, 12, 0, 0, 0, time.UTC)
}
// loadVerdicts puts into server's DNSBL, for each client listedBy names,
// a verdict of each zone SWWAF_DNSBL_ZONES names, fetched at
// verdictsFetched, as reputation.json would at start: one that lists the
// client from each zone listedBy gives for it, and one that does not from
// each other zone.
func loadVerdicts(server *proxy.Server, listedBy map[string][]string) {
verdicts := make([]reputation.Verdict, 0, len(listedBy)*len(server.DNSBL.Zones()))
for client, zones := range listedBy {
for _, zone := range server.DNSBL.Zones() {
verdicts = append(verdicts, reputation.Verdict{
Zone: zone, Client: netip.MustParseAddr(client),
Listed: slices.Contains(zones, zone), Fetched: verdictsFetched(),
})
}
}
server.DNSBL.Load(verdicts)
}
// loadLists puts copies of lists into server's lists, by URL, each with
// its lines, fetched at listsFetched, as reputation.json would at start.
func loadLists(t *testing.T, server *proxy.Server, copies map[string][]string) {
t.Helper()
lists := make([]reputation.List, 0, len(copies))
for listURL, lines := range copies {
lists = append(lists, reputation.List{
URL: listURL, Fetched: listsFetched(), Lines: lines,
})
}
err := server.Lists.Load(lists)
if err != nil {
t.Fatalf("load the lists: %v", err)
}
}
// wantReputation checks the URLs of the blocklists the log line names in
// its reputation.
func wantReputation(t *testing.T, line logLine, want ...string) {
t.Helper()
if !slices.Equal(line.Reputation, want) {
t.Errorf("log line has reputation %v, want %v", line.Reputation, want)
}
}
+27 -138
View File
@@ -3,7 +3,6 @@ package proxy
import ( import (
"context" "context"
"errors" "errors"
"io"
"net/http" "net/http"
"net/http/httptrace" "net/http/httptrace"
"net/http/httputil" "net/http/httputil"
@@ -16,9 +15,6 @@ import (
"sync/atomic" "sync/atomic"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog" "sneak.berlin/go/smallwebwaf/internal/requestlog"
) )
@@ -52,21 +48,7 @@ type request struct {
client netip.Addr client netip.Addr
peer netip.Addr peer netip.Addr
peerTrusted bool peerTrusted bool
// lookedUp is true once the client's AS number and country have been start time.Time
// looked up, whether or not an answer was there, and lookupAnswer is
// what the lookup gave then, the zero Answer while GeoJS had given none.
lookedUp bool
lookupAnswer lookup.Answer
// counted is true for a request the rate limits counted, whose bytes
// the byte limits count once it has ended. limitPercent and
// bytesPercent are then its client's limit percentages for the rate
// limits and for the byte limits.
counted bool
limitPercent, bytesPercent percentage
// blocklisted is true once a blocklist is found to list the client,
// and dnsblListed once a DNSBL zone's verdict is.
blocklisted, dnsblListed bool
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
@@ -77,9 +59,6 @@ type request struct {
refused atomic.Pointer[refusal] refused atomic.Pointer[refusal]
// complete is true once the app's whole answer has been passed on. // complete is true once the app's whole answer has been passed on.
complete bool complete bool
// upgraded is the connection to the app once the app has switched
// protocols, as for a WebSocket, and nil otherwise.
upgraded *upgradedConn
// mu guards what follows. The timeouts run on goroutines of their // mu guards what follows. The timeouts run on goroutines of their
// own, and the transport starts and stops them, and notes the times // own, and the transport starts and stops them, and notes the times
@@ -213,17 +192,14 @@ func (rq *request) check(ctx context.Context) *refusal {
// checkClient runs the checks on the request's client, and returns the // checkClient runs the checks on the request's client, and returns the
// action of the first that refuses the request, or "" when none does. A // action of the first that refuses the request, or "" when none does. A
// client in SWWAF_ALLOW_NETS skips them, and is not looked up. For any // client in SWWAF_ALLOW_NETS skips them. For any other client,
// other client, SWWAF_DENY_NETS comes first, then a ban on its netblock, // SWWAF_DENY_NETS comes first, then a ban on its netblock, so that a
// so that a client either refuses is not looked up, then the lookup of // client either refuses is not looked up, and then the country lists; a
// its AS number and country, then the country lists, then the blocklists, // request any of them refuses is not counted for the rate limits. Then
// and then the DNSBL zones' verdicts; a request any of them refuses is not // come the rate limits, unless the client is in
// counted for the rate limits. Then come the rate limits, unless the // SWWAF_RATE_LIMIT_EXEMPT_NETS or the request's path is exempt under
// client is in SWWAF_RATE_LIMIT_EXEMPT_NETS or the request's path is // SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that every other request is counted,
// exempt under SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that every other request // and last the rule files. ctx is the request's own context.
// 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) {
@@ -240,29 +216,13 @@ func (rq *request) checkClient(ctx context.Context) string {
return requestlog.ActionBanned return requestlog.ActionBanned
} }
rq.lookUp(ctx) if rq.countryDenied(ctx) {
if rq.countryDenied() {
return requestlog.ActionCountryDenied return requestlog.ActionCountryDenied
} }
if rq.blocklistDenied() { exempt := isInside(rq.client, cfg.RateLimitExemptNets) ||
return requestlog.ActionDenied pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
} if !exempt && rq.limitBroken(now) {
if rq.dnsblDenied(ctx) {
return requestlog.ActionDenied
}
rq.counted = !isInside(rq.client, cfg.RateLimitExemptNets) &&
!pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
if rq.counted {
rq.limitPercent, rq.bytesPercent = rq.limitPercentages()
rq.line.LimitPercent, rq.line.LimitPercentSetting = rq.limitPercent.logged()
rq.line.BytesPercent, rq.line.BytesPercentSetting = rq.bytesPercent.logged()
}
if rq.counted && rq.limitBroken(now) {
return requestlog.ActionRateLimited return requestlog.ActionRateLimited
} }
@@ -327,9 +287,7 @@ func (rq *request) forward(ctx context.Context) {
// rewrite makes the request the app receives: the client's request, // rewrite makes the request the app receives: the client's request,
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and // unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
// the request's id set, without any X-Client-ASN or X-Client-Country the // the request's id set.
// client sent, whatever SWWAF_ADD_LOOKUP_HEADERS says, and, while it is
// set, with the client's AS number and country in them.
func (rq *request) rewrite(pr *httputil.ProxyRequest) { func (rq *request) rewrite(pr *httputil.ProxyRequest) {
upstream := rq.h.config.UpstreamURL upstream := rq.h.config.UpstreamURL
pr.Out.URL.Scheme = upstream.Scheme pr.Out.URL.Scheme = upstream.Scheme
@@ -339,12 +297,6 @@ func (rq *request) rewrite(pr *httputil.ProxyRequest) {
pr.Out.URL.RawQuery = pr.In.URL.RawQuery pr.Out.URL.RawQuery = pr.In.URL.RawQuery
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted) setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
pr.Out.Header.Set(requestIDHeader, rq.line.RequestID) pr.Out.Header.Set(requestIDHeader, rq.line.RequestID)
pr.Out.Header.Del(asnHeader)
pr.Out.Header.Del(countryHeader)
if rq.h.config.AddLookupHeaders {
setLookupHeaders(pr.Out.Header, rq.line.ASN, rq.line.Country)
}
} }
// modifyResponse looks at the app's answer before ReverseProxy passes it // modifyResponse looks at the app's answer before ReverseProxy passes it
@@ -355,18 +307,11 @@ func (rq *request) modifyResponse(res *http.Response) error {
if res.StatusCode == http.StatusSwitchingProtocols { if res.StatusCode == http.StatusSwitchingProtocols {
// An upgraded connection, such as a WebSocket, is not cut by the // An upgraded connection, such as a WebSocket, is not cut by the
// timeouts. ReverseProxy writes this answer straight to the // timeouts. ReverseProxy writes this answer straight to the
// connection it takes over, not through rq.out, and then copies // connection it takes over, not through rq.out.
// what passes each way through res.Body, the connection to the app.
rq.stopTimers() rq.stopTimers()
rq.out.status = res.StatusCode rq.out.status = res.StatusCode
rq.line.Websocket = true rq.line.Websocket = true
conn, ok := res.Body.(io.ReadWriteCloser)
if ok {
rq.upgraded = &upgradedConn{ReadWriteCloser: conn}
res.Body = rq.upgraded
}
return nil return nil
} }
@@ -469,7 +414,10 @@ func (rq *request) finish() {
line.ResponseContentType = header.Get("Content-Type") line.ResponseContentType = header.Get("Content-Type")
line.CacheControl = header.Get("Cache-Control") line.CacheControl = header.Get("Cache-Control")
line.Location = header.Get("Location") line.Location = header.Get("Location")
line.RequestBytes = rq.requestBytes()
if rq.body != nil {
line.RequestBytes = rq.body.bytes.Load()
}
// limit is the setting whose size or time limit the request passed. // limit is the setting whose size or time limit the request passed.
var limit string var limit string
@@ -526,83 +474,24 @@ func timing(start, end time.Time) *float64 {
} }
// addToHistory adds the request, which has ended, to its client's // addToHistory adds the request, which has ended, to its client's
// history, and then the lookup's answer about the client, as // history.
// answerAtTheEnd gives it, to that history and to the notes of the bans
// on its netblock: an answer may have come before either was there, and
// one from GeoJS that comes later is added when it comes.
func (rq *request) addToHistory() { func (rq *request) addToHistory() {
var requestBytes int64
if rq.body != nil {
requestBytes = rq.body.bytes.Load()
}
forwarded := !rq.upstreamStart.IsZero() forwarded := !rq.upstreamStart.IsZero()
rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{ rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{
Country: rq.line.Country,
Forwarded: forwarded, Forwarded: forwarded,
Refused: !forwarded && rq.refused.Load() != nil, Refused: !forwarded && rq.refused.Load() != nil,
Status: rq.out.status, Status: rq.out.status,
RequestBytes: rq.requestBytes(), RequestBytes: requestBytes,
ResponseBytes: rq.out.bytes, ResponseBytes: rq.out.bytes,
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit, BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
}) })
answer, found := rq.answerAtTheEnd()
if found {
rq.h.addLookup(answer)
}
}
// countAnomalies counts the request, which has ended, and its bytes, as
// countedBytes gives them, for the anomaly thresholds, whatever was done
// with it: a request refused, one from a client in SWWAF_ALLOW_NETS or
// SWWAF_RATE_LIMIT_EXEMPT_NETS, and one for a path in
// SWWAF_RATE_LIMIT_EXEMPT_PATHS are counted too. It is counted for its
// client's AS number when answerAtTheEnd gives one. With every anomaly
// threshold off, the default, it does nothing.
func (rq *request) countAnomalies() {
if !anomalyThresholdsSet(rq.h.config) {
return
}
answer, _ := rq.answerAtTheEnd()
rq.h.anomalies.Count(rq.h.now(), anomaly.Request{
Client: rq.client,
ClientGroup: clientGroup(rq.client),
ASN: answer.ASN,
ASName: answer.ASName,
Country: answer.Country,
Bytes: rq.countedBytes(),
})
}
// anomalyThresholdsSet reports whether any anomaly threshold is set.
func anomalyThresholdsSet(cfg *config.Config) bool {
off := anomaly.Thresholds{}
return cfg.AnomalyClient != off || cfg.AnomalyNet != off || cfg.AnomalyASN != off ||
cfg.AnomalyTotal != off || cfg.AnomalyWatch != off
}
// answerAtTheEnd returns, for a client that was looked up, the lookup's
// answer about it as the request ends, and whether there is one: the
// lookup database's, which was there at once, or the one GeoJS has given
// by then, which a request does not wait for unless a setting needs it.
func (rq *request) answerAtTheEnd() (lookup.Answer, bool) {
if !rq.lookedUp {
return lookup.Answer{}, false
}
if rq.h.config.LookupSource == "file" {
return rq.lookupAnswer, true
}
return rq.h.geojs.Kept(clientGroup(rq.client))
}
// requestBytes is how many bytes of the request's body have been read.
func (rq *request) requestBytes() int64 {
if rq.body == nil {
return 0
}
return rq.body.bytes.Load()
} }
// clientRequestDeadline is when the client must have sent its whole // clientRequestDeadline is when the client must have sent its whole
+1 -4
View File
@@ -119,10 +119,7 @@ func wantFullLine(t *testing.T, line logLine) {
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound, ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
CacheControl: "no-store", Location: "/elsewhere", CacheControl: "no-store", Location: "/elsewhere",
Action: requestlog.ActionForward, Action: requestlog.ActionForward,
// Its 3 bytes in and 5 out, each way counted by default. Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
Counts: ratelimit.Counts{
Minute: 1, Hour: 1, Day: 1, MinuteBytes: 8, HourBytes: 8, DayBytes: 8,
},
}) })
if !reflect.DeepEqual(line.Line, want) { if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want) t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
+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",instance="app",rule_id="blocked"}`, 1) `smallwebwaf_rule_matches_total{action="block",rule_id="blocked"}`, 1)
wantMetric(t, metrics, wantMetric(t, metrics,
`smallwebwaf_rule_matches_total{action="ban",instance="app",rule_id="probe"}`, 1) `smallwebwaf_rule_matches_total{action="ban",rule_id="probe"}`, 1)
wantMetric(t, metrics, `smallwebwaf_rules_loaded{instance="app"}`, 2) wantMetric(t, metrics, "smallwebwaf_rules_loaded", 2)
wantMetric(t, metrics, `smallwebwaf_requests_total{action="rule_blocked",`+ wantMetric(t, metrics,
`instance="app",status_class="4xx"}`, 1) `smallwebwaf_requests_total{action="rule_blocked",status_class="4xx"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack",instance="app"}`, 1) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack"}`, 1)
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 0) wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 0)
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 1) wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1)
} }
// writeRules writes content as a rule file into a new directory, and // writeRules writes content as a rule file into a new directory, and
+5 -4
View File
@@ -10,9 +10,8 @@ import (
// checkRules checks the request against the rules of the rule files at // checkRules checks the request against the rules of the rule files at
// now, notes the ids of those it matches in the log line, and returns the // now, notes the ids of those it matches in the log line, and returns the
// action of the rule that refuses it, ActionRuleBlocked for a block rule // action of the rule that refuses it, ActionRuleBlocked for a block rule
// and ActionBanned for a ban rule, or "" when none does. A ban rule bans // and ActionBanned for a ban rule, or "" when none does. In enforce mode
// the client's netblock for a clear sign of attack, or in observe mode // a ban rule bans the client's netblock for a clear sign of attack.
// raises the alert for the ban it would have made.
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)
@@ -30,7 +29,9 @@ 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:
rq.banForAttack(now, last) if !rq.h.config.Observe {
rq.banForAttack(now, last)
}
return requestlog.ActionBanned return requestlog.ActionBanned
default: default:
+4 -39
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{
{Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100}, {Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100},
{Forwarded: true, Status: 101}, {Forwarded: true, Status: 101},
{Forwarded: true, Status: 304, RequestBytes: 5}, {Forwarded: true, Status: 304, RequestBytes: 5},
{Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true}, {Country: "FR", Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true},
{Forwarded: true, Status: 502, ResponseBytes: 12}, {Forwarded: true, Status: 502, ResponseBytes: 12},
// Closed without an answer: refused, and no response. // Closed without an answer: refused, and no response.
{Refused: true, Status: 0}, {Refused: true, Status: 0},
@@ -33,6 +33,8 @@ func TestHistoryKeepsEveryRequest(t *testing.T) {
want := ratelimit.History{ want := ratelimit.History{
FirstSeen: start, FirstSeen: start,
LastSeen: start.Add(6 * time.Minute), LastSeen: start.Add(6 * time.Minute),
Country: "FR",
LookedUp: start.Add(3 * time.Minute),
Requests: 7, Requests: 7,
Forwarded: 4, Forwarded: 4,
Refused: 2, Refused: 2,
@@ -50,43 +52,6 @@ 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()
+88 -205
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
// and bytes counted over a minute, an hour and a day, as the "Counting // counted over a minute, an hour and a day, as the "Counting method"
// method" section of SPEC.md describes, which tell when a request takes // section of SPEC.md describes, which tell when a request takes the client
// the client over a rate limit or a byte limit, and each client's history // over a rate limit, and each client's history since it was first seen.
// since it was first seen. At most 20,000 clients are kept, in memory, and // At most 20,000 clients are kept, in memory, and written to clients.json
// written to clients.json and read from it by the state package. // and read from it by the state package.
package ratelimit package ratelimit
import ( import (
@@ -23,30 +23,19 @@ const maxClients = 20000
const day = 24 * time.Hour const day = 24 * time.Hour
// The kinds of limits, as the metrics name them.
const (
// KindRequests is a rate limit, on a client's requests.
KindRequests = "requests"
// KindBytes is a byte limit, on a client's bytes.
KindBytes = "bytes"
)
// Limits are the most requests a client may make in a minute, an hour and // Limits are the most requests a client may make in a minute, an hour and
// a day, and the most bytes. Zero is no limit. // a day. Zero is no limit.
type Limits struct { type Limits struct {
PerMinute int64 PerMinute int64
PerHour int64 PerHour int64
PerDay int64 PerDay int64
BytesPerMinute int64
BytesPerHour int64
BytesPerDay int64
} }
// Limiter counts each client's requests and bytes against the limits, and // Limiter counts each client's requests against the limits, and keeps
// keeps its history. It is safe for concurrent use. // its history. It is safe for concurrent use.
type Limiter struct { type Limiter struct {
// windows are the minute, the hour and the day, in the order of // windows are the minute, the hour and the day, in the order of
// Client.buckets and Client.byteBuckets. // Client.buckets.
windows [3]window windows [3]window
mu sync.Mutex mu sync.Mutex
@@ -54,23 +43,17 @@ type Limiter struct {
} }
// Client is a client in the table, as clients.json holds it: its buckets // Client is a client in the table, as clients.json holds it: its buckets
// of requests and of bytes in each window, and its history. // in each window, and its history.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type Client struct { type Client struct {
Client netip.Prefix `json:"client"` Client netip.Prefix `json:"client"`
Minute Buckets `json:"minute"` Minute Buckets `json:"minute"`
Hour Buckets `json:"hour"` Hour Buckets `json:"hour"`
Day Buckets `json:"day"` Day Buckets `json:"day"`
MinuteBytes Buckets `json:"minute_bytes"` History History `json:"history"`
HourBytes Buckets `json:"hour_bytes"`
DayBytes Buckets `json:"day_bytes"`
History History `json:"history"`
} }
// Buckets are a client's two buckets in one window: the requests, or the // Buckets are a client's two buckets in one window: the requests in the
// bytes, in the bucket under way, which began at Start, and in the bucket // bucket under way, which began at Start, and in the bucket before it.
// before it.
type Buckets struct { type Buckets struct {
Start time.Time `json:"start"` Start time.Time `json:"start"`
Current int64 `json:"current"` Current int64 `json:"current"`
@@ -83,12 +66,8 @@ type Buckets struct {
type History struct { type History struct {
FirstSeen time.Time `json:"first_seen"` FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"` LastSeen time.Time `json:"last_seen"`
// ASN, ASName and Country are the client's AS number, AS name and // Country is the client's country as it was last looked up, and
// country as last looked up, each empty when the lookup could not // LookedUp when that was; both are empty while it never was.
// find it, and LookedUp is when the lookup gave that answer; all are
// empty while the client never was looked up.
ASN string `json:"asn,omitempty"`
ASName string `json:"as_name,omitempty"`
Country string `json:"country,omitempty"` Country string `json:"country,omitempty"`
LookedUp time.Time `json:"looked_up,omitzero"` LookedUp time.Time `json:"looked_up,omitzero"`
// Requests are all the client's requests: Forwarded those passed to // Requests are all the client's requests: Forwarded those passed to
@@ -118,12 +97,14 @@ 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 or a byte limit. // Limit is its requests that broke a rate 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
@@ -136,8 +117,7 @@ type Request struct {
// and of its response. // and of its response.
RequestBytes int64 RequestBytes int64
ResponseBytes int64 ResponseBytes int64
// BrokeLimit is true for a request that broke a rate limit or a byte // BrokeLimit is true for a request that broke a rate limit.
// limit.
BrokeLimit bool BrokeLimit bool
} }
@@ -150,74 +130,62 @@ func New(limits Limits) *Limiter {
return &Limiter{ return &Limiter{
windows: [3]window{ windows: [3]window{
{ {name: "minute", length: time.Minute, limit: limits.PerMinute},
name: "minute", length: time.Minute, {name: "hour", length: time.Hour, limit: limits.PerHour},
limit: limits.PerMinute, byteLimit: limits.BytesPerMinute, {name: "day", length: day, limit: limits.PerDay},
},
{
name: "hour", length: time.Hour,
limit: limits.PerHour, byteLimit: limits.BytesPerHour,
},
{
name: "day", length: day,
limit: limits.PerDay, byteLimit: limits.BytesPerDay,
},
}, },
clients: clients, clients: clients,
} }
} }
// Hit is a request that takes a client over a rate limit, or whose bytes // Hit is a request that takes a client over a rate limit.
// take it over a byte limit.
type Hit struct { type Hit struct {
// Kind is KindRequests for a rate limit, KindBytes for a byte limit.
Kind string
// Window is "minute", "hour" or "day". // Window is "minute", "hour" or "day".
Window string Window string
// Limit is the window's limit, as the client's percentage of it. // Limit is the window's limit.
Limit int64 Limit int64
// Count is the client's requests, or bytes, counted in the window, // Requests is the client's requests counted in the window, this one
// this request's included. // included.
Count float64 Requests float64
} }
// Counts are a client's requests and bytes in the minute, the hour and // Counts are a client's requests in the minute, the hour and the day that
// the day that end at a request, that request's included. // end at a request, that request included.
//
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
type Counts struct { type Counts struct {
Minute float64 `json:"minute"` Minute float64 `json:"minute"`
Hour float64 `json:"hour"` Hour float64 `json:"hour"`
Day float64 `json:"day"` Day float64 `json:"day"`
MinuteBytes float64 `json:"minute_bytes"`
HourBytes float64 `json:"hour_bytes"`
DayBytes float64 `json:"day_bytes"`
} }
// Count counts a request from client at now, in every window, whether or // Count counts a request from client at now, in every window, whether or
// not it is refused, and returns the client's counts in each window. It // not it is refused, and returns the client's requests in each window. It
// reports whether the request takes the client over a rate limit, of // reports whether the request takes the client over a limit, and the
// which the client gets the percentage percent, rounded down, and the hit: // window whose limit it goes over, the shortest if it is over several.
// the window whose limit it goes over, the shortest if it is over func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) {
// several. A limit that is off stays off. l.mu.Lock()
func (l *Limiter) Count( defer l.mu.Unlock()
client netip.Prefix, now time.Time, percent int64,
) (Counts, Hit, bool) { var (
return l.count(client, now, 1, 0, percent) requests [3]float64
hit Hit
)
for i, b := range l.get(client).buckets() {
w := l.windows[i]
requests[i] = b.add(now, w.length)
if hit.Window == "" && w.limit > 0 && requests[i] > float64(w.limit) {
hit = Hit{Window: w.name, Limit: w.limit, Requests: requests[i]}
}
}
counts := Counts{Minute: requests[0], Hour: requests[1], Day: requests[2]}
return counts, hit, hit.Window != ""
} }
// CountBytes counts bytes, those of a request from client that has ended, // Reset sets client's counts in every window back to zero. Its history
// at now, in every window, and returns the client's counts in each window. // keeps its totals.
// It reports whether the bytes take the client over a byte limit, of which
// the client gets the percentage percent, and the hit, as Count does.
func (l *Limiter) CountBytes(
client netip.Prefix, now time.Time, bytes, percent int64,
) (Counts, Hit, bool) {
return l.count(client, now, 0, bytes, percent)
}
// Reset sets client's counts of requests and of bytes in every window
// back to zero. Its history keeps its totals.
func (l *Limiter) Reset(client netip.Prefix) { func (l *Limiter) Reset(client netip.Prefix) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
@@ -225,7 +193,6 @@ func (l *Limiter) Reset(client netip.Prefix) {
c, seen := l.clients.Peek(client) c, seen := l.clients.Peek(client)
if seen { if seen {
c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{} c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{}
c.MinuteBytes, c.HourBytes, c.DayBytes = Buckets{}, Buckets{}, Buckets{}
} }
} }
@@ -242,6 +209,11 @@ func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) {
h.LastSeen = now h.LastSeen = now
if r.Country != "" {
h.Country = r.Country
h.LookedUp = now
}
h.Requests++ h.Requests++
if r.Forwarded { if r.Forwarded {
h.Forwarded++ h.Forwarded++
@@ -260,25 +232,6 @@ 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 {
@@ -359,11 +312,12 @@ func (l *Limiter) Load(clients []Client, now time.Time) {
l.clients.Purge() l.clients.Purge()
for _, c := range clients { for _, c := range clients {
for i, w := range l.windows { for i, b := range c.buckets() {
for _, b := range []*Buckets{c.buckets()[i], c.byteBuckets()[i]} { // The window that ends at now covers neither bucket once it
if b.Passed(now, w.length) { // begins after the bucket under way has ended.
*b = Buckets{} length := l.windows[i].length
} if !now.Add(-length).Before(b.Start.Add(length)) {
*b = Buckets{}
} }
} }
@@ -371,51 +325,6 @@ func (l *Limiter) Load(clients []Client, now time.Time) {
} }
} }
// count adds requests and bytes from client at now to its buckets in
// every window, and returns its counts. A limit is broken only by what is
// added to it, so that a request whose bytes are counted after another of
// the client's requests broke a rate limit does not break it too. The
// client gets the percentage percent of each limit.
func (l *Limiter) count(
client netip.Prefix, now time.Time, requests, bytes, percent int64,
) (Counts, Hit, bool) {
l.mu.Lock()
defer l.mu.Unlock()
c := l.get(client)
requestBuckets, byteBuckets := c.buckets(), c.byteBuckets()
var (
requestCounts, byteCounts [3]float64
hit Hit
)
for i, w := range l.windows {
requestCounts[i] = requestBuckets[i].Add(now, w.length, requests)
byteCounts[i] = byteBuckets[i].Add(now, w.length, bytes)
limit, byteLimit := percentOf(w.limit, percent), percentOf(w.byteLimit, percent)
switch {
case hit.Window != "":
case requests > 0 && w.limit > 0 && requestCounts[i] > float64(limit):
hit = Hit{
Kind: KindRequests, Window: w.name, Limit: limit, Count: requestCounts[i],
}
case bytes > 0 && w.byteLimit > 0 && byteCounts[i] > float64(byteLimit):
hit = Hit{
Kind: KindBytes, Window: w.name, Limit: byteLimit, Count: byteCounts[i],
}
}
}
counts := Counts{
Minute: requestCounts[0], Hour: requestCounts[1], Day: requestCounts[2],
MinuteBytes: byteCounts[0], HourBytes: byteCounts[1], DayBytes: byteCounts[2],
}
return counts, hit, hit.Window != ""
}
// get returns client's entry in the table, a new one if it has none, and // get returns client's entry in the table, a new one if it has none, and
// makes it the most recently seen. // makes it the most recently seen.
func (l *Limiter) get(client netip.Prefix) *Client { func (l *Limiter) get(client netip.Prefix) *Client {
@@ -428,48 +337,30 @@ func (l *Limiter) get(client netip.Prefix) *Client {
return c return c
} }
// buckets returns c's buckets of requests in the minute, the hour and the // buckets returns c's buckets in the minute, the hour and the day.
// day.
func (c *Client) buckets() [3]*Buckets { func (c *Client) buckets() [3]*Buckets {
return [3]*Buckets{&c.Minute, &c.Hour, &c.Day} return [3]*Buckets{&c.Minute, &c.Hour, &c.Day}
} }
// byteBuckets returns c's buckets of bytes in the minute, the hour and the // window is a length of time over which requests are counted, and the
// day. // most requests a client may make in it.
func (c *Client) byteBuckets() [3]*Buckets {
return [3]*Buckets{&c.MinuteBytes, &c.HourBytes, &c.DayBytes}
}
// window is a length of time over which requests and bytes are counted,
// and the most requests and the most bytes a client may have in it.
type window struct { type window struct {
name string name string
length time.Duration length time.Duration
limit int64 limit int64
byteLimit int64
} }
// percentOf returns the percentage percent of limit, rounded down. It is // add counts a request at now in a window of length, and returns the
// written as limit's hundreds times percent, plus the rest's share, since // client's requests in the window that ends at now: those in the bucket
// limit*percent can overflow for a byte limit. // under way, and those in the bucket before it weighted by how much of
func percentOf(limit, percent int64) int64 { // that bucket the window still covers.
const hundred = 100
return limit/hundred*percent + limit%hundred*percent/hundred
}
// Add counts n requests, or n bytes, at now in a window of length, and
// returns the count in the window that ends at now: what is in the bucket
// under way, and what is in the bucket before it weighted by how much of
// that bucket the window still covers. With n zero it counts nothing, and
// returns the count. The anomaly counters count in Buckets too.
// //
// Concurrent requests can be counted out of order, so now can be a moment // Concurrent requests can be counted out of order, so now can be a moment
// before the bucket under way began; such a request is counted in that // before the bucket under way began; such a request is counted in that
// bucket. A request dated more than a second before it means the clock // bucket. A request dated more than a second before it means the clock
// was set back, and the buckets start afresh: otherwise the bucket before // was set back, and the buckets start afresh: otherwise the bucket before
// would keep its full weight until the clock caught up. // would keep its full weight until the clock caught up.
func (b *Buckets) Add(now time.Time, length time.Duration, n int64) float64 { func (b *Buckets) add(now time.Time, length time.Duration) float64 {
if now.Before(b.Start.Add(-time.Second)) { if now.Before(b.Start.Add(-time.Second)) {
*b = Buckets{} *b = Buckets{}
} }
@@ -486,7 +377,7 @@ func (b *Buckets) Add(now time.Time, length time.Duration, n int64) float64 {
b.Current = 0 b.Current = 0
} }
b.Current += n b.Current++
elapsed := max(now.Sub(b.Start), 0) elapsed := max(now.Sub(b.Start), 0)
covered := 1 - float64(elapsed)/float64(length) covered := 1 - float64(elapsed)/float64(length)
@@ -494,14 +385,6 @@ func (b *Buckets) Add(now time.Time, length time.Duration, n int64) float64 {
return float64(b.Previous)*covered + float64(b.Current) return float64(b.Previous)*covered + float64(b.Current)
} }
// Passed reports whether the window of length that ends at now covers
// neither of b's buckets: it begins after the bucket under way has ended.
// What they hold then counts no more, and a state file read at now drops
// it.
func (b *Buckets) Passed(now time.Time, length time.Duration) bool {
return !now.Add(-length).Before(b.Start.Add(length))
}
// add counts a response with status in its class. A status of 0, for // add counts a response with status in its class. A status of 0, for
// nothing sent, is not a response. // nothing sent, is not a response.
func (r *Responses) add(status int) { func (r *Responses) add(status int) {
+6 -184
View File
@@ -1,7 +1,6 @@
package ratelimit_test package ratelimit_test
import ( import (
"math"
"net/netip" "net/netip"
"testing" "testing"
"time" "time"
@@ -12,10 +11,6 @@ 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"
@@ -67,180 +62,22 @@ func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) {
start := midnight() start := midnight()
for range limit { for range limit {
_, _, over := limiter.Count(client, start, whole) _, _, over := limiter.Count(client, start)
if over { if over {
t.Fatal("a request within the limit is over it") t.Fatal("a request within the limit is over it")
} }
} }
// Over both limits; the minute's is named, with the four requests. // Over both limits; the minute's is named, with the four requests.
_, hit, over := limiter.Count(client, start, whole) _, hit, over := limiter.Count(client, start)
want := ratelimit.Hit{ want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1}
Kind: ratelimit.KindRequests, Window: minute, Limit: limit, Count: limit + 1,
}
if !over || hit != want { if !over || hit != want {
t.Errorf("request over the limit gives %+v and %t, want %+v and true", t.Errorf("request over the limit gives %+v and %t, want %+v and true",
hit, over, want) hit, over, want)
} }
} }
func TestClientGetsItsPercentageOfEachLimitRoundedDown(t *testing.T) {
t.Parallel()
limiter := ratelimit.New(ratelimit.Limits{PerMinute: 5, BytesPerDay: math.MaxInt64})
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()
@@ -249,14 +86,14 @@ func TestCountGivesTheRequestsInEachWindow(t *testing.T) {
start := midnight() start := midnight()
for range 3 { for range 3 {
limiter.Count(client, start, whole) limiter.Count(client, start)
} }
// A quarter into the next hour, the minute has only this request. The // A quarter into the next hour, the minute has only this request. The
// hour still covers three quarters of the bucket before, with its three // hour still covers three quarters of the bucket before, with its three
// requests, which count 2.25, and this one: 3.25. The day covers all // requests, which count 2.25, and this one: 3.25. The day covers all
// four. // four.
counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4), whole) counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4))
want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4} want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4}
if counts != want { if counts != want {
@@ -424,24 +261,9 @@ func wantCount(
) { ) {
t.Helper() t.Helper()
_, hit, _ := limiter.Count(client, now, whole) _, hit, _ := limiter.Count(client, now)
if hit.Window != want { if hit.Window != want {
t.Errorf("request from %s at %s is over %q, want %q", t.Errorf("request from %s at %s is over %q, want %q",
client, now.Format(time.RFC3339), hit.Window, want) client, now.Format(time.RFC3339), hit.Window, want)
} }
} }
// wantBytesCount counts bytes from client at now, and checks the kind of
// the limit they break, "" for none.
func wantBytesCount(
t *testing.T, limiter *ratelimit.Limiter, client netip.Prefix, now time.Time,
bytes int64, want string,
) {
t.Helper()
_, hit, _ := limiter.CountBytes(client, now, bytes, whole)
if hit.Kind != want {
t.Errorf("%d bytes from %s at %s break a limit on %q, want %q",
bytes, client, now.Format(time.RFC3339), hit.Kind, want)
}
}
+7 -14
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(), whole) limiter.Count(netip.MustParsePrefix(want[i]), midnight())
} }
snapshot := limiter.Snapshot() snapshot := limiter.Snapshot()
@@ -63,8 +63,7 @@ func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) {
start := midnight() start := midnight()
limiter := ratelimit.New(ratelimit.Limits{}) limiter := ratelimit.New(ratelimit.Limits{})
limiter.Count(client, start, whole) limiter.Count(client, start)
limiter.CountBytes(client, start, 5, whole)
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true}) limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
loaded := func(now time.Time) ratelimit.Client { loaded := func(now time.Time) ratelimit.Client {
@@ -77,25 +76,19 @@ 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, of requests and of bytes, which are emptied; the // minute's buckets, which are emptied; the hour's and the day's stay,
// hour's and the day's stay, and so does the history. // and so does the history.
got := loaded(start.Add(2 * time.Minute)) got := loaded(start.Add(2 * time.Minute))
if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 || if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 ||
got.Day.Current != 1 || got.History.Requests != 1 { got.Day.Current != 1 || got.History.Requests != 1 {
t.Errorf("loaded two minutes on as %+v", got) t.Errorf("loaded two minutes on as %+v", got)
} }
if got.MinuteBytes != (ratelimit.Buckets{}) || got.HourBytes.Current != 5 ||
got.DayBytes.Current != 5 {
t.Errorf("loaded two minutes on with buckets of bytes %+v, %+v and %+v",
got.MinuteBytes, got.HourBytes, got.DayBytes)
}
// A moment before, the window still covers some of the earlier one. // A moment before, the window still covers some of the earlier one.
got = loaded(start.Add(2*time.Minute - time.Nanosecond)) got = loaded(start.Add(2*time.Minute - time.Nanosecond))
if got.Minute.Current != 1 || got.MinuteBytes.Current != 5 { if got.Minute.Current != 1 {
t.Errorf("loaded just under two minutes on with minute buckets %+v and %+v", t.Errorf("loaded just under two minutes on with minute buckets %+v",
got.Minute, got.MinuteBytes) got.Minute)
} }
} }
-335
View File
@@ -1,335 +0,0 @@
package reputation
import (
"cmp"
"context"
"encoding/hex"
"errors"
"fmt"
"log/slog"
"net"
"net/netip"
"slices"
"strings"
"sync"
"time"
"github.com/hashicorp/golang-lru/v2/simplelru"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config"
)
const (
// maxVerdicts is how many verdicts are kept. Past it, the one fetched
// longest ago is dropped.
maxVerdicts = 100000
// maxQueries is how many queries may be under way at once. Past it, a
// zone is not asked about a client until the client's next request, so
// that a swarm of new addresses cannot fill the memory.
maxQueries = 1000
// failureDelay is how long a zone is not asked again after a query to
// it fails, so that a zone refusing queries is not asked on every
// request.
failureDelay = time.Minute
)
var (
errAsk = errors.New("ask the zone")
errRefused = errors.New("the zone refused the query")
errNotListing = errors.New("the answer is outside 127.0.0.0/8")
)
// Verdict is what a zone said about a client, as reputation.json holds
// it: the zone, the client's address, whether the zone lists it, and when
// the zone answered.
type Verdict struct {
Zone string `json:"zone"`
Client netip.Addr `json:"client"`
Listed bool `json:"listed"`
Fetched time.Time `json:"fetched"`
}
// DNSBLParams are what NewDNSBL needs.
type DNSBLParams struct {
// Zones are the DNSBL zones clients are asked about in
// (SWWAF_DNSBL_ZONES).
Zones []string
// Resolver is the resolver they are asked through
// (SWWAF_DNSBL_RESOLVER), or, while it is the zero AddrPort, the
// host's, as /etc/resolv.conf names it.
Resolver netip.AddrPort
// CacheTTL is how long a verdict is used after it was fetched
// (SWWAF_REPUTATION_CACHE_TTL), and Timeout how long a query may take
// (SWWAF_REPUTATION_TIMEOUT).
CacheTTL time.Duration
Timeout time.Duration
// Now tells the time, normally time.Now in UTC.
Now func() time.Time
// ProcessLog receives each query that fails, and why.
ProcessLog *slog.Logger
// Alerts receive a source_failure alert for each query that fails.
Alerts *alerts.Queue
}
// DNSBL asks the DNSBL zones about clients, in the background, and keeps
// their verdicts. It is safe for concurrent use.
type DNSBL struct {
params DNSBLParams
resolver *net.Resolver
mu sync.Mutex
// verdicts are by query. Each is added as it is fetched and never moved
// up, so that the one fetched longest ago is the first dropped.
verdicts *simplelru.LRU[query, Verdict]
// asking are the queries under way.
asking map[query]bool
// queries and failures count, by zone, the queries made and those that
// failed, and retryAt is when a zone whose last query failed may be
// asked again.
queries map[string]int
failures map[string]int
retryAt map[string]time.Time
}
// query is a client's address, to ask a zone about.
type query struct {
zone string
client netip.Addr
}
// NewDNSBL returns a DNSBL with no verdict yet.
func NewDNSBL(params DNSBLParams) *DNSBL {
verdicts, err := simplelru.NewLRU[query, Verdict](maxVerdicts, nil)
if err != nil {
panic(err) // NewLRU fails only for a size below one
}
resolver := &net.Resolver{}
if params.Resolver.IsValid() {
// Dial is used by Go's own resolver alone.
resolver.PreferGo = true
resolver.Dial = func(ctx context.Context, network, _ string) (net.Conn, error) {
var dialer net.Dialer
return dialer.DialContext(ctx, network, params.Resolver.String())
}
}
return &DNSBL{
params: params,
resolver: resolver,
verdicts: verdicts,
asking: map[query]bool{},
queries: map[string]int{},
failures: map[string]int{},
retryAt: map[string]time.Time{},
}
}
// Zones returns the zones, in the order SWWAF_DNSBL_ZONES names them.
func (d *DNSBL) Zones() []string {
return slices.Clone(d.params.Zones)
}
// ListedBy returns the zones whose verdict on addr, a client's address,
// lists it, in the order SWWAF_DNSBL_ZONES names them, each with its key
// masked, as config.MaskZoneKey masks it, since they go to the request
// log, the alerts and the metrics. A verdict is used until CacheTTL has
// passed since it was fetched. Each zone without one is asked about addr
// in the background, unless a query about addr to it is under way, the
// zone is left alone after a failure, or maxQueries are under way;
// ListedBy never waits for a query. ctx is the context of the client's
// request, and a query goes on after the request ends.
func (d *DNSBL) ListedBy(ctx context.Context, addr netip.Addr) []string {
d.mu.Lock()
defer d.mu.Unlock()
now := d.params.Now()
var listedBy []string
for _, zone := range d.params.Zones {
q := query{zone: zone, client: addr}
kept, found := d.verdicts.Peek(q)
switch {
case found && now.Sub(kept.Fetched) < d.params.CacheTTL:
if kept.Listed {
listedBy = append(listedBy, config.MaskZoneKey(zone))
}
case !d.asking[q] && !now.Before(d.retryAt[zone]) && len(d.asking) < maxQueries:
d.asking[q] = true
d.queries[zone]++
go d.ask(context.WithoutCancel(ctx), q)
}
}
return listedBy
}
// Queries returns how many queries were made to zone.
func (d *DNSBL) Queries(zone string) int {
d.mu.Lock()
defer d.mu.Unlock()
return d.queries[zone]
}
// Failures returns how many queries to zone failed.
func (d *DNSBL) Failures(zone string) int {
d.mu.Lock()
defer d.mu.Unlock()
return d.failures[zone]
}
// Snapshot returns every verdict still in use, sorted by client, then by
// zone, as reputation.json lists them.
func (d *DNSBL) Snapshot() []Verdict {
d.mu.Lock()
now := d.params.Now()
verdicts := make([]Verdict, 0, d.verdicts.Len())
for _, kept := range d.verdicts.Values() {
if now.Sub(kept.Fetched) < d.params.CacheTTL {
verdicts = append(verdicts, kept)
}
}
d.mu.Unlock()
slices.SortFunc(verdicts, func(a, b Verdict) int {
return cmp.Or(a.Client.Compare(b.Client), strings.Compare(a.Zone, b.Zone))
})
return verdicts
}
// Load keeps verdicts, read from reputation.json, in place of those it
// keeps, but for those of a zone SWWAF_DNSBL_ZONES does not name, and,
// past maxVerdicts, those fetched longest ago. One fetched CacheTTL ago or
// more is neither used nor written, as for any verdict.
func (d *DNSBL) Load(verdicts []Verdict) {
verdicts = slices.Clone(verdicts)
slices.SortStableFunc(verdicts, func(a, b Verdict) int {
return a.Fetched.Compare(b.Fetched)
})
d.mu.Lock()
defer d.mu.Unlock()
d.verdicts.Purge()
for _, kept := range verdicts {
if slices.Contains(d.params.Zones, kept.Zone) {
d.verdicts.Add(query{zone: kept.Zone, client: kept.Client}, kept)
}
}
}
// ask asks q's zone about q's client, keeps the verdict, and notes the
// query as no longer under way. A query that fails gives no verdict: it
// is counted, logged and raised as a source_failure alert, which show the
// zone with its key masked, and the zone is not asked again for
// failureDelay.
func (d *DNSBL) ask(ctx context.Context, q query) {
listed, err := d.lookUp(ctx, q)
now := d.params.Now()
d.mu.Lock()
delete(d.asking, q)
if err == nil {
d.verdicts.Add(q, Verdict{
Zone: q.zone, Client: q.client, Listed: listed, Fetched: now,
})
} else {
d.failures[q.zone]++
d.retryAt[q.zone] = now.Add(failureDelay)
}
d.mu.Unlock()
if err != nil {
const failed = "asking a DNSBL zone failed"
shown := config.MaskZoneKey(q.zone)
// Raised before it is logged, so that the alert is there once the
// log line is.
d.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventSourceFailure,
Reason: failed,
Detail: map[string]any{"source": shown, "error": err.Error()},
})
d.params.ProcessLog.Warn(failed, "zone", shown, "error", err.Error())
}
}
// lookUp asks q's zone about q's client through the resolver, and returns
// whether the zone lists it, as readAnswer reads the answer. No such name
// is a client the zone does not list. A query not answered within Timeout
// fails.
func (d *DNSBL) lookUp(ctx context.Context, q query) (bool, error) {
ctx, cancel := context.WithTimeout(ctx, d.params.Timeout)
defer cancel()
answer, err := d.resolver.LookupNetIP(ctx, "ip4", queryName(q.zone, q.client))
var dnsErr *net.DNSError
switch {
case err == nil:
return readAnswer(answer)
case errors.As(err, &dnsErr) && dnsErr.IsNotFound:
return false, nil
case errors.As(err, &dnsErr):
// The error names the name asked about, which holds the client's
// address, which is not to be logged: only what went wrong is kept.
return false, fmt.Errorf("%w: %s", errAsk, dnsErr.Err)
default:
return false, fmt.Errorf("%w: %w", errAsk, err)
}
}
// queryName returns the name a zone is asked about addr by, as RFC 5782
// builds it: the four numbers of an IPv4 address, or the 32 hex digits of
// an IPv6 address, in reverse order, each followed by a dot, then the zone
// and a dot, which makes it a full name, to which the resolver adds no
// search domain of /etc/resolv.conf.
func queryName(zone string, addr netip.Addr) string {
parts := strings.Split(addr.String(), ".")
if addr.Is6() {
parts = strings.Split(hex.EncodeToString(addr.AsSlice()), "")
}
slices.Reverse(parts)
return strings.Join(parts, ".") + "." + zone + "."
}
// readAnswer reads the addresses a zone answered with. An address in
// 127.0.0.0/8 lists the client, as RFC 5782 has zones answer, but one in
// 127.255.255.0/24 is how Spamhaus refuses a query, such as one sent
// through a public resolver or one past its limit, and is a failure. So is
// an address outside 127.0.0.0/8, such as a resolver gives that answers
// even for names that do not exist.
func readAnswer(answer []netip.Addr) (bool, error) {
listing := netip.MustParsePrefix("127.0.0.0/8")
refusal := netip.MustParsePrefix("127.255.255.0/24")
for _, addr := range answer {
switch {
case refusal.Contains(addr):
return false, fmt.Errorf("%w: %s", errRefused, addr)
case !listing.Contains(addr):
return false, fmt.Errorf("%w: %s", errNotListing, addr)
}
}
return len(answer) > 0, nil
}
-737
View File
@@ -1,737 +0,0 @@
package reputation_test
import (
"bytes"
"context"
"encoding/binary"
"errors"
"io"
"log/slog"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"reflect"
"slices"
"strings"
"sync"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/reputation"
)
// The tests of the DNSBL zones run in synctest bubbles, as those of the
// lists do, and the resolver the zones are asked through is a stand-in
// reached through an in-memory connection, net.Pipe's, for the same
// reason. They run one at a time, none in parallel with another test of
// this package: Go's resolver counts the queries under way in one
// sync.WaitGroup for the whole process, and the process fails when
// queries from two bubbles, or from a bubble and from outside one, are
// under way at once. TestMain has the resolver make its configuration,
// which it makes on its first query, outside every bubble, since the
// configuration holds a channel, which the bubble it was made in would
// keep to itself.
const (
// zone and otherZone are the DNSBL zones the tests name.
zone = "dnsbl.example"
otherZone = "other.example"
// cacheTTL is the tests' SWWAF_REPUTATION_CACHE_TTL, and timeout their
// SWWAF_REPUTATION_TIMEOUT: a second, the least time /etc/resolv.conf
// can have Go's resolver wait for one server, so that it is the
// DNSBL's own timeout that ends a query, whatever that file says.
cacheTTL = 24 * time.Hour
timeout = time.Second
// listed and unlisted are clients zone is asked about by the names
// listedName and unlistedName, and most tests have zone list the first
// alone, by answering with listing.
listed = "192.0.2.99"
unlisted = "192.0.2.100"
listedName = "99.2.0.192." + zone + "."
unlistedName = "100.2.0.192." + zone + "."
listing = "127.0.0.2"
)
// The DNS response codes the stand-in answers with, besides no error.
const (
serverFailure = 2
noSuchName = 3
refused = 5
)
var errNoNetwork = errors.New("the test dials nothing")
func TestMain(m *testing.M) {
// A query that fails at once, as nothing is dialled for it.
resolver := &net.Resolver{
PreferGo: true,
Dial: func(context.Context, string, string) (net.Conn, error) {
return nil, errNoNetwork
},
}
_, _ = resolver.LookupNetIP(context.Background(), "ip4", "warm-up.invalid.")
m.Run()
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestZonesListOrNotClientsByTheirIPv4AndIPv6Addresses(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
// The addresses of the examples of RFC 5782, and the names it
// gives for them.
const (
v4 = "192.0.2.99"
v6 = "2001:db8:1:2:3:4:567:89ab"
// v6Name is the hex digits of v6, in reverse order.
v6Name = "b.a.9.8.7.6.5.0.4.0.0.0.3.0.0.0.2.0.0.0.1.0.0.0.8.b.d.0.1.0.0.2."
)
resolver := &resolverStandIn{answers: map[string]answer{
"99.2.0.192." + zone + ".": {addrs: []string{listing}},
v6Name + otherZone + ".": {addrs: []string{"127.0.0.4", "127.0.0.10"}},
"99.2.0.192." + otherZone + ".": {},
}}
dnsbl := newDNSBL(resolver, dnsblParams(zone, otherZone))
// Neither client has a verdict yet, so neither is listed, and each
// zone is asked about each.
wantZones(t, dnsbl, v4)
wantZones(t, dnsbl, v6)
synctest.Wait()
wantZones(t, dnsbl, v4, zone)
wantZones(t, dnsbl, v6, otherZone)
wantAsked(t, resolver,
"99.2.0.192."+zone+".", "99.2.0.192."+otherZone+".",
v6Name+zone+".", v6Name+otherZone+".")
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestListedByNeverWaitsForAQuery(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
dnsbl := newDNSBL(&resolverStandIn{hanging: true}, dnsblParams(zone))
began := time.Now()
// The second, while the first's query is under way, starts none.
wantZones(t, dnsbl, listed)
wantZones(t, dnsbl, listed)
if waited := time.Since(began); waited != 0 {
t.Errorf("waited %s for the query, want no wait", waited)
}
synctest.Wait()
wantQueries(t, dnsbl, 1, 0)
waitForTheResolver()
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestVerdictUsedUntilTheCacheTTLHasPassedSinceItWasFetched(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
resolver := &resolverStandIn{answers: map[string]answer{
listedName: {addrs: []string{listing}},
}}
dnsbl := newDNSBL(resolver, dnsblParams(zone))
wantZones(t, dnsbl, listed)
wantZones(t, dnsbl, unlisted)
synctest.Wait()
// The zone lists the other client from now on, but the verdicts
// kept are used, and the zone is not asked again, until the TTL
// has passed.
resolver.set(listedName, answer{rcode: noSuchName})
resolver.set(unlistedName, answer{addrs: []string{listing}})
time.Sleep(cacheTTL - time.Nanosecond)
wantZones(t, dnsbl, listed, zone)
wantZones(t, dnsbl, unlisted)
synctest.Wait()
wantQueries(t, dnsbl, 2, 0)
// Then neither verdict is used, and both clients are asked about
// again.
time.Sleep(time.Nanosecond)
wantZones(t, dnsbl, listed)
wantZones(t, dnsbl, unlisted)
synctest.Wait()
wantQueries(t, dnsbl, 4, 0)
wantZones(t, dnsbl, listed)
wantZones(t, dnsbl, unlisted, zone)
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestQueryNotAnsweredWithinTheTimeoutFails(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
queue := newQueue()
p := dnsblParams(zone)
p.Alerts = queue
dnsbl := newDNSBL(&resolverStandIn{hanging: true}, p)
wantZones(t, dnsbl, listed)
time.Sleep(timeout - time.Nanosecond)
synctest.Wait()
wantQueries(t, dnsbl, 1, 0)
time.Sleep(time.Nanosecond)
synctest.Wait()
wantQueries(t, dnsbl, 1, 1)
if got := waiting(queue); len(got) != 1 ||
got[0].Detail["error"] != "ask the zone: i/o timeout" {
t.Errorf("alerts waiting %+v, want the timeout's", got)
}
if verdicts := dnsbl.Snapshot(); len(verdicts) != 0 {
t.Errorf("verdicts %+v, want none", verdicts)
}
waitForTheResolver()
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestZoneThatFailsOrRefusesGivesNoVerdictAndIsLeftAloneForAMinute(t *testing.T) {
for _, tc := range []struct {
name string
answer answer
error string
}{
{
"a server failure", answer{rcode: serverFailure},
"ask the zone: server misbehaving",
},
{"a refusal", answer{rcode: refused}, "ask the zone: server misbehaving"},
{
"an answer in 127.255.255.0/24, with which Spamhaus refuses a query",
answer{addrs: []string{"127.255.255.254"}},
"the zone refused the query: 127.255.255.254",
},
{
"an answer outside 127.0.0.0/8, as for a name that does not exist",
answer{addrs: []string{"192.0.2.1"}},
"the answer is outside 127.0.0.0/8: 192.0.2.1",
},
} {
//nolint:paralleltest // one at a time, as the comment at the top of this file says
t.Run(tc.name, func(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
var log bytes.Buffer
queue := newQueue()
p := dnsblParams(zone)
p.Alerts = queue
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
dnsbl := newDNSBL(&resolverStandIn{answers: map[string]answer{
listedName: tc.answer,
}}, p)
// The failure gives no verdict, and the zone is not asked
// again within a minute of it.
wantZones(t, dnsbl, listed)
synctest.Wait()
time.Sleep(time.Minute - time.Nanosecond)
wantZones(t, dnsbl, listed)
synctest.Wait()
wantQueries(t, dnsbl, 1, 1)
time.Sleep(time.Nanosecond)
wantZones(t, dnsbl, listed)
synctest.Wait()
wantQueries(t, dnsbl, 2, 2)
if verdicts := dnsbl.Snapshot(); len(verdicts) != 0 {
t.Errorf("verdicts %+v, want none", verdicts)
}
// One alert for the first failure; the cooldown holds back
// the second.
wantAlert(t, queue, alerts.Alert{
Time: time.Now().Add(-time.Minute),
Event: alerts.EventSourceFailure,
Reason: "asking a DNSBL zone failed",
Detail: map[string]any{"source": zone, "error": tc.error},
})
if !strings.Contains(log.String(), `"msg":"asking a DNSBL zone failed",`+
`"zone":"`+zone+`","error":"`+tc.error) {
t.Errorf("logged\n%s\nwant the failures", log.String())
}
})
})
}
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestAtMost1000QueriesUnderWay(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
dnsbl := newDNSBL(&resolverStandIn{hanging: true}, dnsblParams(zone))
client := netip.MustParseAddr("198.18.0.0")
for range 1001 {
dnsbl.ListedBy(t.Context(), client)
client = client.Next()
}
synctest.Wait()
wantQueries(t, dnsbl, 1000, 0)
waitForTheResolver()
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestMetricsCountEachZonesQueriesAndThoseThatFailed(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
resolver := &resolverStandIn{answers: map[string]answer{
listedName: {addrs: []string{listing}},
"99.2.0.192." + otherZone + ".": {rcode: serverFailure},
}}
dnsbl := newDNSBL(resolver, dnsblParams(zone, otherZone))
m := metrics.New(1, "app")
m.AddReputation(reputation.New(params()), dnsbl)
wantZones(t, dnsbl, listed)
synctest.Wait()
scraped := httptest.NewRecorder()
m.ServeHTTP(scraped, httptest.NewRequestWithContext(t.Context(), http.MethodGet,
"/", http.NoBody))
for series, want := range map[string]string{
"queries_total" + `{instance="app",source="` + zone + `"}`: "1",
"failures_total" + `{instance="app",source="` + zone + `"}`: "0",
"queries_total" + `{instance="app",source="` + otherZone + `"}`: "1",
"failures_total" + `{instance="app",source="` + otherZone + `"}`: "1",
} {
line := "\nsmallwebwaf_reputation_" + series + " " + want + "\n"
if !strings.Contains(scraped.Body.String(), line) {
t.Errorf("metrics\n%s\nwant%s", scraped.Body.String(), line)
}
}
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestZoneKeyIsMaskedInTheVerdictsTheFailuresAndTheMetrics(t *testing.T) {
const (
key = "abcdefghijklmnopqrstuvwxyz"
keyed = key + ".xbl.dq.spamhaus.net"
masked = "********.xbl.dq.spamhaus.net"
)
synctest.Test(t, func(t *testing.T) {
var log bytes.Buffer
queue := newQueue()
p := dnsblParams(keyed)
p.Alerts = queue
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
dnsbl := newDNSBL(&resolverStandIn{answers: map[string]answer{
"99.2.0.192." + keyed + ".": {addrs: []string{listing}},
"100.2.0.192." + keyed + ".": {rcode: serverFailure},
}}, p)
m := metrics.New(1, "app")
m.AddReputation(reputation.New(params()), dnsbl)
// Both clients are asked about before either answer comes, so that
// the failure does not keep the zone from the other query.
wantZones(t, dnsbl, listed)
wantZones(t, dnsbl, unlisted)
synctest.Wait()
wantZones(t, dnsbl, listed, masked)
if got := waiting(queue); len(got) != 1 || got[0].Detail["source"] != masked {
t.Errorf("alerts waiting %+v, want the failure's, from %s", got, masked)
}
scraped := httptest.NewRecorder()
m.ServeHTTP(scraped, httptest.NewRequestWithContext(t.Context(), http.MethodGet,
"/", http.NoBody))
for name, shown := range map[string]string{
"the log": log.String(), "the metrics": scraped.Body.String(),
} {
if strings.Contains(shown, key) || !strings.Contains(shown, masked) {
t.Errorf("%s shows the key, or does not name the zone:\n%s", name, shown)
}
}
})
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestVerdictsKeptAcrossARestart(t *testing.T) {
synctest.Test(t, func(t *testing.T) {
resolver := &resolverStandIn{answers: map[string]answer{
listedName: {addrs: []string{listing}},
}}
dnsbl := newDNSBL(resolver, dnsblParams(zone))
fetched := time.Now()
wantZones(t, dnsbl, unlisted)
wantZones(t, dnsbl, listed)
synctest.Wait()
kept := dnsbl.Snapshot()
want := []reputation.Verdict{
{Zone: zone, Client: netip.MustParseAddr(listed), Listed: true, Fetched: fetched},
{Zone: zone, Client: netip.MustParseAddr(unlisted), Fetched: fetched},
}
if !reflect.DeepEqual(kept, want) {
t.Errorf("verdicts %+v, want %+v", kept, want)
}
// Restarted an hour later with what reputation.json keeps, it uses
// the verdicts, and asks the zone nothing, until the TTL has passed
// since they were fetched.
time.Sleep(time.Hour)
restarted := &resolverStandIn{}
again := newDNSBL(restarted, dnsblParams(zone))
again.Load(kept)
wantZones(t, again, listed, zone)
wantZones(t, again, unlisted)
synctest.Wait()
wantAsked(t, restarted)
time.Sleep(cacheTTL - time.Hour)
wantZones(t, again, listed)
synctest.Wait()
wantAsked(t, restarted, listedName)
})
}
func TestNeitherAVerdictOfAZoneNotNamedNorOnePastItsTTLIsKept(t *testing.T) {
t.Parallel()
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
p := dnsblParams(zone)
p.Now = func() time.Time { return now }
dnsbl := reputation.NewDNSBL(p)
// The last verdict still in use, one fetched a TTL ago, and one of a
// zone SWWAF_DNSBL_ZONES does not name.
inUse := reputation.Verdict{
Zone: zone, Client: netip.MustParseAddr(listed), Listed: true,
Fetched: now.Add(-cacheTTL + time.Nanosecond),
}
stale := reputation.Verdict{
Zone: zone, Client: netip.MustParseAddr(unlisted), Fetched: now.Add(-cacheTTL),
}
notNamed := reputation.Verdict{
Zone: otherZone, Client: netip.MustParseAddr(listed), Listed: true, Fetched: now,
}
dnsbl.Load([]reputation.Verdict{notNamed, stale, inUse})
if got := dnsbl.Snapshot(); !reflect.DeepEqual(got, []reputation.Verdict{inUse}) {
t.Errorf("verdicts %+v, want only %+v", got, inUse)
}
}
func TestAtMost100000VerdictsKeptTheOneFetchedLongestAgoDroppedFirst(t *testing.T) {
t.Parallel()
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
p := dnsblParams(zone)
p.Now = func() time.Time { return now }
dnsbl := reputation.NewDNSBL(p)
// 100,001 verdicts, listed by client, as reputation.json lists them,
// each fetched a millisecond before the one before it: the last is one
// too many.
const count = 100001
verdicts := make([]reputation.Verdict, 0, count)
client := netip.MustParseAddr("198.18.0.0")
for i := range count {
verdicts = append(verdicts, reputation.Verdict{
Zone: zone, Client: client, Fetched: now.Add(-time.Duration(i) * time.Millisecond),
})
client = client.Next()
}
dnsbl.Load(verdicts)
got := dnsbl.Snapshot()
if len(got) != count-1 || !slices.Contains(got, verdicts[0]) ||
slices.Contains(got, verdicts[count-1]) {
t.Errorf("%d verdicts kept, want all but the one fetched longest ago", len(got))
}
}
//nolint:paralleltest // one at a time, as the comment at the top of this file says
func TestQueriesGoToTheResolverSWWAFDNSBLResolverNames(t *testing.T) {
resolver := &resolverStandIn{answers: map[string]answer{
listedName: {addrs: []string{listing}},
}}
conn, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen: %v", err)
}
served := make(chan struct{})
go func() {
resolver.serveUDP(conn)
close(served)
}()
t.Cleanup(func() {
_ = conn.Close()
<-served
})
p := dnsblParams(zone)
p.Resolver = netip.MustParseAddrPort(conn.LocalAddr().String())
// On the real clock: the stand-in answers at once, so only a test
// process held up for a whole minute would see the query fail.
p.Timeout = time.Minute
isListed, err := reputation.NewDNSBL(p).LookUp(zone, netip.MustParseAddr(listed))
if err != nil || !isListed {
t.Errorf("listed %t (%v), want true", isListed, err)
}
wantAsked(t, resolver, listedName)
}
// resolverStandIn is a stand-in for the resolver the zones are asked
// through. It answers each query by the name asked about, as answers
// gives, with no such name for a name answers does not give, and not at
// all while hanging. It notes each name asked about.
type resolverStandIn struct {
mu sync.Mutex
answers map[string]answer
hanging bool
names []string
}
// answer is how the stand-in answers a name: with an A record of each of
// addrs, or with the response code rcode, unless it is 0, for no error.
type answer struct {
addrs []string
rcode uint16
}
// What the stand-in reads of a query, and writes in its reply.
const (
// headerLength is the length of a DNS message's header, which the
// question follows: its id, its flags, and how many questions,
// answers and other records it holds, two bytes each.
headerLength = 12
// typeAndClass is the length of the type and the class that end a
// question, after its name.
typeAndClass = 4
// replyFlags mark a reply to a query that asked for recursion, which
// is available, with no error. The response code goes in their last
// four bits.
replyFlags = 0x8180
// maxMessage is the longest query read over UDP.
maxMessage = 1232
)
// set has the stand-in answer name with given.
func (s *resolverStandIn) set(name string, given answer) {
s.mu.Lock()
defer s.mu.Unlock()
s.answers[name] = given
}
// dial connects Go's resolver to the stand-in through an in-memory
// connection, on which it sends each query, and reads each reply, after
// its length, as over TCP.
func (s *resolverStandIn) dial(context.Context, string, string) (net.Conn, error) {
client, server := net.Pipe()
go s.serve(server)
return client, nil
}
// serve answers the queries that come on conn until the resolver closes
// it.
func (s *resolverStandIn) serve(conn net.Conn) {
defer func() {
_ = conn.Close()
}()
for {
var length [2]byte
_, err := io.ReadFull(conn, length[:])
if err != nil {
return
}
message := make([]byte, binary.BigEndian.Uint16(length[:]))
_, err = io.ReadFull(conn, message)
if err != nil {
return
}
reply, answered := s.reply(message)
if !answered {
continue // the resolver gives up, and closes conn
}
//nolint:gosec // a reply of a few dozen bytes
_, err = conn.Write(append(binary.BigEndian.AppendUint16(nil, uint16(len(reply))),
reply...))
if err != nil {
return
}
}
}
// serveUDP answers the queries that come on conn, each in a datagram, as
// a resolver does, until conn is closed.
func (s *resolverStandIn) serveUDP(conn net.PacketConn) {
message := make([]byte, maxMessage)
for {
n, from, err := conn.ReadFrom(message)
if err != nil {
return
}
reply, answered := s.reply(message[:n])
if answered {
_, _ = conn.WriteTo(reply, from)
}
}
}
// reply returns the stand-in's reply to message, a query, and false for
// none, while it hangs. It notes the name asked about.
func (s *resolverStandIn) reply(message []byte) ([]byte, bool) {
// The name is labels, each after its length, ended by a length of 0.
var labels []string
end := headerLength
for message[end] != 0 {
length := int(message[end])
labels = append(labels, string(message[end+1:end+1+length]))
end += 1 + length
}
end += 1 + typeAndClass
name := strings.Join(labels, ".") + "."
s.mu.Lock()
s.names = append(s.names, name)
given, found := s.answers[name]
hanging := s.hanging
s.mu.Unlock()
if hanging {
return nil, false
}
if !found {
given = answer{rcode: noSuchName}
}
// The query's id, the flags, one question, the answers, and no other
// records, then the question, as asked.
reply := slices.Clone(message[:2])
reply = binary.BigEndian.AppendUint16(reply, replyFlags|given.rcode)
reply = binary.BigEndian.AppendUint16(reply, 1)
//nolint:gosec // a handful of answers
reply = binary.BigEndian.AppendUint16(reply, uint16(len(given.addrs)))
reply = append(reply, 0, 0, 0, 0)
reply = append(reply, message[headerLength:end]...)
// An A record starts with the name asked about, by a pointer to it in
// the question, then its type, A, its class, IN, how long it may be
// kept, 60 seconds, and the length of its address, 4 bytes.
record := []byte{0xc0, headerLength, 0, 1, 0, 1, 0, 0, 0, 60, 0, 4}
for _, addr := range given.addrs {
reply = append(reply, record...)
reply = append(reply, netip.MustParseAddr(addr).AsSlice()...)
}
return reply, true
}
// dnsblParams returns the DNSBLParams of zones, with the tests' cache TTL
// and timeout, by the bubble's clock, with alerts to a queue that sends
// none.
func dnsblParams(zones ...string) reputation.DNSBLParams {
return reputation.DNSBLParams{
Zones: zones,
CacheTTL: cacheTTL,
Timeout: timeout,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: newQueue(),
}
}
// waitForTheResolver waits, on the bubble's clock, an hour, until Go's
// resolver has given up on every stand-in that does not answer: it waits
// for a server as long as /etc/resolv.conf has it wait, a few seconds,
// even after the query was given up, and a bubble cannot end before it.
func waitForTheResolver() {
time.Sleep(time.Hour)
}
// newDNSBL returns the DNSBL of p, asking resolver.
func newDNSBL(resolver *resolverStandIn, p reputation.DNSBLParams) *reputation.DNSBL {
dnsbl := reputation.NewDNSBL(p)
dnsbl.SetDial(resolver.dial)
return dnsbl
}
// wantZones checks the zones whose verdict dnsbl says lists client, as a
// request from client finds them.
func wantZones(t *testing.T, dnsbl *reputation.DNSBL, client string, want ...string) {
t.Helper()
got := dnsbl.ListedBy(t.Context(), netip.MustParseAddr(client))
if !slices.Equal(got, want) {
t.Errorf("%s is listed by %v, want %v", client, got, want)
}
}
// wantQueries checks how many queries dnsbl made to zone, and how many of
// them failed.
func wantQueries(t *testing.T, dnsbl *reputation.DNSBL, queries, failures int) {
t.Helper()
if dnsbl.Queries(zone) != queries || dnsbl.Failures(zone) != failures {
t.Errorf("%d queries and %d failures, want %d and %d", dnsbl.Queries(zone),
dnsbl.Failures(zone), queries, failures)
}
}
// wantAsked checks the names the stand-in was asked about, in any order.
func wantAsked(t *testing.T, resolver *resolverStandIn, want ...string) {
t.Helper()
resolver.mu.Lock()
got := slices.Sorted(slices.Values(resolver.names))
resolver.mu.Unlock()
slices.Sort(want)
if !slices.Equal(got, want) {
t.Errorf("asked about %v, want %v", got, want)
}
}
-27
View File
@@ -1,27 +0,0 @@
package reputation
import (
"context"
"net"
"net/http"
"net/netip"
)
// SetTransport has l's fetches go through transport instead of the
// network.
func (l *Lists) SetTransport(transport http.RoundTripper) {
l.httpClient.Transport = transport
}
// SetDial has d's queries go through dial instead of the network.
func (d *DNSBL) SetDial(
dial func(ctx context.Context, network, address string) (net.Conn, error),
) {
d.resolver = &net.Resolver{PreferGo: true, Dial: dial}
}
// LookUp asks zone about addr at once, as a query in the background does,
// and returns whether zone lists addr.
func (d *DNSBL) LookUp(zone string, addr netip.Addr) (bool, error) {
return d.lookUp(context.Background(), query{zone: zone, client: addr})
}
-501
View File
@@ -1,501 +0,0 @@
// Package reputation fetches the lists the settings name by URL: the
// blocklists of SWWAF_BLOCKLIST_URLS, and the file of AS:percent lines
// SWWAF_ASN_LIMIT_PERCENT_URL names. It keeps the last good copy of each,
// whole, comment lines included, which is used while a fetch fails, and
// when each was last tried. It also asks the DNSBL zones of
// SWWAF_DNSBL_ZONES about clients, and keeps their verdicts. The state
// package writes all of these to reputation.json and reads them from it,
// so that a restart keeps them too.
package reputation
import (
"context"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"net/netip"
"slices"
"strings"
"sync"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/config"
)
const (
// maxListBytes is the most of a list that is read. A longer one is a
// failure, so that a wrong URL cannot fill the memory.
maxListBytes = 16 << 20
// fetchTimeout bounds one fetch of a list.
fetchTimeout = time.Minute
// mappedBits is the length of ::ffff:0.0.0.0/96, the netblock of every
// IPv4-mapped address.
mappedBits = 96
)
var (
errStatus = errors.New("the server answered")
errTooLong = errors.New("the list is longer than 16 MiB")
errNotNetblock = errors.New("is not an address or a netblock, such as 192.0.2.0/24")
errNotASNPercent = errors.New(
"is not an AS number, : and a percentage, such as AS64496:50")
)
// List is a list as reputation.json holds it: the URL it is fetched from,
// when it was last tried, the fetch failed or not, and its last good copy:
// when that was fetched, and its lines, as fetched, comment lines
// included, both left out while no fetch of it has succeeded.
type List struct {
URL string `json:"url"`
Tried time.Time `json:"tried"`
Fetched time.Time `json:"fetched,omitzero"`
Lines []string `json:"lines,omitzero"`
}
// Params are what New needs.
type Params struct {
// BlocklistURLs are the blocklists (SWWAF_BLOCKLIST_URLS), and
// ASNLimitPercentURL the file of AS:percent lines
// (SWWAF_ASN_LIMIT_PERCENT_URL), "" while it is unset.
BlocklistURLs []string
ASNLimitPercentURL string
// Refresh is how long after a list was last fetched or tried it is
// fetched again (SWWAF_BLOCKLIST_REFRESH).
Refresh time.Duration
// Now tells the time, normally time.Now in UTC.
Now func() time.Time
// ProcessLog receives each fetch of a list, and why one failed.
ProcessLog *slog.Logger
// Alerts receive a source_failure alert for each fetch that fails.
Alerts *alerts.Queue
}
// Lists are the lists Params names, each with its last good copy. They
// are safe for concurrent use.
type Lists struct {
params Params
httpClient *http.Client
mu sync.Mutex
// lists are by URL, one for each URL Params names.
lists map[string]*list
}
// list is one list: what reputation.json keeps of it, its last try, zero
// before the first, and its last good copy, what that copy says, and how
// many fetches of it failed.
type list struct {
kept List
entries entries
failures int
}
// entries are what the lines of a copy say: for a blocklist, the netblocks
// it names, with the lengths among them, and for the file of AS:percent
// lines, the percentage it gives each AS number.
type entries struct {
netblocks map[netip.Prefix]bool
lengths []int
percents map[string]int64
}
// New returns the lists, without a copy of any yet.
func New(params Params) *Lists {
l := &Lists{params: params, httpClient: &http.Client{}, lists: map[string]*list{}}
for _, listURL := range l.URLs() {
l.lists[listURL] = &list{kept: List{URL: listURL}}
}
return l
}
// URLs returns the URL of every list: the blocklists' in the order
// SWWAF_BLOCKLIST_URLS names them, then SWWAF_ASN_LIMIT_PERCENT_URL.
func (l *Lists) URLs() []string {
urls := slices.Clone(l.params.BlocklistURLs)
if l.params.ASNLimitPercentURL != "" {
urls = append(urls, l.params.ASNLimitPercentURL)
}
return urls
}
// ListedBy returns the URLs of the blocklists whose copy lists addr, in
// the order SWWAF_BLOCKLIST_URLS names them.
func (l *Lists) ListedBy(addr netip.Addr) []string {
l.mu.Lock()
defer l.mu.Unlock()
var listedBy []string
for _, listURL := range l.params.BlocklistURLs {
if l.lists[listURL].entries.contain(addr) {
listedBy = append(listedBy, listURL)
}
}
return listedBy
}
// ASNLimitPercent returns the percentage the copy of the file of
// AS:percent lines gives asn, and whether it lists asn.
func (l *Lists) ASNLimitPercent(asn string) (int64, bool) {
if l.params.ASNLimitPercentURL == "" {
return 0, false
}
l.mu.Lock()
defer l.mu.Unlock()
percent, listed := l.lists[l.params.ASNLimitPercentURL].entries.percents[asn]
return percent, listed
}
// Fetched returns when the copy in use of the list at listURL was
// fetched, or zero while there is none.
func (l *Lists) Fetched(listURL string) time.Time {
l.mu.Lock()
defer l.mu.Unlock()
return l.lists[listURL].kept.Fetched
}
// Failures returns how many fetches of the list at listURL failed.
func (l *Lists) Failures(listURL string) int {
l.mu.Lock()
defer l.mu.Unlock()
return l.lists[listURL].failures
}
// Run fetches each list once Refresh has passed since it was last fetched
// or tried, the later of the two, until ctx is done. A list never tried is
// fetched at once, and so is one whose last try or copy, read from
// reputation.json, is that old.
func (l *Lists) Run(ctx context.Context) {
if len(l.lists) == 0 {
return
}
for ctx.Err() == nil {
next := l.fetchDue(ctx)
timer := time.NewTimer(next.Sub(l.params.Now()))
select {
case <-ctx.Done():
case <-timer.C:
}
timer.Stop()
}
}
// Snapshot returns each list that has been tried, with its copy, if it
// has one, sorted by URL, as reputation.json lists them.
func (l *Lists) Snapshot() []List {
l.mu.Lock()
tried := make([]List, 0, len(l.lists))
for _, held := range l.lists {
if !held.kept.Tried.IsZero() {
tried = append(tried, held.kept)
}
}
l.mu.Unlock()
slices.SortFunc(tried, func(a, b List) int {
return strings.Compare(a.URL, b.URL)
})
return tried
}
// Load puts lists, read from reputation.json, in place of the last tries
// and copies held. A list Params does not name is dropped. A copy with a
// line that parse refuses is an error, and then nothing changes.
func (l *Lists) Load(lists []List) error {
found := make(map[string]entries, len(lists))
for _, kept := range lists {
if _, named := l.lists[kept.URL]; !named {
continue
}
read, err := l.parse(kept.URL, kept.Lines)
if err != nil {
return fmt.Errorf("the copy of %s: %w", kept.URL, err)
}
found[kept.URL] = read
}
l.mu.Lock()
defer l.mu.Unlock()
for listURL, held := range l.lists {
held.kept, held.entries = List{URL: listURL}, entries{}
}
for _, kept := range lists {
read, named := found[kept.URL]
if named {
l.lists[kept.URL].kept, l.lists[kept.URL].entries = kept, read
}
}
return nil
}
// fetchDue fetches each list that is due, one after another, and returns
// when the next is due. Once ctx has ended, it starts none, since a fetch
// cut off is noted as a try.
func (l *Lists) fetchDue(ctx context.Context) time.Time {
var next time.Time
for _, listURL := range l.URLs() {
due := l.due(listURL)
if ctx.Err() == nil && !l.params.Now().Before(due) {
l.fetch(ctx, listURL)
due = l.due(listURL)
}
if next.IsZero() || due.Before(next) {
next = due
}
}
return next
}
// due returns when the list at listURL is to be fetched: Refresh after it
// was last fetched or tried, the later of the two.
func (l *Lists) due(listURL string) time.Time {
l.mu.Lock()
defer l.mu.Unlock()
held := l.lists[listURL]
last := held.kept.Fetched
if held.kept.Tried.After(last) {
last = held.kept.Tried
}
return last.Add(l.params.Refresh)
}
// fetch fetches the list at listURL, and notes the try. A good copy takes
// the place of the one held. A failure leaves that in use, and is counted,
// logged and raised as a source_failure alert. A fetch cut off as ctx
// ends, as smallwebwaf stops, is no failure, but is still noted as a try,
// so that a restart waits for it: the server may have had its request.
func (l *Lists) fetch(ctx context.Context, listURL string) {
lines, err := l.get(ctx, listURL)
var found entries
if err == nil {
found, err = l.parse(listURL, lines)
}
cutOff := err != nil && ctx.Err() != nil
now := l.params.Now()
l.mu.Lock()
held := l.lists[listURL]
held.kept.Tried = now
if err == nil {
held.kept.Fetched, held.kept.Lines = now, lines
held.entries = found
} else if !cutOff {
held.failures++
}
l.mu.Unlock()
if cutOff {
return
}
if err != nil {
const failed = "fetching a list failed"
// Raised before it is logged, so that the alert is there once the
// log line is.
l.params.Alerts.Raise(alerts.Alert{
Event: alerts.EventSourceFailure,
Reason: failed,
Detail: map[string]any{"source": listURL, "error": err.Error()},
})
l.params.ProcessLog.Warn(failed, "url", listURL, "error", err.Error())
return
}
l.params.ProcessLog.Info("fetched a list", "url", listURL, "lines", len(lines))
}
// get fetches the list at listURL, and returns its lines. An answer other
// than 200, or a list longer than maxListBytes, is a failure.
func (l *Lists) get(ctx context.Context, listURL string) ([]string, error) {
ctx, cancel := context.WithTimeout(ctx, fetchTimeout)
defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, listURL, http.NoBody)
if err != nil {
return nil, fmt.Errorf("make the request: %w", err)
}
res, err := l.httpClient.Do(req)
if err != nil {
// Do's error names the URL, which the log line and the alert name
// already: only what went wrong is kept.
return nil, fmt.Errorf("fetch the list: %w", errors.Unwrap(err))
}
defer func() {
_ = res.Body.Close()
}()
if res.StatusCode != http.StatusOK {
return nil, fmt.Errorf("%w %s", errStatus, res.Status)
}
body, err := io.ReadAll(io.LimitReader(res.Body, maxListBytes+1))
if err != nil {
return nil, fmt.Errorf("read the list: %w", err)
}
if len(body) > maxListBytes {
return nil, errTooLong
}
lines := []string{}
for line := range strings.Lines(string(body)) {
lines = append(lines, strings.TrimSuffix(line, "\n"))
}
return lines, nil
}
// parse reads the lines of the list at listURL: those of a blocklist, or
// of the file of AS:percent lines. Anything after a ; or a # on a line is
// left out, and so is a line left blank. Any other line that does not read
// is an error naming it by its number.
func (l *Lists) parse(listURL string, lines []string) (entries, error) {
if listURL == l.params.ASNLimitPercentURL {
return parsePercents(lines)
}
return parseNetblocks(lines)
}
// parseNetblocks reads a blocklist's lines, each an address or a netblock
// as the settings take them.
func parseNetblocks(lines []string) (entries, error) {
found := entries{netblocks: map[netip.Prefix]bool{}}
for i, line := range lines {
text := withoutComment(line)
if text == "" {
continue
}
netblock, ok := parseNetblock(text)
if !ok {
return entries{}, fmt.Errorf("line %d %w", i+1, errNotNetblock)
}
found.netblocks[netblock] = true
if !slices.Contains(found.lengths, netblock.Bits()) {
found.lengths = append(found.lengths, netblock.Bits())
}
}
return found, nil
}
// parseNetblock reads text, a line of a blocklist, and reports whether it
// is an address or a netblock as the settings take them. A client's IPv4
// address is checked as IPv4, never IPv4-mapped, so an IPv4-mapped line,
// such as ::ffff:192.0.2.0/120, is read as the IPv4 address or netblock it
// stands for, 192.0.2.0/24, and a mapped netblock shorter than /96, which
// stands for none, is refused.
func parseNetblock(text string) (netip.Prefix, bool) {
netblock, err := config.ParseNetblock(text)
if err != nil {
return netip.Prefix{}, false
}
// The address as written: ParseNetblock's has the bits past the
// netblock's length cleared, the ::ffff among them below /96.
written, _, _ := strings.Cut(text, "/")
if addr, _ := netip.ParseAddr(written); !addr.Is4In6() {
return netblock, true
}
if netblock.Bits() < mappedBits {
return netip.Prefix{}, false
}
return netip.PrefixFrom(netblock.Addr().Unmap(), netblock.Bits()-mappedBits), true
}
// parsePercents reads the lines of the file of AS:percent lines, each an
// AS number, : and a percentage, as SWWAF_ASN_LIMIT_PERCENT takes them. An
// AS number listed more than once gets the lowest of its percentages.
func parsePercents(lines []string) (entries, error) {
found := entries{percents: map[string]int64{}}
for i, line := range lines {
text := withoutComment(line)
if text == "" {
continue
}
asnText, percentText, _ := strings.Cut(text, ":")
asn, asnErr := config.ParseASN(asnText)
percent, percentErr := config.ParsePercent(percentText)
if asnErr != nil || percentErr != nil {
return entries{}, fmt.Errorf("line %d %w", i+1, errNotASNPercent)
}
earlier, listed := found.percents[asn]
if !listed || percent < earlier {
found.percents[asn] = percent
}
}
return found, nil
}
// withoutComment returns line without anything after a ; or a #, and
// without the spaces around what is left.
func withoutComment(line string) string {
text, _, _ := strings.Cut(line, ";")
text, _, _ = strings.Cut(text, "#")
return strings.TrimSpace(text)
}
// contain reports whether the netblocks of a blocklist's copy hold addr:
// whether addr, cut to one of their lengths, is one of them.
func (e entries) contain(addr netip.Addr) bool {
for _, length := range e.lengths {
netblock, err := addr.Prefix(length)
if err == nil && e.netblocks[netblock] {
return true
}
}
return false
}
-610
View File
@@ -1,610 +0,0 @@
package reputation_test
import (
"bytes"
"context"
"fmt"
"io"
"log/slog"
"net/http"
"net/netip"
"net/url"
"reflect"
"slices"
"strings"
"sync"
"testing"
"testing/synctest"
"time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/reputation"
)
// The tests run in a synctest bubble, where the time package runs on a
// clock of the test's own: time.Sleep moves it on at once, and
// synctest.Wait returns once Run waits for the next list to be due, so
// that every fetch due by then has been made. The stand-in for the
// servers the lists are fetched from answers without the network, since a
// fetch waiting on the network would keep that clock from moving on.
const (
// dropURL and torURL are the blocklists, and asnURL the file of
// AS:percent lines.
dropURL = "https://lists.example/drop.txt"
torURL = "https://lists.example/tor.txt"
asnURL = "https://lists.example/asn.txt"
// refresh is the tests' SWWAF_BLOCKLIST_REFRESH, and cooldown their
// SWWAF_ALERT_COOLDOWN, longer than it.
refresh = 24 * time.Hour
cooldown = 48 * time.Hour
// drop is a blocklist as the Spamhaus DROP list is written, with an
// address and a netblock in each of its comments, which list nothing.
drop = "; Spamhaus DROP List 2026/10/07 - (c) 2026 The Spamhaus Project SLL\n" +
"; Last-Modified: Wed, 07 Oct 2026 00:00:00 GMT ; 192.0.2.1\n" +
"# 198.51.100.0/24\n" +
"\n" +
"203.0.113.0/24 ; SBL1\n" +
" 192.0.2.9 # one address\n" +
"2001:db8:1::/48 ; SBL2\n"
)
func TestListedAddressesAndNetblocksWithTheCommentsLeftOut(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{dropURL: drop}}
lists := start(t, servers, params(dropURL))
for addr, want := range map[string][]string{
"203.0.113.0": {dropURL},
"203.0.113.255": {dropURL},
"192.0.2.9": {dropURL},
"2001:db8:1::7": {dropURL},
"203.0.114.0": nil,
"192.0.2.8": nil,
"192.0.2.1": nil,
"198.51.100.7": nil,
"2001:db8:2::7": nil,
} {
wantListedBy(t, lists, addr, want...)
}
})
}
func TestIPv4MappedLineListsTheIPv4AddressOrNetblockItStandsFor(t *testing.T) {
t.Parallel()
now := time.Now()
lists := reputation.New(params(dropURL))
err := lists.Load([]reputation.List{{
URL: dropURL, Tried: now, Fetched: now,
Lines: []string{"::ffff:192.0.2.9", "::ffff:203.0.113.0/120"},
}})
if err != nil {
t.Fatalf("load: %v", err)
}
for addr, want := range map[string][]string{
"192.0.2.9": {dropURL},
"203.0.113.255": {dropURL},
"192.0.2.8": nil,
"203.0.114.0": nil,
} {
wantListedBy(t, lists, addr, want...)
}
// A mapped netblock shorter than /96 stands for no IPv4 one.
err = lists.Load([]reputation.List{{
URL: dropURL, Tried: now, Fetched: now, Lines: []string{"::ffff:198.51.100.0/88"},
}})
const want = "the copy of " + dropURL +
": line 1 is not an address or a netblock, such as 192.0.2.0/24"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
}
func TestClientIsListedByEachBlocklistThatListsIt(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{
dropURL: "203.0.113.0/24\n", torURL: "203.0.113.9\n",
}}
lists := start(t, servers, params(torURL, dropURL))
// In the order SWWAF_BLOCKLIST_URLS names them.
wantListedBy(t, lists, "203.0.113.9", torURL, dropURL)
wantListedBy(t, lists, "203.0.113.8", dropURL)
})
}
func TestListFetchedAgainOnceRefreshHasPassed(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{dropURL: "203.0.113.9\n"}}
lists := start(t, servers, params(dropURL))
began := time.Now()
wantFetches(t, servers, 1)
servers.set(dropURL, "203.0.113.10\n")
time.Sleep(refresh - time.Nanosecond)
wantFetches(t, servers, 1)
wantListedBy(t, lists, "203.0.113.9", dropURL)
time.Sleep(time.Nanosecond)
wantFetches(t, servers, 2)
wantListedBy(t, lists, "203.0.113.9")
wantListedBy(t, lists, "203.0.113.10", dropURL)
if fetched := lists.Fetched(dropURL); !fetched.Equal(began.Add(refresh)) {
t.Errorf("the copy in use was fetched at %s, want %s", fetched,
began.Add(refresh))
}
})
}
func TestFailedFetchKeepsTheLastGoodCopyAndAlertsOncePerCooldown(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
// fail has the stand-in answer the fetches after the first so that
// they fail with error.
fail func(servers *standIn)
error string
}{
{
"an answer other than 200",
func(servers *standIn) { servers.set(dropURL, "") },
"the server answered 503 Service Unavailable",
},
{
"a line that does not read",
func(servers *standIn) { servers.set(dropURL, "203.0.113.10\n<html>\n") },
"line 2 is not an address or a netblock, such as 192.0.2.0/24",
},
{
"a list longer than 16 MiB",
func(servers *standIn) {
servers.set(dropURL, strings.Repeat("#\n", 8<<20+1))
},
"the list is longer than 16 MiB",
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
var log bytes.Buffer
servers := &standIn{lists: map[string]string{dropURL: drop}}
queue := newQueue()
p := params(dropURL)
p.ProcessLog = slog.New(slog.NewJSONHandler(&log, nil))
p.Alerts = queue
lists := start(t, servers, p)
kept := lists.Snapshot()
tc.fail(servers)
// Each failure is tried again once refresh has passed since it.
for range 2 {
time.Sleep(refresh)
synctest.Wait()
}
wantFetches(t, servers, 3)
wantListedBy(t, lists, "203.0.113.9", dropURL)
want := kept[0]
want.Tried = time.Now()
if got := lists.Snapshot(); !reflect.DeepEqual(got, []reputation.List{want}) {
t.Errorf("lists %+v, want the first copy, last tried now, %+v", got, want)
}
if lists.Failures(dropURL) != 2 {
t.Errorf("%d failures, want 2", lists.Failures(dropURL))
}
// One alert for the first failure; the cooldown holds back the
// second.
wantAlert(t, queue, alerts.Alert{
Time: time.Now().Add(-refresh),
Event: alerts.EventSourceFailure,
Reason: "fetching a list failed",
Detail: map[string]any{"source": dropURL, "error": tc.error},
})
if !strings.Contains(log.String(), `"msg":"fetching a list failed",`+
`"url":"`+dropURL+`","error":"`+tc.error) {
t.Errorf("logged\n%s\nwant the failures", log.String())
}
})
})
}
}
func TestFetchNotDoneWithinAMinuteFails(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{}, hanging: true}
lists := start(t, servers, params(dropURL))
time.Sleep(time.Minute - time.Nanosecond)
synctest.Wait()
if lists.Failures(dropURL) != 0 {
t.Errorf("%d failures before a minute, want none", lists.Failures(dropURL))
}
time.Sleep(time.Nanosecond)
synctest.Wait()
if lists.Failures(dropURL) != 1 {
t.Errorf("%d failures after a minute, want 1", lists.Failures(dropURL))
}
})
}
func TestKeptCopyIsFetchedAgainOnceRefreshHasPassedSinceItWasFetched(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{
dropURL: "203.0.113.10\n", torURL: "198.51.100.10\n",
}}
lists := reputation.New(params(dropURL, torURL))
lists.SetTransport(servers)
// drop.txt was fetched an hour ago, and tor.txt a refresh ago, as
// reputation.json says at start.
err := lists.Load([]reputation.List{
{URL: dropURL, Fetched: time.Now().Add(-time.Hour), Lines: []string{"203.0.113.9"}},
{URL: torURL, Fetched: time.Now().Add(-refresh), Lines: []string{"198.51.100.9"}},
})
if err != nil {
t.Fatalf("load: %v", err)
}
run(t, lists)
wantFetches(t, servers, 1)
wantListedBy(t, lists, "203.0.113.9", dropURL)
wantListedBy(t, lists, "198.51.100.10", torURL)
time.Sleep(refresh - time.Hour - time.Nanosecond)
wantFetches(t, servers, 1)
time.Sleep(time.Nanosecond)
wantFetches(t, servers, 2)
wantListedBy(t, lists, "203.0.113.10", dropURL)
})
}
func TestRestartWaitsRefreshAfterTheLastTryEvenOneThatFailed(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
// drop.txt is fetched, and a refresh later the fetch downloads it
// whole but fails on a line that does not read.
servers := &standIn{lists: map[string]string{dropURL: "198.51.100.1\n"}}
lists := start(t, servers, params(dropURL))
servers.set(dropURL, "198.51.100.2\n<html>\n")
time.Sleep(refresh)
wantFetches(t, servers, 2)
// Restarted with what reputation.json keeps, it waits a refresh
// after the failed try, as it does while it runs.
restarted := &standIn{lists: map[string]string{dropURL: "198.51.100.2\n"}}
again := reputation.New(params(dropURL))
again.SetTransport(restarted)
err := again.Load(lists.Snapshot())
if err != nil {
t.Fatalf("load: %v", err)
}
run(t, again)
wantFetches(t, restarted, 0)
wantListedBy(t, again, "198.51.100.1", dropURL)
time.Sleep(refresh - time.Nanosecond)
wantFetches(t, restarted, 0)
time.Sleep(time.Nanosecond)
wantFetches(t, restarted, 1)
wantListedBy(t, again, "198.51.100.2", dropURL)
})
}
func TestFetchCutOffAsItStopsIsNoFailureButARestartWaitsForIt(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
// Stopped 30 seconds into the fetch of drop.txt, before tor.txt's.
servers := &standIn{lists: map[string]string{}, hanging: true}
queue := newQueue()
p := params(dropURL, torURL)
p.Alerts = queue
lists := reputation.New(p)
lists.SetTransport(servers)
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
lists.Run(ctx)
close(stopped)
}()
time.Sleep(30 * time.Second)
stop()
<-stopped
wantFetches(t, servers, 1)
if lists.Failures(dropURL) != 0 || len(waiting(queue)) != 0 {
t.Errorf("%d failures and alerts %+v, want none", lists.Failures(dropURL),
waiting(queue))
}
// Restarted an hour later with what reputation.json keeps, it fetches
// tor.txt, never tried, at once, and drop.txt a refresh after its
// cut-off try.
time.Sleep(time.Hour)
restarted := &standIn{lists: map[string]string{
dropURL: "203.0.113.7\n", torURL: "198.51.100.7\n",
}}
again := reputation.New(params(dropURL, torURL))
again.SetTransport(restarted)
err := again.Load(lists.Snapshot())
if err != nil {
t.Fatalf("load: %v", err)
}
run(t, again)
wantFetches(t, restarted, 1)
wantListedBy(t, again, "198.51.100.7", torURL)
time.Sleep(refresh - time.Hour - time.Nanosecond)
wantFetches(t, restarted, 1)
time.Sleep(time.Nanosecond)
wantFetches(t, restarted, 2)
wantListedBy(t, again, "203.0.113.7", dropURL)
})
}
func TestASNLimitPercentFileGivesEachASNumberItsLowestPercentage(t *testing.T) {
t.Parallel()
synctest.Test(t, func(t *testing.T) {
servers := &standIn{lists: map[string]string{
asnURL: "# hosting networks\nAS14061:50 ; DigitalOcean\nas16276:25\n\n" +
"AS14061:10\nAS14061:30\n",
}}
p := params()
p.ASNLimitPercentURL = asnURL
lists := start(t, servers, p)
for asn, want := range map[string]int64{"AS14061": 10, "AS16276": 25} {
percent, listed := lists.ASNLimitPercent(asn)
if !listed || percent != want {
t.Errorf("%s has %d (listed %t), want %d", asn, percent, listed, want)
}
}
if _, listed := lists.ASNLimitPercent("AS64496"); listed {
t.Error("AS64496 is listed")
}
// A line that does not read fails the fetch.
servers.set(asnURL, "AS14061:50\nAS16276\n")
time.Sleep(refresh)
synctest.Wait()
if lists.Failures(asnURL) != 1 {
t.Errorf("%d failures, want 1", lists.Failures(asnURL))
}
})
}
func TestLoadDropsCopiesOfListsNotNamedAndRefusesOnesThatDoNotRead(t *testing.T) {
t.Parallel()
fetched := time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)
kept := reputation.List{
URL: dropURL, Tried: fetched, Fetched: fetched, Lines: []string{"203.0.113.9"},
}
lists := reputation.New(params(dropURL))
err := lists.Load([]reputation.List{kept, {
URL: torURL, Tried: fetched, Fetched: fetched, Lines: []string{"198.51.100.9"},
}})
if err != nil {
t.Fatalf("load: %v", err)
}
if got := lists.Snapshot(); !reflect.DeepEqual(got, []reputation.List{kept}) {
t.Errorf("copies %+v, want only %+v", got, kept)
}
err = lists.Load([]reputation.List{{URL: dropURL, Fetched: fetched, Lines: []string{
"; DROP", "203.0.113.300",
}}})
const want = "the copy of " + dropURL +
": line 2 is not an address or a netblock, such as 192.0.2.0/24"
if err == nil || err.Error() != want {
t.Errorf("error %v, want %s", err, want)
}
if got := lists.Snapshot(); !reflect.DeepEqual(got, []reputation.List{kept}) {
t.Errorf("copies %+v after the error, want %+v still", got, kept)
}
}
// standIn is a stand-in for the servers the lists are fetched from. It
// notes the URL of each fetch.
type standIn struct {
mu sync.Mutex
// lists are what it answers with, by URL; it answers a URL it has no
// list for with 503, and none at all while hanging.
lists map[string]string
hanging bool
fetches []string
}
// RoundTrip has the stand-in answer req, in place of the network. A fetch
// abandoned before the stand-in answers fails, as over the network.
func (s *standIn) RoundTrip(req *http.Request) (*http.Response, error) {
s.mu.Lock()
s.fetches = append(s.fetches, req.URL.String())
list, found := s.lists[req.URL.String()]
hanging := s.hanging
s.mu.Unlock()
if hanging {
<-req.Context().Done()
return nil, req.Context().Err()
}
status := http.StatusOK
if !found {
status = http.StatusServiceUnavailable
}
return &http.Response{
StatusCode: status,
Status: fmt.Sprintf("%d %s", status, http.StatusText(status)),
Header: http.Header{},
Body: io.NopCloser(strings.NewReader(list)),
Request: req,
}, nil
}
// set has the stand-in answer listURL with list, or with 503 for "".
func (s *standIn) set(listURL, list string) {
s.mu.Lock()
defer s.mu.Unlock()
if list == "" {
delete(s.lists, listURL)
return
}
s.lists[listURL] = list
}
// params returns the Params of the blocklists at urls, refreshed every
// refresh, by the bubble's clock, with alerts to a queue that sends none.
func params(urls ...string) reputation.Params {
return reputation.Params{
BlocklistURLs: urls,
Refresh: refresh,
Now: time.Now,
ProcessLog: slog.New(slog.DiscardHandler),
Alerts: newQueue(),
}
}
// newQueue returns a queue of alerts to a webhook that is never sent
// them, with a cooldown of cooldown.
func newQueue() *alerts.Queue {
return alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: cooldown,
Now: time.Now,
})
}
// start returns the lists of p, fetched through servers by Run, which runs
// until the test ends, once Run has fetched those due at start.
func start(t *testing.T, servers *standIn, p reputation.Params) *reputation.Lists {
t.Helper()
lists := reputation.New(p)
lists.SetTransport(servers)
run(t, lists)
return lists
}
// run runs lists' Run until the test ends, and waits until it has fetched
// the lists due.
func run(t *testing.T, lists *reputation.Lists) {
t.Helper()
ctx, stop := context.WithCancel(t.Context())
stopped := make(chan struct{})
go func() {
lists.Run(ctx)
close(stopped)
}()
t.Cleanup(func() {
stop()
<-stopped
})
synctest.Wait()
}
// wantFetches waits until Run has made the fetches due, and checks how
// many the servers have had.
func wantFetches(t *testing.T, servers *standIn, want int) {
t.Helper()
synctest.Wait()
servers.mu.Lock()
got := len(servers.fetches)
servers.mu.Unlock()
if got != want {
t.Errorf("%d fetches, want %d", got, want)
}
}
// wantListedBy checks the URLs of the blocklists lists says list addr.
func wantListedBy(t *testing.T, lists *reputation.Lists, addr string, want ...string) {
t.Helper()
got := lists.ListedBy(netip.MustParseAddr(addr))
if !slices.Equal(got, want) {
t.Errorf("%s is listed by %v, want %v", addr, got, want)
}
}
// waiting returns the alerts waiting in queue.
func waiting(queue *alerts.Queue) []alerts.Alert {
return queue.Snapshot().Waiting[alerts.DestinationWebhook]
}
// wantAlert checks that want is the one alert waiting in queue, and that
// the cooldown has held back one repeat of it.
func wantAlert(t *testing.T, queue *alerts.Queue, want alerts.Alert) {
t.Helper()
got := waiting(queue)
if len(got) != 1 || !reflect.DeepEqual(got[0], want) || queue.Suppressed() != 1 {
t.Errorf("alerts waiting %+v, %d held back, want only %+v and 1", got,
queue.Suppressed(), want)
}
}
+10 -29
View File
@@ -35,8 +35,7 @@ const (
// rule. // rule.
ActionRuleBlocked = "rule_blocked" ActionRuleBlocked = "rule_blocked"
// ActionDenied is a request refused because its client is in // ActionDenied is a request refused because its client is in
// SWWAF_DENY_NETS, in a blocklist while SWWAF_BLOCKLIST_ACTION is deny, // SWWAF_DENY_NETS.
// or listed by a DNSBL zone while SWWAF_REPUTATION_ACTION is deny.
ActionDenied = "denied" ActionDenied = "denied"
// ActionCountryDenied is a request refused for its client's country. // ActionCountryDenied is a request refused for its client's country.
ActionCountryDenied = "country_denied" ActionCountryDenied = "country_denied"
@@ -46,7 +45,7 @@ const (
) )
// OffenceLimit is the offence a request line names for a request that // OffenceLimit is the offence a request line names for a request that
// broke a rate limit, or whose bytes broke a byte limit. // broke a rate limit.
const OffenceLimit = "limit" const OffenceLimit = "limit"
// timeLayout is RFC 3339 with milliseconds. // timeLayout is RFC 3339 with milliseconds.
@@ -81,14 +80,11 @@ type Line struct {
// Request detail. RequestID is the X-Request-ID a trusted proxy sent, // Request detail. RequestID is the X-Request-ID a trusted proxy sent,
// or a new one, and is sent on to the app. ForwardedFor is the // or a new one, and is sent on to the app. ForwardedFor is the
// X-Forwarded-For header as received. ClientGroup is the netblock the // X-Forwarded-For header as received. ClientGroup is the netblock the
// client is counted as. ASN, ASName and Country are the client's AS // client is counted as.
// number, AS name and country, as looked up.
RequestID string `json:"request_id"` RequestID string `json:"request_id"`
PeerIP string `json:"peer_ip"` PeerIP string `json:"peer_ip"`
ForwardedFor string `json:"forwarded_for,omitempty"` ForwardedFor string `json:"forwarded_for,omitempty"`
ClientGroup string `json:"client_group"` ClientGroup string `json:"client_group"`
ASN string `json:"asn"`
ASName string `json:"as_name"`
Country string `json:"country"` Country string `json:"country"`
ContentType string `json:"content_type,omitempty"` ContentType string `json:"content_type,omitempty"`
// ContentLength is the length of its body the request announced. // ContentLength is the length of its body the request announced.
@@ -118,29 +114,14 @@ type Line struct {
// ActionBanned, ActionCountryDenied, ActionRateLimited or // ActionBanned, ActionCountryDenied, ActionRateLimited or
// ActionRuleBlocked. // ActionRuleBlocked.
WouldAction string `json:"would_action,omitempty"` WouldAction string `json:"would_action,omitempty"`
// LimitPercent and LimitPercentSetting are, for a request the rate // Counts are the client's requests as the rate limits counted them
// limits counted whose client a biased threshold gives a percentage of // with this one, for a request they counted.
// the rate limits below 100, that percentage and the setting that gave
// it. BytesPercent and BytesPercentSetting are the same for the byte
// limits.
LimitPercent *int64 `json:"limit_percent,omitempty"`
LimitPercentSetting string `json:"limit_percent_setting,omitempty"`
BytesPercent *int64 `json:"bytes_percent,omitempty"`
BytesPercentSetting string `json:"bytes_percent_setting,omitempty"`
// Counts are, for a request the rate limits counted, the client's
// requests as they counted them with this one, and its bytes as the
// byte limits counted them, with this request's once it has ended if
// they count them.
Counts ratelimit.Counts `json:"counts,omitzero"` Counts ratelimit.Counts `json:"counts,omitzero"`
// RuleIDs are the ids of the rule file rules the request matched. // RuleIDs are the ids of the rule file rules the request matched.
RuleIDs []string `json:"rule_ids,omitempty"` RuleIDs []string `json:"rule_ids,omitempty"`
// LimitHit is the window whose limit the request went over, named as // LimitHit is the window whose rate limit the request went over:
// Counts names its count: minute, hour or day for a rate limit, and // minute, hour or day.
// minute_bytes, hour_bytes or day_bytes for a byte limit.
LimitHit string `json:"limit_hit,omitempty"` LimitHit string `json:"limit_hit,omitempty"`
// Reputation are the URLs of the blocklists that list the client, then
// the DNSBL zones whose verdict lists it, their keys masked.
Reputation []string `json:"reputation,omitempty"`
// Offence is the offence the request was held as, OffenceLimit. // Offence is the offence the request was held as, OffenceLimit.
Offence string `json:"offence,omitempty"` Offence string `json:"offence,omitempty"`
// BanExpires is when the ban the request made, or was refused under, // BanExpires is when the ban the request made, or was refused under,
@@ -190,8 +171,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, and instanceName, SWWAF_INSTANCE_NAME, as instance. // as a request line's.
func NewProcessLogger(w io.Writer, instanceName string) *slog.Logger { func NewProcessLogger(w io.Writer) *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 {
@@ -202,5 +183,5 @@ func NewProcessLogger(w io.Writer, instanceName string) *slog.Logger {
}, },
}) })
return slog.New(handler).With("type", "process", "instance", instanceName) return slog.New(handler).With("type", "process")
} }
+4 -5
View File
@@ -65,12 +65,12 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
} }
} }
func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) { func TestProcessLinesAreMarkedProcess(t *testing.T) {
t.Parallel() t.Parallel()
var out bytes.Buffer var out bytes.Buffer
requestlog.NewProcessLogger(&out, "fsn1app1/gitea").Info("starting", "version", "v1") requestlog.NewProcessLogger(&out).Info("starting", "version", "v1")
var fields map[string]any var fields map[string]any
@@ -79,9 +79,8 @@ func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
t.Fatalf("decode %q: %v", out.String(), err) t.Fatalf("decode %q: %v", out.String(), err)
} }
if fields["type"] != "process" || fields["instance"] != "fsn1app1/gitea" || if fields["type"] != "process" || fields["msg"] != "starting" ||
fields["msg"] != "starting" || fields["level"] != "INFO" || fields["level"] != "INFO" || fields["version"] != "v1" {
fields["version"] != "v1" {
t.Errorf("process line %v", fields) t.Errorf("process line %v", fields)
} }
+14 -28
View File
@@ -22,7 +22,6 @@ 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"
) )
@@ -101,8 +100,6 @@ 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
@@ -129,7 +126,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
} }
@@ -226,21 +223,13 @@ func (f *Files) readAfterChanges(
} }
// readAgain reads the rule files again, in place of the rules loaded, or // readAgain reads the rule files again, in place of the rules loaded, or
// logs the error that keeps the rules as they were, and raises a // logs the error that keeps the rules as they were.
// file_error alert for it, for the file it is in.
func (f *Files) readAgain() { func (f *Files) readAgain() {
rules, path, err := read(f.params.Dir) rules, err := read(f.params.Dir)
if err != nil { if err != nil {
const kept = "a rule file has an error, and the rules stay as they were" f.params.ProcessLog.Error(
"a rule file has an error, and the rules stay as they were",
// Raised before it is logged, so that the alert is there once the "error", err.Error())
// 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
} }
@@ -257,14 +246,13 @@ func (f *Files) logRead(count int) {
} }
// read returns the rules of every rule file in dir, in the order of the // read returns the rules of every rule file in dir, in the order of the
// files' names, and then of their lines, or an error, with the path of the // files' names, and then of their lines. A file whose name starts with a
// rule file it is in, or dir. A file whose name starts with a dot, such as // dot, such as an editor's lock file .#50-app.rules, is not a rule file,
// an editor's lock file .#50-app.rules, is not a rule file, as a shell's // as a shell's *.rules would not match it.
// *.rules would not match it. func read(dir string) ([]Rule, error) {
func read(dir string) ([]Rule, string, error) {
entries, err := os.ReadDir(dir) entries, err := os.ReadDir(dir)
if err != nil { if err != nil {
return nil, dir, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err) return nil, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err)
} }
var rules []Rule var rules []Rule
@@ -278,15 +266,13 @@ func read(dir string) ([]Rule, string, error) {
continue continue
} }
path := filepath.Join(dir, name) rules, err = readFile(filepath.Join(dir, name), rules, places)
rules, err = readFile(path, rules, places)
if err != nil { if err != nil {
return nil, path, err return nil, err
} }
} }
return rules, "", nil return rules, nil
} }
// readFile appends the rules of the rule file at path to rules. places // readFile appends the rules of the rule file at path to rules. places
+7 -32
View File
@@ -7,15 +7,12 @@ 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"
) )
@@ -340,7 +337,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 }
@@ -369,7 +366,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, queue := watch(t, dir) files, lines := 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.
@@ -383,27 +380,13 @@ func TestBrokenEditKeepsTheRulesAsTheyWere(t *testing.T) {
t.Errorf("logged %v, want an error %q", line, want) t.Errorf("logged %v, want an error %q", line, want)
} }
// The error is raised as a file_error alert too, for the file.
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, and raises no alert. // Once mended, the file is read again.
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) {
@@ -521,8 +504,7 @@ 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, and the alerts waiting in a // the process log in the processLog returned.
// 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)
@@ -530,12 +512,6 @@ 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
} }
@@ -555,9 +531,8 @@ 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. It returns the alerts' queue as // waits until it watches the directory.
// well. func watch(t *testing.T, dir string) (*rules.Files, processLog) {
func watch(t *testing.T, dir string) (*rules.Files, processLog, *alerts.Queue) {
t.Helper() t.Helper()
params, lines := newParams(dir) params, lines := newParams(dir)
@@ -582,7 +557,7 @@ func watch(t *testing.T, dir string) (*rules.Files, processLog, *alerts.Queue) {
lines.waitFor(t, watching) lines.waitFor(t, watching)
return files, lines, params.Alerts return files, lines
} }
// 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,8 +13,6 @@ 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
@@ -93,7 +91,6 @@ 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)
+53 -146
View File
@@ -1,7 +1,6 @@
// Package smallwebwaf runs the smallwebwaf process: it reads the settings, // Package smallwebwaf runs the smallwebwaf process: it reads the settings,
// the rule files, the lookup database and the state files, serves requests // the rule files and the state files, serves requests until it is told to
// until it is told to stop, and then stops in an orderly way, writing the // stop, and then stops in an orderly way, writing the state files.
// state files.
package smallwebwaf package smallwebwaf
import ( import (
@@ -16,7 +15,6 @@ 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"
@@ -65,12 +63,11 @@ func Main(version string) int {
}) })
} }
// Run reads the settings, the rule files, the lookup database and the // Run reads the settings, the rule files and the state files, then serves
// state files, then serves requests until ctx is done. It returns the // requests until ctx is done. It returns the process's exit status, 1
// process's exit status, 1 when smallwebwaf cannot start. // when smallwebwaf cannot start.
func Run(ctx context.Context, params Params) int { func Run(ctx context.Context, params Params) int {
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 {
@@ -88,22 +85,16 @@ 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, cfg.InstanceName) processLog = requestlog.NewProcessLogger(stdout)
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())
@@ -111,18 +102,32 @@ func Run(ctx context.Context, params Params) int {
return 1 return 1
} }
server, err := newServer(cfg, stdout, processLog, now, ruleFiles, alertQueue) // The state files give times in UTC.
if err != nil { now := func() time.Time { return time.Now().UTC() }
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 := loadStateFiles(cfg, server, alertQueue, now, processLog) files, err := state.Load(state.Params{
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())
@@ -142,92 +147,7 @@ func Run(ctx context.Context, params Params) int {
"address", listener.Addr().String(), "address", listener.Addr().String(),
"settings", cfg) "settings", cfg)
return serve(ctx, server, listener, files, ruleFiles, alertQueue, processLog) return serve(ctx, server.Server, listener, files, ruleFiles, 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,
Lists: server.Lists,
DNSBL: server.DNSBL,
Alerts: alertQueue,
Anomalies: server.Anomalies,
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
@@ -268,16 +188,12 @@ func startSending(
} }
// serve serves requests on listener, writes the state files as they are // serve serves requests on listener, writes the state files as they are
// due, takes in an admin's edits of them, reads the rule files again as // due, takes in an admin's edits of them, and reads the rule files again
// they change, and the lookup database when it is replaced, fetches the // as they change, until ctx is done. Then it gives the requests in
// lists the settings name by URL as they are due, and sends the alerts, // progress shutdownTimeout to finish, and writes every state file.
// 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 *proxy.Server, listener net.Listener, ctx context.Context, server *http.Server, listener net.Listener,
files *state.Files, ruleFiles *rules.Files, alertQueue *alerts.Queue, files *state.Files, ruleFiles *rules.Files, processLog *slog.Logger,
processLog *slog.Logger,
) int { ) int {
served := make(chan error, 1) served := make(chan error, 1)
@@ -288,16 +204,24 @@ func serve(
writing, stopWriting := context.WithCancel(ctx) writing, stopWriting := context.WithCancel(ctx)
defer stopWriting() defer stopWriting()
written := inBackground(func() { files.Run(writing) }) written := make(chan struct{})
watched := inBackground(func() { files.Watch(writing) }) watched := make(chan struct{})
rulesWatched := inBackground(func() { ruleFiles.Watch(writing) }) rulesWatched := make(chan struct{})
lookupFileWatched := inBackground(func() {
if server.LookupFile != nil { go func() {
server.LookupFile.Watch(writing) files.Run(writing)
} close(written)
}) }()
listsFetched := inBackground(func() { server.Lists.Run(writing) })
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:
@@ -329,8 +253,7 @@ 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, and no alert is being sent, so that alerts.json keeps every // files. Every request has ended too, but for two kinds
// 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
@@ -339,9 +262,6 @@ func serve(
<-written <-written
<-watched <-watched
<-rulesWatched <-rulesWatched
<-lookupFileWatched
<-listsFetched
<-alertsSent
err = files.WriteAll() err = files.WriteAll()
if err != nil { if err != nil {
@@ -354,16 +274,3 @@ 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
}
+11 -641
View File
@@ -14,11 +14,9 @@ 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"
) )
@@ -38,22 +36,12 @@ const (
stateWriteDelay = "SWWAF_STATE_WRITE_DELAY" stateWriteDelay = "SWWAF_STATE_WRITE_DELAY"
stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL" stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL"
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
rulesDir = "SWWAF_RULES_DIR" rulesDir = "SWWAF_RULES_DIR"
lookupSource = "SWWAF_LOOKUP_SOURCE" adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
lookupDBPath = "SWWAF_LOOKUP_DB_PATH"
instanceName = "SWWAF_INSTANCE_NAME"
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name
// adminSecret is the SWWAF_ADMIN_TOKEN the tests set. // adminSecret is the SWWAF_ADMIN_TOKEN the tests set.
adminSecret = "fedcba9876543210fedcba9876543210" adminSecret = "fedcba9876543210fedcba9876543210"
// instance is the SWWAF_INSTANCE_NAME the tests set where they look at
// it.
instance = "fsn1app1/gitea"
// greeting is what the tests' app answers. // greeting is what the tests' app answers.
greeting = "hello from the app" greeting = "hello from the app"
// placed is the client the tests' lookup databases place.
placed = "203.0.113.9"
) )
// output collects what smallwebwaf writes on stdout. // output collects what smallwebwaf writes on stdout.
@@ -110,16 +98,12 @@ func (o *output) text() string {
} }
// run runs smallwebwaf with the settings in env until ctx is done, and // run runs smallwebwaf with the settings in env until ctx is done, and
// returns its exit status. SWWAF_LOOKUP_SOURCE is off unless env sets it, // returns its exit status.
// so that no test sends GeoJS its clients' addresses.
func run(ctx context.Context, env map[string]string, out *output) int { func run(ctx context.Context, env map[string]string, out *output) int {
return smallwebwaf.Run(ctx, smallwebwaf.Params{ return smallwebwaf.Run(ctx, smallwebwaf.Params{
Version: testVersion, Version: testVersion,
LookupEnv: func(name string) (string, bool) { LookupEnv: func(name string) (string, bool) {
value, ok := env[name] value, ok := env[name]
if !ok && name == "SWWAF_LOOKUP_SOURCE" {
return "off", true
}
return value, ok return value, ok
}, },
@@ -132,10 +116,7 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
out := &output{} out := &output{}
status := run(t.Context(), map[string]string{ status := run(t.Context(), map[string]string{"SWWAF_REQUEST_MAX_BYTES": "lots"}, out)
"SWWAF_REQUEST_MAX_BYTES": "lots",
instanceName: instance,
}, out)
if status != 1 { if status != 1 {
t.Errorf("exit status %d, want 1", status) t.Errorf("exit status %d, want 1", status)
} }
@@ -143,8 +124,7 @@ func TestInvalidSettingStopsTheStart(t *testing.T) {
line := out.line(t, "msg", "invalid setting") line := out.line(t, "msg", "invalid setting")
message, _ := line["error"].(string) message, _ := line["error"].(string)
if line["type"] != "process" || line["instance"] != instance || if line["type"] != "process" || line["level"] != "ERROR" ||
line["level"] != "ERROR" ||
!strings.HasPrefix(message, "SWWAF_REQUEST_MAX_BYTES: ") { !strings.HasPrefix(message, "SWWAF_REQUEST_MAX_BYTES: ") {
t.Errorf("start refused with %v", line) t.Errorf("start refused with %v", line)
} }
@@ -155,7 +135,7 @@ func TestShortTokenStopsTheStartUnshown(t *testing.T) {
const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use
for _, name := range []string{adminToken, metricsToken} { for _, name := range []string{adminToken, "SWWAF_METRICS_TOKEN"} {
t.Run(name, func(t *testing.T) { t.Run(name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -249,46 +229,6 @@ 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()
@@ -501,212 +441,6 @@ func TestRulesDirThatDoesNotExistStopsTheStart(t *testing.T) {
": no such file or directory") ": no such file or directory")
} }
func TestLookupDatabaseReplacedWhileRunningTakesEffect(t *testing.T) {
t.Parallel()
const token = "0123456789abcdef0123456789abcdef"
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
writeLookupDatabase(t, path, "DE")
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: t.TempDir(),
rulesDir: t.TempDir(),
trustedProxies: localhost + "/32",
metricsToken: token,
instanceName: instance,
lookupSource: "file",
lookupDBPath: path,
"SWWAF_DENIED_COUNTRIES": "kp",
// The requests sent until a replacement takes effect, and those
// for the metrics, must not break a rate limit, whose ban would
// refuse them too.
rateLimitExemptNets: placed + "," + localhost,
}
began := time.Now()
// Each replacement is written beside the file and renamed over it, as
// a refresh is.
replacement := path + ".new"
var lastRead float64
runUntilStopped(t, env, func(url string) {
wantStatus(t, url, placed, http.StatusOK)
// smallwebwaf reads the file again once it watches its directory,
// which may be after the first replacement. Once that has taken
// effect, only the watch can show the next. Each takes as long as
// it takes, so that a slow test process cannot fail the test.
writeLookupDatabase(t, replacement, "KP")
rename(t, replacement, path)
for statusFrom(t, url, placed) != http.StatusForbidden {
time.Sleep(pollInterval)
}
writeLookupDatabase(t, replacement, "DE")
rename(t, replacement, path)
for statusFrom(t, url, placed) != http.StatusOK {
time.Sleep(pollInterval)
}
// One that cannot be read leaves the file read before in use.
err := os.WriteFile(replacement, []byte("not a lookup database\n"), 0o600)
if err != nil {
t.Fatalf("write %s: %v", replacement, err)
}
rename(t, replacement, path)
metrics := metricsWith(t, url+"_smallwebwaf/metrics", token,
`smallwebwaf_lookup_database_read_failures_total{instance="fsn1app1/gitea"} 1`)
lastRead = seriesValue(t, metrics,
`smallwebwaf_lookup_database_last_read_timestamp_seconds{instance="fsn1app1/gitea"}`)
wantStatus(t, url, placed, http.StatusOK)
})
if lastRead < float64(began.Unix()) || lastRead > float64(time.Now().Unix()) {
t.Errorf("the lookup database was last read at %v, want a time in the test", lastRead)
}
}
func TestBlocklistTriesAndCopiesKeptInReputationJSONAcrossRestarts(t *testing.T) {
t.Parallel()
const (
token = "0123456789abcdef0123456789abcdef"
torPath = "/tor.txt"
dropPath = "/drop.txt"
// copyright is the DROP list's date and copyright line.
copyright = "; Spamhaus DROP List 2026/10/07 - (c) 2026 The Spamhaus Project SLL"
)
failing, torFetches := new(atomic.Bool), new(atomic.Int32)
lists := map[string]string{
torPath: "198.51.100.0/24\n",
dropPath: copyright + "\n" + placed + " ; SBL1\n",
}
server := httptest.NewServer(http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == torPath {
torFetches.Add(1)
}
if failing.Load() {
w.WriteHeader(http.StatusServiceUnavailable)
return
}
_, _ = io.WriteString(w, lists[r.URL.Path])
}))
t.Cleanup(server.Close)
torURL, dropURL := server.URL+torPath, server.URL+dropPath
dir := t.TempDir()
env := map[string]string{
listenAddr: localhost + ":0",
upstreamURL: startApp(t),
stateDir: dir,
rulesDir: t.TempDir(),
trustedProxies: localhost + "/32",
metricsToken: token,
instanceName: instance,
"SWWAF_BLOCKLIST_URLS": torURL,
// The requests sent until the list takes effect, and those for the
// metrics, must not break a rate limit, whose ban would refuse them
// too.
rateLimitExemptNets: placed + "," + localhost,
}
failures := `smallwebwaf_reputation_failures_total{instance="fsn1app1/gitea",` +
`source="` + torURL + `"} `
// tor.txt cannot be fetched at first, which is counted.
failing.Store(true)
runUntilStopped(t, env, func(url string) {
metricsWith(t, url+"_smallwebwaf/metrics", token, failures+"1")
})
// Restarted with drop.txt named after it and the server answering, tor.txt
// waits SWWAF_BLOCKLIST_REFRESH after its failed try, kept in reputation.json,
// while drop.txt, never tried, is fetched at once. Lists are fetched in the
// order named, so once drop.txt refuses the client, tor.txt has had its turn.
failing.Store(false)
env["SWWAF_BLOCKLIST_URLS"] = torURL + "," + dropURL
out := runUntilStopped(t, env, func(url string) {
for statusFrom(t, url, placed) != http.StatusForbidden {
time.Sleep(pollInterval)
}
})
wantDeniedByList(t, out.line(t, "action", "denied"), dropURL)
if fetches := torFetches.Load(); fetches != 1 {
t.Errorf("tor.txt fetched %d times, want once, before the restart", fetches)
}
// After another restart, with the server failing, the copy of drop.txt
// kept in reputation.json, its copyright line included, refuses the
// client from the first request.
failing.Store(true)
out = runUntilStopped(t, env, func(url string) {
wantStatus(t, url, placed, http.StatusForbidden)
})
wantDeniedByList(t, out.line(t, "type", "request"), dropURL)
path := filepath.Join(dir, "reputation.json")
kept, err := os.ReadFile(path) //nolint:gosec // a file in the test's directory
if err != nil || !strings.Contains(string(kept), `"`+copyright+`"`) {
t.Errorf("reputation.json holds\n%s\nwant the copy with %q (%v)", kept, copyright,
err)
}
}
// wantDeniedByList checks that the request log line is of a request the
// blocklist at listURL refused.
func wantDeniedByList(t *testing.T, line map[string]any, listURL string) {
t.Helper()
reputation, _ := line["reputation"].([]any)
if line["action"] != "denied" || len(reputation) != 1 || reputation[0] != listURL {
t.Errorf("request log line %v, want one denied for %s", line, listURL)
}
}
func TestLookupDatabaseThatCannotBeReadStopsTheStart(t *testing.T) {
t.Parallel()
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
// If it starts instead, it is stopped after waitLimit.
ctx, stop := context.WithTimeout(t.Context(), waitLimit)
defer stop()
out := &output{}
status := run(ctx, map[string]string{
listenAddr: localhost + ":0", stateDir: t.TempDir(), rulesDir: t.TempDir(),
lookupSource: "file", lookupDBPath: path,
}, out)
if status != 1 {
t.Fatalf("exit status %d, want 1", status)
}
want := "SWWAF_LOOKUP_DB_PATH cannot be read: open " + path +
": no such file or directory"
line := out.line(t, "msg", "cannot use the lookup database")
if line["error"] != want {
t.Errorf("start refused with %v, want the error %q", line, want)
}
}
func TestEveryLineIsAlsoSentToTheRemoteLogEndpoint(t *testing.T) { func TestEveryLineIsAlsoSentToTheRemoteLogEndpoint(t *testing.T) {
t.Parallel() t.Parallel()
@@ -782,8 +516,7 @@ func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
rulesDir: t.TempDir(), rulesDir: t.TempDir(),
"SWWAF_LOG_REMOTE_URL": "syslog+tls://" + endpoint.Addr().String(), "SWWAF_LOG_REMOTE_URL": "syslog+tls://" + endpoint.Addr().String(),
"SWWAF_LOG_REMOTE_BUFFER": "1", "SWWAF_LOG_REMOTE_BUFFER": "1",
metricsToken: token, "SWWAF_METRICS_TOKEN": token,
instanceName: instance,
} }
out := runUntilStopped(t, env, func(url string) { out := runUntilStopped(t, env, func(url string) {
@@ -793,17 +526,16 @@ func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
// last. // last.
metrics := metricsText(t, url+"_smallwebwaf/metrics", token) metrics := metricsText(t, url+"_smallwebwaf/metrics", token)
for _, series := range []string{ for _, series := range []string{
`smallwebwaf_remote_log_lines_sent_total{instance="fsn1app1/gitea"} 0`, "smallwebwaf_remote_log_lines_sent_total 0",
`smallwebwaf_remote_log_buffer_depth{instance="fsn1app1/gitea"} 1`, "smallwebwaf_remote_log_buffer_depth 1",
} { } {
if !strings.Contains(metrics, "\n"+series+"\n") { if !strings.Contains(metrics, "\n"+series+"\n") {
t.Errorf("no %q in the metrics:\n%s", series, metrics) t.Errorf("no %q in the metrics:\n%s", series, metrics)
} }
} }
const dropped = "\nsmallwebwaf_remote_log_lines_dropped_total" + if strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total 0\n") ||
`{instance="fsn1app1/gitea"} ` !strings.Contains(metrics, "\nsmallwebwaf_remote_log_lines_dropped_total ") {
if strings.Contains(metrics, dropped+"0\n") || !strings.Contains(metrics, dropped) {
t.Errorf("no line dropped in the metrics:\n%s", metrics) t.Errorf("no line dropped in the metrics:\n%s", metrics)
} }
@@ -813,170 +545,6 @@ func TestStalledRemoteLogEndpointHoldsUpNoRequest(t *testing.T) {
}) })
out.line(t, "type", "request") out.line(t, "type", "request")
// While the lines are sent, the process's lines give the instance name
// too.
line := out.line(t, "msg", "starting")
if line["instance"] != instance {
t.Errorf("start logged with instance %v, want %s", line["instance"], instance)
}
}
func TestBanIsAlertedAndAnAlertNotSentIsKeptAcrossARestart(t *testing.T) {
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) {
@@ -1120,7 +688,7 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) {
"SWWAF_REQUEST_MAX_BYTES": "100M", "SWWAF_REQUEST_MAX_BYTES": "100M",
"SWWAF_RESPONSE_MAX_BYTES": "5G", "SWWAF_RESPONSE_MAX_BYTES": "5G",
"SWWAF_ALLOW_NETS": "", "SWWAF_ALLOW_NETS": "",
rateLimitExemptNets: "", "SWWAF_RATE_LIMIT_EXEMPT_NETS": "",
"SWWAF_DENY_NETS": "", "SWWAF_DENY_NETS": "",
"SWWAF_RATE_LIMIT_PER_MINUTE": "1000", "SWWAF_RATE_LIMIT_PER_MINUTE": "1000",
"SWWAF_RATE_LIMIT_PER_HOUR": "10000", "SWWAF_RATE_LIMIT_PER_HOUR": "10000",
@@ -1238,70 +806,6 @@ func metricsText(t *testing.T, url, token string) string {
return string(body) return string(body)
} }
// metricsWith asks for the metrics at url with token until they hold each
// of series, as long as that takes, so that a slow test process cannot
// fail the test, and returns them.
func metricsWith(t *testing.T, url, token string, series ...string) string {
t.Helper()
for {
metrics := metricsText(t, url, token)
missing := slices.ContainsFunc(series, func(one string) bool {
return !strings.Contains(metrics, one)
})
if !missing {
return metrics
}
time.Sleep(pollInterval)
}
}
// seriesValue returns the value of series, such as
// name{instance="app"}, in metrics.
func seriesValue(t *testing.T, metrics, series string) float64 {
t.Helper()
for line := range strings.Lines(metrics) {
value, found := strings.CutPrefix(strings.TrimSpace(line), series+" ")
if !found {
continue
}
number, err := strconv.ParseFloat(value, 64)
if err != nil {
t.Fatalf("%s has the value %q: %v", series, value, err)
}
return number
}
t.Fatalf("no %s in the metrics:\n%s", series, metrics)
return 0
}
// rename renames the file at from to, replacing any file there.
func rename(t *testing.T, from, to string) {
t.Helper()
err := os.Rename(from, to)
if err != nil {
t.Fatalf("rename %s to %s: %v", from, to, err)
}
}
// writeLookupDatabase writes a lookup database at path that places the
// client placed in country, and no other address.
func writeLookupDatabase(t *testing.T, path, country string) {
t.Helper()
lookuptest.Write(t, path, map[string]lookuptest.Network{
placed + "/32": {ASN: "AS64496", ASName: "Example Net", Country: country},
})
}
// askAsAdmin sends a request with method to url, with body and // askAsAdmin sends a request with method to url, with body and
// adminSecret, and checks that it is answered 200. // adminSecret, and checks that it is answered 200.
func askAsAdmin(t *testing.T, method, url, body string) { func askAsAdmin(t *testing.T, method, url, body string) {
@@ -1413,140 +917,6 @@ func saveUntilAnswered(t *testing.T, path, content, url, from string, status int
} }
} }
// destination is a stand-in for SWWAF_ALERT_SLACK_WEBHOOK_URL or
// SWWAF_ALERT_NTFY_URL. It notes each request it is sent, and answers
// 200.
type destination struct {
url string
mu sync.Mutex
posts []destinationPost
}
// destinationPost is a request a destination was sent: its headers and
// its body.
type destinationPost struct {
header http.Header
body string
}
// startDestination starts a destination that takes every alert.
func startDestination(t *testing.T) *destination {
t.Helper()
d := &destination{}
server := httptest.NewServer(http.HandlerFunc(
func(_ http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
d.mu.Lock()
d.posts = append(d.posts, destinationPost{
header: r.Header.Clone(), body: string(body),
})
d.mu.Unlock()
}))
t.Cleanup(server.Close)
d.url = server.URL + "/alerts"
return d
}
// firstPost waits until the destination has been sent a request, and
// returns the first. It waits as long as that takes, so that a slow test
// process cannot fail the test.
func (d *destination) firstPost(t *testing.T) destinationPost {
t.Helper()
for {
d.mu.Lock()
if len(d.posts) > 0 {
post := d.posts[0]
d.mu.Unlock()
return post
}
d.mu.Unlock()
time.Sleep(pollInterval)
}
}
// webhook is a stand-in for SWWAF_ALERT_WEBHOOK_URL. It notes each alert
// 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 {
+48 -329
View File
@@ -1,16 +1,11 @@
// Package state keeps smallwebwaf's state in JSON files in // Package state keeps smallwebwaf's state in JSON files in
// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes: // SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes:
// bans.json holds the bans, clients.json each client's counters and // bans.json holds the bans, clients.json each client's counters and
// history, lookups.json GeoJS's answers, reputation.json the last try and // history, and lookups.json GeoJS's answers. Load reads them at start,
// last good copy of each list fetched from a URL and the DNSBL zones' // Watch takes in an admin's edit of one while smallwebwaf runs, and Run
// verdicts, and alerts.json the // and WriteAll write them. The disk is read and written outside the
// cooldowns, the hour under way, the alerts waiting for each destination // parts' locks, which are held only to take a snapshot or to put in what
// and the anomaly counters. Load // a file holds, so that no request waits on the disk.
// reads them at start, Watch takes in an admin's edit of one while
// smallwebwaf runs, and Run and WriteAll write them. The disk is read and
// written outside the parts' locks, which are held only to take a
// snapshot or to put in what a file holds, so that no request waits on
// the disk.
package state package state
import ( import (
@@ -22,22 +17,17 @@ 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/anomaly"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/reputation"
) )
// version is the version of the files' format, the only one read. // version is the version of the files' format, the only one read.
@@ -49,23 +39,16 @@ const fileMode = 0o600
// The state files' names. // The state files' names.
const ( const (
bansJSON = "bans.json" bansJSON = "bans.json"
clientsJSON = "clients.json" clientsJSON = "clients.json"
lookupsJSON = "lookups.json" lookupsJSON = "lookups.json"
reputationJSON = "reputation.json"
alertsJSON = "alerts.json"
) )
var ( var (
errVersion = errors.New("unknown version") errVersion = errors.New("unknown version")
// errMissing is for an entry without a field it needs. // errMissing is for an entry without a field it needs.
errMissing = errors.New("has no") errMissing = errors.New("has no")
errCause = errors.New("is not limit, attack or admin") errCause = errors.New("is not limit, attack or admin")
errDestination = errors.New("is not webhook, slack or ntfy")
errScope = errors.New("is not client, net, asn, total or watch")
errWaitingList = errors.New(`waiting is a list, but now lists the alerts by ` +
`destination: put the list under "webhook", as "waiting": {"webhook": [...]}, ` +
`or remove the file`)
) )
// Params are what Load needs. // Params are what Load needs.
@@ -77,16 +60,10 @@ type Params struct {
// is (SWWAF_STATE_COUNTER_INTERVAL). // is (SWWAF_STATE_COUNTER_INTERVAL).
WriteDelay time.Duration WriteDelay time.Duration
CounterInterval time.Duration CounterInterval time.Duration
// Ledger, Limiter, GeoJS, Lists, DNSBL, Alerts and Anomalies hold the // Ledger, Limiter and GeoJS hold the state.
// state. Alerts also receive a file_error alert for an edit set aside, Ledger *bans.Ledger
// and for a write that fails while smallwebwaf runs. Limiter *ratelimit.Limiter
Ledger *bans.Ledger GeoJS *lookup.GeoJS
Limiter *ratelimit.Limiter
GeoJS *lookup.GeoJS
Lists *reputation.Lists
DNSBL *reputation.DNSBL
Alerts *alerts.Queue
Anomalies *anomaly.Counters
// Now tells the time by which the counters' buckets run out, normally // Now tells the time by which the counters' buckets run out, normally
// time.Now in UTC. // time.Now in UTC.
Now func() time.Time Now func() time.Time
@@ -144,25 +121,6 @@ type lookupsFile struct {
Lookups []lookup.Answer `json:"lookups"` Lookups []lookup.Answer `json:"lookups"`
} }
// reputationFile is reputation.json, indented for an admin to read and
// edit, so that each line of a list's copy is on a line of its own.
type reputationFile struct {
Version int `json:"version"`
Lists []reputation.List `json:"lists"`
Verdicts []reputation.Verdict `json:"verdicts"`
}
// alertsFile is alerts.json, indented for an admin to read and edit.
//
//nolint:tagliatelle // the state files use snake_case, as the request log does
type alertsFile struct {
Version int `json:"version"`
Cooldowns []alerts.Cooldown `json:"cooldowns"`
Hour alerts.Hour `json:"hour"`
Waiting map[string][]alerts.Alert `json:"waiting"`
AnomalyCounters []anomaly.Counter `json:"anomaly_counters"`
}
// stateFile is the struct of a state file. Once the file is decoded, its // 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
@@ -173,10 +131,10 @@ type stateFile interface {
} }
// Load checks that files can be written in Dir, and reads the state files // Load checks that files can be written in Dir, and reads the state files
// in it into the parts of Params that hold the state. A missing file is // in it into the ledger, the limiter and GeoJS. A missing file is empty
// empty state, as on a first start. A file that does not parse, has an // state, as on a first start. A file that does not parse, has an unknown
// unknown version, or has an entry without a field it needs, is an error // version, or has an entry without a field it needs, is an error that
// that names the file and, where the JSON decoder tells it, the line and // names the file and, where the JSON decoder tells it, the line and
// column, or else the entry. // column, or else the entry.
func Load(params Params) (*Files, error) { func Load(params Params) (*Files, error) {
err := checkWritable(params.Dir) err := checkWritable(params.Dir)
@@ -189,26 +147,23 @@ func Load(params Params) (*Files, error) {
bansRead, bansErr := f.read(bansJSON) bansRead, bansErr := f.read(bansJSON)
clientsRead, clientsErr := f.read(clientsJSON) clientsRead, clientsErr := f.read(clientsJSON)
lookupsRead, lookupsErr := f.read(lookupsJSON) lookupsRead, lookupsErr := f.read(lookupsJSON)
reputationRead, reputationErr := f.read(reputationJSON)
alertsRead, alertsErr := f.read(alertsJSON)
err = errors.Join(bansErr, clientsErr, lookupsErr, reputationErr, alertsErr) err = errors.Join(bansErr, clientsErr, lookupsErr)
if err != nil { if err != nil {
return nil, err return nil, err
} }
params.ProcessLog.Info("read the state files", "directory", params.Dir, params.ProcessLog.Info("read the state files", "directory", params.Dir,
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead, "bans", bansRead, "clients", clientsRead, "lookups", lookupsRead)
"lists", reputationRead, "alerts_waiting", alertsRead)
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, raised as a file_error alert, and // done. A write that fails is logged, and the file is written again at
// the file is written again at its next write. Each write takes in an // its next write. Each write takes in an admin's edit of its file first,
// admin's edit of its file first, as writeFile describes. // 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()
@@ -226,13 +181,9 @@ func (f *Files) Run(ctx context.Context) {
case <-bansDue: case <-bansDue:
bansDue = nil bansDue = nil
f.logFailure(bansJSON, f.writeFile(bansJSON)) f.logFailure(f.writeFile(bansJSON))
case <-interval.C: case <-interval.C:
for _, name := range []string{ f.logFailure(f.WriteAll())
bansJSON, clientsJSON, lookupsJSON, reputationJSON, alertsJSON,
} {
f.logFailure(name, f.writeFile(name))
}
} }
} }
} }
@@ -241,7 +192,7 @@ func (f *Files) Run(ctx context.Context) {
// fails does not keep the others from being written. // fails does not keep the others from being written.
func (f *Files) WriteAll() error { func (f *Files) WriteAll() error {
return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON), return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON),
f.writeFile(lookupsJSON), f.writeFile(reputationJSON), f.writeFile(alertsJSON)) f.writeFile(lookupsJSON))
} }
// 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
@@ -276,7 +227,7 @@ func (f *Files) Watch(ctx context.Context) {
return return
case event := <-watcher.Events: case event := <-watcher.Events:
switch name := filepath.Base(event.Name); name { switch name := filepath.Base(event.Name); name {
case bansJSON, clientsJSON, lookupsJSON, reputationJSON, alertsJSON: case bansJSON, clientsJSON, lookupsJSON:
f.fileChanged(name) f.fileChanged(name)
} }
case err = <-watcher.Errors: case err = <-watcher.Errors:
@@ -286,22 +237,11 @@ func (f *Files) Watch(ctx context.Context) {
} }
} }
// logFailure logs a write of the state file name that failed, and raises // logFailure logs a write that failed.
// a file_error alert for it. func (f *Files) logFailure(err error) {
func (f *Files) logFailure(name string, err error) {
if err != nil { if err != nil {
const failed = "writing the state files failed" f.params.ProcessLog.Error("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())
} }
} }
@@ -422,28 +362,6 @@ func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
f.params.GeoJS.Load(file.Lookups) f.params.GeoJS.Load(file.Lookups)
entries = len(file.Lookups) entries = len(file.Lookups)
case reputationJSON:
var file reputationFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
err = f.params.Lists.Load(file.Lists)
if err != nil {
return 0, fmt.Errorf("%s: %w", path, err)
}
f.params.DNSBL.Load(file.Verdicts)
entries = len(file.Lists)
case alertsJSON:
waiting, err := f.takeInAlerts(path, data)
if err != nil {
return 0, err
}
entries = waiting
} }
f.sums[name] = sha256.Sum256(data) f.sums[name] = sha256.Sum256(data)
@@ -451,41 +369,6 @@ func (f *Files) takeIn(name string, data []byte, edit bool) (int, error) {
return entries, nil return entries, nil
} }
// takeInAlerts parses data, what alerts.json, at path, holds, puts it
// into the alerts and the anomaly counters, in place of what they held,
// and returns how many alerts wait in it, as takeIn describes.
func (f *Files) takeInAlerts(path string, data []byte) (int, error) {
// waiting was a list, of the alerts waiting for the webhook, before
// alerts went to Slack and ntfy too.
var written struct {
Waiting json.RawMessage `json:"waiting"`
}
if json.Unmarshal(data, &written) == nil &&
bytes.HasPrefix(written.Waiting, []byte("[")) {
return 0, fmt.Errorf("%s: %w", path, errWaitingList)
}
var file alertsFile
err := parse(path, data, &file)
if err != nil {
return 0, err
}
f.params.Alerts.Load(alerts.State{
Cooldowns: file.Cooldowns, Hour: file.Hour, Waiting: file.Waiting,
})
f.params.Anomalies.Load(file.AnomalyCounters, f.params.Now())
entries := 0
for _, waiting := range file.Waiting {
entries += len(waiting)
}
return entries, nil
}
// writeFile writes the state file name from what smallwebwaf holds. An // writeFile writes the state file name from what smallwebwaf holds. An
// edit made since smallwebwaf last read or wrote the file is taken in // edit made since smallwebwaf last read or wrote the file is taken in
// first, so that it is not overwritten, or set aside if it does not // first, so that it is not overwritten, or set aside if it does not
@@ -530,9 +413,8 @@ 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, and raises a file_error alert for it. If the // the file the error is. If the rename fails, the edit is left as it is,
// rename fails, the edit is left as it is, and the error returned is // and the error returned is parseErr joined with the rename's.
// 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)
@@ -541,16 +423,8 @@ func (f *Files) setAside(name string, parseErr error) error {
return errors.Join(parseErr, err) return errors.Join(parseErr, err)
} }
const setAside = "set aside an edit of a state file that does not parse" f.params.ProcessLog.Error("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
@@ -561,39 +435,21 @@ func (f *Files) setAside(name string, parseErr error) error {
func (f *Files) encode(name string) ([]byte, error) { func (f *Files) encode(name string) ([]byte, error) {
switch name { switch name {
case bansJSON: case bansJSON:
return encodeIndented(bansFile{ file := bansFile{Version: version, Bans: BanEntries(f.params.Ledger.Snapshot())}
Version: version, Bans: BanEntries(f.params.Ledger.Snapshot()),
}) data, err := json.MarshalIndent(file, "", " ")
if err != nil {
return nil, err
}
return append(data, '\n'), nil
case clientsJSON: case clientsJSON:
return encodeOnePerLine("clients", f.params.Limiter.Snapshot()) return encodeOnePerLine("clients", f.params.Limiter.Snapshot())
case lookupsJSON: default: // lookups.json
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot()) return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
case reputationJSON:
return encodeIndented(reputationFile{
Version: version, Lists: f.params.Lists.Snapshot(),
Verdicts: f.params.DNSBL.Snapshot(),
})
default: // alerts.json
held := f.params.Alerts.Snapshot()
return encodeIndented(alertsFile{
Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour,
Waiting: held.Waiting, AnomalyCounters: f.params.Anomalies.Snapshot(),
})
} }
} }
// encodeIndented encodes file, a state file's struct, indented for an
// admin to read and edit.
func encodeIndented(file any) ([]byte, error) {
data, err := json.MarshalIndent(file, "", " ")
if err != nil {
return nil, err
}
return append(data, '\n'), nil
}
// BanEntries returns held as bans.json lists them, an empty list for // BanEntries returns held as bans.json lists them, an empty list for
// none. // none.
func BanEntries(held []bans.Ban) []BanEntry { func BanEntries(held []bans.Ban) []BanEntry {
@@ -675,8 +531,8 @@ func (f *bansFile) check(data []byte) error {
} }
// check refuses a client without its address, which would count nobody's // check refuses a client without its address, which would count nobody's
// requests, or with requests or bytes in a window but no start, which // requests, or with requests in a window but no start, which would drop
// would drop them and give the client a fresh allowance. // them and give the client a fresh allowance.
func (f *clientsFile) check([]byte) error { func (f *clientsFile) check([]byte) error {
for i, client := range f.Clients { for i, client := range f.Clients {
switch { switch {
@@ -688,12 +544,6 @@ func (f *clientsFile) check([]byte) error {
return missing(i, "hour.start") return missing(i, "hour.start")
case countsWithoutStart(client.Day): case countsWithoutStart(client.Day):
return missing(i, "day.start") return missing(i, "day.start")
case countsWithoutStart(client.MinuteBytes):
return missing(i, "minute_bytes.start")
case countsWithoutStart(client.HourBytes):
return missing(i, "hour_bytes.start")
case countsWithoutStart(client.DayBytes):
return missing(i, "day_bytes.start")
} }
} }
@@ -731,139 +581,8 @@ func (f *lookupsFile) check(data []byte) error {
return nil return nil
} }
// check refuses a list without its URL, which would name no list, or the // countsWithoutStart reports whether b holds requests but no start, which
// time it was last tried, which would have it fetched at once, and a copy // places them in time.
// of it without the time it was fetched, or without its lines, which hold
// the list. It refuses a verdict without its zone or its client, which
// would be about no one, whether the zone lists the client, or the time
// it was fetched, which would drop it. A verdict's listed is false for a
// client the zone does not list, which Verdicts cannot tell from a
// missing one, so each listed is read again as written.
func (f *reputationFile) check(data []byte) error {
for i, kept := range f.Lists {
switch {
case kept.URL == "":
return missing(i, "url")
case kept.Tried.IsZero():
return missing(i, "tried")
case kept.Fetched.IsZero() && kept.Lines != nil:
return missing(i, "fetched")
case kept.Lines == nil && !kept.Fetched.IsZero():
return missing(i, "lines")
}
}
var written struct {
Verdicts []struct {
Listed *bool `json:"listed"`
} `json:"verdicts"`
}
err := json.Unmarshal(data, &written)
if err != nil {
return err
}
for i, verdict := range f.Verdicts {
switch {
case verdict.Zone == "":
return fmt.Errorf("verdicts %w", missing(i, "zone"))
case !verdict.Client.IsValid():
return fmt.Errorf("verdicts %w", missing(i, "client"))
case written.Verdicts[i].Listed == nil:
return fmt.Errorf("verdicts %w", missing(i, "listed"))
case verdict.Fetched.IsZero():
return fmt.Errorf("verdicts %w", missing(i, "fetched"))
}
}
return nil
}
// check refuses a cooldown without its event or when its alert was sent,
// which would hold back no repeat, alerts waiting for a destination with
// another name than webhook, slack or ntfy, most likely misspelt, an
// alert waiting without its event or its time, and an anomaly counter as
// checkAnomalyCounters does.
func (f *alertsFile) check([]byte) error {
for i, cooldown := range f.Cooldowns {
switch {
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 checkAnomalyCounters(f.AnomalyCounters)
}
// checkAnomalyCounters refuses an anomaly counter whose scope is not
// client, net, asn, total or watch, most likely misspelt, and one without
// a field it needs, as missingFromCounter tells.
func checkAnomalyCounters(counters []anomaly.Counter) error {
for i, counter := range counters {
if !slices.Contains(anomaly.Scopes(), counter.Scope) {
return fmt.Errorf("anomaly_counters entry %d's scope %q %w", i+1,
counter.Scope, errScope)
}
field := missingFromCounter(counter)
if field != "" {
return fmt.Errorf("anomaly_counters %w", missing(i, field))
}
}
return nil
}
// missingFromCounter returns the first field counter, an anomaly counter,
// needs and has not, or "" when it has them all: what tells it from the
// others in its scope, without which it would never be counted again, the
// netblock of a client, net or watch counter, the AS number of an asn one
// and the name of a watch one; and the start of a window in which it has
// requests or bytes, without which they would be dropped.
func missingFromCounter(counter anomaly.Counter) string {
scope := counter.Scope
switch {
case scope != anomaly.ScopeASN && scope != anomaly.ScopeTotal &&
!counter.Netblock.IsValid():
return "netblock"
case scope == anomaly.ScopeASN && counter.ASN == "":
return "asn"
case scope == anomaly.ScopeWatch && counter.Name == "":
return "name"
case countsWithoutStart(counter.Minute):
return "minute.start"
case countsWithoutStart(counter.Hour):
return "hour.start"
case countsWithoutStart(counter.MinuteBytes):
return "minute_bytes.start"
case countsWithoutStart(counter.HourBytes):
return "hour_bytes.start"
default:
return ""
}
}
// countsWithoutStart reports whether b holds requests, or bytes, but no
// start, which places them in time.
func countsWithoutStart(b ratelimit.Buckets) bool { 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)
} }
+35 -713
View File
@@ -10,10 +10,8 @@ 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"
@@ -21,31 +19,18 @@ import (
"testing/synctest" "testing/synctest"
"time" "time"
"sneak.berlin/go/smallwebwaf/internal/alerts"
"sneak.berlin/go/smallwebwaf/internal/anomaly"
"sneak.berlin/go/smallwebwaf/internal/bans" "sneak.berlin/go/smallwebwaf/internal/bans"
"sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/lookup"
"sneak.berlin/go/smallwebwaf/internal/metrics" "sneak.berlin/go/smallwebwaf/internal/metrics"
"sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/reputation"
"sneak.berlin/go/smallwebwaf/internal/state" "sneak.berlin/go/smallwebwaf/internal/state"
) )
const ( const (
// The state files. // The state files.
bansJSON = "bans.json" bansJSON = "bans.json"
clientsJSON = "clients.json" clientsJSON = "clients.json"
lookupsJSON = "lookups.json" lookupsJSON = "lookups.json"
reputationJSON = "reputation.json"
alertsJSON = "alerts.json"
// blocklistURL and torURL are the blocklists the tests' lists name, and
// dnsblZone the DNSBL zone of the tests' verdicts.
blocklistURL = "https://lists.example/drop.txt"
torURL = "https://lists.example/tor.txt"
dnsblZone = "dnsbl.example"
// The AS number and AS name the tests' clients are looked up in.
asn = "AS64496"
asName = "Example Net"
// 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"
@@ -53,9 +38,6 @@ 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.
@@ -69,8 +51,6 @@ 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",
@@ -105,150 +85,6 @@ 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
}
]
},
"anomaly_counters": [
{
"scope": "asn",
"asn": "AS64496",
"hour_bytes": {
"start": "2026-10-06T00:00:00Z",
"current": 8,
"previous": 0
}
},
{
"scope": "net",
"netblock": "203.0.113.0/24",
"minute": {
"start": "2026-10-06T00:00:00Z",
"current": 1,
"previous": 0
}
},
{
"scope": "total",
"minute": {
"start": "2026-10-06T00:00:00Z",
"current": 1,
"previous": 0
},
"minute_bytes": {
"start": "2026-10-06T00:00:00Z",
"current": 8,
"previous": 0
}
},
{
"scope": "watch",
"netblock": "203.0.113.0/24",
"name": "office",
"hour": {
"start": "2026-10-06T00:00:00Z",
"current": 1,
"previous": 0
}
}
]
}
`
// filledReputationJSON is reputation.json holding the blocklists' last
// tries and the copy of one, with its comment line, and two verdicts of a
// DNSBL zone, as fill puts them in.
const filledReputationJSON = `{
"version": 1,
"lists": [
{
"url": "https://lists.example/drop.txt",
"tried": "2026-10-06T00:00:00Z",
"fetched": "2026-10-05T23:00:00Z",
"lines": [
"; Spamhaus DROP List 2026/10/05 - (c) 2026 The Spamhaus Project SLL",
"203.0.113.0/24 ; SBL1",
"2001:db8::/32 ; SBL2"
]
},
{
"url": "https://lists.example/tor.txt",
"tried": "2026-10-06T00:00:00Z"
}
],
"verdicts": [
{
"zone": "dnsbl.example",
"client": "203.0.113.9",
"listed": true,
"fetched": "2026-10-05T23:00:00Z"
},
{
"zone": "dnsbl.example",
"client": "2001:db8::1",
"listed": false,
"fetched": "2026-10-05T22:00:00Z"
}
]
}
`
func TestFilesWrittenAndReadBack(t *testing.T) { func TestFilesWrittenAndReadBack(t *testing.T) {
t.Parallel() t.Parallel()
@@ -275,141 +111,13 @@ 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.Lists.Snapshot(), before.Lists.Snapshot(); !reflect.DeepEqual(
got, want) {
t.Errorf("%s read back\n%+v\nwant\n%+v", reputationJSON, got, want)
}
wantEqual(t, reputationJSON, after.DNSBL.Snapshot(), before.DNSBL.Snapshot())
if got, want := after.Alerts.Snapshot(), before.Alerts.Snapshot(); !reflect.DeepEqual(
got, want) {
t.Errorf("%s read back\n%+v\nwant\n%+v", alertsJSON, got, want)
}
wantEqual(t, alertsJSON, after.Anomalies.Snapshot(), before.Anomalies.Snapshot())
// Each one-per-line file lists its entries by client, and nothing // Each one-per-line file lists its entries by client, and nothing
// but the five files is left in the directory. // but the three 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, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON) wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
}
func TestReputationJSONKeepsEachCopyWholeOneLineToALine(t *testing.T) {
t.Parallel()
dir := t.TempDir()
params := newParams(dir)
fill(params)
files := load(t, params)
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
got := readFile(t, filepath.Join(dir, reputationJSON))
if got != filledReputationJSON {
t.Errorf("reputation.json\n%s\nwant\n%s", got, filledReputationJSON)
}
}
func TestAlertsJSONIsIndentedWithTheCooldownsTheHourAndTheAlertsWaiting(t *testing.T) {
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 TestAnomalyCountersKeptInAlertsJSONAcrossARestart(t *testing.T) {
t.Parallel()
dir := t.TempDir()
request := anomaly.Request{
Client: netip.MustParseAddr("203.0.113.9"),
ClientGroup: netip.MustParsePrefix("203.0.113.9/32"),
}
// The whole service may have two requests a minute.
withThreshold := func() state.Params {
params := newParams(dir)
params.Anomalies = anomaly.New(anomaly.Params{
Total: anomaly.Thresholds{RequestsPerMinute: 2}, Alerts: params.Alerts,
})
return params
}
before := withThreshold()
files := load(t, before)
for range 2 {
before.Anomalies.Count(midnight(), request)
}
err := files.WriteAll()
if err != nil {
t.Fatalf("write: %v", err)
}
// After the restart, the third request in the minute is over it.
after := withThreshold()
load(t, after)
after.Anomalies.Count(midnight(), request)
waiting := after.Alerts.Snapshot().Waiting[alerts.DestinationWebhook]
if len(waiting) != 1 || waiting[0].Event != alerts.EventAnomaly ||
waiting[0].Detail["count"] != float64(3) {
t.Errorf("alerts wait %+v, want one for 3 requests", waiting)
}
} }
func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) { func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) {
@@ -438,11 +146,8 @@ 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.Lists.Snapshot()) != 0 || len(params.GeoJS.Snapshot()) != 0 {
len(params.DNSBL.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")
} }
} }
@@ -484,31 +189,6 @@ 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`,
},
{
"an anomaly counter of an unknown scope", alertsJSON,
`{"version": 1, "anomaly_counters": [{"scope": "total"}, ` +
`{"scope": "nett", "netblock": "203.0.113.0/24"}]}`,
`: anomaly_counters entry 2's scope "nett" is not client, net, asn, total ` +
`or watch`,
},
{
"a copy of a list with a line that does not read", reputationJSON,
`{"version": 1, "lists": [{"url": "` + blocklistURL + `", ` +
`"tried": "2026-10-06T00:00:00Z", "fetched": "2026-10-06T00:00:00Z", ` +
`"lines": ["; DROP", "203.0.113.300"]}]}`,
`: the copy of ` + blocklistURL + `: line 2 is not an address or a netblock, ` +
`such as 192.0.2.0/24`,
},
} { } {
t.Run(tc.name, func(t *testing.T) { t.Run(tc.name, func(t *testing.T) {
t.Parallel() t.Parallel()
@@ -572,12 +252,6 @@ 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 + `}]}`,
@@ -604,167 +278,6 @@ func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
} }
} }
func TestReputationJSONEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) {
t.Parallel()
const (
drop = `"url": "` + blocklistURL + `", `
tried = `"tried": "2026-10-06T00:00:00Z", `
fetched = `"fetched": "2026-10-06T00:00:00Z"`
// verdictZone, verdictClient and listed start a verdict, which
// fetched ends.
verdictZone = `"zone": "` + dnsblZone + `", `
verdictClient = `"client": "198.51.100.7", `
listed = `"listed": false, `
)
for _, tc := range []struct {
name, content string
// want is what the error says after the file's path.
want string
}{
{
"a list without its URL",
`{"version": 1, "lists": [{` + tried + fetched + `, "lines": []}]}`,
`: entry 1 has no "url"`,
},
{
"a list without the time it was last tried",
`{"version": 1, "lists": [{` + drop + fetched + `, "lines": []}]}`,
`: entry 1 has no "tried"`,
},
{
"a copy of a list without the time it was fetched",
`{"version": 1, "lists": [{` + drop + tried + `"lines": []}]}`,
`: entry 1 has no "fetched"`,
},
{
// An empty list has no lines, which is not having none.
"a copy of a list without its lines",
`{"version": 1, "lists": [{` + drop + tried + fetched + `, "lines": []}, ` +
`{"url": "` + torURL + `", ` + tried + fetched + `}]}`,
`: entry 2 has no "lines"`,
},
{
"a verdict without its zone",
`{"version": 1, "verdicts": [{` + verdictClient + listed + fetched + `}]}`,
`: verdicts entry 1 has no "zone"`,
},
{
"a verdict without its client",
`{"version": 1, "verdicts": [{` + verdictZone + listed + fetched + `}]}`,
`: verdicts entry 1 has no "client"`,
},
{
// A client the zone does not list has a listed of false, which is
// not having none.
"a verdict without whether the zone lists the client",
`{"version": 1, "verdicts": [{` + verdictZone + verdictClient + listed +
fetched + `}, {` + verdictZone + verdictClient + fetched + `}]}`,
`: verdicts entry 2 has no "listed"`,
},
{
"a verdict without the time it was fetched",
`{"version": 1, "verdicts": [{` + verdictZone + verdictClient +
`"listed": true}]}`,
`: verdicts entry 1 has no "fetched"`,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
wantRefused(t, reputationJSON, tc.content, tc.want)
})
}
}
func 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"`,
},
{
// The whole service's counter needs nothing to tell it apart.
"an anomaly counter of a netblock without it",
`{"version": 1, "anomaly_counters": [{"scope": "total"}, {"scope": "net"}]}`,
`: anomaly_counters entry 2 has no "netblock"`,
},
{
"an anomaly counter of a client without its netblock",
`{"version": 1, "anomaly_counters": [{"scope": "client"}]}`,
`: anomaly_counters entry 1 has no "netblock"`,
},
{
"an anomaly counter of an AS number without it",
`{"version": 1, "anomaly_counters": [{"scope": "asn"}]}`,
`: anomaly_counters entry 1 has no "asn"`,
},
{
"an anomaly counter of a named netblock without its name",
`{"version": 1, "anomaly_counters": [` +
`{"scope": "watch", "netblock": "203.0.113.0/24"}]}`,
`: anomaly_counters entry 1 has no "name"`,
},
{
"an anomaly counter of a named netblock without its netblock",
`{"version": 1, "anomaly_counters": [{"scope": "watch", "name": "office"}]}`,
`: anomaly_counters entry 1 has no "netblock"`,
},
{
"an anomaly counter with bytes in a window without its start",
`{"version": 1, "anomaly_counters": [` +
`{"scope": "total", "hour_bytes": {"current": 5}}]}`,
`: anomaly_counters entry 1 has no "hour_bytes.start"`,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
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()
@@ -781,9 +294,7 @@ func TestBanWithAnotherCauseStopsTheStart(t *testing.T) {
func TestUnknownVersionStopsTheStart(t *testing.T) { func TestUnknownVersionStopsTheStart(t *testing.T) {
t.Parallel() t.Parallel()
for _, file := range []string{ for _, file := range []string{bansJSON, clientsJSON, lookupsJSON} {
bansJSON, clientsJSON, lookupsJSON, reputationJSON, 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()
@@ -833,12 +344,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)
@@ -885,9 +396,8 @@ func TestEveryFileWrittenEveryCounterInterval(t *testing.T) {
time.Sleep(time.Nanosecond) time.Sleep(time.Nanosecond)
synctest.Wait() synctest.Wait()
wantFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON) wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
removeFiles(t, dir, alertsJSON, bansJSON, clientsJSON, lookupsJSON, removeFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
reputationJSON)
} }
}) })
} }
@@ -930,57 +440,6 @@ 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()
@@ -1003,7 +462,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(), whole) params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight())
err = files.WriteAll() err = files.WriteAll()
if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") { if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") {
@@ -1037,8 +496,8 @@ func TestWritesAreCountedInTheMetrics(t *testing.T) {
} }
const ( const (
ofBans = `{file="bans.json",instance="app"}` ofBans = `{file="bans.json"}`
ofClients = `{file="clients.json",instance="app"}` ofClients = `{file="clients.json"}`
) )
got := scrape(t, params) got := scrape(t, params)
@@ -1106,7 +565,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, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON) wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON)
wantWriteFailed(t, params, bansJSON) wantWriteFailed(t, params, bansJSON)
} }
@@ -1178,43 +637,6 @@ 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()}})
edit(t, dir, reputationJSON, `{"version": 1, "lists": [{"url": "`+blocklistURL+`", `+
`"tried": "2026-10-06T00:00:00Z", "fetched": "2026-10-06T00:00:00Z", `+
`"lines": ["198.51.100.7"]}], "verdicts": [{"zone": "`+dnsblZone+`", `+
`"client": "198.51.100.7", "listed": true, "fetched": "2026-10-06T00:00:00Z"}]}`)
wantTakenIn(t, lines, dir, reputationJSON)
listedBy := params.Lists.ListedBy(client.Addr())
if !slices.Equal(listedBy, []string{blocklistURL}) {
t.Errorf("%s taken in lists %s on %v, want on the blocklist", reputationJSON,
client.Addr(), listedBy)
}
wantEqual(t, reputationJSON, params.DNSBL.Snapshot(), []reputation.Verdict{{
Zone: dnsblZone, Client: client.Addr(), Listed: true, Fetched: midnight(),
}})
// 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) {
@@ -1286,7 +708,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")
} }
@@ -1295,7 +717,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")
} }
@@ -1385,7 +807,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")
} }
@@ -1422,7 +844,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, alertsJSON, bansJSON, clientsJSON, lookupsJSON, reputationJSON) wantFiles(t, dir, 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.
@@ -1439,15 +861,7 @@ func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) {
t.Errorf("set aside with %v", line) t.Errorf("set aside with %v", line)
} }
// It is raised as a file_error alert, with the same file and error. wantFiles(t, dir, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON)
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,
reputationJSON)
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)
@@ -1480,7 +894,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",instance="app"}`, 2) `smallwebwaf_state_file_edits_taken_in_total{file="bans.json"}`, 2)
} }
func TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) { func TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) {
@@ -1545,7 +959,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",instance="app"}`, 1) `smallwebwaf_state_file_edits_set_aside_total{file="bans.json"}`, 1)
} }
func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) { func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) {
@@ -1576,21 +990,10 @@ 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, the lists, two blocklists, are // hold nothing yet. GeoJS is never asked.
// never fetched, and the alerts, at most two an hour, are
// never sent. The anomaly counters count the scopes fill counts, with
// thresholds fill does not reach.
func newParams(dir string) state.Params { func newParams(dir string) state.Params {
discard := slog.New(slog.DiscardHandler) discard := slog.New(slog.DiscardHandler)
m := metrics.New(1, "app") m := metrics.New(1)
queue := alerts.New(alerts.Params{
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
Events: alerts.Events(),
Cooldown: 15 * time.Minute,
MaxPerHour: 2,
Instance: "fsn1app1/gitea",
Now: midnight,
})
return state.Params{ return state.Params{
Dir: dir, Dir: dir,
@@ -1607,118 +1010,39 @@ 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,
}), }),
Lists: reputation.New(reputation.Params{
BlocklistURLs: []string{blocklistURL, torURL}, Refresh: 24 * time.Hour,
Now: midnight, ProcessLog: discard, Alerts: queue,
}),
DNSBL: reputation.NewDNSBL(reputation.DNSBLParams{
Zones: []string{dnsblZone}, CacheTTL: 24 * time.Hour, Timeout: time.Second,
Now: midnight, ProcessLog: discard, Alerts: queue,
}),
Alerts: queue,
Anomalies: anomaly.New(anomaly.Params{
Net: anomaly.Thresholds{RequestsPerMinute: 1000},
ASN: anomaly.Thresholds{BytesPerHour: 1 << 30},
Total: anomaly.Thresholds{RequestsPerMinute: 1000, BytesPerMinute: 1 << 30},
Watch: anomaly.Thresholds{RequestsPerHour: 1000},
NetV4Prefix: 24,
NetV6Prefix: 48,
NamedNetblocks: []anomaly.NamedNetblock{{Name: "office", Netblock: office()}},
Alerts: queue,
}),
Now: midnight, Now: midnight,
ProcessLog: discard, ProcessLog: discard,
Metrics: m, Metrics: m,
} }
} }
// office is the named netblock of the anomaly counters of newParams.
func office() netip.Prefix {
return netip.MustParsePrefix("203.0.113.0/24")
}
// fill puts a permanent ban an admin made, a ban for a broken limit and // 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, // one for a clear sign of attack, clients with counts and histories, and
// GeoJS answers, the blocklists' last tries and the copy of one, and two // GeoJS answers into the parts of params.
// verdicts of a DNSBL zone, as filledReputationJSON holds them, and alerts
// and anomaly counters, 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{ params.Ledger.BanForLimit(client, now, bans.Notes{Country: "DE", Limit: 1})
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, whole) params.Limiter.Count(netip.MustParsePrefix(c), now)
} }
params.Limiter.CountBytes(client, now, 8, whole)
params.Limiter.AddToHistory(client, now, ratelimit.Request{ params.Limiter.AddToHistory(client, now, ratelimit.Request{
Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5, Country: "DE", 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),
}, },
}) })
// The copy of drop.txt was fetched an hour ago, and the fetches of it
// and of tor.txt tried since failed.
err := params.Lists.Load([]reputation.List{{
URL: blocklistURL, Tried: now, Fetched: now.Add(-time.Hour),
Lines: []string{
"; Spamhaus DROP List 2026/10/05 - (c) 2026 The Spamhaus Project SLL",
"203.0.113.0/24 ; SBL1",
"2001:db8::/32 ; SBL2",
},
}, {URL: torURL, Tried: now}})
if err != nil {
panic(err) // the copy reads
}
params.DNSBL.Load([]reputation.Verdict{
{
Zone: dnsblZone, Client: netip.MustParseAddr("2001:db8::1"),
Fetched: now.Add(-2 * time.Hour),
},
{Zone: dnsblZone, Client: client.Addr(), Listed: true, Fetched: now.Add(-time.Hour)},
})
// 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",
})
params.Anomalies.Count(now, anomaly.Request{
Client: client.Addr(), ClientGroup: client, ASN: asn, Bytes: 8,
})
} }
// permanentBan is the ban permanentBansJSON holds. // permanentBan is the ban permanentBansJSON holds.
@@ -1729,8 +1053,6 @@ 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",
@@ -1776,13 +1098,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))
} }
@@ -2028,8 +1350,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",instance="app"}, or // smallwebwaf_state_file_writes_total{file="bans.json"}, or fails the test
// fails the test if there is no such series. // 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()
@@ -2058,7 +1380,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 + `",instance="app"}` file := `{file="` + name + `"}`
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)
+5 -7
View File
@@ -7,10 +7,9 @@
# request bans for good, that `sv stop` stops smallwebwaf in order, that # request bans for good, that `sv stop` stops smallwebwaf in order, that
# `docker stop` stops the container without having to kill it, and that # `docker stop` stops the container without having to kill it, and that
# a new container on the same volume still refuses the banned client. The # a new container on the same volume still refuses the banned client. The
# containers run with SWWAF_LOOKUP_SOURCE=off, so that no address is sent # containers, the volume and both images are removed however the script
# to GeoJS. The containers, the volume and both images are removed however # ends. Building the app needs network access, for nixpkgs' binary cache.
# the script ends. Building the app needs network access, for nixpkgs' # script/check does not run this.
# binary cache. script/check does not run this.
set -eu set -eu
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
@@ -64,13 +63,12 @@ logged() {
} }
# start_container: run the app's container, with the state files on the # start_container: run the app's container, with the state files on the
# volume, a rate limit of one request a minute and no client looked up, # volume and a rate limit of one request a minute, and wait until it is
# and wait until it is healthy. # healthy.
start_container() { start_container() {
docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \ docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \
--volume "$VOLUME:/var/lib/smallwebwaf" \ --volume "$VOLUME:/var/lib/smallwebwaf" \
--env SWWAF_RATE_LIMIT_PER_MINUTE=1 \ --env SWWAF_RATE_LIMIT_PER_MINUTE=1 \
--env SWWAF_LOOKUP_SOURCE=off \
"$APP_IMAGE" >/dev/null "$APP_IMAGE" >/dev/null
wait_for "the health check did not pass" healthy wait_for "the health check did not pass" healthy
address="$(docker port "$CONTAINER" 8080/tcp)" address="$(docker port "$CONTAINER" 8080/tcp)"
-19
View File
@@ -1,19 +0,0 @@
#!/bin/sh
# script/tidy: write go.mod and go.sum as `go mod tidy` writes them, which
# the test phase of the Dockerfile checks. This builds the Dockerfile's
# tidy-files stage, which holds the two files alone, and --output writes
# them into the working tree. The build makes no image, so it has no tag.
# --no-cache because `go mod tidy` asks the module proxy, whose answers a
# cached layer would repeat.
set -eu
ROOT="$(cd "$(dirname "$0")/.." && pwd -P)"
main() {
cd "$ROOT"
docker build --no-cache \
--target tidy-files \
--output "type=local,dest=$ROOT" .
}
main "$@"